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)