From a6fbd8f011cb114bc85389205460f720efdc9624 Mon Sep 17 00:00:00 2001 From: Ming-Jer Lee Date: Thu, 27 Aug 2026 15:20:11 -0700 Subject: [PATCH] feat: add clgraph detect and clgraph init MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Step 2 of the plugin work. The dialect cannot be guessed safely, so it has to be asked — and an MCP server on stdio has no channel to ask through. Asking therefore lives in the CLI and, later, in /clgraph:setup. clgraph detect proposes a SQL directory and a dialect with the evidence attached, ranked: a dbt profile's adapter type (authoritative, high confidence), then dialect-specific syntax markers, then parse scoring. It writes nothing and decides nothing. --json makes it agent-readable. clgraph init records the answer in clgraph.toml, or [tool.clgraph] with --into-pyproject, and adds .clgraph/ to .gitignore. Run bare it prompts with detection pre-filling the defaults; with --yes it never prompts and fails instead, which is how agents and CI invoke it. A missing dialect is an error either way. Two things testing turned up, both fixed here: Parse scoring counted only raised ParseErrors, but sqlglot far more often degrades unsupported syntax into an exp.Command node without erroring — so a wrong dialect scored as a clean parse. Both now count as failures. Even so, scoring stays weak: Snowflake-only syntax like QUALIFY and IFF() parses cleanly under all eight candidates. TestParseScoringIsWeak pins that, since it is the concrete reason this design refuses to guess. sqlglot logged warnings while probing wrong dialects, and in any process with logging configured to stdout those landed inside the JSON that `clgraph detect --json` emits. Scoring now silences sqlglot and restores the prior level. --- CHANGELOG.md | 22 ++ README.md | 41 ++- src/clgraph/cli.py | 172 ++++++++++++ src/clgraph/config_writer.py | 252 ++++++++++++++++++ src/clgraph/detect.py | 496 +++++++++++++++++++++++++++++++++++ tests/test_cli.py | 129 +++++++++ tests/test_config_writer.py | 223 ++++++++++++++++ tests/test_detect.py | 329 +++++++++++++++++++++++ 8 files changed, 1661 insertions(+), 3 deletions(-) create mode 100644 src/clgraph/config_writer.py create mode 100644 src/clgraph/detect.py create mode 100644 tests/test_config_writer.py create mode 100644 tests/test_detect.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 8c14df0..dd7dce1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -19,6 +19,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 directly instead of going through `python -m clgraph.mcp`. - `clgraph.toml` / `[tool.clgraph]` project configuration, with keys `sql_dir` and `dialect`. +- `clgraph detect` - proposes which directory holds your SQL and which + dialect it is, with the evidence and a confidence rating attached. + Evidence is ranked: a dbt profile's adapter `type:` (high confidence), + then dialect-specific syntax markers, then parse scoring. It writes + nothing and decides nothing. `--json` for machine-readable output. +- `clgraph init` - records the answer in `clgraph.toml`, or in + `[tool.clgraph]` with `--into-pyproject`, and adds `.clgraph/` to + `.gitignore`. Prompts when run bare, using detection to pre-fill; with + `--yes` it never prompts and fails instead, which is how agents and CI + invoke it. A missing dialect is an error, never a guess. +- `clgraph.detect` and `clgraph.config_writer` public API: `detect()`, + `Detection`, `Confidence`, `write_config()`, `known_dialects()`. ### Changed @@ -36,6 +48,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 in clients as a failed server with no explanation. - `--pipeline` is now optional for `clgraph mcp` and `clgraph-mcp`. +### Fixed + +- Dialect parse scoring counted only raised `ParseError`s, but sqlglot + more often degrades unsupported syntax into an `exp.Command` node + without erroring. Both now count as a failed parse, so scoring can + actually distinguish dialects. +- sqlglot warnings emitted while probing candidate dialects no longer + reach stdout, where they corrupted `clgraph detect --json` for any + caller parsing it. + ## [0.0.8] - 2026-08-11 ### Added diff --git a/README.md b/README.md index 3394fe0..2e3df6a 100644 --- a/README.md +++ b/README.md @@ -775,8 +775,43 @@ pip install 'clgraph[mcp]' #### Project Configuration +Run `clgraph init` and it will find your SQL, propose a dialect with its +reasoning, and write the config: + +```bash +clgraph init +``` + +To see what it would propose without writing anything: + +```bash +clgraph detect +``` + +``` +SQL directories: + models (34 files, dbt) + +Dialect: snowflake (confidence: high) + - profiles.yml declares adapter type 'snowflake' +``` + +Confidence is the part to read. `high` comes from a dbt profile's own +adapter type and can be accepted as-is; anything lower is a real question, +because **the dialect is never guessed** — sqlglot parses most of a corpus +under the wrong grammar without complaining, and the resulting lineage +graph is plausible and wrong. + +For scripted or CI use, pass everything explicitly and it will never +prompt: + +```bash +clgraph init --sql-dir models/ --dialect snowflake --yes +``` + The server works out what to index from the project itself, so hosts can -launch it with no arguments. Put this in `clgraph.toml` at your project root: +launch it with no arguments. `clgraph init` writes this for you, or write +it by hand at your project root: ```toml @@ -799,8 +834,8 @@ what to do instead of showing a dead connection. **The dialect is never guessed.** Pointing clgraph at SQL files without saying which dialect they are is an error, not a silent fall back to -BigQuery — sqlglot parses most of a corpus under the wrong grammar without -complaining, and the resulting lineage graph is plausible and wrong. +BigQuery. Run `clgraph init` to be asked once and have the answer +recorded. #### Claude Desktop Configuration diff --git a/src/clgraph/cli.py b/src/clgraph/cli.py index c276cd5..a1fc2cc 100644 --- a/src/clgraph/cli.py +++ b/src/clgraph/cli.py @@ -255,6 +255,178 @@ def _print_diff_summary(diff_result): console.print(f" [yellow]~ {cd.full_name} ({cd.field_name})[/yellow]") +@app.command() +def detect( + path: Annotated[ + Path, + typer.Argument(help="Project directory to inspect"), + ] = Path("."), + as_json: Annotated[ + bool, + typer.Option("--json", help="Emit machine-readable JSON"), + ] = False, +): + """Propose which SQL to index and which dialect it is. + + Reports evidence and a confidence level; writes nothing and decides + nothing. `clgraph init` records the answer. + + Confidence is what matters: "high" comes from a dbt profile's own + adapter type and can be accepted as-is. Anything lower is a genuine + question, because the wrong dialect produces a lineage graph that + looks right and is not. + """ + from clgraph.detect import detect as run_detection + + if not path.exists(): + typer.echo(f"Error: path does not exist: {path}", err=True) + raise typer.Exit(code=1) + + result = run_detection(path) + + if as_json: + typer.echo(json.dumps(result.to_dict(), indent=2)) + return + + _print_detection(result) + + +def _print_detection(result) -> None: + """Render a Detection for a human reader.""" + if not result.sql_dirs: + typer.echo("No SQL files found.") + else: + typer.echo("SQL directories:") + for candidate in result.sql_dirs: + relative = candidate.to_dict(result.root)["path"] + typer.echo(f" {relative} ({candidate.file_count} files, {candidate.kind})") + + guess = result.dialect + typer.echo("") + if guess.best is None: + typer.echo("Dialect: no evidence found.") + else: + typer.echo(f"Dialect: {guess.best} (confidence: {guess.confidence})") + for reason in guess.evidence: + typer.echo(f" - {reason}") + + if guess.confidence != "high": + typer.echo("") + typer.echo("Confirm the dialect before indexing; clgraph will not guess it.") + + +@app.command() +def init( + sql_dir: Annotated[ + Optional[str], + typer.Option("--sql-dir", help="Directory of SQL files, or a JSON pipeline"), + ] = None, + dialect: Annotated[ + Optional[str], + typer.Option(help="SQL dialect. Required for SQL files; never inferred."), + ] = None, + project_dir: Annotated[ + Path, + typer.Option("--project-dir", help="Project root to configure"), + ] = Path("."), + into_pyproject: Annotated[ + bool, + typer.Option("--into-pyproject", help="Write [tool.clgraph] instead of clgraph.toml"), + ] = False, + force: Annotated[ + bool, + typer.Option("--force", help="Overwrite existing clgraph configuration"), + ] = False, + yes: Annotated[ + bool, + typer.Option("--yes", "-y", help="Never prompt; fail instead of asking"), + ] = False, +): + """Configure this project for clgraph. + + Run bare, it detects candidates and asks. With --yes it never prompts, + which is how agents and CI invoke it: everything must be passed + explicitly, and a missing dialect is an error rather than a guess. + """ + from clgraph.config_writer import write_config + from clgraph.detect import detect as run_detection + + detection = run_detection(project_dir) + + resolved_sql_dir = sql_dir or _choose_sql_dir(detection, yes) + resolved_dialect = dialect or _choose_dialect(detection, resolved_sql_dir, yes) + + try: + result = write_config( + project_dir, + sql_dir=resolved_sql_dir, + dialect=resolved_dialect, + into_pyproject=into_pyproject, + force=force, + ) + except (ValueError, FileExistsError) as err: + typer.echo(f"Error: {err}", err=True) + raise typer.Exit(code=1) from err + + typer.echo(f"Wrote {result.describe()}") + if result.gitignore_updated: + typer.echo("Added .clgraph/ to .gitignore") + typer.echo("Next: clgraph index") + + +def _choose_sql_dir(detection, yes: bool) -> str: + """Pick the SQL directory, prompting only when allowed.""" + candidates = detection.sql_dirs + + if yes: + if not candidates: + typer.echo( + "Error: no SQL directory found and --sql-dir was not given.", + err=True, + ) + raise typer.Exit(code=1) + return candidates[0].to_dict(detection.root)["path"] + + if not candidates: + return typer.prompt("Where are your SQL files?") + + default = candidates[0].to_dict(detection.root)["path"] + typer.echo("SQL directories found:") + for candidate in candidates: + relative = candidate.to_dict(detection.root)["path"] + typer.echo(f" {relative} ({candidate.file_count} files, {candidate.kind})") + return typer.prompt("Which directory holds your SQL?", default=default) + + +def _choose_dialect(detection, sql_dir: str, yes: bool) -> Optional[str]: + """Pick the dialect, prompting only when allowed. + + A JSON pipeline carries its own dialect, so it is never asked about. + """ + if sql_dir.endswith(".json"): + return None + + guess = detection.dialect + + if yes: + typer.echo( + "Error: --dialect is required with --yes. clgraph does not guess: " + "the wrong dialect yields a plausible but incorrect lineage graph.", + err=True, + ) + if guess.best: + typer.echo(f"Detection suggests: {guess.best} ({guess.confidence})", err=True) + raise typer.Exit(code=1) + + if guess.best: + typer.echo(f"Detected dialect: {guess.best} (confidence: {guess.confidence})") + for reason in guess.evidence: + typer.echo(f" - {reason}") + return typer.prompt("Which SQL dialect?", default=guess.best) + + return typer.prompt("Which SQL dialect?") + + @app.command() def mcp( pipeline: Annotated[ diff --git a/src/clgraph/config_writer.py b/src/clgraph/config_writer.py new file mode 100644 index 0000000..b2d590a --- /dev/null +++ b/src/clgraph/config_writer.py @@ -0,0 +1,252 @@ +""" +Record the setup answers as project configuration. + +This is the counterpart to :mod:`clgraph.mcp.config`, which reads what +this writes. Detection proposes, a human confirms, and this module makes +the answer durable so nothing downstream ever has to guess again. + +Two surfaces are supported because plenty of dbt and warehouse repos have +no ``pyproject.toml`` at all: ``clgraph.toml`` at the project root is the +default and what the docs show, and ``[tool.clgraph]`` is available for +Python projects that would rather keep configuration in one file. +""" + +import logging +import re +from dataclasses import dataclass +from pathlib import Path +from typing import FrozenSet, Optional + +from sqlglot.dialects.dialect import Dialect + +from .mcp.config import CACHE_RELPATH, CONFIG_FILENAME, PYPROJECT_FILENAME + +logger = logging.getLogger(__name__) + +GITIGNORE_FILENAME = ".gitignore" +GITIGNORE_ENTRY = ".clgraph/" +GITIGNORE_COMMENT = "# clgraph lineage index (rebuild with: clgraph index)" + +PYPROJECT_TABLE = "[tool.clgraph]" + +_CLGRAPH_TABLE_RE = re.compile( + r"^\[tool\.clgraph\]\s*$.*?(?=^\[|\Z)", + re.MULTILINE | re.DOTALL, +) + + +@dataclass(frozen=True) +class InitResult: + """What ``write_config`` did, so callers can report it accurately.""" + + config_path: Path + target: str + sql_dir: str + dialect: Optional[str] + gitignore_updated: bool + + def describe(self) -> str: + """One-line summary for CLI and agent output.""" + dialect = self.dialect or "dialect from pipeline file" + return f"{self.target}: sql_dir={self.sql_dir}, {dialect}" + + +def write_config( + project_dir: Path, + sql_dir: str, + dialect: Optional[str], + *, + into_pyproject: bool = False, + force: bool = False, +) -> InitResult: + """Write project configuration and ignore the index directory. + + Args: + project_dir: Project root. Config is written here. + sql_dir: SQL directory or JSON pipeline, relative to project_dir. + An absolute path inside the project is stored relative. + dialect: SQL dialect. Required unless sql_dir is a .json pipeline, + which carries its own. + into_pyproject: Write ``[tool.clgraph]`` instead of clgraph.toml. + force: Overwrite existing clgraph configuration. + + Returns: + An InitResult describing what was written. + + Raises: + ValueError: If the dialect is unknown or sql_dir does not exist. + Validation runs before anything is written, so a rejected call + leaves the project untouched. + FileExistsError: If configuration already exists and force is not + set. + """ + root = Path(project_dir) + relative_sql_dir = _validate(root, sql_dir, dialect) + + if into_pyproject: + config_path = _write_pyproject(root, relative_sql_dir, dialect, force=force) + target = PYPROJECT_FILENAME + else: + config_path = _write_clgraph_toml(root, relative_sql_dir, dialect, force=force) + target = CONFIG_FILENAME + + return InitResult( + config_path=config_path, + target=target, + sql_dir=relative_sql_dir, + dialect=dialect, + gitignore_updated=ensure_gitignore(root), + ) + + +# ============================================================================= +# Validation +# ============================================================================= + + +def _validate(root: Path, sql_dir: str, dialect: Optional[str]) -> str: + """Check inputs and normalise sql_dir. Writes nothing.""" + target = _resolve_sql_dir(root, sql_dir) + + if not target.exists(): + raise ValueError( + f"sql_dir does not exist: {target}\n" + f"Run 'clgraph detect' to see which directories hold SQL." + ) + + if target.suffix != ".json": + _validate_dialect(dialect) + + return _relative_sql_dir(root, target) + + +def _validate_dialect(dialect: Optional[str]) -> None: + """Reject anything sqlglot cannot actually parse with.""" + if not dialect: + raise ValueError( + "A dialect is required for SQL files. clgraph does not guess: " + "the wrong dialect yields a plausible but incorrect lineage graph.\n" + f"Known dialects include: {_dialect_examples()}" + ) + + if dialect.lower() not in known_dialects(): + raise ValueError( + f"Unknown dialect: {dialect!r}\nKnown dialects include: {_dialect_examples()}" + ) + + +def known_dialects() -> FrozenSet[str]: + """Every dialect name sqlglot recognises.""" + return frozenset(name for name in Dialect.classes if name) + + +def _dialect_examples() -> str: + """A readable subset for error messages, not the full registry.""" + common = ("bigquery", "snowflake", "postgres", "redshift", "duckdb", "databricks", "spark") + available = known_dialects() + return ", ".join(name for name in common if name in available) + ", ..." + + +def _resolve_sql_dir(root: Path, sql_dir: str) -> Path: + candidate = Path(sql_dir.rstrip("/\\")).expanduser() + return candidate if candidate.is_absolute() else root / candidate + + +def _relative_sql_dir(root: Path, target: Path) -> str: + """Store paths relative to the project so the config stays portable.""" + try: + return target.resolve().relative_to(root.resolve()).as_posix() + except ValueError: + return target.as_posix() + + +# ============================================================================= +# Writers +# ============================================================================= + + +def _write_clgraph_toml(root: Path, sql_dir: str, dialect: Optional[str], *, force: bool) -> Path: + config_path = root / CONFIG_FILENAME + + if config_path.exists() and not force: + raise FileExistsError(f"{config_path} already exists. Pass --force to overwrite it.") + + config_path.write_text(_render_toml(sql_dir, dialect)) + logger.info("Wrote %s", config_path) + return config_path + + +def _render_toml(sql_dir: str, dialect: Optional[str]) -> str: + lines = [ + "# clgraph project configuration", + "# Written by `clgraph init`. Read by the CLI and the MCP server.", + "", + f'sql_dir = "{sql_dir}"', + ] + if dialect: + lines.append(f'dialect = "{dialect}"') + return "\n".join(lines) + "\n" + + +def _write_pyproject(root: Path, sql_dir: str, dialect: Optional[str], *, force: bool) -> Path: + config_path = root / PYPROJECT_FILENAME + existing = config_path.read_text() if config_path.exists() else "" + + if _CLGRAPH_TABLE_RE.search(existing): + if not force: + raise FileExistsError( + f"{config_path} already has a {PYPROJECT_TABLE} table. Pass --force to replace it." + ) + updated = _CLGRAPH_TABLE_RE.sub(_render_pyproject_table(sql_dir, dialect), existing) + else: + separator = "" if not existing or existing.endswith("\n\n") else "\n" + updated = existing + separator + _render_pyproject_table(sql_dir, dialect) + + config_path.write_text(updated) + logger.info("Wrote %s to %s", PYPROJECT_TABLE, config_path) + return config_path + + +def _render_pyproject_table(sql_dir: str, dialect: Optional[str]) -> str: + lines = [PYPROJECT_TABLE, f'sql_dir = "{sql_dir}"'] + if dialect: + lines.append(f'dialect = "{dialect}"') + return "\n".join(lines) + "\n" + + +# ============================================================================= +# .gitignore +# ============================================================================= + + +def ensure_gitignore(root: Path) -> bool: + """Ignore the index directory, exactly once. + + The index is a build artifact: it goes stale on every SQL edit, and a + committed one shows up as a spurious diff in every branch. CI rebuilds + it with ``clgraph index --check``. + + Returns: + True if the entry was added, False if it was already covered. + """ + gitignore = root / GITIGNORE_FILENAME + existing = gitignore.read_text() if gitignore.exists() else "" + + if _already_ignored(existing): + return False + + # A blank separator line only makes sense when there is content above. + if not existing: + addition = f"{GITIGNORE_COMMENT}\n{GITIGNORE_ENTRY}\n" + else: + prefix = "" if existing.endswith("\n") else "\n" + addition = f"{prefix}\n{GITIGNORE_COMMENT}\n{GITIGNORE_ENTRY}\n" + gitignore.write_text(existing + addition) + logger.info("Added %s to %s", GITIGNORE_ENTRY, gitignore) + return True + + +def _already_ignored(contents: str) -> bool: + """Whether any existing line already covers the index directory.""" + wanted = {GITIGNORE_ENTRY, GITIGNORE_ENTRY.rstrip("/"), str(CACHE_RELPATH)} + return any(line.strip() in wanted for line in contents.splitlines()) diff --git a/src/clgraph/detect.py b/src/clgraph/detect.py new file mode 100644 index 0000000..f9ba766 --- /dev/null +++ b/src/clgraph/detect.py @@ -0,0 +1,496 @@ +""" +Propose a SQL source and dialect for a project. + +Detection gathers evidence and rates its own confidence. It never +decides, and it never writes anything: ``clgraph init`` and +``/clgraph:setup`` take what this returns, show it to a human, and record +the answer. + +That split exists because the dialect cannot be guessed safely. sqlglot +parses most of a corpus under the wrong grammar without complaining, so a +bad guess yields a lineage graph that looks right and is not -- the worst +failure mode for something an agent will trust. Confidence is the signal +setup keys on: HIGH is offered pre-selected, anything lower is a real +question. + +Evidence, strongest first: + + 1. A dbt profile's ``type:`` -- the warehouse's own config. HIGH. + 2. Syntax markers unique to one dialect (QUALIFY, SAFE_CAST, ...). + 3. Parse scoring: how many statements fail per candidate dialect. + +Scoring is deliberately last, and weak. Snowflake-only syntax such as +QUALIFY and IFF() parses cleanly under all eight candidates, so scoring +separates almost nothing there; it earns its place mainly on BigQuery's +backticked paths and other syntax sqlglot genuinely cannot read. The +markers do most of the discriminating. See TestParseScoringIsWeak. +""" + +import logging +import re +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +import sqlglot +from sqlglot import exp + +logger = logging.getLogger(__name__) + +try: + import yaml + + YAML_AVAILABLE = True +except ImportError: # pyyaml is an optional extra + YAML_AVAILABLE = False + + +class Confidence: + """How much weight setup should give the proposed dialect.""" + + HIGH = "high" + MEDIUM = "medium" + LOW = "low" + NONE = "none" + + +#: Dialects considered during scoring. Curated rather than exhaustive -- +#: every extra candidate costs a full parse pass over the corpus. +CANDIDATE_DIALECTS: Tuple[str, ...] = ( + "bigquery", + "snowflake", + "postgres", + "redshift", + "duckdb", + "mysql", + "spark", + "trino", +) + +#: Directories never worth scanning for a project's own SQL. +IGNORED_DIRS = frozenset( + { + ".git", + ".hg", + ".svn", + ".venv", + "venv", + "env", + "node_modules", + "__pycache__", + ".clgraph", + ".mypy_cache", + ".pytest_cache", + ".ruff_cache", + "dist", + "build", + "target", + "site-packages", + ".tox", + "htmlcov", + } +) + +#: Regexes that only appear in one dialect's SQL, with the label shown as +#: evidence. Kept small and strongly discriminating on purpose: a marker +#: that fires on two dialects is worse than no marker. +DIALECT_MARKERS: Dict[str, Tuple[Tuple[str, str], ...]] = { + "snowflake": ( + (r"\bQUALIFY\b", "QUALIFY clause"), + (r"\bIFF\s*\(", "IFF() function"), + (r"::\s*VARIANT\b", "::VARIANT cast"), + (r"\bLATERAL\s+FLATTEN\b", "LATERAL FLATTEN"), + (r"\bCURRENT_WAREHOUSE\s*\(", "CURRENT_WAREHOUSE()"), + ), + "bigquery": ( + (r"\bSAFE_CAST\s*\(", "SAFE_CAST()"), + (r"\bSTRUCT\s*<", "STRUCT<> type"), + (r"`[\w-]+\.[\w-]+\.[\w-]+`", "backticked project.dataset.table"), + (r"_PARTITIONTIME\b", "_PARTITIONTIME pseudo-column"), + (r"\bSELECT\s+\*\s+EXCEPT\s*\(", "SELECT * EXCEPT()"), + ), + "postgres": ( + (r"\bSERIAL\b", "SERIAL type"), + (r"\bILIKE\b", "ILIKE operator"), + (r"\bRETURNING\b", "RETURNING clause"), + (r"\bgenerate_series\s*\(", "generate_series()"), + ), + "redshift": ( + (r"\bDISTKEY\b", "DISTKEY"), + (r"\bSORTKEY\b", "SORTKEY"), + (r"\bENCODE\s+\w+", "column ENCODE"), + ), + "duckdb": ( + (r"\bread_csv_auto\s*\(", "read_csv_auto()"), + (r"\bread_parquet\s*\(", "read_parquet()"), + ), + "mysql": ( + (r"\bAUTO_INCREMENT\b", "AUTO_INCREMENT"), + (r"\bENGINE\s*=", "ENGINE="), + ), + "spark": ( + (r"\bUSING\s+DELTA\b", "USING DELTA"), + (r"\bLATERAL\s+VIEW\b", "LATERAL VIEW"), + ), +} + +#: dbt adapter names that map onto a sqlglot dialect. Adapters absent here +#: (custom vendor adapters) are reported as unknown rather than guessed. +DBT_TYPE_TO_DIALECT = { + "bigquery": "bigquery", + "snowflake": "snowflake", + "postgres": "postgres", + "redshift": "redshift", + "duckdb": "duckdb", + "databricks": "spark", + "spark": "spark", + "trino": "trino", + "athena": "trino", + "mysql": "mysql", +} + +MAX_FILES_SCORED = 40 +"""Cap on files pushed through the parse-scoring pass, for large repos.""" + + +@dataclass(frozen=True) +class SqlDirCandidate: + """A directory that holds SQL, and how much of it.""" + + path: Path + file_count: int + kind: str # "sql_dir" | "dbt" + + def to_dict(self, root: Path) -> Dict[str, Any]: + return { + "path": _relative(self.path, root), + "file_count": self.file_count, + "kind": self.kind, + } + + +@dataclass(frozen=True) +class DialectGuess: + """A proposed dialect, with the reasoning attached.""" + + best: Optional[str] + confidence: str + evidence: Tuple[str, ...] + scores: Dict[str, int] + + def to_dict(self) -> Dict[str, Any]: + return { + "best": self.best, + "confidence": self.confidence, + "evidence": list(self.evidence), + "scores": dict(self.scores), + } + + +@dataclass(frozen=True) +class Detection: + """Everything detection has to say about a project.""" + + root: Path + sql_dirs: Tuple[SqlDirCandidate, ...] + dialect: DialectGuess + + def to_dict(self) -> Dict[str, Any]: + return { + "sql_dirs": [c.to_dict(self.root) for c in self.sql_dirs], + "dialect": self.dialect.to_dict(), + } + + +def detect(project_dir: Path) -> Detection: + """Inspect a project and propose what to index. + + Args: + project_dir: Root of the project to inspect. + + Returns: + A Detection. Writes nothing, decides nothing. + """ + root = Path(project_dir) + candidates = _find_sql_dirs(root) + sql_files = _sample_sql_files(candidates) + + return Detection(root=root, sql_dirs=candidates, dialect=_guess_dialect(root, sql_files)) + + +# ============================================================================= +# SQL directories +# ============================================================================= + + +def _find_sql_dirs(root: Path) -> Tuple[SqlDirCandidate, ...]: + """Find directories containing .sql files, most populous first.""" + dbt_roots = _dbt_project_roots(root) + counts: Dict[Path, int] = {} + + for sql_file in root.rglob("*.sql"): + if _is_ignored(sql_file, root): + continue + counts[sql_file.parent] = counts.get(sql_file.parent, 0) + 1 + + candidates = [ + SqlDirCandidate( + path=directory, + file_count=count, + kind="dbt" if _under_any(directory, dbt_roots) else "sql_dir", + ) + for directory, count in counts.items() + ] + candidates.sort(key=lambda c: (-c.file_count, str(c.path))) + return tuple(candidates) + + +def _is_ignored(path: Path, root: Path) -> bool: + """Whether any path segment below root is a noise directory.""" + try: + parts = path.relative_to(root).parts + except ValueError: + return True + return any(part in IGNORED_DIRS for part in parts) + + +def _dbt_project_roots(root: Path) -> Tuple[Path, ...]: + """Directories containing a dbt_project.yml.""" + return tuple( + found.parent for found in root.rglob("dbt_project.yml") if not _is_ignored(found, root) + ) + + +def _under_any(directory: Path, roots: Tuple[Path, ...]) -> bool: + return any(directory == r or r in directory.parents for r in roots) + + +def _sample_sql_files(candidates: Tuple[SqlDirCandidate, ...]) -> List[Path]: + """Collect SQL files to analyse, capped for large repositories.""" + files: List[Path] = [] + for candidate in candidates: + for sql_file in sorted(candidate.path.glob("*.sql")): + files.append(sql_file) + if len(files) >= MAX_FILES_SCORED: + return files + return files + + +# ============================================================================= +# Dialect +# ============================================================================= + + +def _guess_dialect(root: Path, sql_files: List[Path]) -> DialectGuess: + """Weigh the evidence and propose a dialect.""" + if not sql_files: + return DialectGuess(best=None, confidence=Confidence.NONE, evidence=(), scores={}) + + corpus = _read_corpus(sql_files) + scores = _parse_scores(corpus) + + profile = _dialect_from_dbt_profile(root) + if profile is not None: + dialect, evidence = profile + return DialectGuess( + best=dialect, + confidence=Confidence.HIGH, + evidence=(evidence,), + scores=scores, + ) + + marker_hits = _marker_hits(corpus) + if marker_hits: + ranked = sorted(marker_hits.items(), key=lambda kv: (-len(kv[1]), kv[0])) + best, reasons = ranked[0] + tied = len(ranked) > 1 and len(ranked[1][1]) == len(reasons) + return DialectGuess( + best=best, + confidence=Confidence.LOW if tied else Confidence.MEDIUM, + evidence=tuple(f"{best}: {reason}" for reason in reasons), + scores=scores, + ) + + return _from_parse_scores_alone(scores) + + +def _read_corpus(sql_files: List[Path]) -> str: + """Concatenate readable SQL files into one blob.""" + chunks = [] + for sql_file in sql_files: + try: + chunks.append(sql_file.read_text(errors="replace")) + except OSError as err: + logger.warning("Could not read %s: %s", sql_file, err) + return "\n".join(chunks) + + +def _marker_hits(corpus: str) -> Dict[str, List[str]]: + """Which dialect-specific markers appear, by dialect.""" + hits: Dict[str, List[str]] = {} + for dialect, markers in DIALECT_MARKERS.items(): + found = [label for pattern, label in markers if re.search(pattern, corpus, re.IGNORECASE)] + if found: + hits[dialect] = found + return hits + + +@contextmanager +def _quiet_sqlglot(): + """Silence sqlglot while we parse under deliberately wrong dialects. + + Scoring works by trying every candidate, so failures are the signal, + not an incident. Left unsilenced, sqlglot's warnings land on stdout in + any process that configured logging there -- which corrupts + `clgraph detect --json` for the agent reading it. + """ + loggers = [logging.getLogger("sqlglot"), logging.getLogger("sqlglot.parser")] + previous = [(log, log.level) for log in loggers] + for log in loggers: + log.setLevel(logging.ERROR) + try: + yield + finally: + for log, level in previous: + log.setLevel(level) + + +def _parse_scores(corpus: str) -> Dict[str, int]: + """Count parse failures per candidate dialect. Lower is better. + + A raised ParseError is only half the story: sqlglot more often + degrades unsupported syntax into an ``exp.Command`` node rather than + erroring, so counting exceptions alone would score a wrong dialect as + a clean parse. Both count as failures here. + """ + statements = [s for s in corpus.split(";") if s.strip()] + scores: Dict[str, int] = {} + + with _quiet_sqlglot(): + for dialect in CANDIDATE_DIALECTS: + scores[dialect] = sum( + 1 for statement in statements if not _parses_cleanly(statement, dialect) + ) + + return scores + + +def _parses_cleanly(statement: str, dialect: str) -> bool: + """Whether sqlglot fully understood a statement in this dialect.""" + try: + parsed = sqlglot.parse(statement, dialect=dialect) + except Exception: + return False + + return bool(parsed) and not any( + node is None or isinstance(node, exp.Command) or node.find(exp.Command) for node in parsed + ) + + +def _from_parse_scores_alone(scores: Dict[str, int]) -> DialectGuess: + """Fall back to parse scoring when no marker fired. + + Portable SQL parses cleanly everywhere, which is a tie, which is LOW. + Saying so is the honest answer -- and the one that makes setup ask. + """ + if not scores: + return DialectGuess(best=None, confidence=Confidence.NONE, evidence=(), scores={}) + + best_score = min(scores.values()) + winners = sorted(d for d, s in scores.items() if s == best_score) + + if len(winners) == 1: + return DialectGuess( + best=winners[0], + confidence=Confidence.MEDIUM, + evidence=(f"{winners[0]}: only dialect that parsed the corpus cleanly",), + scores=scores, + ) + + return DialectGuess( + best=winners[0], + confidence=Confidence.LOW, + evidence=( + f"{len(winners)} dialects parse this corpus equally well " + f"({', '.join(winners)}); nothing distinguishes them", + ), + scores=scores, + ) + + +def _dialect_from_dbt_profile(root: Path) -> Optional[Tuple[str, str]]: + """Read the adapter type from a dbt profile, if there is one. + + Returns (dialect, evidence) or None. An adapter with no sqlglot + equivalent returns None rather than a guess. + """ + for profiles in _profile_candidates(root): + adapter = _adapter_type(profiles) + if adapter is None: + continue + dialect = DBT_TYPE_TO_DIALECT.get(adapter) + if dialect is None: + logger.info("dbt adapter '%s' has no known sqlglot dialect", adapter) + continue + return dialect, f"{profiles.name} declares adapter type '{adapter}'" + return None + + +def _profile_candidates(root: Path) -> List[Path]: + """Project-local profiles.yml first, then the user's ~/.dbt one.""" + found = [p for p in root.rglob("profiles.yml") if not _is_ignored(p, root)] + home_profile = Path.home() / ".dbt" / "profiles.yml" + if home_profile.exists(): + found.append(home_profile) + return found + + +def _adapter_type(profiles: Path) -> Optional[str]: + """Extract an output's ``type:`` from a dbt profiles.yml.""" + try: + text = profiles.read_text(errors="replace") + except OSError as err: + logger.warning("Could not read %s: %s", profiles, err) + return None + + if YAML_AVAILABLE: + adapter = _adapter_from_yaml(text) + if adapter is not None: + return adapter + + match = re.search(r"^\s*type:\s*([\w-]+)\s*$", text, re.MULTILINE) + return match.group(1).lower() if match else None + + +def _adapter_from_yaml(text: str) -> Optional[str]: + """Walk parsed YAML for the first outputs..type value.""" + try: + data = yaml.safe_load(text) + except Exception as err: # malformed YAML is not fatal; regex still tries + logger.warning("Could not parse dbt profile YAML: %s", err) + return None + + if not isinstance(data, dict): + return None + + for profile in data.values(): + if not isinstance(profile, dict): + continue + outputs = profile.get("outputs") + if not isinstance(outputs, dict): + continue + target = profile.get("target") + ordered = [outputs[target]] if target in outputs else list(outputs.values()) + for output in ordered: + if isinstance(output, dict) and isinstance(output.get("type"), str): + return output["type"].lower() + return None + + +def _relative(path: Path, root: Path) -> str: + """Render path relative to root when possible, for readable output.""" + try: + return str(path.relative_to(root)) + except ValueError: + return str(path) diff --git a/tests/test_cli.py b/tests/test_cli.py index 3f0aa98..6f4c201 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -224,3 +224,132 @@ def test_analyze_help(self): output = _strip_ansi(result.stdout) assert "--dialect" in output assert "--format" in output + + +class TestDetectCommand: + """A6 — proposing a source and dialect, deciding nothing.""" + + def test_detect_reports_sql_directory(self, sql_dir): + result = runner.invoke(app, ["detect", sql_dir]) + + assert result.exit_code == 0 + assert "2" in result.output # two .sql files in the fixture + + def test_detect_json_is_machine_readable(self, sql_dir): + result = runner.invoke(app, ["detect", sql_dir, "--json"]) + + assert result.exit_code == 0 + payload = json.loads(result.stdout) + assert "sql_dirs" in payload + assert "dialect" in payload + assert "confidence" in payload["dialect"] + + def test_detect_reports_low_confidence_for_portable_sql(self, sql_dir): + """The fixture SQL is plain ANSI — nothing should be asserted.""" + result = runner.invoke(app, ["detect", sql_dir, "--json"]) + + assert json.loads(result.stdout)["dialect"]["confidence"] in {"low", "medium"} + + def test_detect_writes_nothing(self, sql_dir): + before = set(Path(sql_dir).rglob("*")) + + runner.invoke(app, ["detect", sql_dir]) + + assert set(Path(sql_dir).rglob("*")) == before + + def test_detect_help(self): + result = runner.invoke(app, ["detect", "--help"]) + assert result.exit_code == 0 + assert "--json" in _strip_ansi(result.stdout) + + +class TestInitCommand: + """A7 — recording the answer.""" + + def test_init_writes_config(self, sql_dir): + project = Path(sql_dir) + + result = runner.invoke( + app, + [ + "init", + "--sql-dir", + sql_dir, + "--dialect", + "bigquery", + "--yes", + "--project-dir", + sql_dir, + ], + ) + + assert result.exit_code == 0 + assert (project / "clgraph.toml").exists() + + def test_init_refuses_unknown_dialect(self, sql_dir): + result = runner.invoke( + app, + [ + "init", + "--sql-dir", + sql_dir, + "--dialect", + "nonsense", + "--yes", + "--project-dir", + sql_dir, + ], + ) + + assert result.exit_code != 0 + assert "dialect" in result.output.lower() + + def test_init_requires_dialect_in_non_interactive_mode(self, sql_dir): + """--yes with no dialect must fail rather than pick one.""" + result = runner.invoke( + app, ["init", "--sql-dir", sql_dir, "--yes", "--project-dir", sql_dir] + ) + + assert result.exit_code != 0 + + def test_init_refuses_to_overwrite_without_force(self, sql_dir): + args = [ + "init", + "--sql-dir", + sql_dir, + "--dialect", + "bigquery", + "--yes", + "--project-dir", + sql_dir, + ] + runner.invoke(app, args) + + result = runner.invoke(app, args) + + assert result.exit_code != 0 + assert "force" in result.output.lower() + + def test_init_adds_gitignore_entry(self, sql_dir): + runner.invoke( + app, + [ + "init", + "--sql-dir", + sql_dir, + "--dialect", + "bigquery", + "--yes", + "--project-dir", + sql_dir, + ], + ) + + assert ".clgraph/" in (Path(sql_dir) / ".gitignore").read_text() + + def test_init_help(self): + result = runner.invoke(app, ["init", "--help"]) + assert result.exit_code == 0 + output = _strip_ansi(result.stdout) + assert "--dialect" in output + assert "--into-pyproject" in output diff --git a/tests/test_config_writer.py b/tests/test_config_writer.py new file mode 100644 index 0000000..741d258 --- /dev/null +++ b/tests/test_config_writer.py @@ -0,0 +1,223 @@ +""" +Tests for clgraph.config_writer — recording the setup answers. + +Writing config is the step that turns a proposal into a decision. It +validates before it writes, refuses to clobber an existing answer without +being told to, and produces a file that clgraph.mcp.config can read back +— that round trip is the contract that matters. +""" + +import pytest + +from clgraph.config_writer import write_config +from clgraph.mcp.config import resolve + + +def write(path, content): + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(content) + return path + + +@pytest.fixture +def project(tmp_path): + write(tmp_path / "models" / "a.sql", "SELECT 1") + return tmp_path + + +# ============================================================================= +# clgraph.toml +# ============================================================================= + + +class TestWriteClgraphToml: + def test_writes_a_readable_config(self, project): + result = write_config(project, sql_dir="models", dialect="snowflake") + + assert result.config_path == project / "clgraph.toml" + assert result.target == "clgraph.toml" + + def test_round_trips_through_resolve(self, project): + """The contract: what init writes, the MCP server reads back.""" + write_config(project, sql_dir="models", dialect="snowflake") + + source = resolve(project, env={}) + + assert source.origin == "clgraph.toml" + assert source.path == project / "models" + assert source.dialect == "snowflake" + assert source.is_complete is True + + def test_refuses_to_overwrite_without_force(self, project): + write_config(project, sql_dir="models", dialect="snowflake") + + with pytest.raises(FileExistsError): + write_config(project, sql_dir="models", dialect="postgres") + + def test_force_overwrites(self, project): + write(project / "other" / "b.sql", "SELECT 1") + write_config(project, sql_dir="models", dialect="snowflake") + + write_config(project, sql_dir="other", dialect="postgres", force=True) + + assert resolve(project, env={}).dialect == "postgres" + + +# ============================================================================= +# pyproject.toml +# ============================================================================= + + +class TestWriteIntoPyproject: + def test_appends_tool_clgraph_section(self, project): + write(project / "pyproject.toml", '[project]\nname = "demo"\n') + + result = write_config(project, sql_dir="models", dialect="postgres", into_pyproject=True) + + assert result.target == "pyproject.toml" + assert resolve(project, env={}).origin == "pyproject.toml" + + def test_leaves_existing_content_intact(self, project): + original = '[project]\nname = "demo"\nversion = "1.0"\n\n[tool.ruff]\nline-length = 100\n' + write(project / "pyproject.toml", original) + + write_config(project, sql_dir="models", dialect="postgres", into_pyproject=True) + + updated = (project / "pyproject.toml").read_text() + assert original in updated + assert "[tool.clgraph]" in updated + + def test_refuses_existing_clgraph_table_without_force(self, project): + write( + project / "pyproject.toml", + '[project]\nname = "demo"\n\n[tool.clgraph]\nsql_dir = "old"\ndialect = "mysql"\n', + ) + + with pytest.raises(FileExistsError): + write_config(project, sql_dir="models", dialect="postgres", into_pyproject=True) + + def test_creates_pyproject_if_absent(self, project): + write_config(project, sql_dir="models", dialect="postgres", into_pyproject=True) + + assert (project / "pyproject.toml").exists() + assert resolve(project, env={}).dialect == "postgres" + + +# ============================================================================= +# Validation — refusing bad input before it reaches config +# ============================================================================= + + +class TestValidation: + def test_unknown_dialect_is_rejected(self, project): + with pytest.raises(ValueError, match="dialect"): + write_config(project, sql_dir="models", dialect="not_a_warehouse") + + def test_error_names_some_valid_dialects(self, project): + with pytest.raises(ValueError) as excinfo: + write_config(project, sql_dir="models", dialect="nope") + + assert "snowflake" in str(excinfo.value) + + def test_missing_sql_dir_is_rejected(self, project): + with pytest.raises(ValueError, match="does not exist"): + write_config(project, sql_dir="nowhere", dialect="postgres") + + def test_empty_dialect_for_sql_dir_is_rejected(self, project): + with pytest.raises(ValueError): + write_config(project, sql_dir="models", dialect="") + + def test_nothing_is_written_when_validation_fails(self, project): + before = set(project.rglob("*")) + + with pytest.raises(ValueError): + write_config(project, sql_dir="models", dialect="nope") + + assert set(project.rglob("*")) == before + + +# ============================================================================= +# .gitignore +# ============================================================================= + + +class TestGitignore: + def test_adds_clgraph_entry(self, project): + result = write_config(project, sql_dir="models", dialect="postgres") + + assert result.gitignore_updated is True + assert ".clgraph/" in (project / ".gitignore").read_text() + + def test_creates_gitignore_when_absent(self, project): + write_config(project, sql_dir="models", dialect="postgres") + + assert (project / ".gitignore").exists() + + def test_entry_is_added_only_once(self, project): + write_config(project, sql_dir="models", dialect="postgres") + write_config(project, sql_dir="models", dialect="postgres", force=True) + write_config(project, sql_dir="models", dialect="postgres", force=True) + + contents = (project / ".gitignore").read_text() + assert contents.count(".clgraph/") == 1 + + def test_preserves_existing_gitignore_content(self, project): + write(project / ".gitignore", "*.pyc\n.venv/\n") + + write_config(project, sql_dir="models", dialect="postgres") + + contents = (project / ".gitignore").read_text() + assert "*.pyc" in contents + assert ".venv/" in contents + assert ".clgraph/" in contents + + def test_reports_no_update_when_already_ignored(self, project): + write(project / ".gitignore", ".clgraph/\n") + + result = write_config(project, sql_dir="models", dialect="postgres") + + assert result.gitignore_updated is False + + +# ============================================================================= +# Paths +# ============================================================================= + + +class TestPathHandling: + def test_sql_dir_is_stored_relative_to_the_project(self, project): + """An absolute tmp path in a committed config helps nobody.""" + write_config(project, sql_dir=str(project / "models"), dialect="postgres") + + assert 'sql_dir = "models"' in (project / "clgraph.toml").read_text() + + def test_trailing_slash_is_accepted(self, project): + write_config(project, sql_dir="models/", dialect="postgres") + + assert resolve(project, env={}).path == project / "models" + + def test_json_pipeline_needs_no_dialect(self, project): + write(project / "pipeline.json", '{"queries": {}}') + + write_config(project, sql_dir="pipeline.json", dialect=None) + + source = resolve(project, env={}) + assert source.kind == "json" + assert source.is_complete is True + + +class TestGitignoreFormatting: + def test_fresh_gitignore_has_no_leading_blank_line(self, project): + write_config(project, sql_dir="models", dialect="postgres") + + contents = (project / ".gitignore").read_text() + + assert not contents.startswith("\n") + assert contents.startswith("#") + + def test_separator_blank_line_when_appending(self, project): + write(project / ".gitignore", "*.pyc\n") + + write_config(project, sql_dir="models", dialect="postgres") + + assert "*.pyc\n\n#" in (project / ".gitignore").read_text() diff --git a/tests/test_detect.py b/tests/test_detect.py new file mode 100644 index 0000000..5a0e2b8 --- /dev/null +++ b/tests/test_detect.py @@ -0,0 +1,329 @@ +""" +Tests for clgraph.detect — proposing a SQL source and dialect. + +Detection never decides. It gathers evidence and rates its own +confidence, so setup can pre-select a strong answer and ask a real +question about a weak one. The tests that matter most here are the ones +asserting it stays quiet when the evidence is thin. +""" + +import pytest + +from clgraph.detect import Confidence, detect + + +def write(path, content): + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(content) + return path + + +SNOWFLAKE_SQL = """ +CREATE TABLE analytics.revenue AS +SELECT + order_id, + IFF(amount > 100, 'large', 'small') AS bucket, + amount +FROM raw.orders +QUALIFY ROW_NUMBER() OVER (PARTITION BY order_id ORDER BY ts DESC) = 1 +""" + +BIGQUERY_SQL = """ +CREATE TABLE `proj.analytics.revenue` AS +SELECT + order_id, + SAFE_CAST(amount AS FLOAT64) AS amount +FROM `proj.raw.orders` +""" + +POSTGRES_SQL = """ +CREATE TABLE analytics.revenue AS +SELECT + order_id, + amount::numeric AS amount +FROM raw.orders +WHERE status ILIKE 'paid%' +""" + +PLAIN_SQL = """ +CREATE TABLE analytics.revenue AS +SELECT order_id, amount FROM raw.orders +""" + + +# ============================================================================= +# SQL directory candidates +# ============================================================================= + + +class TestSqlDirDiscovery: + def test_finds_directory_with_sql_files(self, tmp_path): + write(tmp_path / "sql" / "a.sql", PLAIN_SQL) + write(tmp_path / "sql" / "b.sql", PLAIN_SQL) + + result = detect(tmp_path) + + assert [c.path for c in result.sql_dirs] == [tmp_path / "sql"] + assert result.sql_dirs[0].file_count == 2 + + def test_ranks_directories_by_file_count(self, tmp_path): + write(tmp_path / "small" / "a.sql", PLAIN_SQL) + for name in "abc": + write(tmp_path / "big" / f"{name}.sql", PLAIN_SQL) + + result = detect(tmp_path) + + assert [c.path.name for c in result.sql_dirs] == ["big", "small"] + + def test_ignores_noise_directories(self, tmp_path): + write(tmp_path / "sql" / "a.sql", PLAIN_SQL) + write(tmp_path / ".venv" / "lib" / "x.sql", PLAIN_SQL) + write(tmp_path / "node_modules" / "y.sql", PLAIN_SQL) + write(tmp_path / ".git" / "z.sql", PLAIN_SQL) + + result = detect(tmp_path) + + assert [c.path.name for c in result.sql_dirs] == ["sql"] + + def test_no_sql_anywhere_yields_no_candidates(self, tmp_path): + write(tmp_path / "README.md", "# nothing here") + + assert detect(tmp_path).sql_dirs == () + + def test_dbt_models_directory_is_marked_dbt(self, tmp_path): + write(tmp_path / "dbt_project.yml", "name: analytics\n") + write(tmp_path / "models" / "stg_users.sql", PLAIN_SQL) + + result = detect(tmp_path) + + assert result.sql_dirs[0].kind == "dbt" + + def test_plain_directory_is_marked_sql_dir(self, tmp_path): + write(tmp_path / "sql" / "a.sql", PLAIN_SQL) + + assert detect(tmp_path).sql_dirs[0].kind == "sql_dir" + + +# ============================================================================= +# Dialect evidence +# ============================================================================= + + +class TestDbtProfileEvidence: + def test_profiles_yml_target_type_is_high_confidence(self, tmp_path): + write(tmp_path / "dbt_project.yml", "name: analytics\nprofile: analytics\n") + write( + tmp_path / "profiles.yml", + "analytics:\n" + " target: prod\n" + " outputs:\n" + " prod:\n" + " type: snowflake\n" + " account: xy12345\n", + ) + write(tmp_path / "models" / "a.sql", PLAIN_SQL) + + guess = detect(tmp_path).dialect + + assert guess.best == "snowflake" + assert guess.confidence == Confidence.HIGH + assert any("profiles.yml" in e for e in guess.evidence) + + def test_profile_type_beats_conflicting_syntax_markers(self, tmp_path): + """The warehouse config is authoritative; markers are a heuristic.""" + write(tmp_path / "dbt_project.yml", "name: analytics\n") + write( + tmp_path / "profiles.yml", + "analytics:\n outputs:\n prod:\n type: postgres\n", + ) + write(tmp_path / "models" / "a.sql", SNOWFLAKE_SQL) + + assert detect(tmp_path).dialect.best == "postgres" + + def test_unknown_profile_type_is_not_reported_as_a_dialect(self, tmp_path): + write(tmp_path / "dbt_project.yml", "name: analytics\n") + write( + tmp_path / "profiles.yml", + "analytics:\n outputs:\n prod:\n type: some_vendor_db\n", + ) + write(tmp_path / "models" / "a.sql", PLAIN_SQL) + + assert detect(tmp_path).dialect.best != "some_vendor_db" + + +class TestSyntaxMarkerEvidence: + def test_snowflake_markers(self, tmp_path): + write(tmp_path / "sql" / "a.sql", SNOWFLAKE_SQL) + + guess = detect(tmp_path).dialect + + assert guess.best == "snowflake" + assert guess.confidence == Confidence.MEDIUM + + def test_bigquery_markers(self, tmp_path): + write(tmp_path / "sql" / "a.sql", BIGQUERY_SQL) + + assert detect(tmp_path).dialect.best == "bigquery" + + def test_postgres_markers(self, tmp_path): + write(tmp_path / "sql" / "a.sql", POSTGRES_SQL) + + assert detect(tmp_path).dialect.best == "postgres" + + def test_evidence_names_the_marker_that_fired(self, tmp_path): + write(tmp_path / "sql" / "a.sql", SNOWFLAKE_SQL) + + evidence = " ".join(detect(tmp_path).dialect.evidence).lower() + + assert "qualify" in evidence + + +class TestParseScoring: + def test_scores_are_reported_per_dialect(self, tmp_path): + write(tmp_path / "sql" / "a.sql", SNOWFLAKE_SQL) + + scores = detect(tmp_path).dialect.scores + + assert "snowflake" in scores + assert "bigquery" in scores + assert all(isinstance(v, int) for v in scores.values()) + + def test_snowflake_corpus_never_scores_bigquery_best(self, tmp_path): + """The regression guard for the whole silent-wrong-dialect class.""" + for name in "abc": + write(tmp_path / "sql" / f"{name}.sql", SNOWFLAKE_SQL) + + assert detect(tmp_path).dialect.best != "bigquery" + + def test_bigquery_corpus_never_scores_snowflake_best(self, tmp_path): + for name in "abc": + write(tmp_path / "sql" / f"{name}.sql", BIGQUERY_SQL) + + assert detect(tmp_path).dialect.best != "snowflake" + + +class TestLowConfidence: + def test_portable_sql_is_low_confidence(self, tmp_path): + """Plain ANSI SQL parses everywhere. Saying so is the honest answer.""" + write(tmp_path / "sql" / "a.sql", PLAIN_SQL) + + guess = detect(tmp_path).dialect + + assert guess.confidence == Confidence.LOW + + def test_no_sql_means_no_guess(self, tmp_path): + write(tmp_path / "README.md", "# nothing") + + guess = detect(tmp_path).dialect + + assert guess.best is None + assert guess.confidence == Confidence.NONE + + def test_conflicting_markers_lower_confidence(self, tmp_path): + write(tmp_path / "sql" / "snow.sql", SNOWFLAKE_SQL) + write(tmp_path / "sql" / "pg.sql", POSTGRES_SQL) + + assert detect(tmp_path).dialect.confidence != Confidence.HIGH + + +# ============================================================================= +# Serialization +# ============================================================================= + + +class TestSerialization: + def test_to_dict_is_json_serializable(self, tmp_path): + import json + + write(tmp_path / "sql" / "a.sql", SNOWFLAKE_SQL) + + payload = json.loads(json.dumps(detect(tmp_path).to_dict())) + + assert payload["dialect"]["best"] == "snowflake" + assert payload["sql_dirs"][0]["file_count"] == 1 + assert isinstance(payload["sql_dirs"][0]["path"], str) + + def test_paths_are_relative_to_the_project(self, tmp_path): + """Absolute tmp paths would be noise in agent output.""" + write(tmp_path / "sql" / "a.sql", PLAIN_SQL) + + assert detect(tmp_path).to_dict()["sql_dirs"][0]["path"] == "sql" + + +class TestDetectionIsReadOnly: + def test_detect_writes_nothing(self, tmp_path): + write(tmp_path / "sql" / "a.sql", SNOWFLAKE_SQL) + before = set(tmp_path.rglob("*")) + + detect(tmp_path) + + assert set(tmp_path.rglob("*")) == before + + +@pytest.mark.parametrize( + "confidence", + [Confidence.HIGH, Confidence.MEDIUM, Confidence.LOW, Confidence.NONE], +) +def test_confidence_values_are_strings(confidence): + """Setup branches on these, and they land in JSON.""" + assert isinstance(confidence, str) + + +class TestParseScoringIsWeak: + """Documents why markers carry the weight, not parse scoring. + + sqlglot is permissive by design: Snowflake-only syntax parses cleanly + under every candidate dialect. That permissiveness is exactly why a + wrong dialect yields a plausible-but-wrong graph instead of an error, + and why detection must not lean on scoring alone. If a future sqlglot + gets stricter this test will fail loudly, which is the point. + """ + + def test_snowflake_syntax_parses_under_every_dialect(self): + from clgraph.detect import CANDIDATE_DIALECTS, _parse_scores + + scores = _parse_scores(SNOWFLAKE_SQL) + + assert set(scores) == set(CANDIDATE_DIALECTS) + assert all(failures == 0 for failures in scores.values()) + + def test_backticked_bigquery_paths_do_discriminate(self): + """The one marker family scoring actually catches.""" + from clgraph.detect import _parse_scores + + scores = _parse_scores(BIGQUERY_SQL) + + assert scores["bigquery"] < scores["postgres"] + + +class TestOutputPurity: + """`clgraph detect --json` must be parseable, always. + + Parse scoring deliberately probes wrong dialects, and sqlglot logs a + warning each time. In any process that has configured logging to + stdout, those warnings land in the middle of the JSON an agent is + trying to read. This is the regression test for that. + """ + + def test_scoring_emits_no_log_records(self, caplog): + import logging + + from clgraph.detect import _parse_scores + + with caplog.at_level(logging.WARNING, logger="sqlglot"): + _parse_scores(BIGQUERY_SQL) + + assert [r for r in caplog.records if r.name.startswith("sqlglot")] == [] + + def test_logger_levels_are_restored(self): + import logging + + from clgraph.detect import _parse_scores + + sqlglot_logger = logging.getLogger("sqlglot") + sqlglot_logger.setLevel(logging.DEBUG) + + _parse_scores(BIGQUERY_SQL) + + assert sqlglot_logger.level == logging.DEBUG + sqlglot_logger.setLevel(logging.NOTSET)