From b16af763be5cfe4acaa158c0dd1a91d6e8482d0d Mon Sep 17 00:00:00 2001 From: Fszta <36471574+Fszta@users.noreply.github.com> Date: Sun, 27 Sep 2026 23:00:15 +0200 Subject: [PATCH 1/5] fix: resolve the mypy errors on main - service.py: type the partial-edges reason list as List[str] so the dominant-reason pick returns str, and expose get_model_config on the ProjectMetadataProvider protocol (both concrete registries already implement it) so the LineageAndMetadataProvider seam matches its use. - sql_parser.py: narrow _phantom_token_reason's return type to the two Literal reasons it actually emits, matching UnresolvedColumnEdge. - tests: pass str paths to ManifestReader, add the fake provider's get_model_config, and type-ignore the path-injected _build fixture imports (deliberately outside mypy's module map). Co-Authored-By: Claude Fable 5 --- parrant/lineage/provider.py | 8 ++++++++ parrant/lineage/service.py | 4 +++- parrant/parser/sql_parser.py | 4 +++- tests/e2e/test_unresolved_edges_e2e.py | 2 +- tests/unit/artifacts/test_manifest.py | 4 ++-- tests/unit/registry/test_unresolved_edges.py | 2 +- tests/unit/test_lineage_provider.py | 3 +++ 7 files changed, 21 insertions(+), 6 deletions(-) diff --git a/parrant/lineage/provider.py b/parrant/lineage/provider.py index d6846e8..e4f8a5b 100644 --- a/parrant/lineage/provider.py +++ b/parrant/lineage/provider.py @@ -196,6 +196,14 @@ def get_column_dbt_meta(self, model: str, column: str) -> Dict[str, Any]: Same contract as :meth:`get_model_dbt_meta`, scoped to one column. Absent ⇒ ``{}``. """ + def get_model_config(self, model: str) -> Dict[str, Any]: + """The node's resolved dbt ``config`` dict for a model (case-insensitive), or ``{}``. + + Manifest-sourced and metadata-agnostic, exactly like :meth:`get_model_dbt_meta`: + every key (``materialized``, ``grants``, ``tags``, …) is exposed generically, none + privileged, values surfaced raw. Absent config ⇒ ``{}`` — never guessed. + """ + @runtime_checkable class LineageAndMetadataProvider(LineageProvider, ProjectMetadataProvider, Protocol): diff --git a/parrant/lineage/service.py b/parrant/lineage/service.py index 925f511..e73f69c 100644 --- a/parrant/lineage/service.py +++ b/parrant/lineage/service.py @@ -164,7 +164,9 @@ def _partial_edges_reason(registry: LineageAndMetadataProvider, name: str) -> st model = registry.get_model(name) except ModelNotFoundError: return "unresolved_edge" - reasons = [edge.get("reason") for edge in _unresolved_edges(model) if edge.get("reason")] + reasons: List[str] = [ + str(edge["reason"]) for edge in _unresolved_edges(model) if edge.get("reason") + ] if not reasons: return "unresolved_edge" counts = Counter(reasons) diff --git a/parrant/parser/sql_parser.py b/parrant/parser/sql_parser.py index 4cdb923..0d9c703 100644 --- a/parrant/parser/sql_parser.py +++ b/parrant/parser/sql_parser.py @@ -735,7 +735,9 @@ def _expand_flatten_sources( return out @staticmethod - def _phantom_token_reason(token: str, flatten_aliases: Set[str]) -> Optional[str]: + def _phantom_token_reason( + token: str, flatten_aliases: Set[str] + ) -> Optional[Literal["pivot_output", "phantom_alias"]]: """Classify a source token as a fabricated edge, or ``None`` if it is genuine. Returns ``"pivot_output"`` for a quoted pivot literal, ``"phantom_alias"`` for a diff --git a/tests/e2e/test_unresolved_edges_e2e.py b/tests/e2e/test_unresolved_edges_e2e.py index 1cfc35a..14a8bc5 100644 --- a/tests/e2e/test_unresolved_edges_e2e.py +++ b/tests/e2e/test_unresolved_edges_e2e.py @@ -39,7 +39,7 @@ _FIXTURE_DIR = Path(__file__).parents[1] / "fixtures" / "unresolved_edges" sys.path.insert(0, str(_FIXTURE_DIR)) -import _build # noqa: E402 (path-injected fixture builder) +import _build # type: ignore[import-not-found] # noqa: E402 (path-injected fixture builder) # Materialize the abstract manifest + catalog once into a tmp dir for the whole module. _TMP = tempfile.TemporaryDirectory() diff --git a/tests/unit/artifacts/test_manifest.py b/tests/unit/artifacts/test_manifest.py index 760f990..a869af1 100644 --- a/tests/unit/artifacts/test_manifest.py +++ b/tests/unit/artifacts/test_manifest.py @@ -268,7 +268,7 @@ def test_source_dependencies_without_identifier(tmp_path): with open(manifest_path, "w") as f: json.dump(manifest_data, f) - reader = ManifestReader(manifest_path) + reader = ManifestReader(str(manifest_path)) reader.load() upstream = reader.get_model_upstream() @@ -348,7 +348,7 @@ def test_manifest_normalizes_exposure_dependencies(tmp_path: Path) -> None: with open(manifest_path, "w") as f: json.dump(manifest_data, f) - reader = ManifestReader(manifest_path) + reader = ManifestReader(str(manifest_path)) reader.load() exposure_deps = reader.get_exposure_dependencies() diff --git a/tests/unit/registry/test_unresolved_edges.py b/tests/unit/registry/test_unresolved_edges.py index d92c41c..59eaa47 100644 --- a/tests/unit/registry/test_unresolved_edges.py +++ b/tests/unit/registry/test_unresolved_edges.py @@ -17,7 +17,7 @@ _FIXTURE_DIR = Path(__file__).parents[2] / "fixtures" / "unresolved_edges" sys.path.insert(0, str(_FIXTURE_DIR)) -import _build # noqa: E402 (path-injected fixture builder) +import _build # type: ignore[import-not-found] # noqa: E402 (path-injected fixture builder) # Materialize the manifest + catalog once into a tmp dir for the whole module. _TMP = tempfile.TemporaryDirectory() diff --git a/tests/unit/test_lineage_provider.py b/tests/unit/test_lineage_provider.py index 2582a39..83203af 100644 --- a/tests/unit/test_lineage_provider.py +++ b/tests/unit/test_lineage_provider.py @@ -175,6 +175,9 @@ def get_model_dbt_meta(self, model: str) -> Dict[str, Any]: def get_column_dbt_meta(self, model: str, column: str) -> Dict[str, Any]: return {} + def get_model_config(self, model: str) -> Dict[str, Any]: + return {} + # --- fixtures -------------------------------------------------------------- From 226a311883e6ae06526ad175613eacfe2f93611d Mon Sep 17 00:00:00 2001 From: Fszta <36471574+Fszta@users.noreply.github.com> Date: Sun, 27 Sep 2026 23:00:26 +0200 Subject: [PATCH 2/5] chore: point mypy at the parrant package [tool.mypy].files still named dbt_column_lineage, the pre-rename package, so bare mypy checked nothing under the source tree. Bare mypy now matches poetry run type-check. Co-Authored-By: Claude Fable 5 --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 5e8b577..b098b40 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -74,7 +74,7 @@ build-backend = "poetry.core.masonry.api" [tool.mypy] python_version = "3.10" -files = ["dbt_column_lineage", "tests"] +files = ["parrant", "tests"] show_error_codes = true pretty = true explicit_package_bases = true From 50912b1bda09b2f3eb40e29a48d8edfa3ec64470 Mon Sep 17 00:00:00 2001 From: Fszta <36471574+Fszta@users.noreply.github.com> Date: Sun, 27 Sep 2026 23:04:11 +0200 Subject: [PATCH 3/5] chore: pin the format and type toolchain Pin black 25.12.0, ruff 0.16.3, and mypy 1.19.1 exactly in the dev dependencies (the versions the lockfile had already resolved to) and lift the .pre-commit-config.yaml revs (black 24.10.0 / ruff v0.8.4 / mypy v1.15.0) to the same versions, so local runs, hooks, and CI share one definition of clean. Also make that definition explicit for ruff: adopt 0.16's default rule set with a policy ignore list for the rules this codebase violates on purpose (fail-safe blind excepts, returncode-forwarding subprocess wrappers). Co-Authored-By: Claude Fable 5 --- .pre-commit-config.yaml | 8 +++++--- poetry.lock | 2 +- pyproject.toml | 28 +++++++++++++++++++++++++--- 3 files changed, 31 insertions(+), 7 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index a4c5c70..23fe1d0 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -14,21 +14,23 @@ repos: - id: mixed-line-ending args: ['--fix=lf'] + # black / ruff / mypy revs are pinned to the SAME versions as pyproject.toml's dev + # dependencies (see the comment there) — bump them together. - repo: https://github.com/psf/black - rev: 24.10.0 + rev: 25.12.0 hooks: - id: black language_version: python3 args: ['--line-length=100'] - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.8.4 + rev: v0.16.3 hooks: - id: ruff args: [--fix, --unsafe-fixes, --exit-non-zero-on-fix] - repo: https://github.com/pre-commit/mirrors-mypy - rev: v1.15.0 + rev: v1.19.1 hooks: - id: mypy args: [--show-error-codes, --pretty, --explicit-package-bases] diff --git a/poetry.lock b/poetry.lock index 512fd27..740ca7e 100644 --- a/poetry.lock +++ b/poetry.lock @@ -2614,4 +2614,4 @@ type = ["pytest-mypy"] [metadata] lock-version = "2.1" python-versions = ">=3.10" -content-hash = "741e87981e54553c2c74546cdadb50c91eb422cc3c0c8f10d38d468b8e8558cc" +content-hash = "2cfb513f81ad4e1e284b772b8edc2c99de374059187fe493c36a11f584845efd" diff --git a/pyproject.toml b/pyproject.toml index b098b40..031f527 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -41,9 +41,11 @@ dbt-duckdb = ">=1.9.2,<2.0.0" dbt-core = ">=1.9.0,<2.0.0" duckdb = ">=1.2.1,<2.0.0" pytest = ">=8.3.4,<10.0.0" -mypy = ">=1.15.0,<2.0.0" -black = ">=24.10.0" -ruff = ">=0.8.4" +# Pinned exactly and mirrored by .pre-commit-config.yaml revs — one definition of +# "clean" across local runs, hooks, and CI. Bump all four places together. +mypy = "1.19.1" +black = "25.12.0" +ruff = "0.16.3" mkdocs-material = ">=9.5.0,<10.0.0" mkdocs-click = ">=0.8.1,<0.9.0" requests = ">=2.32.0,<3.0.0" @@ -68,6 +70,26 @@ line-length = 100 [tool.ruff] line-length = 100 +[tool.ruff.lint] +# Ruff 0.16's broad default rule set, minus rules that flag patterns this codebase uses +# deliberately. Do not silence a rule here to get a diff through — these are policy: +# BLE001/S110/S112 — broad/silent `except Exception` is the documented fail-safe +# contract (degrade honestly, never raise) across the registry/service/parser. +# PLW1510 — the scripts/ wrappers intentionally forward `subprocess.run` returncodes. +# PLW1508/G201/TRY002/RUF034/RUF059/EXE002 — style opinions not worth code churn. +ignore = [ + "BLE001", + "S110", + "S112", + "PLW1510", + "PLW1508", + "G201", + "TRY002", + "RUF034", + "RUF059", + "EXE002", +] + [build-system] requires = ["poetry-core>=2.0.0,<3.0.0"] build-backend = "poetry.core.masonry.api" From 6f6df9870609577b369fde8223e3eefa4cc02e84 Mon Sep 17 00:00:00 2001 From: Fszta <36471574+Fszta@users.noreply.github.com> Date: Sun, 27 Sep 2026 23:04:11 +0200 Subject: [PATCH 4/5] ci: gate on type-check and format-check Add mypy (poetry run type-check), black --check (line length 100), and ruff check (no --fix) steps to the test workflow, before the test tiers. The format gates go green with the follow-up one-time repo-wide style commit on this branch. Co-Authored-By: Claude Fable 5 --- .github/workflows/test.yml | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 1ffaef4..47282bb 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -56,6 +56,15 @@ jobs: - name: Install dependencies run: poetry install --no-interaction + - name: Type check + run: poetry run type-check + + - name: Check formatting (black) + run: poetry run black --check --line-length=100 . + + - name: Lint (ruff) + run: poetry run ruff check . + - name: Run unit tests run: poetry run test-unit From 9d5ec4f2dddac2ae68492925c3191ca4b60b4722 Mon Sep 17 00:00:00 2001 From: Fszta <36471574+Fszta@users.noreply.github.com> Date: Sun, 27 Sep 2026 23:05:25 +0200 Subject: [PATCH 5/5] style: one-time repo-wide format Mechanical output of the pinned toolchain over the whole repo: black 25.12.0 plus ruff 0.16.3 --fix --unsafe-fixes (1894 fixes, mostly UP006/UP045 pyupgrade rewrites). No hand edits. Verified idempotent (a second run is a no-op) and behavior-preserving: mypy clean and the full suite identical before and after (760 passed, 1 skipped). Regenerable: if this commit conflicts at merge time, drop it and rerun on merged main: poetry run format && poetry run format && poetry run test (first run applies fixes and exits non-zero via --exit-non-zero-on-fix; the second must be a no-op) Co-Authored-By: Claude Fable 5 --- .github/scripts/gen_changelog.py | 11 +- parrant/artifacts/adapter_mapping.py | 11 +- parrant/artifacts/catalog.py | 9 +- parrant/artifacts/exceptions.py | 7 +- parrant/artifacts/manifest.py | 85 +++--- parrant/artifacts/registry.py | 124 ++++---- parrant/cli/main.py | 114 ++++---- parrant/lineage/backtest.py | 88 +++--- parrant/lineage/changeset.py | 128 ++++----- parrant/lineage/ci.py | 34 ++- parrant/lineage/display/__init__.py | 4 +- parrant/lineage/display/backtest.py | 6 +- parrant/lineage/display/base.py | 13 +- parrant/lineage/display/dot.py | 32 +-- parrant/lineage/display/html/explore.py | 187 +++++++------ parrant/lineage/display/json.py | 21 +- parrant/lineage/display/markdown.py | 80 +++--- parrant/lineage/display/text.py | 14 +- parrant/lineage/policy.py | 149 +++++----- parrant/lineage/policy_init.py | 46 +-- parrant/lineage/provider.py | 36 +-- parrant/lineage/semantic_diff.py | 4 +- parrant/lineage/service.py | 192 ++++++------- parrant/lineage/sqlglot_provider.py | 10 +- parrant/lineage/verdict.py | 60 ++-- parrant/metabase/artifact.py | 5 +- parrant/metabase/cli.py | 25 +- parrant/metabase/client.py | 33 +-- parrant/metabase/extract.py | 47 ++-- parrant/metabase/join.py | 14 +- parrant/metabase/pmbql.py | 16 +- parrant/metabase/reach.py | 49 ++-- parrant/metabase/resolvers.py | 86 +++--- parrant/metabase/warehouse_meta.py | 36 ++- parrant/models/schema.py | 264 +++++++++--------- parrant/parser/sql_parser.py | 227 +++++++-------- parrant/parser/sql_parser_utils.py | 25 +- scripts/run_tests.py | 5 +- tests/e2e/conftest.py | 9 +- tests/e2e/test_lineage_api.py | 37 +-- tests/e2e/test_unresolved_edges_e2e.py | 22 +- tests/fixtures/unresolved_edges/_build.py | 20 +- tests/integration/conftest.py | 11 +- .../test_compiled_sql_fallback_integration.py | 3 +- .../integration/test_exposures_integration.py | 165 ++++++----- .../test_json_output_integration.py | 1 - .../test_lineage_explorer_integration.py | 6 +- .../integration/test_manifest_integration.py | 2 - .../test_predicate_impact_integration.py | 4 +- .../integration/test_registry_integration.py | 8 +- tests/live/seed.py | 32 +-- tests/resources/dbt_test_project/setup.py | 9 +- tests/unit/artifacts/test_adapter_mapping.py | 2 +- tests/unit/artifacts/test_catalog.py | 1 + tests/unit/artifacts/test_manifest.py | 1 + tests/unit/artifacts/test_meta_index.py | 4 +- tests/unit/lineage/test_backtest.py | 1 - tests/unit/lineage/test_backtest_display.py | 24 +- tests/unit/lineage/test_changeset.py | 36 ++- .../lineage/test_confidence_completeness.py | 4 +- tests/unit/lineage/test_config_axis.py | 42 +-- tests/unit/lineage/test_inferred_meta.py | 9 +- .../lineage/test_manifest_seeded_universe.py | 3 +- .../lineage/test_markdown_confidence_cap.py | 4 +- tests/unit/lineage/test_metabase_markdown.py | 8 +- tests/unit/lineage/test_opaque_models.py | 7 +- tests/unit/lineage/test_policy_engine.py | 26 +- tests/unit/lineage/test_policy_init.py | 1 - tests/unit/lineage/test_resolution.py | 4 +- tests/unit/lineage/test_selection.py | 22 +- .../test_star_passthrough_catalog_missing.py | 9 +- .../test_unresolved_edges_propagation.py | 6 +- tests/unit/lineage/test_verdict.py | 23 +- tests/unit/metabase/_fixtures.py | 26 +- tests/unit/metabase/test_incremental.py | 26 +- tests/unit/metabase/test_join_reach.py | 2 +- tests/unit/metabase/test_pmbql.py | 6 +- tests/unit/metabase/test_resolvers.py | 2 +- tests/unit/parser/test_sql_parser.py | 9 +- tests/unit/registry/test_coverage.py | 2 +- tests/unit/registry/test_registry.py | 124 ++++---- tests/unit/registry/test_unresolved_edges.py | 6 +- tests/unit/test_lineage_provider.py | 58 ++-- 83 files changed, 1565 insertions(+), 1559 deletions(-) diff --git a/.github/scripts/gen_changelog.py b/.github/scripts/gen_changelog.py index 3811737..25e0aee 100755 --- a/.github/scripts/gen_changelog.py +++ b/.github/scripts/gen_changelog.py @@ -58,7 +58,12 @@ def main() -> int: m = COMMIT_RE.match(subj) if not m: continue - typ, scope, bang, desc = (m.group("type"), m.group("scope"), m.group("bang"), m.group("desc")) + typ, scope, bang, desc = ( + m.group("type"), + m.group("scope"), + m.group("bang"), + m.group("desc"), + ) # Release/CI plumbing is never user-facing: drop it even when a commit is # mis-typed as feat/fix (e.g. `fix(ci): ...`) instead of `ci: ...`. if scope in ("ci", "release"): @@ -72,9 +77,7 @@ def main() -> int: buckets[typ].append(bullet) if prev_tag: - header = ( - f"## [{new_version}]({repo_url}/compare/{prev_tag}...v{new_version}) ({date})\n" - ) + header = f"## [{new_version}]({repo_url}/compare/{prev_tag}...v{new_version}) ({date})\n" else: header = f"## {new_version} ({date})\n" diff --git a/parrant/artifacts/adapter_mapping.py b/parrant/artifacts/adapter_mapping.py index 505937d..5577a9f 100644 --- a/parrant/artifacts/adapter_mapping.py +++ b/parrant/artifacts/adapter_mapping.py @@ -16,7 +16,6 @@ """ import logging -from typing import Dict, Optional, Set logger = logging.getLogger(__name__) @@ -27,7 +26,7 @@ # bigquery, redshift, databricks, postgres, duckdb, spark, trino, presto, # athena, clickhouse, ...) an explicit identity entry is optional -- the # fallthrough in normalize_adapter returns the lowercased name unchanged. -ADAPTER_TO_DIALECT: Dict[str, str] = { +ADAPTER_TO_DIALECT: dict[str, str] = { # --- adapters whose dbt name differs from the sqlglot dialect name --- # The T-SQL family: dbt reports sqlserver/synapse/fabric, sqlglot uses "tsql". "sqlserver": "tsql", @@ -47,7 +46,7 @@ } -def _known_sqlglot_dialects() -> Set[str]: +def _known_sqlglot_dialects() -> set[str]: """Return the set of dialect names sqlglot actually supports. Derived from sqlglot at runtime so the check stays correct as sqlglot is @@ -62,14 +61,14 @@ def _known_sqlglot_dialects() -> Set[str]: return set() -_KNOWN_DIALECTS: Set[str] = _known_sqlglot_dialects() +_KNOWN_DIALECTS: set[str] = _known_sqlglot_dialects() # Track adapters we have already warned about so the warning fires once per # unresolved adapter instead of on every column parsed. -_warned_adapters: Set[str] = set() +_warned_adapters: set[str] = set() -def normalize_adapter(adapter_name: Optional[str]) -> Optional[str]: +def normalize_adapter(adapter_name: str | None) -> str | None: """Normalize a dbt adapter name to a sqlglot dialect name. If adapter_name is None or empty, returns it unchanged. diff --git a/parrant/artifacts/catalog.py b/parrant/artifacts/catalog.py index 43e93a1..d423c8e 100644 --- a/parrant/artifacts/catalog.py +++ b/parrant/artifacts/catalog.py @@ -1,13 +1,14 @@ -from pathlib import Path -from typing import Dict, Any import json +from pathlib import Path +from typing import Any + from parrant.models.schema import Model class CatalogReader: def __init__(self, catalog_path: str): self.catalog_path = Path(catalog_path) - self.catalog: Dict[str, Any] = {} + self.catalog: dict[str, Any] = {} def load(self) -> None: if not self.catalog_path.exists(): @@ -15,7 +16,7 @@ def load(self) -> None: with open(self.catalog_path, "r") as f: self.catalog = json.load(f) - def get_models_nodes(self) -> Dict[str, Model]: + def get_models_nodes(self) -> dict[str, Model]: models = {} nodes = self.catalog.get("nodes", {}) sources = self.catalog.get("sources", {}) diff --git a/parrant/artifacts/exceptions.py b/parrant/artifacts/exceptions.py index ae6c2b8..966e7e2 100644 --- a/parrant/artifacts/exceptions.py +++ b/parrant/artifacts/exceptions.py @@ -1,15 +1,14 @@ class RegistryError(Exception): """Base exception for all registry-related errors.""" - pass + class ModelNotFoundError(RegistryError): """Raised when a requested model is not found.""" - pass + class RegistryNotLoadedError(RegistryError): """Raised when trying to access registry before loading data.""" - pass + class RegistryLoadError(Exception): """Base exception for registry loading errors.""" - pass diff --git a/parrant/artifacts/manifest.py b/parrant/artifacts/manifest.py index e45283e..07e1f38 100644 --- a/parrant/artifacts/manifest.py +++ b/parrant/artifacts/manifest.py @@ -1,20 +1,19 @@ import json import os import re -from typing import Dict, List, Optional, Set, Any from pathlib import Path +from typing import Any from parrant.artifacts.adapter_mapping import normalize_adapter from parrant.models.schema import TestNode - # Matches the quoted name(s) inside a dbt ``ref(...)`` expression, e.g. # ``ref('stg_accounts')`` or ``ref('my_pkg', 'stg_accounts')``. The *last* quoted # token is the model name (the first, when present, is the package). _REF_QUOTED_RE = re.compile(r"""['"]([^'"]+)['"]""") -def _model_name_from_ref(ref_expr: Optional[str]) -> Optional[str]: +def _model_name_from_ref(ref_expr: str | None) -> str | None: """Extract the model name from a dbt ``ref(...)`` expression string. Returns ``None`` when nothing quoted can be found (e.g. a ``source(...)`` target @@ -28,7 +27,7 @@ def _model_name_from_ref(ref_expr: Optional[str]) -> Optional[str]: return matches[-1].lower() -def _model_name_from_unique_id(unique_id: Optional[str]) -> Optional[str]: +def _model_name_from_unique_id(unique_id: str | None) -> str | None: """Return the lowercased model name from a ``model..`` unique_id.""" if not unique_id: return None @@ -39,13 +38,13 @@ def _model_name_from_unique_id(unique_id: Optional[str]) -> Optional[str]: class ManifestReader: - def __init__(self, manifest_path: Optional[str] = None): + def __init__(self, manifest_path: str | None = None): self.manifest_path = Path(manifest_path) if manifest_path else None - self.manifest: Dict[str, Any] = {} + self.manifest: dict[str, Any] = {} # Lazily-built index of on-disk compiled SQL keyed by filename (e.g. ``orders.sql``), # used to recover a model's compiled SQL when the manifest's ``original_file_path`` # has drifted from the ``target/compiled`` layout (a model moved between builds). - self._compiled_index: Optional[Dict[str, List[Path]]] = None + self._compiled_index: dict[str, list[Path]] | None = None def load(self) -> None: if not self.manifest_path or not self.manifest_path.exists(): @@ -53,21 +52,21 @@ def load(self) -> None: with open(self.manifest_path, "r") as f: self.manifest = json.load(f) - def get_adapter(self) -> Optional[str]: + def get_adapter(self) -> str | None: adapter_name = self.manifest.get("metadata", {}).get("adapter_type") return normalize_adapter(adapter_name) - def _find_node(self, model_name: str) -> Optional[Dict[str, Any]]: + def _find_node(self, model_name: str) -> dict[str, Any] | None: """Find a node in the manifest by model name.""" if not self.manifest: return None model_name_lower = model_name.lower() - for _, node in self.manifest.get("nodes", {}).items(): + for node in self.manifest.get("nodes", {}).values(): if node.get("name", "").lower() == model_name_lower: return dict(node) return None - def get_model_dependencies(self) -> Dict[str, Set[str]]: + def get_model_dependencies(self) -> dict[str, set[str]]: """Return a dictionary of model dependencies with full model names. Returns: @@ -81,11 +80,11 @@ def get_model_dependencies(self) -> Dict[str, Set[str]]: dependencies[model_id] = depends_on return dependencies - def get_model_upstream(self) -> Dict[str, Set[str]]: + def get_model_upstream(self) -> dict[str, set[str]]: """Get upstream dependencies for each model.""" - upstream: Dict[str, Set[str]] = {} + upstream: dict[str, set[str]] = {} - for _, node in self.manifest.get("nodes", {}).items(): + for node in self.manifest.get("nodes", {}).values(): resource_type = node.get("resource_type") if resource_type in ("model", "snapshot"): model_name = node.get("name") @@ -116,9 +115,9 @@ def get_model_upstream(self) -> Dict[str, Set[str]]: return upstream - def get_model_downstream(self) -> Dict[str, Set[str]]: + def get_model_downstream(self) -> dict[str, set[str]]: """Return a dictionary of model downstream dependencies.""" - downstream: Dict[str, Set[str]] = {} + downstream: dict[str, set[str]] = {} upstream_deps = self.get_model_upstream() @@ -130,7 +129,7 @@ def get_model_downstream(self) -> Dict[str, Set[str]]: return downstream - def _resolve_compiled_file(self, node: Dict[str, Any]) -> Optional[Path]: + def _resolve_compiled_file(self, node: dict[str, Any]) -> Path | None: """Locate the on-disk compiled SQL file for a node. Many real manifests are produced without embedded ``compiled_code`` (e.g. @@ -175,9 +174,7 @@ def _resolve_compiled_file(self, node: Dict[str, Any]) -> Optional[Path]: return self._recover_compiled_by_name(Path(original_file_path).name, package_name) return None - def _recover_compiled_by_name( - self, filename: str, package_name: Optional[str] - ) -> Optional[Path]: + def _recover_compiled_by_name(self, filename: str, package_name: str | None) -> Path | None: """Find an on-disk compiled file by its ``.sql`` name, unambiguously. Prefers a single match under the model's own package dir; otherwise accepts a single @@ -195,11 +192,11 @@ def _recover_compiled_by_name( return scoped[0] return matches[0] if len(matches) == 1 else None - def _compiled_basename_index(self) -> Dict[str, List[Path]]: + def _compiled_basename_index(self) -> dict[str, list[Path]]: """Lazily index ``target/compiled/**/*.sql`` by filename → list of paths.""" if self._compiled_index is not None: return self._compiled_index - index: Dict[str, List[Path]] = {} + index: dict[str, list[Path]] = {} if self.manifest_path: compiled_dir = self.manifest_path.parent / "compiled" if compiled_dir.is_dir(): @@ -208,7 +205,7 @@ def _compiled_basename_index(self) -> Dict[str, List[Path]]: self._compiled_index = index return index - def get_compiled_sql(self, model_name: str) -> Optional[str]: + def get_compiled_sql(self, model_name: str) -> str | None: """Get compiled SQL for a model. Prefers SQL embedded in the manifest, falling back to the compiled file on @@ -233,8 +230,8 @@ def get_compiled_sql(self, model_name: str) -> Optional[str]: @staticmethod def _merged_meta( - top_meta: Optional[Dict[str, Any]], config_meta: Optional[Dict[str, Any]] - ) -> Dict[str, Any]: + top_meta: dict[str, Any] | None, config_meta: dict[str, Any] | None + ) -> dict[str, Any]: """Merge a node's two dbt meta locations, ``config.meta`` winning over top-level. dbt exposes user-authored meta at both ``node.meta`` (legacy) and @@ -242,14 +239,14 @@ def _merged_meta( ``config`` value is authoritative (it is what dbt itself resolves). Neither present yields an empty dict — meta is *absent*, never guessed. """ - merged: Dict[str, Any] = {} + merged: dict[str, Any] = {} if isinstance(top_meta, dict): merged.update(top_meta) if isinstance(config_meta, dict): merged.update(config_meta) return merged - def get_model_meta(self, model_name: str) -> Dict[str, Any]: + def get_model_meta(self, model_name: str) -> dict[str, Any]: """Merged user-authored dbt ``meta`` for a model (``config.meta`` over ``meta``). This is arbitrary consumer metadata — ANY key an author declared — captured @@ -262,7 +259,7 @@ def get_model_meta(self, model_name: str) -> Dict[str, Any]: config = node.get("config") or {} return self._merged_meta(node.get("meta"), config.get("meta")) - def get_model_config(self, model_name: str) -> Dict[str, Any]: + def get_model_config(self, model_name: str) -> dict[str, Any]: """The node's resolved dbt ``config`` dict for a model (``node.config``). This is the generic dbt config surface — ``grants``, ``materialized``, ``tags``, @@ -276,7 +273,7 @@ def get_model_config(self, model_name: str) -> Dict[str, Any]: config = node.get("config") or {} return dict(config) if isinstance(config, dict) else {} - def get_column_meta(self, model_name: str) -> Dict[str, Dict[str, Any]]: + def get_column_meta(self, model_name: str) -> dict[str, dict[str, Any]]: """Per-column merged user meta for a model, keyed by lowercased column name. Each column's meta merges ``columns..config.meta`` over ``columns..meta`` @@ -287,7 +284,7 @@ def get_column_meta(self, model_name: str) -> Dict[str, Dict[str, Any]]: node = self._find_node(model_name) if not node: return {} - result: Dict[str, Dict[str, Any]] = {} + result: dict[str, dict[str, Any]] = {} for col_name, col_data in (node.get("columns") or {}).items(): col_data = col_data or {} col_config = col_data.get("config") or {} @@ -296,7 +293,7 @@ def get_column_meta(self, model_name: str) -> Dict[str, Dict[str, Any]]: ) return result - def get_model_path(self, model_name: str) -> Optional[str]: + def get_model_path(self, model_name: str) -> str | None: """Get the path to the model from the manifest.""" node = self._find_node(model_name) if not node: @@ -304,27 +301,27 @@ def get_model_path(self, model_name: str) -> Optional[str]: return node.get("path") - def get_model_language(self, model_name: str) -> Optional[str]: + def get_model_language(self, model_name: str) -> str | None: """Get the language of a model from the manifest.""" node = self._find_node(model_name) if not node: return None return node.get("language") - def get_model_resource_path(self, model_name: str) -> Optional[str]: + def get_model_resource_path(self, model_name: str) -> str | None: """Get the original file path of a model from the manifest.""" node = self._find_node(model_name) if not node: return None return node.get("original_file_path") - def get_node(self, node_id: str) -> Optional[Dict[str, Any]]: + def get_node(self, node_id: str) -> dict[str, Any] | None: node = self.manifest.get("nodes", {}).get(node_id) if node is None: return None return dict(node) - def get_tests(self) -> List[TestNode]: + def get_tests(self) -> list[TestNode]: """Read dbt test nodes (``resource_type == "test"``) from the manifest. We never run the tests; we read what they *declare*. For each test we extract: @@ -339,7 +336,7 @@ def get_tests(self) -> List[TestNode]: unknown field set to ``None`` (never guessed), so the reverse index can report coverage honestly. """ - tests: List[TestNode] = [] + tests: list[TestNode] = [] for node_id, node in self.manifest.get("nodes", {}).items(): if node.get("resource_type") != "test": @@ -375,8 +372,8 @@ def get_tests(self) -> List[TestNode]: if len(model_deps) == 1: target_model = model_deps[0] - referenced_model: Optional[str] = None - referenced_column: Optional[str] = None + referenced_model: str | None = None + referenced_column: str | None = None if test_name == "relationships": referenced_model = _model_name_from_ref(kwargs.get("to")) field = kwargs.get("field") @@ -397,7 +394,7 @@ def get_tests(self) -> List[TestNode]: return tests - def get_exposures(self) -> Dict[str, Dict[str, Any]]: + def get_exposures(self) -> dict[str, dict[str, Any]]: """Get all exposures from the manifest. Returns: @@ -405,15 +402,15 @@ def get_exposures(self) -> Dict[str, Dict[str, Any]]: """ return self.manifest.get("exposures", {}) - def get_exposure_dependencies(self) -> Dict[str, Set[str]]: + def get_exposure_dependencies(self) -> dict[str, set[str]]: """Get model dependencies for each exposure. Returns: Dict[str, Set[str]]: Key is exposure name, value is set of model names it depends on """ - exposure_deps: Dict[str, Set[str]] = {} + exposure_deps: dict[str, set[str]] = {} - for exposure_id, exposure_data in self.manifest.get("exposures", {}).items(): + for exposure_data in self.manifest.get("exposures", {}).values(): exposure_name = exposure_data.get("name") if not exposure_name: continue @@ -440,13 +437,13 @@ def get_exposure_dependencies(self) -> Dict[str, Set[str]]: return exposure_deps - def get_model_exposures(self) -> Dict[str, Set[str]]: + def get_model_exposures(self) -> dict[str, set[str]]: """Get exposures that depend on each model. Returns: Dict[str, Set[str]]: Key is model name, value is set of exposure names that depend on it """ - model_exposures: Dict[str, Set[str]] = {} + model_exposures: dict[str, set[str]] = {} exposure_deps = self.get_exposure_dependencies() diff --git a/parrant/artifacts/registry.py b/parrant/artifacts/registry.py index ae7cfc1..09b4f2a 100644 --- a/parrant/artifacts/registry.py +++ b/parrant/artifacts/registry.py @@ -1,23 +1,23 @@ -from typing import Any, Dict, List, Optional, Set, Tuple -from dataclasses import dataclass, field import logging +from dataclasses import dataclass, field +from typing import Any from parrant.artifacts.catalog import CatalogReader +from parrant.artifacts.exceptions import ( + ModelNotFoundError, + RegistryError, + RegistryNotLoadedError, +) from parrant.artifacts.manifest import ManifestReader from parrant.models.schema import ( - Model, Column, - SQLParseResult, ColumnLineage, - Exposure, Coverage, + Exposure, + Model, + SQLParseResult, TestNode, ) -from parrant.artifacts.exceptions import ( - ModelNotFoundError, - RegistryNotLoadedError, - RegistryError, -) from parrant.parser import SQLColumnParser logger = logging.getLogger(__name__) @@ -43,17 +43,17 @@ class ParseStats: # them is preserved from the manifest dependency graph. opaque: int = 0 skipped_no_sql: int = 0 - failed_model_names: List[str] = field(default_factory=list) - opaque_model_names: List[str] = field(default_factory=list) - skipped_model_names: List[str] = field(default_factory=list) + failed_model_names: list[str] = field(default_factory=list) + opaque_model_names: list[str] = field(default_factory=list) + skipped_model_names: list[str] = field(default_factory=list) @dataclass class RegistryState: """Immutable state of the registry.""" - models: Dict[str, Model] - exposures: Dict[str, Exposure] + models: dict[str, Model] + exposures: dict[str, Exposure] is_loaded: bool = False @@ -62,14 +62,14 @@ def __init__( self, catalog_path: str, manifest_path: str, - adapter_override: Optional[str] = None, + adapter_override: str | None = None, ): self._catalog_reader = CatalogReader(catalog_path) self._manifest_reader = ManifestReader(manifest_path) self._state = RegistryState(models={}, exposures={}, is_loaded=False) - self._sql_parser: Optional[SQLColumnParser] = None - self._dialect: Optional[str] = None - self._adapter_override: Optional[str] = adapter_override + self._sql_parser: SQLColumnParser | None = None + self._dialect: str | None = None + self._adapter_override: str | None = adapter_override self._parse_stats: ParseStats = ParseStats() # Names of model-like nodes that have a real catalog entry (data types known). # A manifest node absent from this set is "catalog-missing": still analyzable via @@ -77,32 +77,32 @@ def __init__( self._catalog_backed_model_names: set = set() # Lazily-built reverse index: upstream column -> models that reference it ONLY in a # predicate (filter/join), i.e. a row-set dependency rather than a value one. - self._filter_dependents: Optional[Dict[str, set]] = None + self._filter_dependents: dict[str, set] | None = None # Reverse index built at load time: (model, column) -> tests targeting that column. # Keys are lowercased to match the codebase's case-insensitive model/column naming. - self._column_tests: Dict[Tuple[str, str], List[TestNode]] = {} + self._column_tests: dict[tuple[str, str], list[TestNode]] = {} # Reverse index for the *referenced* side of relationships tests: (model, column) -> # relationships tests pointing AT that column via ``to=``/``field=``. Removing this # parent key breaks the child's relationships test just as surely as removing the # child column does, so it is a distinct provable-break lookup. - self._referenced_tests: Dict[Tuple[str, str], List[TestNode]] = {} + self._referenced_tests: dict[tuple[str, str], list[TestNode]] = {} # Tests we could not attribute to a (model, column) pair — kept for coverage honesty # (counted, never guessed at). See :meth:`get_unattributable_test_count`. - self._unattributable_tests: List[TestNode] = [] + self._unattributable_tests: list[TestNode] = [] # Every test node's unique_id present in this manifest. Lets the verdict classifier # confirm a base test STILL EXISTS in head before flagging it broken — so a rename # that updates the test's yml (new unique_id) is not a false break. - self._test_unique_ids: Set[str] = set() + self._test_unique_ids: set[str] = set() # model (lowercased) -> every test that breaks if the whole model is removed: those # attached to it AND relationships tests referencing it. Column-level recovery can # miss a model's tested columns, but a wholly-removed model breaks all of its tests. - self._model_tests: Dict[str, List[TestNode]] = {} + self._model_tests: dict[str, list[TestNode]] = {} @property def is_loaded(self) -> bool: return self._state.is_loaded - def _initialize_models(self) -> Dict[str, Model]: + def _initialize_models(self) -> dict[str, Model]: """Initialize the model universe from the *manifest*, enriched by the catalog. The manifest is the source of truth for which models exist: it lists every @@ -124,7 +124,7 @@ def _initialize_models(self) -> Dict[str, Model]: except Exception as e: raise RegistryError(f"Failed to initialize models: {e}") - models: Dict[str, Model] = {} + models: dict[str, Model] = {} catalog_backed: set = set() # 1) Seed the universe from manifest model-like nodes (model/snapshot/seed). @@ -168,7 +168,7 @@ def _initialize_models(self) -> Dict[str, Model]: raise RegistryError("No models found in manifest or catalog") return models - def _apply_dependencies(self, models: Dict[str, Model]) -> None: + def _apply_dependencies(self, models: dict[str, Model]) -> None: """Apply upstream and downstream dependencies to models.""" try: upstream_deps = self._manifest_reader.get_model_upstream() @@ -176,7 +176,7 @@ def _apply_dependencies(self, models: Dict[str, Model]) -> None: model_exposures = self._manifest_reader.get_model_exposures() manifest_sources = self._manifest_reader.manifest.get("sources", {}) - for source_id, source_node in manifest_sources.items(): + for source_node in manifest_sources.values(): source_name = source_node.get("source_name") source_identifier = ( source_node.get("identifier", "").lower() @@ -202,7 +202,7 @@ def _apply_dependencies(self, models: Dict[str, Model]) -> None: except Exception as e: raise RegistryError(f"Failed to apply dependencies: {e}") - def _apply_descriptions(self, models: Dict[str, Model]) -> None: + def _apply_descriptions(self, models: dict[str, Model]) -> None: """Populate model and column descriptions from the dbt-authored docs. The *manifest* is the primary source: dbt records the docs a person wrote in @@ -234,7 +234,7 @@ def _apply_descriptions(self, models: Dict[str, Model]) -> None: if manifest_desc: column.description = manifest_desc - def _apply_meta(self, models: Dict[str, Model]) -> None: + def _apply_meta(self, models: dict[str, Model]) -> None: """Attach arbitrary dbt ``meta`` from the manifest onto models and columns. User-authored meta (ANY key) is namespaced under ``Model.metadata["dbt_meta"]`` @@ -269,7 +269,7 @@ def _apply_meta(self, models: Dict[str, Model]) -> None: if col_meta: column.metadata = col_meta - def _load_exposures(self) -> Dict[str, Exposure]: + def _load_exposures(self) -> dict[str, Exposure]: """Load exposures from manifest.""" exposures = {} exposure_data = self._manifest_reader.get_exposures() @@ -311,7 +311,7 @@ def _build_test_index(self) -> None: self._test_unique_ids = set() self._model_tests = {} - def _attach_to_model(model_name: Optional[str], test: TestNode) -> None: + def _attach_to_model(model_name: str | None, test: TestNode) -> None: if model_name is None: return bucket = self._model_tests.setdefault(model_name.lower(), []) @@ -337,14 +337,14 @@ def _attach_to_model(model_name: Optional[str], test: TestNode) -> None: key = (test.target_model.lower(), test.target_column.lower()) self._column_tests.setdefault(key, []).append(test) - def get_column_tests(self, model: str, column: str) -> List[TestNode]: + def get_column_tests(self, model: str, column: str) -> list[TestNode]: """Return the dbt tests targeting ``model.column`` (case-insensitive). Returns an empty list for an unknown (model, column) pair or one with no tests. """ return list(self._column_tests.get((model.lower(), column.lower()), [])) - def get_tests_referencing(self, model: str, column: str) -> List[TestNode]: + def get_tests_referencing(self, model: str, column: str) -> list[TestNode]: """Return relationships tests whose *referenced* (parent) side is ``model.column``. These break when the parent key is removed/renamed, distinct from the tests that @@ -353,7 +353,7 @@ def get_tests_referencing(self, model: str, column: str) -> List[TestNode]: """ return list(self._referenced_tests.get((model.lower(), column.lower()), [])) - def get_model_tests(self, model: str) -> List[TestNode]: + def get_model_tests(self, model: str) -> list[TestNode]: """Every test that breaks if ``model`` is removed wholesale (case-insensitive). Includes tests attached to the model and relationships tests referencing it — used @@ -362,7 +362,7 @@ def get_model_tests(self, model: str) -> List[TestNode]: """ return list(self._model_tests.get(model.lower(), [])) - def get_test_unique_ids(self) -> Set[str]: + def get_test_unique_ids(self) -> set[str]: """All dbt test unique_ids present in this manifest. The verdict classifier intersects a base test against this head set to confirm the @@ -379,11 +379,11 @@ def get_unattributable_test_count(self) -> int: """ return len(self._unattributable_tests) - def get_unattributable_tests(self) -> List[TestNode]: + def get_unattributable_tests(self) -> list[TestNode]: """The test nodes whose (model, column) target could not be attributed.""" return list(self._unattributable_tests) - def _process_lineage(self, models: Dict[str, Model]) -> None: + def _process_lineage(self, models: dict[str, Model]) -> None: """Process and apply column lineage to models.""" logger = logging.getLogger(__name__) @@ -501,7 +501,7 @@ def _mark_opaque(self, model: Model) -> None: model.upstream = set(model.upstream or set()) | set(manifest_upstream) def _apply_column_lineage( - self, model: Model, parse_result: SQLParseResult, models: Dict[str, Model] + self, model: Model, parse_result: SQLParseResult, models: dict[str, Model] ) -> None: """Apply parsed lineage to model columns. @@ -534,7 +534,7 @@ def _apply_column_lineage( self._declare_unresolved_edges(model, parse_result, models) def _declare_unresolved_edges( - self, model: Model, parse_result: SQLParseResult, models: Dict[str, Model] + self, model: Model, parse_result: SQLParseResult, models: dict[str, Model] ) -> None: """Finalize the model's unresolved-edge markers and stamp them onto its metadata. @@ -565,7 +565,7 @@ def _declare_unresolved_edges( The finalized set is stored on ``model.metadata["unresolved_edges"]`` as a list of dicts, mirroring ``star_sources`` — the complete, uncapped machine surface the resolution/confidence pass consumes. """ - markers: List[Dict[str, Any]] = [] + markers: list[dict[str, Any]] = [] # 1. Parser markers: stamp the model name. for edge in parse_result.unresolved_edges: @@ -584,7 +584,7 @@ def _declare_unresolved_edges( for col_name, column in model.columns.items(): for lineage in column.lineage or []: - kept: Set[str] = set() + kept: set[str] = set() for token in lineage.source_columns: reason = self._registry_phantom_reason( token, phantom_bases, declared_upstreams, models @@ -609,10 +609,10 @@ def _declare_unresolved_edges( def _registry_phantom_reason( self, token: str, - phantom_bases: Set[str], - declared_upstreams: Set[str], - models: Dict[str, Model], - ) -> Optional[str]: + phantom_bases: set[str], + declared_upstreams: set[str], + models: dict[str, Model], + ) -> str | None: """Classify a source token against registry ground truth, or ``None`` to keep it. * ``unexpandable_star`` — qualifier is a leaked ``select *`` base, not a declared upstream. @@ -649,10 +649,10 @@ def _token_column(token: str) -> str: return tail.strip().strip('"').lower() @staticmethod - def _dedupe_edge_records(records: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + def _dedupe_edge_records(records: list[dict[str, Any]]) -> list[dict[str, Any]]: """Order-stable de-dup of marker dicts on (column, reason, detail).""" - seen: Set[Tuple[Any, Any, Any]] = set() - unique: List[Dict[str, Any]] = [] + seen: set[tuple[Any, Any, Any]] = set() + unique: list[dict[str, Any]] = [] for record in records: key = (record.get("column"), record.get("reason"), record.get("detail")) if key not in seen: @@ -660,7 +660,7 @@ def _dedupe_edge_records(records: List[Dict[str, Any]]) -> List[Dict[str, Any]]: unique.append(record) return unique - def _process_star_references(self, models: Dict[str, Model]) -> None: + def _process_star_references(self, models: dict[str, Model]) -> None: """Process star references between models.""" for model in models.values(): if not model.metadata or "star_sources" not in model.metadata: @@ -689,7 +689,7 @@ def _apply_star_columns(self, target: Model, source_name: str, source: Model) -> catalog-missing branch of :meth:`_apply_column_lineage`. """ catalog_missing = bool(target.metadata and target.metadata.get("catalog_missing")) - for col_name, source_col in source.columns.items(): + for col_name in source.columns: if col_name not in target.columns: if not catalog_missing: continue @@ -746,13 +746,13 @@ def load(self) -> None: except Exception as e: raise RegistryError(f"Failed to load registry: {e}") - def get_models(self) -> Dict[str, Model]: + def get_models(self) -> dict[str, Model]: """Get all models in the registry.""" if not self.is_loaded: raise RegistryNotLoadedError("Registry must be loaded before accessing models") return self._state.models - def get_dialect(self) -> Optional[str]: + def get_dialect(self) -> str | None: """Return the resolved SQL dialect (adapter), or ``None`` when unknown. Public accessor over the dialect the registry already computes at load time @@ -771,7 +771,7 @@ def get_model(self, model_name: str) -> Model: raise ModelNotFoundError(f"Model '{model_name}' not found") return model - def get_model_dbt_meta(self, model: str) -> Dict[str, Any]: + def get_model_dbt_meta(self, model: str) -> dict[str, Any]: """Arbitrary user-authored dbt ``meta`` for a model (case-insensitive). Reads the meta namespaced under ``Model.metadata["dbt_meta"]`` by @@ -784,7 +784,7 @@ def get_model_dbt_meta(self, model: str) -> Dict[str, Any]: return {} return dict(model_obj.metadata.get("dbt_meta") or {}) - def get_model_config(self, model: str) -> Dict[str, Any]: + def get_model_config(self, model: str) -> dict[str, Any]: """The node's resolved dbt ``config`` dict for a model (case-insensitive). Reads the config namespaced under ``Model.metadata["dbt_config"]`` by @@ -798,7 +798,7 @@ def get_model_config(self, model: str) -> Dict[str, Any]: return {} return dict(model_obj.metadata.get("dbt_config") or {}) - def get_column_dbt_meta(self, model: str, column: str) -> Dict[str, Any]: + def get_column_dbt_meta(self, model: str, column: str) -> dict[str, Any]: """Arbitrary user-authored dbt ``meta`` for a column (case-insensitive). Reads ``Column.metadata`` populated by :meth:`_apply_meta`. Returns an empty dict @@ -812,7 +812,7 @@ def get_column_dbt_meta(self, model: str, column: str) -> Dict[str, Any]: return {} return dict(col.metadata) - def get_exposures(self) -> Dict[str, Exposure]: + def get_exposures(self) -> dict[str, Exposure]: """Get all exposures in the registry.""" if not self.is_loaded: raise RegistryNotLoadedError("Registry must be loaded before accessing exposures") @@ -899,7 +899,7 @@ def is_catalog_backed(self, model_name: str) -> bool: """Whether a model has a real catalog entry (known column types).""" return model_name.lower() in self._catalog_backed_model_names - def get_manifest_downstream(self) -> Dict[str, set]: + def get_manifest_downstream(self) -> dict[str, set]: """Manifest-level downstream child map, covering every model (not just catalog ones).""" return self._manifest_reader.get_model_downstream() @@ -912,7 +912,7 @@ def get_filter_dependents(self, source_column: str) -> set: is already reported as a value impact), so this stays the purely-predicate set. """ if self._filter_dependents is None: - index: Dict[str, set] = {} + index: dict[str, set] = {} for name, model in self.get_models().items(): projected: set = set() for column in model.columns.values(): @@ -931,7 +931,7 @@ def _check_loaded(self) -> None: if not self._state.models: raise RegistryNotLoadedError("Registry must be loaded before accessing models") - def _find_compiled_sql(self, model_name: str) -> Optional[str]: + def _find_compiled_sql(self, model_name: str) -> str | None: """Find compiled SQL for a model from manifest or target file.""" self._check_loaded() model_name_lower = model_name.lower() @@ -953,7 +953,7 @@ def _find_compiled_sql(self, model_name: str) -> Optional[str]: compiled_sql = f.read() model.compiled_sql = compiled_sql return compiled_sql - except (FileNotFoundError, IOError): + except (OSError, FileNotFoundError): pass return None diff --git a/parrant/cli/main.py b/parrant/cli/main.py index 764643e..7b96fc0 100644 --- a/parrant/cli/main.py +++ b/parrant/cli/main.py @@ -1,9 +1,10 @@ import json +import logging import sys from pathlib import Path +from typing import Any + import click -import logging -from typing import Any, Dict, List, Optional from parrant.lineage.changeset import ( ChangesetBuilder, @@ -14,6 +15,11 @@ git_changed_models, scope_changes_to_models, ) +from parrant.lineage.display import DotDisplay, JsonDisplay, TextDisplay +from parrant.lineage.display.base import LineageStaticDisplay +from parrant.lineage.display.html.explore import LineageExplorer +from parrant.lineage.display.markdown import render_changeset_markdown +from parrant.lineage.service import LineageSelector, LineageService from parrant.lineage.verdict import ( applied_overrides, break_is_overridden, @@ -21,12 +27,6 @@ decide_verdict, ineffective_overrides, ) -from parrant.lineage.display import TextDisplay, DotDisplay, JsonDisplay -from parrant.lineage.display.html.explore import LineageExplorer -from parrant.lineage.display.markdown import render_changeset_markdown -from parrant.lineage.service import LineageService, LineageSelector -from parrant.lineage.display.base import LineageStaticDisplay - logging.basicConfig(level=logging.INFO, format="%(levelname)s - %(message)s") @@ -124,12 +124,12 @@ def cli( format: str, output: str, port: int, - adapter: Optional[str], - base_manifest: Optional[str], - base_catalog: Optional[str], - git_base: Optional[str], - policy_path: Optional[str], - metabase_path: Optional[str], + adapter: str | None, + base_manifest: str | None, + base_catalog: str | None, + git_base: str | None, + policy_path: str | None, + metabase_path: str | None, no_overrides: bool, ) -> None: """Parrant - column-level lineage and change-impact for dbt (parry breaks, warrant safe).""" @@ -238,7 +238,7 @@ def cli( if format == "dot": display.save() else: - available_columns = ", ".join(model.columns.keys()) + ", ".join(model.columns.keys()) click.echo( f"Error: Column '{selector.column}' not found in model '{selector.model}'", err=True, @@ -262,21 +262,21 @@ def cli( click.echo(f" {downstream}") except Exception as e: - click.echo(f"Error: {str(e)}", err=True) + click.echo(f"Error: {e!s}", err=True) sys.exit(1) def _build_explore_change_context( head_service: LineageService, *, - adapter: Optional[str], - base_manifest: Optional[str], - base_catalog: Optional[str], - git_base: Optional[str], - policy_path: Optional[str], - metabase_path: Optional[str], + adapter: str | None, + base_manifest: str | None, + base_catalog: str | None, + git_base: str | None, + policy_path: str | None, + metabase_path: str | None, no_overrides: bool = False, -) -> Optional[Dict[str, Any]]: +) -> dict[str, Any] | None: """Assemble the changeset report the explorer surfaces, or ``None`` when no change source was supplied (pure-explore mode). @@ -289,10 +289,10 @@ def _build_explore_change_context( return None honor_overrides = not no_overrides - stale_overrides: List[Dict[str, object]] = [] - override_warnings: List[str] = [] + stale_overrides: list[dict[str, object]] = [] + override_warnings: list[str] = [] - base_service: Optional[LineageService] = None + base_service: LineageService | None = None if base_manifest: resolved_base_catalog = base_catalog if not resolved_base_catalog: @@ -363,7 +363,7 @@ def _build_explore_change_context( ) by_change_list = aggregated.get("by_change") if isinstance(aggregated, dict) else None summary_obj = report.get("summary", {}) - summary: Dict[str, Any] = summary_obj if isinstance(summary_obj, dict) else {} + summary: dict[str, Any] = summary_obj if isinstance(summary_obj, dict) else {} report["verdict"] = decide_verdict(breaks, summary, changes, by_change=by_change_list) # Mirror the impact() gate: unexcused (blocking) breaks only; excused ones surface as # allow-break override records so the explorer shows the same signals as CI. @@ -527,22 +527,22 @@ def _build_explore_change_context( def impact( manifest: str, catalog: str, - base_manifest: Optional[str], - base_catalog: Optional[str], - git_base: Optional[str], - scope_git: Optional[str], + base_manifest: str | None, + base_catalog: str | None, + git_base: str | None, + scope_git: str | None, format: str, - adapter: Optional[str], + adapter: str | None, explain: bool, no_overrides: bool, ci: bool, fail_on: str, - policy_path: Optional[str], - metabase_path: Optional[str], + policy_path: str | None, + metabase_path: str | None, emit_selector: bool, - github_token: Optional[str], - repo: Optional[str], - pr_number: Optional[int], + github_token: str | None, + repo: str | None, + pr_number: int | None, ) -> None: """Diff-driven impact: assess the blast radius of a whole change (PR). @@ -582,11 +582,11 @@ def impact( honor_overrides = not no_overrides # Override side-outputs (populated by the changeset build below): stale directives # (no matching change) and parse warnings (malformed pragmas). Empty under --no-overrides. - stale_overrides: List[Dict[str, object]] = [] - override_warnings: List[str] = [] + stale_overrides: list[dict[str, object]] = [] + override_warnings: list[str] = [] - base_service: Optional[LineageService] = None - changes: List[ColumnChange] + base_service: LineageService | None = None + changes: list[ColumnChange] # Whether structural checks (added/removed/type_changed) could run. They need a # real catalog on both sides; the two-manifest path decides this from the builder # below. The git-diff fallback is a separate, self-evident coarse mode, so it is @@ -690,7 +690,7 @@ def impact( ) by_change_list = aggregated.get("by_change") if isinstance(aggregated, dict) else None summary_obj = report.get("summary", {}) - summary: Dict[str, Any] = summary_obj if isinstance(summary_obj, dict) else {} + summary: dict[str, Any] = summary_obj if isinstance(summary_obj, dict) else {} report["verdict"] = decide_verdict(breaks, summary, changes, by_change=by_change_list) # a break excused by an allow-break override is DEMOTED — it must not keep the # gate armed. Split the breaks so report/gate reflect only the UNEXCUSED (blocking) @@ -760,7 +760,7 @@ def impact( click.echo(render_changeset_markdown(report, explain=explain)) except Exception as e: - click.echo(f"Error: {str(e)}", err=True) + click.echo(f"Error: {e!s}", err=True) sys.exit(1) # Selector emission is a pure side-channel to $GITHUB_OUTPUT: independent of --ci, it posts @@ -852,15 +852,15 @@ def policy_group() -> None: def policy_test( manifest: str, catalog: str, - adapter: Optional[str], + adapter: str | None, policy_path: str, - git_range: Optional[str], - last: Optional[int], - changesets_dir: Optional[str], - repo_dir: Optional[str], + git_range: str | None, + last: int | None, + changesets_dir: str | None, + repo_dir: str | None, fmt: str, fail_on: str, - baseline_path: Optional[str], + baseline_path: str | None, ) -> None: """Backtest a candidate policy over git history (or a saved changeset corpus). @@ -901,8 +901,8 @@ def policy_test( ) sys.exit(1) - report: Optional[BacktestReport] = None - baseline: Optional[BacktestReport] = None + report: BacktestReport | None = None + baseline: BacktestReport | None = None try: # Resolve the policy up front so a broken file fails loudly (PolicyConfigError -> exit 1), # never treated as "no policy". @@ -939,7 +939,7 @@ def policy_test( else: click.echo(render_backtest_table(report)) except Exception as exc: - click.echo(f"Error: {str(exc)}", err=True) + click.echo(f"Error: {exc!s}", err=True) sys.exit(1) # The gate exit is OUTSIDE the try/except (mirrors impact's CI gate): a tripped gate is a @@ -987,7 +987,7 @@ def policy_test( def policy_init( manifest: str, catalog: str, - adapter: Optional[str], + adapter: str | None, output: str, force: bool, stdout: bool, @@ -1042,7 +1042,7 @@ def _relation_name_resolver(registry: Any): if reader is None or not hasattr(reader, "_find_node"): return None - def _resolve(model_name: str) -> Optional[str]: + def _resolve(model_name: str) -> str | None: node = reader._find_node(model_name) if not node: return None @@ -1055,9 +1055,9 @@ def _resolve(model_name: str) -> Optional[str]: def _run_ci( report: dict, fail_on_value: str, - token: Optional[str], - repo: Optional[str], - pr_number: Optional[int], + token: str | None, + repo: str | None, + pr_number: int | None, explain: bool = False, ) -> None: """Post the sticky PR comment (best-effort) and exit per the severity gate.""" diff --git a/parrant/lineage/backtest.py b/parrant/lineage/backtest.py index 47dc347..f7cd2c8 100644 --- a/parrant/lineage/backtest.py +++ b/parrant/lineage/backtest.py @@ -24,7 +24,7 @@ import subprocess import sys from collections import defaultdict -from typing import Any, Dict, List, Optional, Tuple +from typing import Any from parrant.lineage.changeset import ( ChangeKind, @@ -73,7 +73,7 @@ def _changesets_fidelity_note() -> str: ) -def _changes_for_models(head: LineageService, models: List[str]) -> List[ColumnChange]: +def _changes_for_models(head: LineageService, models: list[str]) -> list[ColumnChange]: """Expand a set of touched models into coarse ``logic_changed`` / ``INDETERMINATE`` changes. Mirrors :func:`build_git_changeset`'s inner loop but takes the already-computed model set, so @@ -81,7 +81,7 @@ def _changes_for_models(head: LineageService, models: List[str]) -> List[ColumnC same diff) instead of shelling out twice. """ head_models = head.registry.get_models() - chosen: Dict[Tuple[str, str], ColumnChange] = {} + chosen: dict[tuple[str, str], ColumnChange] = {} for model_name in models: model = head_models[model_name] for column in sorted(model.columns): @@ -95,7 +95,7 @@ def _changes_for_models(head: LineageService, models: List[str]) -> List[ColumnC return sorted(chosen.values(), key=lambda c: (c.model, c.column)) -def _parent_ref(sha: str, repo_dir: Optional[str]) -> str: +def _parent_ref(sha: str, repo_dir: str | None) -> str: """The first-parent ref to diff ``sha`` against, or the empty tree for a root commit. Squash-merge repos are linear (the spec's stated assumption), so first-parent ``^`` is the @@ -117,12 +117,12 @@ def _parent_ref(sha: str, repo_dir: Optional[str]) -> str: def _replay_point( head_service: LineageService, policy: Policy, - changes: List[ColumnChange], + changes: list[ColumnChange], ref: str, source: str, unmapped: int, - parse_failures: List[str], -) -> Tuple[BacktestPointResult, PolicyVerdict]: + parse_failures: list[str], +) -> tuple[BacktestPointResult, PolicyVerdict]: """Run the existing impact + breaks + policy pipeline for one point. In git-diff / changesets mode there is no base registry, so ``classify_provable_breaks`` @@ -163,9 +163,9 @@ def _replay_point( def _aggregate_rule_stats( - points_verdicts: List[Tuple[BacktestPointResult, PolicyVerdict]], + points_verdicts: list[tuple[BacktestPointResult, PolicyVerdict]], policy: Policy, -) -> List[BacktestRuleStat]: +) -> list[BacktestRuleStat]: """Per-rule aggregate across the whole range. Rows: every ``policy.rules`` id (always present, so a rule that never fired reads as a @@ -177,11 +177,11 @@ def _aggregate_rule_stats( policy_ids = [rule.id for rule in policy.rules] policy_id_set = set(policy_ids) - fired_total: Dict[str, int] = defaultdict(int) - fired_unknown: Dict[str, int] = defaultdict(int) - block_prs: Dict[str, int] = defaultdict(int) - warn_prs: Dict[str, int] = defaultdict(int) - extra_order: List[str] = [] + fired_total: dict[str, int] = defaultdict(int) + fired_unknown: dict[str, int] = defaultdict(int) + block_prs: dict[str, int] = defaultdict(int) + warn_prs: dict[str, int] = defaultdict(int) + extra_order: list[str] = [] for _point, verdict in points_verdicts: block_here: set[str] = set() @@ -202,7 +202,7 @@ def _aggregate_rule_stats( for rid in warn_here: warn_prs[rid] += 1 - stats: List[BacktestRuleStat] = [] + stats: list[BacktestRuleStat] = [] for rid in policy_ids + extra_order: total = fired_total.get(rid, 0) stats.append( @@ -218,7 +218,7 @@ def _aggregate_rule_stats( return stats -def _load_changesets_dir(path: str) -> List[Tuple[str, List[ColumnChange]]]: +def _load_changesets_dir(path: str) -> list[tuple[str, list[ColumnChange]]]: """Read ``*.json`` from ``path`` and reconstruct each into ``(filename, changes)``. Accepts either a bare change-list, ``{"changes": [...]}``, or the full changeset report shape @@ -226,7 +226,7 @@ def _load_changesets_dir(path: str) -> List[Tuple[str, List[ColumnChange]]]: """ if not os.path.isdir(path): raise RuntimeError(f"--changesets path is not a directory: '{path}'") - out: List[Tuple[str, List[ColumnChange]]] = [] + out: list[tuple[str, list[ColumnChange]]] = [] for name in sorted(os.listdir(path)): if not name.endswith(".json"): continue @@ -238,7 +238,7 @@ def _load_changesets_dir(path: str) -> List[Tuple[str, List[ColumnChange]]]: return out -def _extract_change_entries(data: Any) -> List[Dict[str, Any]]: +def _extract_change_entries(data: Any) -> list[dict[str, Any]]: if isinstance(data, list): return data if isinstance(data, dict): @@ -254,7 +254,7 @@ def _extract_change_entries(data: Any) -> List[Dict[str, Any]]: ) -def _resolve_git_range(git_range: Optional[str], last: Optional[int]) -> Tuple[str, str]: +def _resolve_git_range(git_range: str | None, last: int | None) -> tuple[str, str]: """Resolve ``(base, head)`` from ``--git-range base..head`` or ``--last N`` sugar.""" if last is not None: if last < 1: @@ -275,11 +275,11 @@ def run_backtest( head_service: LineageService, policy: Policy, *, - git_range: Optional[str] = None, - last: Optional[int] = None, - changesets_dir: Optional[str] = None, - repo_dir: Optional[str] = None, - baseline: Optional[BacktestReport] = None, + git_range: str | None = None, + last: int | None = None, + changesets_dir: str | None = None, + repo_dir: str | None = None, + baseline: BacktestReport | None = None, policy_source: str = "", progress: bool = True, ) -> BacktestReport: @@ -311,21 +311,23 @@ def run_backtest( def _run_git( head_service: LineageService, policy: Policy, - git_range: Optional[str], - last: Optional[int], - repo_dir: Optional[str], - baseline: Optional[BacktestReport], + git_range: str | None, + last: int | None, + repo_dir: str | None, + baseline: BacktestReport | None, policy_source: str, progress: bool, ) -> BacktestReport: base, head = _resolve_git_range(git_range, last) commits = git_rev_list(base, head, repo_dir) - warnings: List[str] = [ - "Merge/non-squash repos: first-parent diffs may double-count changes across a " - "non-linear history (squash-merge repos are linear and unaffected)." + warnings: list[str] = [ + ( + "Merge/non-squash repos: first-parent diffs may double-count changes across a " + "non-linear history (squash-merge repos are linear and unaffected)." + ) ] - points_verdicts: List[Tuple[BacktestPointResult, PolicyVerdict]] = [] + points_verdicts: list[tuple[BacktestPointResult, PolicyVerdict]] = [] skipped = 0 total = len(commits) for i, sha in enumerate(commits, start=1): @@ -375,13 +377,13 @@ def _run_changesets( head_service: LineageService, policy: Policy, changesets_dir: str, - baseline: Optional[BacktestReport], + baseline: BacktestReport | None, policy_source: str, progress: bool, ) -> BacktestReport: corpus = _load_changesets_dir(changesets_dir) - warnings: List[str] = [] - points_verdicts: List[Tuple[BacktestPointResult, PolicyVerdict]] = [] + warnings: list[str] = [] + points_verdicts: list[tuple[BacktestPointResult, PolicyVerdict]] = [] skipped = 0 total = len(corpus) for i, (name, changes) in enumerate(corpus, start=1): @@ -422,14 +424,14 @@ def _assemble_report( *, mode: str, policy_source: str, - base: Optional[str], - head: Optional[str], - points_verdicts: List[Tuple[BacktestPointResult, PolicyVerdict]], + base: str | None, + head: str | None, + points_verdicts: list[tuple[BacktestPointResult, PolicyVerdict]], policy: Policy, skipped: int, - warnings: List[str], + warnings: list[str], fidelity_note: str, - baseline: Optional[BacktestReport], + baseline: BacktestReport | None, ) -> BacktestReport: points = [p for p, _v in points_verdicts] rule_stats = _aggregate_rule_stats(points_verdicts, policy) @@ -463,10 +465,10 @@ def _assemble_report( ) -def _baseline_delta(rule_stats: List[BacktestRuleStat], baseline: BacktestReport) -> Dict[str, Any]: +def _baseline_delta(rule_stats: list[BacktestRuleStat], baseline: BacktestReport) -> dict[str, Any]: """Per-rule would-BLOCK delta vs a saved baseline (for the regression gate + report).""" base_block = {s.rule_id: s.would_block_prs for s in baseline.rule_stats} - per_rule: Dict[str, Dict[str, int]] = {} + per_rule: dict[str, dict[str, int]] = {} for stat in rule_stats: prev = base_block.get(stat.rule_id, 0) per_rule[stat.rule_id] = { @@ -484,7 +486,7 @@ def _baseline_delta(rule_stats: List[BacktestRuleStat], baseline: BacktestReport def backtest_exit_code( report: BacktestReport, fail_on: str, - baseline: Optional[BacktestReport] = None, + baseline: BacktestReport | None = None, ) -> int: """Translate a report into a CI exit code. diff --git a/parrant/lineage/changeset.py b/parrant/lineage/changeset.py index 20018f9..3b5e7c0 100644 --- a/parrant/lineage/changeset.py +++ b/parrant/lineage/changeset.py @@ -16,7 +16,7 @@ import subprocess from dataclasses import dataclass, field from enum import Enum -from typing import Any, Dict, List, Optional, Set, Tuple +from typing import Any from parrant.lineage.provider import LineageProvider from parrant.lineage.semantic_diff import ( @@ -53,7 +53,7 @@ def priority(self) -> int: return _KIND_PRIORITY[self] -_KIND_PRIORITY: Dict[ChangeKind, int] = { +_KIND_PRIORITY: dict[ChangeKind, int] = { ChangeKind.REMOVED: 5, ChangeKind.TYPE_CHANGED: 4, ChangeKind.LOGIC_CHANGED: 3, @@ -75,21 +75,21 @@ class ColumnChange: model: str column: str kind: ChangeKind - detail: Optional[str] = None - semantic: Optional[SemanticChangeKind] = None + detail: str | None = None + semantic: SemanticChangeKind | None = None # Why a ``logic_changed`` column was flagged: the human-readable semantic reason plus the # two compared defining expressions. Populated only for logic changes (structural kinds # leave them ``None``), and surfaced by ``--explain`` / the JSON ``explain`` block. - reason: Optional[str] = None - base_expression: Optional[str] = None - head_expression: Optional[str] = None + reason: str | None = None + base_expression: str | None = None + head_expression: str | None = None # the override pragma acknowledging this change, when one resolved to it. Excluded from # equality/hashing (``compare=False``) so it never perturbs the sort key or dedup, and so a # frozen ``ColumnChange`` stays hashable even though ``OverrideDirective`` (pydantic) is not. - override: Optional[OverrideDirective] = field(default=None, compare=False) + override: OverrideDirective | None = field(default=None, compare=False) - def to_dict(self) -> Dict[str, object]: - payload: Dict[str, object] = { + def to_dict(self) -> dict[str, object]: + payload: dict[str, object] = { "model": self.model, "column": self.column, "kind": self.kind.value, @@ -117,7 +117,7 @@ def to_dict(self) -> Dict[str, object]: return payload -def _normalize_sql(sql: Optional[str]) -> Optional[str]: +def _normalize_sql(sql: str | None) -> str | None: """Normalize compiled SQL so cosmetic reformatting isn't read as a logic change. ``strip_sql_comments`` already removes comments and collapses whitespace runs, @@ -129,7 +129,7 @@ def _normalize_sql(sql: Optional[str]) -> Optional[str]: return strip_sql_comments(sql) -def _registry_dialect(registry: object) -> Optional[str]: +def _registry_dialect(registry: object) -> str | None: """Best-effort SQL dialect from a registry, ``None`` when it exposes no getter. Defensive so a real ``ModelRegistry`` yields its dialect while lightweight test stubs @@ -152,11 +152,11 @@ class OverrideResolution: stays ``List[ColumnChange]`` and existing callers are unaffected. """ - stale: List[Dict[str, object]] = field(default_factory=list) - warnings: List[str] = field(default_factory=list) + stale: list[dict[str, object]] = field(default_factory=list) + warnings: list[str] = field(default_factory=list) -def _stale_record(directive: OverrideDirective) -> Dict[str, object]: +def _stale_record(directive: OverrideDirective) -> dict[str, object]: """The report skeleton for a stale (no matching change) override.""" record = directive.to_record() return record @@ -180,9 +180,9 @@ def _attach_override(change: ColumnChange, directive: OverrideDirective) -> Colu def resolve_overrides( - model_to_sql: Dict[str, Optional[str]], - changes: List[ColumnChange], -) -> Tuple[List[ColumnChange], List[Dict[str, object]], List[str]]: + model_to_sql: dict[str, str | None], + changes: list[ColumnChange], +) -> tuple[list[ColumnChange], list[dict[str, object]], list[str]]: """Attach override pragmas parsed from each model's head SQL to the matching changes. Shared by :class:`ChangesetBuilder` and :func:`build_git_changeset` so both entry points @@ -194,12 +194,12 @@ def resolve_overrides( Names are lowercased to match the ``ColumnChange`` keys. """ result = list(changes) - changes_by_model: Dict[str, List[int]] = {} + changes_by_model: dict[str, list[int]] = {} for idx, change in enumerate(result): changes_by_model.setdefault(change.model.lower(), []).append(idx) - stale: List[Dict[str, object]] = [] - warnings: List[str] = [] + stale: list[dict[str, object]] = [] + warnings: list[str] = [] for model_name, sql in model_to_sql.items(): if not sql: @@ -245,9 +245,9 @@ class _ColumnDiff: """ kind: SemanticChangeKind - reason: Optional[str] - base_expression: Optional[str] - head_expression: Optional[str] + reason: str | None + base_expression: str | None + head_expression: str | None class ChangesetBuilder: @@ -263,7 +263,7 @@ def __init__( self, base: LineageProvider, head: LineageProvider, - dialect: Optional[str] = None, + dialect: str | None = None, honor_overrides: bool = True, ): self.base = base @@ -275,12 +275,12 @@ def __init__( # when True (default), parse override pragmas from head SQL and attach them. # ``--no-overrides`` sets this False to compute the raw gate (audit / the backtest). self.honor_overrides = honor_overrides - self.stale_overrides: List[Dict[str, object]] = [] - self.override_warnings: List[str] = [] + self.stale_overrides: list[dict[str, object]] = [] + self.override_warnings: list[str] = [] - def build(self) -> List[ColumnChange]: + def build(self) -> list[ColumnChange]: # (model, column) -> ColumnChange, keeping the highest-priority kind. - chosen: Dict[Tuple[str, str], ColumnChange] = {} + chosen: dict[tuple[str, str], ColumnChange] = {} def record(change: ColumnChange) -> None: key = (change.model, change.column) @@ -379,13 +379,13 @@ def record(change: ColumnChange) -> None: chosen_changes = self._apply_overrides(chosen_changes) return sorted(chosen_changes, key=lambda c: (c.model, c.column, c.kind.value)) - def _apply_overrides(self, changes: List[ColumnChange]) -> List[ColumnChange]: + def _apply_overrides(self, changes: list[ColumnChange]) -> list[ColumnChange]: """Parse override pragmas from each changed model's head SQL and attach them. Compiled dbt SQL preserves ``--`` comments, so the head compiled SQL is the pragma source. Records stale directives / parse warnings on ``self`` for the report. """ - model_to_sql: Dict[str, Optional[str]] = { + model_to_sql: dict[str, str | None] = { model_name: self._safe_compiled_sql(self.head, model_name) for model_name in {change.model for change in changes} } @@ -426,7 +426,7 @@ def _logic_changed(self, model_name: str) -> bool: return False return base_sql != head_sql - def _logic_changed_columns(self, base_model, head_model) -> Dict[str, "_ColumnDiff"]: + def _logic_changed_columns(self, base_model, head_model) -> dict[str, _ColumnDiff]: """Which output columns changed derivation, each with a semantic classification. The model's compiled SQL differs, but usually only a few columns are responsible. @@ -461,7 +461,7 @@ def _logic_changed_columns(self, base_model, head_model) -> Dict[str, "_ColumnDi for column in head_model.columns } - changed: Dict[str, _ColumnDiff] = {} + changed: dict[str, _ColumnDiff] = {} for column in head_model.columns: base_sig = base_sigs.get(column) head_sig = head_sigs.get(column) @@ -478,7 +478,7 @@ def _logic_changed_columns(self, base_model, head_model) -> Dict[str, "_ColumnDi ) return changed - def _unattributed_logic_fallback(self, model_name: str, head_model) -> Dict[str, "_ColumnDiff"]: + def _unattributed_logic_fallback(self, model_name: str, head_model) -> dict[str, _ColumnDiff]: """Fail-safe for a proven compiled-SQL change that no output column can explain. ``_logic_changed`` proved the compiled SQL differs, but ``_logic_changed_columns`` @@ -514,7 +514,7 @@ def _unattributed_logic_fallback(self, model_name: str, head_model) -> Dict[str, for column in head_model.columns } - def _classify_change(self, base_exprs: List[str], head_exprs: List[str]) -> "_ColumnDiff": + def _classify_change(self, base_exprs: list[str], head_exprs: list[str]) -> _ColumnDiff: """Classify a signature-differing column, keeping the reason and compared expressions. Fail-safe: if any involved defining expression is unparseable we cannot prove *how* @@ -550,14 +550,14 @@ def _classify_change(self, base_exprs: List[str], head_exprs: List[str]) -> "_Co head_expression=" | ".join(head_exprs) or None, ) - def _any_unparseable(self, expressions: List[str]) -> bool: + def _any_unparseable(self, expressions: list[str]) -> bool: return any( canonical_key(expression, self._dialect).startswith(_UNPARSEABLE_PREFIX) for expression in expressions ) @staticmethod - def _column_expressions(model, column_name: str) -> List[str]: + def _column_expressions(model, column_name: str) -> list[str]: """The raw defining expression string(s) of a column's lineage entries (or ``[]``).""" column = model.columns.get(column_name) if column is None: @@ -565,7 +565,7 @@ def _column_expressions(model, column_name: str) -> List[str]: lineage = getattr(column, "lineage", None) or [] return [getattr(entry, "sql_expression", None) or "" for entry in lineage] - def _column_signatures(self, model) -> Dict[str, Tuple]: + def _column_signatures(self, model) -> dict[str, tuple]: """Per-column derivation signature: {column -> sorted lineage fingerprint}. Columns with no parsed lineage are omitted (no signature), so the caller can tell @@ -573,7 +573,7 @@ def _column_signatures(self, model) -> Dict[str, Tuple]: dialect-aware AST canonical key (``canonical_key``), so cosmetic-only differences collapse to the same signature. """ - signatures: Dict[str, Tuple] = {} + signatures: dict[str, tuple] = {} for column_name, column in model.columns.items(): lineage = getattr(column, "lineage", None) or [] if not lineage: @@ -587,7 +587,7 @@ def _column_signatures(self, model) -> Dict[str, Tuple]: return signatures @staticmethod - def _safe_compiled_sql(registry: LineageProvider, model_name: str) -> Optional[str]: + def _safe_compiled_sql(registry: LineageProvider, model_name: str) -> str | None: try: return registry.get_compiled_sql(model_name) except Exception: @@ -596,9 +596,9 @@ def _safe_compiled_sql(registry: LineageProvider, model_name: str) -> Optional[s return None -def _path_to_model_map(head: LineageProvider) -> Dict[str, str]: +def _path_to_model_map(head: LineageProvider) -> dict[str, str]: """Map each model's ``resource_path`` (dbt ``original_file_path``) to its name.""" - mapping: Dict[str, str] = {} + mapping: dict[str, str] = {} for model_name, model in head.get_models().items(): if model.resource_path: mapping[_norm_path(model.resource_path)] = model_name @@ -608,9 +608,9 @@ def _path_to_model_map(head: LineageProvider) -> Dict[str, str]: def git_changed_models( head: LineageProvider, git_base: str, - repo_dir: Optional[str] = None, + repo_dir: str | None = None, git_head: str = "HEAD", -) -> Set[str]: +) -> set[str]: """Return the set of models whose ``.sql`` file changed between ``git_base`` and ``git_head``. Files with no matching model (macros, tests, deleted files) are ignored, so @@ -627,8 +627,8 @@ def git_changed_models_and_unmapped( head: LineageProvider, git_base: str, git_head: str = "HEAD", - repo_dir: Optional[str] = None, -) -> Tuple[Set[str], List[str]]: + repo_dir: str | None = None, +) -> tuple[set[str], list[str]]: """Split changed ``.sql`` files into (models that mapped, paths that did NOT). Reuses the same path->model map and git diff as :func:`git_changed_models` but also returns @@ -637,8 +637,8 @@ def git_changed_models_and_unmapped( (spec honesty invariant). Non-model SQL (macros/tests/snapshots) shows up here too. """ path_to_model = _path_to_model_map(head) - matched: Set[str] = set() - unmapped: List[str] = [] + matched: set[str] = set() + unmapped: list[str] = [] for changed_file in _git_changed_sql_files(git_base, repo_dir, git_head): model = path_to_model.get(_norm_path(changed_file)) if model: @@ -648,7 +648,7 @@ def git_changed_models_and_unmapped( return matched, unmapped -def git_rev_list(base: str, head: str, repo_dir: Optional[str] = None) -> List[str]: +def git_rev_list(base: str, head: str, repo_dir: str | None = None) -> list[str]: """Enumerate commits in ``base..head`` (oldest -> newest) that touch a ``.sql`` file. Each surviving commit is replayed as one changeset (one commit ≈ one squash-merged PR). @@ -669,7 +669,7 @@ def git_rev_list(base: str, head: str, repo_dir: Optional[str] = None) -> List[s return [line.strip() for line in result.stdout.splitlines() if line.strip()] -def changes_from_dicts(entries: List[Dict[str, Any]]) -> List[ColumnChange]: +def changes_from_dicts(entries: list[dict[str, Any]]) -> list[ColumnChange]: """Reconstruct :class:`ColumnChange` objects from the ``changeset.changes`` JSON shape. Accepts the dicts produced by :meth:`ColumnChange.to_dict` (model/column/kind + optional @@ -678,7 +678,7 @@ def changes_from_dicts(entries: List[Dict[str, Any]]) -> List[ColumnChange]: expressions/reason are not needed for policy evaluation. Unknown/malformed entries raise a ``ValueError`` (via the enum constructors) so a corrupt corpus fails loudly. """ - changes: List[ColumnChange] = [] + changes: list[ColumnChange] = [] for entry in entries: semantic_raw = entry.get("semantic") semantic = SemanticChangeKind(semantic_raw) if semantic_raw else None @@ -701,11 +701,11 @@ def changes_from_dicts(entries: List[Dict[str, Any]]) -> List[ColumnChange]: def build_git_changeset( head: LineageProvider, git_base: str, - repo_dir: Optional[str] = None, + repo_dir: str | None = None, honor_overrides: bool = True, - collect: Optional[OverrideResolution] = None, + collect: OverrideResolution | None = None, git_head: str = "HEAD", -) -> List[ColumnChange]: +) -> list[ColumnChange]: """Fallback changeset: diff ``.sql`` model files between ``git_base`` and ``git_head``. When only one manifest is available we cannot diff columns, so every column @@ -726,7 +726,7 @@ def build_git_changeset( return [] head_models = head.get_models() - chosen: Dict[Tuple[str, str], ColumnChange] = {} + chosen: dict[tuple[str, str], ColumnChange] = {} for model_name in changed_models: model = head_models[model_name] for column in sorted(model.columns): @@ -742,7 +742,7 @@ def build_git_changeset( changes = sorted(chosen.values(), key=lambda c: (c.model, c.column)) if honor_overrides: - model_to_sql: Dict[str, Optional[str]] = { + model_to_sql: dict[str, str | None] = { model_name: ChangesetBuilder._safe_compiled_sql(head, model_name) for model_name in changed_models } @@ -753,7 +753,7 @@ def build_git_changeset( return changes -def scope_changes_to_models(changes: List[ColumnChange], models: Set[str]) -> List[ColumnChange]: +def scope_changes_to_models(changes: list[ColumnChange], models: set[str]) -> list[ColumnChange]: """Keep only changes whose model is in ``models``. Used to intersect a precise two-manifest changeset with the set of models @@ -768,8 +768,8 @@ def _norm_path(path: str) -> str: def _git_changed_sql_files( - git_base: str, repo_dir: Optional[str], git_head: str = "HEAD" -) -> List[str]: + git_base: str, repo_dir: str | None, git_head: str = "HEAD" +) -> list[str]: try: result = subprocess.run( ["git", "diff", "--name-only", f"{git_base}...{git_head}", "--", "*.sql"], @@ -786,20 +786,20 @@ def _git_changed_sql_files( def build_changeset_report( source: str, - changes: List[ColumnChange], - aggregated: Dict[str, object], -) -> Dict[str, object]: + changes: list[ColumnChange], + aggregated: dict[str, object], +) -> dict[str, object]: """Assemble the final report: a ``changeset`` block plus the aggregated impact. The impact keys (``summary``, ``affected_models``, ``affected_columns``, ``affected_exposures``) are a superset of the single-column ``impact`` block, so existing consumers keep working; ``changeset`` and ``by_change`` are added. """ - by_kind: Dict[str, int] = {} + by_kind: dict[str, int] = {} for change in changes: by_kind[change.kind.value] = by_kind.get(change.kind.value, 0) + 1 - report: Dict[str, object] = { + report: dict[str, object] = { "changeset": { "source": source, "total_changes": len(changes), diff --git a/parrant/lineage/ci.py b/parrant/lineage/ci.py index 5f87b4c..dd915de 100644 --- a/parrant/lineage/ci.py +++ b/parrant/lineage/ci.py @@ -21,7 +21,7 @@ import os from dataclasses import dataclass from enum import Enum -from typing import Any, Dict, Optional +from typing import Any import requests @@ -54,9 +54,9 @@ def blocks(self) -> bool: def gate_exit_code( - summary: Dict[str, Any], + summary: dict[str, Any], fail_on: FailOn, - policy_verdict: Optional[Any] = None, + policy_verdict: Any | None = None, ) -> int: """Map an aggregated-impact ``summary`` to an exit code under ``fail_on``. @@ -86,7 +86,7 @@ def gate_exit_code( return 0 # FailOn.NONE and any unknown policy: warn only. -def highest_tripped_level(summary: Dict[str, Any]) -> str: +def highest_tripped_level(summary: dict[str, Any]) -> str: """Return the most severe gate level this ``summary`` trips, ignoring policy. Walks the blocking policies from most to least severe and returns the first whose @@ -101,7 +101,7 @@ def highest_tripped_level(summary: Dict[str, Any]) -> str: return FailOn.NONE.value -def write_github_outputs(report: Dict[str, Any]) -> bool: +def write_github_outputs(report: dict[str, Any]) -> bool: """Emit machine-readable results to ``$GITHUB_OUTPUT`` for the composite action. Writes ``affected_models``, ``affected_columns``, ``affected_exposures`` and @@ -134,15 +134,14 @@ def write_github_outputs(report: Dict[str, Any]) -> bool: values["test_set_size"] = len(policy_verdict.get("test_set", []) or []) try: with open(output_path, "a", encoding="utf-8") as handle: - for key, value in values.items(): - handle.write(f"{key}={value}\n") + handle.writelines(f"{key}={value}\n" for key, value in values.items()) except OSError as exc: logger.warning("Could not write GitHub Action outputs: %s", exc) return False return True -def write_selector_outputs(report: Dict[str, Any]) -> bool: +def write_selector_outputs(report: dict[str, Any]) -> bool: """Emit the policy-free rebuild selection to ``$GITHUB_OUTPUT`` for a selective build. Projects ``report["selection"]`` verbatim — it never recomputes — writing exactly two keys: @@ -169,8 +168,7 @@ def write_selector_outputs(report: Dict[str, Any]) -> bool: } try: with open(output_path, "a", encoding="utf-8") as handle: - for key, value in values.items(): - handle.write(f"{key}={value}\n") + handle.writelines(f"{key}={value}\n" for key, value in values.items()) except OSError as exc: logger.warning("Could not write selector GitHub Action outputs: %s", exc) return False @@ -194,7 +192,7 @@ class GitHubContext: api_url: str = _DEFAULT_API -def resolve_pr_number(explicit: Optional[int] = None) -> Optional[int]: +def resolve_pr_number(explicit: int | None = None) -> int | None: """Resolve the PR number, preferring an explicit value over the GH event. In GitHub Actions the ``pull_request`` event payload carries the number at @@ -222,10 +220,10 @@ def resolve_pr_number(explicit: Optional[int] = None) -> Optional[int]: def resolve_context( - token: Optional[str] = None, - repo: Optional[str] = None, - pr_number: Optional[int] = None, -) -> Optional[GitHubContext]: + token: str | None = None, + repo: str | None = None, + pr_number: int | None = None, +) -> GitHubContext | None: """Build a :class:`GitHubContext` from explicit args + the GH Actions env. Returns ``None`` when any of token / repo / PR number is missing, so callers @@ -241,7 +239,7 @@ def resolve_context( return GitHubContext(repo=repo, pr_number=int(number), token=token, api_url=api_url) -def _headers(token: str) -> Dict[str, str]: +def _headers(token: str) -> dict[str, str]: return { "Authorization": f"Bearer {token}", "Accept": "application/vnd.github+json", @@ -250,8 +248,8 @@ def _headers(token: str) -> Dict[str, str]: def _find_comment_id( - session: Any, base_url: str, headers: Dict[str, str], marker: str -) -> Optional[int]: + session: Any, base_url: str, headers: dict[str, str], marker: str +) -> int | None: """Return the id of the existing marked comment, paging through all comments.""" page = 1 while True: diff --git a/parrant/lineage/display/__init__.py b/parrant/lineage/display/__init__.py index d326bbd..de509e5 100644 --- a/parrant/lineage/display/__init__.py +++ b/parrant/lineage/display/__init__.py @@ -1,5 +1,5 @@ -from parrant.lineage.display.text import TextDisplay from parrant.lineage.display.dot import DotDisplay from parrant.lineage.display.json import JsonDisplay +from parrant.lineage.display.text import TextDisplay -__all__ = ['TextDisplay', 'DotDisplay', 'JsonDisplay'] \ No newline at end of file +__all__ = ["DotDisplay", "JsonDisplay", "TextDisplay"] diff --git a/parrant/lineage/display/backtest.py b/parrant/lineage/display/backtest.py index 1fa7384..09c7ed1 100644 --- a/parrant/lineage/display/backtest.py +++ b/parrant/lineage/display/backtest.py @@ -11,8 +11,6 @@ from __future__ import annotations -from typing import List - from parrant.models.schema import BacktestReport, BacktestRuleStat @@ -38,7 +36,7 @@ def _totals_line(report: BacktestReport) -> str: def render_backtest_table(report: BacktestReport) -> str: """A fixed-width text table for terminals — the per-rule aggregate + totals + fidelity note.""" - lines: List[str] = [] + lines: list[str] = [] lines.append(f"Policy backtest [{report.mode}] — policy: {report.policy_source}") if report.base or report.head: lines.append(f"Range: {report.base}..{report.head}") @@ -75,7 +73,7 @@ def render_backtest_table(report: BacktestReport) -> str: def render_backtest_markdown(report: BacktestReport) -> str: """A Markdown report for a CI artifact / PR comment / agent over MCP.""" - lines: List[str] = [] + lines: list[str] = [] lines.append(f"## Policy backtest — `{report.policy_source}`") lines.append("") lines.append(f"- **Mode:** {report.mode}") diff --git a/parrant/lineage/display/base.py b/parrant/lineage/display/base.py index 7c4706d..2d01214 100644 --- a/parrant/lineage/display/base.py +++ b/parrant/lineage/display/base.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import Dict, Union, Set + from parrant.models.schema import Column, ColumnLineage, Coverage @@ -26,25 +26,18 @@ class LineageStaticDisplay(ABC): @abstractmethod def display_column_info(self, column: Column) -> None: """Display basic column information.""" - pass @abstractmethod - def display_upstream(self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]) -> None: + def display_upstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: """Display upstream lineage.""" - pass @abstractmethod - def display_downstream( - self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]] - ) -> None: + def display_downstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: """Display downstream lineage.""" - pass def display_coverage(self, coverage: Coverage) -> None: """Render a coverage statement. Default no-op; text/json override.""" - pass @abstractmethod def save(self) -> None: """Save or finalize the display output.""" - pass diff --git a/parrant/lineage/display/dot.py b/parrant/lineage/display/dot.py index e1bd8df..122fbc6 100644 --- a/parrant/lineage/display/dot.py +++ b/parrant/lineage/display/dot.py @@ -1,14 +1,14 @@ -from typing import Dict, Set, Optional, Any, Union +from typing import Any + from graphviz import Digraph # type: ignore # missing stubs for graphviz -from parrant.models.schema import Column, ColumnLineage -from parrant.lineage.provider import LineageProvider + from parrant.lineage.display.base import LineageStaticDisplay +from parrant.lineage.provider import LineageProvider +from parrant.models.schema import Column, ColumnLineage class DotDisplay(LineageStaticDisplay): - def __init__( - self, output_file: str = "lineage.dot", registry: Optional[LineageProvider] = None - ): + def __init__(self, output_file: str = "lineage.dot", registry: LineageProvider | None = None): self.dot = Digraph(comment="Column Lineage") self.dot.attr(rankdir="LR") self.dot.attr("node", fontname="Helvetica") @@ -16,10 +16,10 @@ def __init__( self.dot.attr(nodesep="1.0") self.dot.attr(ranksep="1.0") self.output_file = output_file - self.models: Dict[str, Any] = {} + self.models: dict[str, Any] = {} self.registry = registry - self.model_columns: Dict[str, Dict[str, str]] = {} - self.edges: Set[tuple[str, str]] = set() + self.model_columns: dict[str, dict[str, str]] = {} + self.edges: set[tuple[str, str]] = set() self.main_model: str = "" self.main_column: str = "" @@ -27,7 +27,7 @@ def display_column_info(self, column: Column) -> None: self._add_column_to_model(column.model_name, column.name, column.data_type) def _add_column_to_model( - self, model_name: str, col_name: str, data_type: Optional[str] = None + self, model_name: str, col_name: str, data_type: str | None = None ) -> None: if model_name not in self.model_columns: self.model_columns[model_name] = {} @@ -75,8 +75,8 @@ def _process_model_chain( self, current_model_name: str, current_col_name: str, - model_refs: Dict[str, Dict[str, ColumnLineage]], - processed: Optional[Set[str]] = None, + model_refs: dict[str, dict[str, ColumnLineage]], + processed: set[str] | None = None, ) -> None: if self.registry is None: return @@ -101,7 +101,7 @@ def _process_model_chain( self._add_edge(current_ref, f"{model_name}.{col_name}") self._process_model_chain(model_name, col_name, model_refs, processed) - def _add_refs(self, refs: Dict[str, Dict[str, ColumnLineage]], direction: str) -> None: + def _add_refs(self, refs: dict[str, dict[str, ColumnLineage]], direction: str) -> None: if not refs: return @@ -113,7 +113,7 @@ def _add_refs(self, refs: Dict[str, Dict[str, ColumnLineage]], direction: str) - for model_name in self.model_columns: self._create_model_subgraph(model_name, model_name == self.main_model) - def display_upstream(self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]) -> None: + def display_upstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: model_refs = { k: v for k, v in refs.items() @@ -121,9 +121,7 @@ def display_upstream(self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[s } self._add_refs(model_refs, direction="upstream") - def display_downstream( - self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]] - ) -> None: + def display_downstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: model_refs = { k: v for k, v in refs.items() diff --git a/parrant/lineage/display/html/explore.py b/parrant/lineage/display/html/explore.py index 5e0df9d..671b94d 100644 --- a/parrant/lineage/display/html/explore.py +++ b/parrant/lineage/display/html/explore.py @@ -1,12 +1,15 @@ -from typing import Dict, Tuple, Union, Set, List, Any, Optional, Mapping, TYPE_CHECKING -from pydantic import BaseModel, Field -from fastapi import FastAPI, Request -from fastapi.templating import Jinja2Templates -from fastapi.staticfiles import StaticFiles -from fastapi.responses import HTMLResponse +import logging +from collections.abc import Mapping from pathlib import Path +from typing import TYPE_CHECKING, Any + import uvicorn -import logging +from fastapi import FastAPI, Request +from fastapi.responses import HTMLResponse +from fastapi.staticfiles import StaticFiles +from fastapi.templating import Jinja2Templates +from pydantic import BaseModel, Field + from parrant.models.schema import Column, ColumnLineage, TestNode if TYPE_CHECKING: @@ -18,8 +21,8 @@ class ColumnInfo(BaseModel): name: str model: str - type: Optional[str] = None - description: Optional[str] = None + type: str | None = None + description: str | None = None class GraphNode(BaseModel): @@ -27,24 +30,24 @@ class GraphNode(BaseModel): label: str type: str model: str - data_type: Optional[str] = None + data_type: str | None = None is_main: bool = False - resource_type: Optional[str] = None + resource_type: str | None = None is_key: bool = False # For row-set (filter/join/QUALIFY) dependents: the predicate the upstream column appears # in — the "why" behind a node that consumes the column without projecting its value. - note: Optional[str] = None + note: str | None = None # dbt tests (not_null / unique / relationships / ...) declared on this column — the # guardrails a reviewer wants to see at a glance. Empty list = untested; None until enriched. - tests: Optional[List[Dict[str, Any]]] = None + tests: list[dict[str, Any]] | None = None # change-context marks (append-only; None in pure-explore mode → today's payload # byte-for-byte). ``semantic`` is the AST-diff class of a CHANGED column node # (equivalent|meaning_changed|indeterminate); ``breaking`` is the fail-safe convenience # (anything not proven equivalent, plus removed/type_changed); ``boundary`` tags a node # that sits past the dbt edge (e.g. "metabase") for the graph's BI band. - semantic: Optional[str] = None - breaking: Optional[bool] = None - boundary: Optional[str] = None + semantic: str | None = None + breaking: bool | None = None + boundary: str | None = None class GraphEdge(BaseModel): @@ -52,15 +55,15 @@ class GraphEdge(BaseModel): target: str type: str = "lineage" # the single amber blast-path edge — set when the edge leaves a breaking column. - breaking: Optional[bool] = None + breaking: bool | None = None class GraphData(BaseModel): - nodes: List[Dict[str, Any]] = Field(default_factory=list) - edges: List[Dict[str, Any]] = Field(default_factory=list) - main_node: Optional[str] = None - column_info: Optional[ColumnInfo] = None - impact_summary: Optional[Dict[str, Any]] = None + nodes: list[dict[str, Any]] = Field(default_factory=list) + edges: list[dict[str, Any]] = Field(default_factory=list) + main_node: str | None = None + column_info: ColumnInfo | None = None + impact_summary: dict[str, Any] | None = None class LineageExplorer: @@ -71,36 +74,36 @@ def __init__(self, host: str = "127.0.0.1", port: int = 8000): self.host = host self.port = port self.data = GraphData() - self.lineage_service: Optional["LineageService"] = None - self._start_model: Optional[str] = None - self._start_column: Optional[str] = None + self.lineage_service: LineageService | None = None + self._start_model: str | None = None + self._start_column: str | None = None # change context (optional). Populated by ``set_change_context`` with the # already-computed changeset report (semantic per changed column, policy verdict, # cross-boundary Metabase reach, coverage honesty). All None here => pure-explore # mode, and every endpoint renders exactly today's payload. - self._change_report: Optional[Dict[str, Any]] = None - self._policy_verdict: Optional[Dict[str, Any]] = None - self._policy_decision: Optional[str] = None - self._metabase_coverage: Optional[Dict[str, Any]] = None + self._change_report: dict[str, Any] | None = None + self._policy_verdict: dict[str, Any] | None = None + self._policy_decision: str | None = None + self._metabase_coverage: dict[str, Any] | None = None # (model, column) -> {"semantic": str|None, "breaking": bool} for the CHANGED columns. - self._change_by_column: Dict[Tuple[str, str], Dict[str, Any]] = {} + self._change_by_column: dict[tuple[str, str], dict[str, Any]] = {} # name -> full metabase exposure entry (source/precision/via_cards/meta). - self._metabase_exposures: Dict[str, Dict[str, Any]] = {} + self._metabase_exposures: dict[str, dict[str, Any]] = {} # (model, column) -> the metabase dashboards THIS change reaches, each as # {name, via_columns, precision} so per-change reach stays column-precise (F4). - self._metabase_reach_by_change: Dict[Tuple[str, str], List[Dict[str, Any]]] = {} + self._metabase_reach_by_change: dict[tuple[str, str], list[dict[str, Any]]] = {} # A MetabaseReach index for STATIC cross-boundary exploration (no changeset needed): # "what Metabase cards/dashboards read this column?" is a lineage question, not an # impact question, so it should answer for any browsed column. Populated on demand # per explored column in pure-explore mode; the changeset path takes precedence. - self._metabase_static_reach: Optional[Any] = None - self._metabase_dashboard_names: Dict[int, str] = {} + self._metabase_static_reach: Any | None = None + self._metabase_dashboard_names: dict[int, str] = {} self._setup_templates_and_routes() @staticmethod - def _is_breaking(kind: Optional[str], semantic: Optional[str]) -> bool: + def _is_breaking(kind: str | None, semantic: str | None) -> bool: """Fail-safe breaking flag for a changed column. Everything is breaking except a purely additive column or a proven-equivalent @@ -112,7 +115,7 @@ def _is_breaking(kind: Optional[str], semantic: Optional[str]) -> bool: return False return semantic != "equivalent" - def set_change_context(self, report: Optional[Dict[str, Any]]) -> None: + def set_change_context(self, report: dict[str, Any] | None) -> None: """Attach an already-computed changeset report so the explorer can surface the product signals (semantic categorization, policy verdict, Metabase reach, coverage). @@ -173,15 +176,15 @@ def set_change_context(self, report: Optional[Dict[str, Any]]) -> None: if reached: self._metabase_reach_by_change[key] = reached - def _change_mark(self, model: Any, column: Any) -> Optional[Dict[str, Any]]: + def _change_mark(self, model: Any, column: Any) -> dict[str, Any] | None: """The semantic/breaking mark for ``model.column`` if it is a changed column.""" if not isinstance(model, str) or not isinstance(column, str): return None return self._change_by_column.get((model, column)) def _enrich_impact_with_change_context( - self, impact_data: Dict[str, Any], model: str, column: str - ) -> Dict[str, Any]: + self, impact_data: dict[str, Any], model: str, column: str + ) -> dict[str, Any]: """Fold the change-context signals onto a single-column impact payload. Additive and guarded: with no change context (pure-explore mode) or a non-dict @@ -264,11 +267,11 @@ async def home(request: Request) -> Any: ) @self.app.get("/api/graph") - async def get_graph_data() -> Dict[str, Any]: + async def get_graph_data() -> dict[str, Any]: return self.data.model_dump() @self.app.get("/api/coverage") - async def get_coverage() -> Dict[str, Any]: + async def get_coverage() -> dict[str, Any]: if not self.lineage_service: return {"error": "Lineage service not initialized"} coverage = self.lineage_service.get_coverage().model_dump() @@ -280,7 +283,7 @@ async def get_coverage() -> Dict[str, Any]: return coverage @self.app.get("/api/policy-verdict") - async def get_policy_verdict() -> Dict[str, Any]: + async def get_policy_verdict() -> dict[str, Any]: # Change-wide decision (block/warn/allow), independent of the selected column. # ``{"decision": None}`` when no policy resolved / no changeset — the frontend # treats that as "feature not active" and renders nothing new. @@ -289,16 +292,16 @@ async def get_policy_verdict() -> Dict[str, Any]: return {"decision": None} @self.app.get("/api/models") - async def get_models() -> List[Dict[str, Any]]: + async def get_models() -> list[dict[str, Any]]: if not self.lineage_service: return [] - model_tree_root: List[Dict[str, Any]] = [] + model_tree_root: list[dict[str, Any]] = [] all_models = self.lineage_service.registry.get_models() all_exposures = self.lineage_service.registry.get_exposures() def insert_into_tree( - tree: List[Dict[str, Any]], path_parts: List[str], model_data: Dict[str, Any] + tree: list[dict[str, Any]], path_parts: list[str], model_data: dict[str, Any] ) -> None: current_level = tree for i, part in enumerate(path_parts): @@ -405,7 +408,7 @@ def insert_into_tree( return model_tree_root @self.app.get("/api/lineage/{model}/{column}") - async def get_lineage(model: str, column: str) -> Dict[str, Any]: + async def get_lineage(model: str, column: str) -> dict[str, Any]: if not self.lineage_service: return {"error": "Lineage service not initialized"} @@ -457,7 +460,7 @@ async def get_lineage(model: str, column: str) -> Dict[str, Any]: return {"error": str(e)} @self.app.get("/api/model/{model_name}/details") - async def get_model_details(model_name: str) -> Dict[str, Any]: + async def get_model_details(model_name: str) -> dict[str, Any]: if not self.lineage_service: return {"error": "Lineage service not initialized"} @@ -475,7 +478,7 @@ async def get_model_details(model_name: str) -> Dict[str, Any]: return {"error": str(e)} @self.app.get("/api/impact-analysis/{model}/{column}") - async def get_impact_analysis(model: str, column: str) -> Dict[str, Any]: + async def get_impact_analysis(model: str, column: str) -> dict[str, Any]: if not self.lineage_service: return {"error": "Lineage service not initialized"} @@ -484,7 +487,7 @@ async def get_impact_analysis(model: str, column: str) -> Dict[str, Any]: try: model_obj = self.lineage_service.registry.get_model(model) except (ValueError, KeyError) as e: - return {"error": f"Model '{model}' not found: {str(e)}"} + return {"error": f"Model '{model}' not found: {e!s}"} # Verify column exists if column not in model_obj.columns: @@ -551,7 +554,7 @@ def _annotate_nodes_with_semantic(self) -> None: """ if not self._change_by_column: return - breaking_ids: Set[str] = set() + breaking_ids: set[str] = set() for node in self.data.nodes: if node.get("type") != "column": continue @@ -568,9 +571,9 @@ def _annotate_nodes_with_semantic(self) -> None: edge["breaking"] = True @staticmethod - def _boundary_exposure_data(entry: Dict[str, Any]) -> Dict[str, Any]: + def _boundary_exposure_data(entry: dict[str, Any]) -> dict[str, Any]: """The graph ``exposure_data`` payload for a reached Metabase dashboard node.""" - exposure_data: Dict[str, Any] = { + exposure_data: dict[str, Any] = { "boundary": "metabase", "type": entry.get("type") or "dashboard", } @@ -590,7 +593,7 @@ def _boundary_exposure_data(entry: Dict[str, Any]) -> Dict[str, Any]: return exposure_data def attach_metabase_static( - self, reach: Any, dashboard_names: Optional[Dict[int, str]] = None + self, reach: Any, dashboard_names: dict[int, str] | None = None ) -> None: """Attach a :class:`MetabaseReach` for STATIC cross-boundary exploration. @@ -615,7 +618,7 @@ def _populate_static_reach(self, model: str, column: str) -> None: ) except Exception: # reach is best-effort; never break the graph over it return - reach_list: List[Dict[str, Any]] = [] + reach_list: list[dict[str, Any]] = [] for entry in entries: name = entry.get("name") if not isinstance(name, str): @@ -659,7 +662,7 @@ def _annotate_boundary_nodes(self) -> None: return # Pass 1 — tag exposure nodes the registry already emitted that are dashboards. - present_exposure_names: Set[str] = set() + present_exposure_names: set[str] = set() for node in self.data.nodes: if node.get("type") != "exposure": continue @@ -740,12 +743,12 @@ def _annotate_boundary_nodes(self) -> None: GraphEdge(source=anchor_id, target=node_id, type="exposure").model_dump() ) - def _downstream_mart_leaf_ids(self) -> List[str]: + def _downstream_mart_leaf_ids(self) -> list[str]: """The terminal downstream dbt-model columns — targets of a lineage edge that are not themselves a source of one — where the blast path exits into the BI layer. Snapshots and sources are excluded so a dashboard fans out of the marts, not a raw table.""" - lineage_sources: Set[str] = set() - lineage_targets: Set[str] = set() + lineage_sources: set[str] = set() + lineage_targets: set[str] = set() for edge in self.data.edges: if edge.get("type", "lineage") != "lineage": continue @@ -756,7 +759,7 @@ def _downstream_mart_leaf_ids(self) -> List[str]: if target is not None: lineage_targets.add(target) by_id = {n.get("id"): n for n in self.data.nodes} - leaves: List[str] = [] + leaves: list[str] = [] for node_id in lineage_targets: if node_id in lineage_sources: continue @@ -769,7 +772,7 @@ def _downstream_mart_leaf_ids(self) -> List[str]: return sorted(leaves) @staticmethod - def _serialize_test(test: TestNode) -> Dict[str, Any]: + def _serialize_test(test: TestNode) -> dict[str, Any]: """Flatten a :class:`TestNode` into the compact dict the frontend renders.""" return { "test_name": test.test_name, @@ -779,9 +782,7 @@ def _serialize_test(test: TestNode) -> Dict[str, Any]: "referenced_column": test.referenced_column, } - def _column_tests_payload( - self, model: Optional[str], column: Optional[str] - ) -> List[Dict[str, Any]]: + def _column_tests_payload(self, model: str | None, column: str | None) -> list[dict[str, Any]]: """The dbt tests declared on ``model.column`` (empty when untested/unknown). Reuses the registry's prebuilt reverse index — never re-parses artifacts. @@ -795,7 +796,7 @@ def _column_tests_payload( return [] return [self._serialize_test(t) for t in tests] - def _enrich_impact_with_tests(self, impact_data: Dict[str, Any]) -> Dict[str, Any]: + def _enrich_impact_with_tests(self, impact_data: dict[str, Any]) -> dict[str, Any]: """Attach the tests covering each affected column to an impact payload. Lets the impact panel show which guarantees a change threatens. Mutates and returns @@ -871,7 +872,7 @@ def _add_rowset_dependents(self) -> None: ) def _enrich_nodes_with_metadata( - self, refs_list: List[Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]] + self, refs_list: list[dict[str, dict[str, ColumnLineage] | set[str]]] ) -> None: """Enrich nodes with metadata like data types and resource types.""" if not self.lineage_service: @@ -917,13 +918,13 @@ def _enrich_nodes_with_metadata( def _queue_additional_nodes( self, - upstream_refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]], - downstream_refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]], - processed: Set[tuple[str, str]], - to_process: List[tuple[str, str]], + upstream_refs: dict[str, dict[str, ColumnLineage] | set[str]], + downstream_refs: dict[str, dict[str, ColumnLineage] | set[str]], + processed: set[tuple[str, str]], + to_process: list[tuple[str, str]], ) -> None: """Queue additional nodes for processing in sorted order for deterministic BFS.""" - new_nodes: List[tuple[str, str]] = [] + new_nodes: list[tuple[str, str]] = [] for refs in [upstream_refs, downstream_refs]: for model_name, columns in sorted(refs.items()): if model_name == "exposures" or not isinstance(columns, dict): @@ -941,9 +942,9 @@ def _queue_additional_nodes( def _add_processed_data( self, - refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]], + refs: dict[str, dict[str, ColumnLineage] | set[str]], direction: str, - main_node_id: Optional[str] = None, + main_node_id: str | None = None, ) -> None: """Process refs and add to graph.""" processed = self._process_refs(refs, direction, main_node_id) @@ -1089,7 +1090,7 @@ def _set_column_info(self, column: Column) -> None: resource_type=resource_type, ) - def _get_model_resource_type(self, model_name: str) -> Optional[str]: + def _get_model_resource_type(self, model_name: str) -> str | None: """Get resource type for a model.""" try: if self.lineage_service: @@ -1109,11 +1110,11 @@ def _add_node( id: str, label: str, model: str, - data_type: Optional[str] = None, + data_type: str | None = None, is_main: bool = False, - resource_type: Optional[str] = None, + resource_type: str | None = None, is_key: bool = False, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Helper to create and add a node.""" node = GraphNode( id=id, @@ -1129,7 +1130,7 @@ def _add_node( self.data.nodes.append(node) return node - def _add_edge(self, source_id: str, target_id: str) -> Dict[str, str]: + def _add_edge(self, source_id: str, target_id: str) -> dict[str, str]: """Helper to create and add an edge.""" edge = GraphEdge(source=source_id, target=target_id, type="lineage").model_dump() @@ -1138,13 +1139,13 @@ def _add_edge(self, source_id: str, target_id: str) -> Dict[str, str]: def _process_refs( self, - refs: Mapping[str, Union[Dict[str, ColumnLineage], Set[str]]], + refs: Mapping[str, dict[str, ColumnLineage] | set[str]], direction: str, - main_node_id: Optional[str] = None, - ) -> Dict[str, List[Dict[str, Any]]]: + main_node_id: str | None = None, + ) -> dict[str, list[dict[str, Any]]]: """Process reference data into nodes and edges.""" - nodes: List[Dict[str, Any]] = [] - edges: List[Dict[str, Any]] = [] + nodes: list[dict[str, Any]] = [] + edges: list[dict[str, Any]] = [] node_ids = set() if "exposures" in refs and isinstance(refs["exposures"], set): @@ -1234,7 +1235,7 @@ def _process_refs( return {"nodes": nodes, "edges": edges} - def _split_qualified_name(self, qualified_name: str) -> Optional[tuple[str, str]]: + def _split_qualified_name(self, qualified_name: str) -> tuple[str, str] | None: """Split a fully qualified name into model and column parts. Returns None if invalid.""" if "." not in qualified_name: return None @@ -1247,9 +1248,9 @@ def _split_qualified_name(self, qualified_name: str) -> Optional[tuple[str, str] def _add_downstream_edges( self, - source_columns: Union[List[str], Set[str]], + source_columns: list[str] | set[str], target_node_id: str, - edges: List[Dict[str, Any]], + edges: list[dict[str, Any]], ) -> None: """Add edges for downstream lineage.""" for source in source_columns: @@ -1263,12 +1264,12 @@ def _add_downstream_edges( def _process_source_columns( self, - source_columns: Union[List[str], Set[str]], + source_columns: list[str] | set[str], target_node_id: str, - refs: Mapping[str, Union[Dict[str, ColumnLineage], Set[str]]], - nodes: List[Dict[str, Any]], - edges: List[Dict[str, Any]], - node_ids: Set[str], + refs: Mapping[str, dict[str, ColumnLineage] | set[str]], + nodes: list[dict[str, Any]], + edges: list[dict[str, Any]], + node_ids: set[str], ) -> None: """Process source columns and create nodes/edges.""" for source in source_columns: @@ -1288,9 +1289,9 @@ def _add_source_node( self, src_model: str, src_col: str, - refs: Mapping[str, Union[Dict[str, ColumnLineage], Set[str]]], - nodes: List[Dict[str, Any]], - node_ids: Set[str], + refs: Mapping[str, dict[str, ColumnLineage] | set[str]], + nodes: list[dict[str, Any]], + node_ids: set[str], ) -> None: """Add a source node to the graph.""" src_node_id = f"col_{src_model}_{src_col}" diff --git a/parrant/lineage/display/json.py b/parrant/lineage/display/json.py index 488bd3d..2842162 100644 --- a/parrant/lineage/display/json.py +++ b/parrant/lineage/display/json.py @@ -1,9 +1,10 @@ import json -from typing import Any, Dict, Optional, Set, Union +from typing import Any import click from parrant.models.schema import Column, ColumnLineage, Coverage + from .base import LineageStaticDisplay # Keys in a lineage refs dict that hold plain string sets rather than @@ -12,8 +13,8 @@ def serialize_refs( - refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]], -) -> Dict[str, Any]: + refs: dict[str, dict[str, ColumnLineage] | set[str]], +) -> dict[str, Any]: """Convert a lineage refs dict into a JSON-serializable structure. Produces a stable shape regardless of which special sets are present:: @@ -25,7 +26,7 @@ def serialize_refs( "exposures": [...], } """ - models: Dict[str, Any] = {} + models: dict[str, Any] = {} sources: list = [] direct_refs: list = [] exposures: list = [] @@ -60,7 +61,7 @@ class JsonDisplay(LineageStaticDisplay): """ def __init__(self) -> None: - self._result: Dict[str, Any] = {} + self._result: dict[str, Any] = {} def display_column_info(self, column: Column) -> None: self._result["model"] = column.model_name @@ -68,7 +69,7 @@ def display_column_info(self, column: Column) -> None: self._result["data_type"] = column.data_type self._result["description"] = column.description - def set_model_description(self, description: Optional[str]) -> None: + def set_model_description(self, description: str | None) -> None: """Attach the selected column's parent model description (its dbt docs). Kept alongside the column's own ``description`` so an agent triaging @@ -76,15 +77,13 @@ def set_model_description(self, description: Optional[str]) -> None: """ self._result["model_description"] = description - def display_upstream(self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]) -> None: + def display_upstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: self._result["upstream"] = serialize_refs(refs) - def display_downstream( - self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]] - ) -> None: + def display_downstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: self._result["downstream"] = serialize_refs(refs) - def set_impact(self, impact: Optional[Dict[str, Any]]) -> None: + def set_impact(self, impact: dict[str, Any] | None) -> None: """Attach impact-analysis results to the JSON document.""" if impact is not None: self._result["impact"] = impact diff --git a/parrant/lineage/display/markdown.py b/parrant/lineage/display/markdown.py index 39a15f5..0db980c 100644 --- a/parrant/lineage/display/markdown.py +++ b/parrant/lineage/display/markdown.py @@ -10,7 +10,7 @@ with per-expression folds (oversized SQL truncated) and low-risk pass-through folded away. """ -from typing import Any, Dict, List, Tuple +from typing import Any # Fold long dashboard lists so a huge blast radius stays scrollable. _MAX_DASHBOARDS_INLINE = 8 @@ -49,7 +49,7 @@ ) -def _structural_checks_skipped(report: Dict[str, Any]) -> bool: +def _structural_checks_skipped(report: dict[str, Any]) -> bool: return not report.get("structural_checks_available", True) @@ -61,7 +61,7 @@ def _kind_label(kind: str) -> str: return _KIND_LABELS.get(kind, kind.replace("_", " ")) -def _truncate_sql(raw: str) -> Tuple[str, str]: +def _truncate_sql(raw: str) -> tuple[str, str]: """Return (sql_to_show, note). Oversized one-liners are head-elided with a pointer.""" raw = raw.strip() if len(raw) <= _MAX_SQL_CHARS: @@ -74,7 +74,7 @@ def _truncate_sql(raw: str) -> Tuple[str, str]: return head + " …", note -def _format_break(b: Dict[str, Any]) -> str: +def _format_break(b: dict[str, Any]) -> str: """One compiler-style diagnostic line for a provable break. ``error[BREAK-TEST]`` — `model.column` breaks the **** test, @@ -90,7 +90,7 @@ def _format_break(b: Dict[str, Any]) -> str: return f"- `error[BREAK-TEST]` {verb} {node} breaks the **{test_name}** test{via}{where}" -def _owner_suffix(exposure: Dict[str, Any]) -> str: +def _owner_suffix(exposure: dict[str, Any]) -> str: """A ' — owner: **Name**' clause routing the exposure to who must sign off. dbt stores an exposure ``owner`` as ``{name, email}``. Surfacing it turns blast @@ -104,14 +104,14 @@ def _owner_suffix(exposure: Dict[str, Any]) -> str: return f" — owner: **{label}**" if label else "" -def _group_by_model(columns: List[Dict[str, Any]]) -> Dict[str, List[Dict[str, Any]]]: - grouped: Dict[str, List[Dict[str, Any]]] = {} +def _group_by_model(columns: list[dict[str, Any]]) -> dict[str, list[dict[str, Any]]]: + grouped: dict[str, list[dict[str, Any]]] = {} for column in columns: grouped.setdefault(column.get("model", "?"), []).append(column) return {model: grouped[model] for model in sorted(grouped)} -def _confidence_floor_clause(confidence: Dict[str, Any]) -> str: +def _confidence_floor_clause(confidence: dict[str, Any]) -> str: """A short clause for the verdict banner when impact is a lower bound.""" if not confidence or confidence.get("level") == "full": return "" @@ -121,7 +121,7 @@ def _confidence_floor_clause(confidence: Dict[str, Any]) -> str: return f" Impact is a **lower bound** — {_plural(n, 'downstream model')} couldn't be analyzed." -def _confidence_reason_words(confidence: Dict[str, Any]) -> str: +def _confidence_reason_words(confidence: dict[str, Any]) -> str: """Plain-language reason models were unanalyzable, for the footer.""" no_column_info = confidence.get("no_column_info", 0) parse_failed = confidence.get("parse_failed", 0) @@ -134,7 +134,7 @@ def _confidence_reason_words(confidence: Dict[str, Any]) -> str: return "" -def _capped_name_lines(names: List[str], label: str) -> Tuple[List[str], bool]: +def _capped_name_lines(names: list[str], label: str) -> tuple[list[str], bool]: """Render up to the display cap of sorted model names, with a "… +N more" line when the source list is longer. Returns the lines and whether names were elided.""" if not names: @@ -149,7 +149,7 @@ def _capped_name_lines(names: List[str], label: str) -> Tuple[List[str], bool]: return lines, truncated -def _render_unanalyzable_names(confidence: Dict[str, Any]) -> List[str]: +def _render_unanalyzable_names(confidence: dict[str, Any]) -> list[str]: """A folded ``
`` disclosure of the reachable models that couldn't be analyzed, capped for readability. Mutates ``confidence`` to set the display-only ``*_truncated`` flags True when it actually elided names (the JSON surface, which is @@ -161,7 +161,7 @@ def _render_unanalyzable_names(confidence: Dict[str, Any]) -> List[str]: if not no_column_info and not parse_failed and not opaque: return [] total = len(no_column_info) + len(parse_failed) + len(opaque) - body: List[str] = [] + body: list[str] = [] nci_lines, nci_truncated = _capped_name_lines(no_column_info, "No column info") pf_lines, pf_truncated = _capped_name_lines(parse_failed, "Parse failed") # Opaque is a deliberate choice, not a failure — labelled distinctly from the others. @@ -195,7 +195,7 @@ def _render_unanalyzable_names(confidence: Dict[str, Any]) -> List[str]: _POLICY_DECISION_ORDER = ["block", "warn", "allow"] -def _reach_sample(reach: List[str]) -> str: +def _reach_sample(reach: list[str]) -> str: """A capped, deterministic ```a`, `b` +N more`` sample of the matched reach.""" if not reach: return "" @@ -206,7 +206,7 @@ def _reach_sample(reach: List[str]) -> str: return shown -def _proof_marker(hit: Dict[str, Any]) -> str: +def _proof_marker(hit: dict[str, Any]) -> str: """The load-bearing honesty distinction for one fired rule. A rule that fired because a fail-safe knob resolved an UNKNOWN (missing meta / an @@ -221,7 +221,7 @@ def _proof_marker(hit: Dict[str, Any]) -> str: return "✓ proven match" -def _policy_hit_line(hit: Dict[str, Any]) -> str: +def _policy_hit_line(hit: dict[str, Any]) -> str: """One "why this verdict" row for a fired rule: the honesty marker (proven vs fail-safe), the rule id, the subject change, and a capped sample of the matched reach. @@ -249,7 +249,7 @@ def _policy_hit_line(hit: Dict[str, Any]) -> str: return line -def _override_applied_line(record: Dict[str, Any]) -> str: +def _override_applied_line(record: dict[str, Any]) -> str: """One line for a honored override: verb, subject, the severity delta, and the reason.""" model = record.get("model", "?") column = record.get("column") @@ -263,7 +263,7 @@ def _override_applied_line(record: Dict[str, Any]) -> str: ) -def _render_overrides_section(report: Dict[str, Any]) -> List[str]: +def _render_overrides_section(report: dict[str, Any]) -> list[str]: """Render the override signals: malformed-pragma warnings (loud, unfolded, FIRST — a dropped pragma must be noticed), honored overrides with their severity delta, ineffective (no-op) overrides with a fix hint, and a folded list of stale overrides to prune. @@ -277,7 +277,7 @@ def _render_overrides_section(report: Dict[str, Any]) -> List[str]: if not (applied or ineffective or stale or warnings): return [] - out: List[str] = [] + out: list[str] = [] # Warnings FIRST and UNFOLDED — the audit invariant (a reasonless/malformed pragma is # dropped, ruling unchanged) is only useful if the author actually notices it did nothing. if warnings: @@ -333,7 +333,7 @@ def _render_overrides_section(report: Dict[str, Any]) -> List[str]: return out -def _render_policy_section(verdict: Dict[str, Any]) -> List[str]: +def _render_policy_section(verdict: dict[str, Any]) -> list[str]: """Render the policy-engine verdict: fired rules grouped by decision (block first), the selective build/test sets, and the notify intents for the consumer's CI to route. @@ -343,7 +343,7 @@ def _render_policy_section(verdict: Dict[str, Any]) -> List[str]: decision = str(verdict.get("decision", "allow")) marker = _POLICY_DECISION_MARKER.get(decision, "🟢") hits = verdict.get("hits") or [] - out: List[str] = [f"### {marker} Policy verdict — {decision.upper()}", ""] + out: list[str] = [f"### {marker} Policy verdict — {decision.upper()}", ""] if decision == "block": # A block must state its EXIT, not just the obstacle: reframe it as @@ -351,10 +351,12 @@ def _render_policy_section(verdict: Dict[str, Any]) -> List[str]: # tripping the rules below — so a reviewer sees the release path, not a dead end. This is # pure messaging over the existing verdict; no override input is consulted. out += [ - "> **Blocked until the change stops tripping the rules below.** This gate re-runs on " - "every push and clears itself — no manual override needed. Clear it by any of: " - "reverting or proving-equivalent the breaking change; evolving the downstream model / " - "schema to absorb it; or stopping it from reaching the flagged object.", + ( + "> **Blocked until the change stops tripping the rules below.** This gate re-runs on " + "every push and clears itself — no manual override needed. Clear it by any of: " + "reverting or proving-equivalent the breaking change; evolving the downstream model / " + "schema to absorb it; or stopping it from reaching the flagged object." + ), "", ] @@ -365,7 +367,7 @@ def _render_policy_section(verdict: Dict[str, Any]) -> List[str]: # The honesty marker on each row states whether the rule PROVED its match or fired on a # fail-safe default, so a fail-safe block never reads as a confident one. out += ["**Why this verdict** — the rules that fired, and on what:", ""] - by_decision: Dict[str, List[Dict[str, Any]]] = {} + by_decision: dict[str, list[dict[str, Any]]] = {} for hit in hits: by_decision.setdefault(str(hit.get("decision", "allow")), []).append(hit) for band in _POLICY_DECISION_ORDER: @@ -409,7 +411,7 @@ def _render_policy_section(verdict: Dict[str, Any]) -> List[str]: # be silently folded into a clean pass. Shown whenever either counter is > 0. unresolved = int(verdict.get("unresolved_reach_count", 0) or 0) skipped = int(verdict.get("skipped_missing_meta", 0) or 0) - coverage_bits: List[str] = [] + coverage_bits: list[str] = [] if skipped: coverage_bits.append(f"{_plural(skipped, 'column')} undecided (missing meta)") if unresolved: @@ -432,7 +434,7 @@ def _truncate_expr(raw: str, limit: int = 120) -> str: return collapsed -def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> str: +def render_changeset_markdown(report: dict[str, Any], explain: bool = False) -> str: """Render a changeset impact report (from ``build_changeset_report``) as Markdown. The compact semantic reason a column was flagged is shown by default, so the default gate @@ -443,7 +445,7 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> changeset = report.get("changeset", {}) summary = report.get("summary", {}) - out: List[str] = ["## Column-level impact of this change", ""] + out: list[str] = ["## Column-level impact of this change", ""] total_changes = changeset.get("total_changes", 0) if not total_changes: @@ -461,7 +463,7 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> changed_nodes = sorted({(c.get("model", "?"), c.get("column", "?")) for c in by_change}) # (model, column) -> the explain block a logic change carried, so `--explain` can annotate # each changed column with WHY it was flagged. Structural changes carry no explain block. - explain_by_node: Dict[Tuple[str, str], Dict[str, Any]] = {} + explain_by_node: dict[tuple[str, str], dict[str, Any]] = {} for change in by_change: block = change.get("explain") if isinstance(block, dict): @@ -499,7 +501,7 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> subject = f"**{_plural(len(changed_nodes), 'column')}**" else: subject = f"**{_plural(total_changes, 'column')}**" - reach_bits: List[str] = [] + reach_bits: list[str] = [] if apps: reach_bits.append(f"**{_plural(len(apps), 'automation')}**") if dashboards: @@ -524,7 +526,7 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> kind_txt = ", ".join(f"{_kind_label(k)}: {v}" for k, v in sorted(by_kind.items())) out.append(f"**Changed:** {_plural(total_changes, 'column')} — {kind_txt}") if changed_nodes: - rows: List[str] = [] + rows: list[str] = [] for model, column in changed_nodes: rows.append(f"- `{model}.{column}`") # The compact semantic reason ("why was this flagged?") shows BY DEFAULT so the @@ -575,10 +577,10 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> out.append("") out.append("| Model | What changes | How |") out.append("|---|---|---|") - folds: List[str] = [] + folds: list[str] = [] for model in review_models: - what_bits: List[str] = [] - how_bits: List[str] = [] + what_bits: list[str] = [] + how_bits: list[str] = [] derived_cols = sorted( derived_by_model.get(model, []), key=lambda c: c.get("column", "") ) @@ -622,7 +624,7 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> # --- ⚠️ Business-facing exposures (apps surfaced above dashboards) -------------------- if exposures: - rollup_bits: List[str] = [] + rollup_bits: list[str] = [] if dashboards: rollup_bits.append(_plural(len(dashboards), "dashboard")) if apps: @@ -630,7 +632,7 @@ def render_changeset_markdown(report: Dict[str, Any], explain: bool = False) -> out.append(f"### ⚠️ Business-facing exposures ({len(exposures)})") out.append("") - def _fmt(exp: Dict[str, Any]) -> str: + def _fmt(exp: dict[str, Any]) -> str: name = exp.get("name", "?") url = exp.get("url") head = f"- **[{name}]({url})**" if url else f"- **{name}**" @@ -696,7 +698,7 @@ def _fmt(exp: Dict[str, Any]) -> str: out += _render_policy_section(policy_verdict) # --- Footer: confidence + coverage (small, plain, honest) ---------------------------- - footer: List[str] = [] + footer: list[str] = [] if _structural_checks_skipped(report): footer.append(_STRUCTURAL_SKIP_NOTE) # Break detection is a lower bound: tests it couldn't attribute to a column are never @@ -707,7 +709,7 @@ def _fmt(exp: Dict[str, Any]) -> str: f"Break detection skipped {_plural(unattributable, 'dbt test')} it couldn't tie " f"to a column (singular/custom tests); a clean ruling is a lower bound." ) - unanalyzable_disclosure: List[str] = [] + unanalyzable_disclosure: list[str] = [] if confidence: if confidence.get("level") == "full": footer.append( @@ -722,7 +724,7 @@ def _fmt(exp: Dict[str, Any]) -> str: # Additional degradation clauses, kept distinct so the reviewer can tell a parser # *failure* apart from a deliberate *choice* not to analyze (opaque, e.g. semantic # views). Both are rebuilt rather than proven safe to skip. - extra_clauses: List[str] = [] + extra_clauses: list[str] = [] if partial_edges: extra_clauses.append(f"{partial_edges} more carried unresolved column edges") if opaque: diff --git a/parrant/lineage/display/text.py b/parrant/lineage/display/text.py index 734e413..844045d 100644 --- a/parrant/lineage/display/text.py +++ b/parrant/lineage/display/text.py @@ -1,6 +1,7 @@ -from typing import Dict, Set, Union import click + from parrant.models.schema import Column, ColumnLineage, Coverage + from .base import LineageStaticDisplay, format_coverage_line @@ -11,7 +12,7 @@ def display_column_info(self, column: Column) -> None: if column.description: click.echo(f"Description: {column.description}") - def display_upstream(self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]) -> None: + def display_upstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: if not refs: return @@ -30,12 +31,10 @@ def display_upstream(self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[s for model_name, columns in refs.items(): if model_name not in ("sources", "direct_refs") and isinstance(columns, dict): click.echo(f" Model {model_name}:") - for col_name, lineage in columns.items(): + for col_name in columns: click.echo(f" {col_name}") - def display_downstream( - self, refs: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]] - ) -> None: + def display_downstream(self, refs: dict[str, dict[str, ColumnLineage] | set[str]]) -> None: if not refs: return @@ -61,7 +60,7 @@ def display_downstream( columns, dict ): click.echo(f" Model {model_name}:") - for col_name, lineage in columns.items(): + for col_name in columns: click.echo(f" {col_name}") def display_coverage(self, coverage: Coverage) -> None: @@ -70,4 +69,3 @@ def display_coverage(self, coverage: Coverage) -> None: def save(self) -> None: """No-op for text display as output is immediate.""" - pass diff --git a/parrant/lineage/policy.py b/parrant/lineage/policy.py index 1b6f408..ed64d61 100644 --- a/parrant/lineage/policy.py +++ b/parrant/lineage/policy.py @@ -32,12 +32,14 @@ from __future__ import annotations import re +from collections.abc import Iterable from enum import Enum -from typing import Any, Dict, Iterable, List, NamedTuple, Optional, Protocol, Set, Tuple +from typing import Any, NamedTuple, Protocol import yaml from parrant.lineage.changeset import ColumnChange +from parrant.lineage.verdict import ineffective_override_record from parrant.models.schema import ( Action, ActionKind, @@ -60,7 +62,6 @@ SemanticChangeKind, StructuralCondition, ) -from parrant.lineage.verdict import ineffective_override_record class PolicyConfigError(Exception): @@ -99,7 +100,7 @@ def _is_unknown(value: Tri) -> bool: return value is Tri.UNKNOWN_MISSING or value is Tri.UNKNOWN_ERROR -def _merge_unknown(values: List[Tri]) -> Tri: +def _merge_unknown(values: list[Tri]) -> Tri: """Collapse the surviving UNKNOWNs to a single cause. ERROR dominates MISSING so a genuine type error is never masked by a fail-open missing-meta default (fail-safe bias).""" if any(v is Tri.UNKNOWN_ERROR for v in values): @@ -107,7 +108,7 @@ def _merge_unknown(values: List[Tri]) -> Tri: return Tri.UNKNOWN_MISSING -def _and(values: List[Tri]) -> Tri: +def _and(values: list[Tri]) -> Tri: """Kleene AND: any FALSE -> FALSE; else any UNKNOWN -> UNKNOWN; else TRUE. Empty -> TRUE.""" if any(v is Tri.FALSE for v in values): return Tri.FALSE @@ -117,7 +118,7 @@ def _and(values: List[Tri]) -> Tri: return Tri.TRUE -def _or(values: List[Tri]) -> Tri: +def _or(values: list[Tri]) -> Tri: """Kleene OR: any TRUE -> TRUE; else any UNKNOWN -> UNKNOWN; else FALSE. Empty -> FALSE.""" if any(v is Tri.TRUE for v in values): return Tri.TRUE @@ -138,7 +139,7 @@ def _not(value: Tri) -> Tri: # --- config loading --------------------------------------------------------- -def load_policy(path: Optional[str]) -> Optional[Policy]: +def load_policy(path: str | None) -> Policy | None: """Resolve and parse the policy file. Resolution order (first found): explicit ``path`` -> ``./parrant.policy.yml`` -> @@ -181,7 +182,7 @@ def parse_policy(raw: Any, source: str = "") -> Policy: raise PolicyConfigError(f"policy '{source}' is invalid: {exc}") from exc -def _resolve_policy_path(path: Optional[str]) -> Optional[str]: +def _resolve_policy_path(path: str | None) -> str | None: """First existing path among explicit -> repo default. The default filename is ``parrant.policy.yml``; the legacy ``dbt-col-lineage.policy.yml`` is @@ -259,7 +260,7 @@ class CombineStrategy(Protocol): def to_element(self, value: Any) -> _Lattice: ... - def combine(self, elements: List[_Lattice]) -> _Lattice: ... + def combine(self, elements: list[_Lattice]) -> _Lattice: ... class _MostRestrictive: @@ -269,7 +270,7 @@ class _MostRestrictive: def to_element(self, value: Any) -> _Lattice: return _Lattice.HIGH if bool(value) else _Lattice.LOW - def combine(self, elements: List[_Lattice]) -> _Lattice: + def combine(self, elements: list[_Lattice]) -> _Lattice: if not elements: return _Lattice.UNKNOWN return max(elements, key=lambda element: element.value) @@ -284,7 +285,7 @@ class _BooleanOr: def to_element(self, value: Any) -> _Lattice: return _Lattice.HIGH if bool(value) else _Lattice.LOW - def combine(self, elements: List[_Lattice]) -> _Lattice: + def combine(self, elements: list[_Lattice]) -> _Lattice: if any(element is _Lattice.HIGH for element in elements): return _Lattice.HIGH if any(element is _Lattice.UNKNOWN for element in elements): @@ -295,7 +296,7 @@ def combine(self, elements: List[_Lattice]) -> _Lattice: # The fold policy per meta key. Unregistered keys fall back to the most-restrictive fold — the # fail-safe direction (an unclassified lineage is treated as "not proven safe"). _DEFAULT_STRATEGY: CombineStrategy = _MostRestrictive() -_COMBINE_STRATEGIES: Dict[str, CombineStrategy] = { +_COMBINE_STRATEGIES: dict[str, CombineStrategy] = { "pii": _MostRestrictive(), "secret": _BooleanOr(), } @@ -306,7 +307,7 @@ def _combine_strategy_for(key: str) -> CombineStrategy: return _COMBINE_STRATEGIES.get(key.lower(), _DEFAULT_STRATEGY) -def _split_ref(source: str) -> Optional[Tuple[str, str]]: +def _split_ref(source: str) -> tuple[str, str] | None: """Split a ``model.column`` lineage ref into ``(model, column)``, both lowercased. Mirrors the service's split: everything before the last ``.`` is the model, the last segment @@ -319,7 +320,7 @@ def _split_ref(source: str) -> Optional[Tuple[str, str]]: return (".".join(parts[:-1]).lower(), parts[-1].lower()) -def _element_to_lookup(element: _Lattice) -> "MetaLookup": +def _element_to_lookup(element: _Lattice) -> MetaLookup: """Map a folded lattice element to a :class:`MetaLookup`. ``UNKNOWN`` -> ``present=False`` so the engine treats an unresolvable inferred value as a missing key (fail-closed).""" if element is _Lattice.HIGH: @@ -349,15 +350,15 @@ class MetaIndex: def __init__(self, registry: Any, metabase_reach: Any = None) -> None: self._registry = registry self._metabase_reach = metabase_reach - self._model_cache: Dict[str, Dict[str, Any]] = {} - self._config_cache: Dict[str, Dict[str, Any]] = {} - self._column_cache: Dict[Tuple[str, str], Dict[str, Any]] = {} - self._exposures: Optional[Dict[str, Any]] = None + self._model_cache: dict[str, dict[str, Any]] = {} + self._config_cache: dict[str, dict[str, Any]] = {} + self._column_cache: dict[tuple[str, str], dict[str, Any]] = {} + self._exposures: dict[str, Any] | None = None # Memoize resolved inferred meta across the (immutable) column DAG, keyed by # (model, column, key) lowercased — folding a diamond visits each node once. - self._inferred_cache: Dict[Tuple[str, str, str], MetaLookup] = {} + self._inferred_cache: dict[tuple[str, str, str], MetaLookup] = {} - def _model_dict(self, model: str) -> Dict[str, Any]: + def _model_dict(self, model: str) -> dict[str, Any]: cached = self._model_cache.get(model) if cached is None: getter = getattr(self._registry, "get_model_dbt_meta", None) @@ -365,7 +366,7 @@ def _model_dict(self, model: str) -> Dict[str, Any]: self._model_cache[model] = cached return cached - def _config_dict(self, model: str) -> Dict[str, Any]: + def _config_dict(self, model: str) -> dict[str, Any]: cached = self._config_cache.get(model) if cached is None: getter = getattr(self._registry, "get_model_config", None) @@ -373,7 +374,7 @@ def _config_dict(self, model: str) -> Dict[str, Any]: self._config_cache[model] = cached return cached - def _column_dict(self, model: str, column: str) -> Dict[str, Any]: + def _column_dict(self, model: str, column: str) -> dict[str, Any]: key = (model, column) cached = self._column_cache.get(key) if cached is None: @@ -382,7 +383,7 @@ def _column_dict(self, model: str, column: str) -> Dict[str, Any]: self._column_cache[key] = cached return cached - def _exposure_dict(self, exposure: str) -> Optional[Any]: + def _exposure_dict(self, exposure: str) -> Any | None: if self._exposures is None: getter = getattr(self._registry, "get_exposures", None) try: @@ -444,8 +445,8 @@ def _inferred_lookup( column: str, key: str, strategy: CombineStrategy, - visiting: Set[Tuple[str, str, str]], - ) -> Tuple[MetaLookup, bool]: + visiting: set[tuple[str, str, str]], + ) -> tuple[MetaLookup, bool]: """Fold ``model.column``'s inferred value, returning ``(result, touched_cycle)``. ``touched_cycle`` is True iff this node's fold DEPENDED on a cycle-guard hit (an upstream @@ -477,7 +478,7 @@ def _inferred_lookup( # Rules 2 & 3: fold the upstream source columns' inferred values. visiting.add(memo_key) - elements: List[_Lattice] = [] + elements: list[_Lattice] = [] touched_cycle = False for src_model, src_column in self._upstream_source_columns(model, column): child, child_touched = self._inferred_lookup( @@ -494,7 +495,7 @@ def _inferred_lookup( self._inferred_cache[memo_key] = result return result, touched_cycle - def _upstream_source_columns(self, model: str, column: str) -> Iterable[Tuple[str, str]]: + def _upstream_source_columns(self, model: str, column: str) -> Iterable[tuple[str, str]]: """The distinct ``(model, column)`` upstream source columns feeding ``model.column``. Reads the provider's ``get_column_lineage`` edges (duck-typed like the meta accessors, so @@ -507,8 +508,8 @@ def _upstream_source_columns(self, model: str, column: str) -> Iterable[Tuple[st edges = getter(model, column) or [] except Exception: return [] - seen: Set[Tuple[str, str]] = set() - out: List[Tuple[str, str]] = [] + seen: set[tuple[str, str]] = set() + out: list[tuple[str, str]] = [] for edge in edges: for source in getattr(edge, "source_columns", None) or []: pair = _split_ref(str(source)) @@ -549,11 +550,11 @@ class ReachedObject(NamedTuple): kind: ReachKind name: str - column: Optional[str] - mechanism: Optional[Mechanism] + column: str | None + mechanism: Mechanism | None -def _to_mechanism(raw: Optional[str]) -> Optional[Mechanism]: +def _to_mechanism(raw: str | None) -> Mechanism | None: if raw is None: return None try: @@ -571,8 +572,8 @@ class ImpactView: :meth:`is_resolved` reports that so the engine can fail-safe. """ - def __init__(self, changeset_impact: Dict[str, Any]) -> None: - self._by_key: Dict[Tuple[str, str, str], Dict[str, Any]] = {} + def __init__(self, changeset_impact: dict[str, Any]) -> None: + self._by_key: dict[tuple[str, str, str], dict[str, Any]] = {} for entry in changeset_impact.get("by_change", []): key = ( str(entry.get("model")), @@ -583,7 +584,7 @@ def __init__(self, changeset_impact: Dict[str, Any]) -> None: if key not in self._by_key: self._by_key[key] = entry - def _entry(self, change: ColumnChange) -> Optional[Dict[str, Any]]: + def _entry(self, change: ColumnChange) -> dict[str, Any] | None: return self._by_key.get((change.model, change.column, change.kind.value)) def is_resolved(self, change: ColumnChange) -> bool: @@ -594,13 +595,13 @@ def reached( self, change: ColumnChange, kind: ReachKind, - mechanism: Optional[List[Mechanism]] = None, - ) -> List[ReachedObject]: + mechanism: list[Mechanism] | None = None, + ) -> list[ReachedObject]: entry = self._entry(change) if not entry or not entry.get("resolved"): return [] wanted = set(mechanism) if mechanism else None - out: List[ReachedObject] = [] + out: list[ReachedObject] = [] if kind is ReachKind.MODEL: for item in entry.get("reached_models", []): mech = _to_mechanism(item.get("mechanism")) @@ -620,7 +621,7 @@ def reached( return out -def build_impact_view(changeset_impact: Dict[str, Any]) -> ImpactView: +def build_impact_view(changeset_impact: dict[str, Any]) -> ImpactView: """Adapt a ``get_changeset_impact`` report dict into an :class:`ImpactView` (pure).""" return ImpactView(changeset_impact) @@ -628,7 +629,7 @@ def build_impact_view(changeset_impact: Dict[str, Any]) -> ImpactView: # --- operator evaluation ---------------------------------------------------- -def _as_list(value: Any) -> Optional[List[Any]]: +def _as_list(value: Any) -> list[Any] | None: if isinstance(value, (list, tuple, set)): return list(value) return None @@ -776,7 +777,7 @@ class _Trace: """Accumulates matched reach object names while a predicate is evaluated for one subject.""" def __init__(self) -> None: - self.matched_reach: List[str] = [] + self.matched_reach: list[str] = [] self.saw_unresolved_reach: bool = False def add(self, name: str) -> None: @@ -795,7 +796,7 @@ def __init__( policy: Policy, meta: MetaIndex, impact: ImpactView, - breaks: List[BreakFinding], + breaks: list[BreakFinding], ) -> None: self._policy = policy self._meta = meta @@ -806,26 +807,26 @@ def __init__( # -- public API ---------------------------------------------------------- - def evaluate(self, changes: List[ColumnChange]) -> PolicyVerdict: + def evaluate(self, changes: list[ColumnChange]) -> PolicyVerdict: """Run every rule against every subject (or once, for aggregate rules); combine per §2.6. Total and deterministic: never raises on rule content (config errors surface at load). An undecidable predicate resolves per the rule's ``MissingMetaPolicy``. """ # index the changeset by (model, column) lowercased so override caps are O(1). - self._change_by_key: Dict[Tuple[str, str], ColumnChange] = { + self._change_by_key: dict[tuple[str, str], ColumnChange] = { (c.model.lower(), c.column.lower()): c for c in changes } - hits: List[RuleHit] = [] + hits: list[RuleHit] = [] build_set: set[str] = set() test_set: set[str] = set() - notifications: List[Notification] = [] + notifications: list[Notification] = [] skipped = 0 unresolved_reach = 0 for rule in self._policy.rules: - subjects: List[Optional[ColumnChange]] + subjects: list[ColumnChange | None] subjects = [None] if rule.scope == "aggregate" else list(changes) for subject in subjects: trace = _Trace() @@ -881,7 +882,7 @@ def evaluate(self, changes: List[ColumnChange]) -> PolicyVerdict: # -- built-in semantic-severity knobs ------------------------------------ - def _semantic_default_hits(self, changes: List[ColumnChange]) -> List[RuleHit]: + def _semantic_default_hits(self, changes: list[ColumnChange]) -> list[RuleHit]: """Synthesize gate contributions from ``defaults.on_meaning_changed`` / ``on_indeterminate``. These are the ergonomic shortcut for "gate on the semantic axis" without authoring a @@ -901,7 +902,7 @@ def _semantic_default_hits(self, changes: List[ColumnChange]) -> List[RuleHit]: SemanticChangeKind.MEANING_CHANGED: (defaults.on_meaning_changed, "on_meaning_changed"), SemanticChangeKind.INDETERMINATE: (defaults.on_indeterminate, "on_indeterminate"), } - hits: List[RuleHit] = [] + hits: list[RuleHit] = [] for change in changes: if change.semantic is None: continue @@ -921,7 +922,7 @@ def _semantic_default_hits(self, changes: List[ColumnChange]) -> List[RuleHit]: # -- fail-safe resolution ------------------------------------------------ - def _resolve(self, result: Tri, rule: Rule) -> Optional[bool]: + def _resolve(self, result: Tri, rule: Rule) -> bool | None: """Resolve a (possibly UNKNOWN) predicate result to fire / not-fire / skip. TRUE -> fire; FALSE -> not fire. An UNKNOWN is routed to the fail-safe knob that matches @@ -1020,7 +1021,7 @@ def _eval_reach(self, cond: ReachCondition, subject: ColumnChange, trace: _Trace return Tri.UNKNOWN_MISSING objects = self._impact.reached(subject, cond.kind, cond.mechanism) n_true = 0 - unknowns: List[Tri] = [] + unknowns: list[Tri] = [] for obj in objects: inner = self._eval_where(cond.where, obj) if inner is Tri.TRUE: @@ -1112,7 +1113,7 @@ def _eval_structural( # -- predicate evaluation (aggregate scope) ------------------------------ def _eval_aggregate( - self, predicate: Predicate, changes: List[ColumnChange], trace: _Trace + self, predicate: Predicate, changes: list[ColumnChange], trace: _Trace ) -> Tri: """Aggregate scope: each leaf is quantified existentially over the whole changeset. @@ -1131,12 +1132,12 @@ def _eval_aggregate( # -- actions ------------------------------------------------------------- def _apply_actions( - self, rule: Rule, subject: Optional[ColumnChange], trace: _Trace - ) -> Tuple[RuleHit, "set[str]", "set[str]", List[Notification]]: + self, rule: Rule, subject: ColumnChange | None, trace: _Trace + ) -> tuple[RuleHit, set[str], set[str], list[Notification]]: build_add: set[str] = set() test_add: set[str] = set() - notes: List[Notification] = [] - action_kinds: List[ActionKind] = [] + notes: list[Notification] = [] + action_kinds: list[ActionKind] = [] for action in rule.action: action_kinds.append(action.type) @@ -1161,7 +1162,7 @@ def _apply_actions( notes, ) - def _collect_nodes(self, action: Action, subject: Optional[ColumnChange]) -> "set[str]": + def _collect_nodes(self, action: Action, subject: ColumnChange | None) -> set[str]: nodes: set[str] = set() if subject is None: return nodes @@ -1173,7 +1174,7 @@ def _collect_nodes(self, action: Action, subject: Optional[ColumnChange]) -> "se return nodes def _build_notification( - self, rule: Rule, action: Action, subject: Optional[ColumnChange], trace: _Trace + self, rule: Rule, action: Action, subject: ColumnChange | None, trace: _Trace ) -> Notification: template = action.message or "" values = { @@ -1191,7 +1192,7 @@ def _build_notification( # -- override caps -------------------------------------------------- - def _apply_override_caps(self, hits: List[RuleHit]) -> None: + def _apply_override_caps(self, hits: list[RuleHit]) -> None: """Cap each subject-scoped hit whose change carries an override (mutates in place). - ``allow-break`` caps a BLOCK to WARN (the only verb that may touch a block). @@ -1232,7 +1233,7 @@ def _apply_override_caps(self, hits: List[RuleHit]) -> None: # -- combination --------------------------------------------------------- - def _combine_decision(self, hits: List[RuleHit]) -> GateDecision: + def _combine_decision(self, hits: list[RuleHit]) -> GateDecision: decision = GateDecision.ALLOW for hit in hits: if hit.decision.severity > decision.severity: @@ -1265,22 +1266,22 @@ def _reached_display(obj: ReachedObject) -> str: _INTERPOLATE_RE = re.compile(r"\{([a-z]+(?:\.[a-z]+)?)\}") -def _interpolate(template: str, values: Dict[str, str]) -> str: +def _interpolate(template: str, values: dict[str, str]) -> str: """Substitute the small safe vocabulary ``{change.*}`` / ``{reach.count}`` / ``{rule.id}``. Unknown tokens are left verbatim (no arbitrary code, no KeyError). """ - def repl(match: "re.Match[str]") -> str: + def repl(match: re.Match[str]) -> str: token = match.group(1) return values.get(token, match.group(0)) return _INTERPOLATE_RE.sub(repl, template) -def _dedup_notifications(notes: List[Notification]) -> List[Notification]: - seen: set[Tuple[str, str, str]] = set() - out: List[Notification] = [] +def _dedup_notifications(notes: list[Notification]) -> list[Notification]: + seen: set[tuple[str, str, str]] = set() + out: list[Notification] = [] for note in notes: key = (note.channel, note.target, note.message) if key not in seen: @@ -1293,13 +1294,13 @@ def _dedup_notifications(notes: List[Notification]) -> List[Notification]: def applied_policy_overrides( - verdict: PolicyVerdict, changes: List[ColumnChange] -) -> List[Dict[str, Any]]: + verdict: PolicyVerdict, changes: list[ColumnChange] +) -> list[dict[str, Any]]: """Honored-override records derived from capped policy hits, in the SAME shape the default gate emits (``applied_overrides``). Cross-references ``ColumnChange.override`` for the verb / source_line / scope that ``RuleHit`` does not carry, so both report paths are uniform.""" by_key = {(c.model.lower(), c.column.lower()): c for c in changes} - records: List[Dict[str, Any]] = [] + records: list[dict[str, Any]] = [] for hit in verdict.hits: if not hit.overridden: continue @@ -1323,9 +1324,9 @@ def applied_policy_overrides( def ineffective_policy_overrides( verdict: PolicyVerdict, - changes: List[ColumnChange], - breaks: Optional[List[BreakFinding]] = None, -) -> List[Dict[str, Any]]: + changes: list[ColumnChange], + breaks: list[BreakFinding] | None = None, +) -> list[dict[str, Any]]: """Override records that landed on a real changed column but capped NO hit — surfaced so an ineffective pragma (e.g. allow-change on a break, or an override where no rule fired) is never silently ignored. Same shape as the default gate's ``ineffective_overrides``.""" @@ -1335,7 +1336,7 @@ def ineffective_policy_overrides( for h in verdict.hits if h.overridden } - records: List[Dict[str, Any]] = [] + records: list[dict[str, Any]] = [] for change in changes: if change.override is None: continue @@ -1347,11 +1348,11 @@ def ineffective_policy_overrides( def evaluate_policy( - changes: List[ColumnChange], - changeset_impact: Dict[str, Any], + changes: list[ColumnChange], + changeset_impact: dict[str, Any], registry: Any, policy: Policy, - breaks: Optional[List[BreakFinding]] = None, + breaks: list[BreakFinding] | None = None, metabase_reach: Any = None, ) -> PolicyVerdict: """One-call helper can wire into ``cli/main.py``: build the indexes and evaluate. diff --git a/parrant/lineage/policy_init.py b/parrant/lineage/policy_init.py index ac281ae..21239ac 100644 --- a/parrant/lineage/policy_init.py +++ b/parrant/lineage/policy_init.py @@ -21,7 +21,7 @@ import os from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any from parrant.lineage.service import LineageService from parrant.models.schema import MetaKeyCoverage, PolicyInitScan @@ -37,7 +37,7 @@ _META_TEMPLATE_CAP = 12 -def _flatten_meta_keys(meta: Dict[str, Any], _prefix: str = "") -> List[str]: +def _flatten_meta_keys(meta: dict[str, Any], _prefix: str = "") -> list[str]: """Yield the dotted leaf-key paths (``a.b.c``) of a (possibly nested) meta dict. Nested dicts are recursed so a nested key still gets an accurate histogram entry AND a @@ -46,7 +46,7 @@ def _flatten_meta_keys(meta: Dict[str, Any], _prefix: str = "") -> List[str]: operator against an intermediate dict is rarely what an author means. An empty dict value is treated as a leaf so it still surfaces as a (present-but-empty) key rather than vanishing. """ - keys: List[str] = [] + keys: list[str] = [] for key, value in meta.items(): dotted = f"{_prefix}{key}" if isinstance(value, dict) and value: @@ -56,7 +56,7 @@ def _flatten_meta_keys(meta: Dict[str, Any], _prefix: str = "") -> List[str]: return keys -def _histogram(counts: Dict[str, int], total: int) -> List[MetaKeyCoverage]: +def _histogram(counts: dict[str, int], total: int) -> list[MetaKeyCoverage]: """Build the coverage rows, sorted most-covered-first then by key (stable, deterministic).""" rows = [MetaKeyCoverage(key=key, n_present=n, total=total) for key, n in counts.items()] rows.sort(key=lambda row: (-row.n_present, row.key)) @@ -76,8 +76,8 @@ def scan_project(registry: Any) -> PolicyInitScan: total_columns = 0 column_test_count = 0 models_with_column_tests = 0 - model_meta_counts: Dict[str, int] = {} - column_meta_counts: Dict[str, int] = {} + model_meta_counts: dict[str, int] = {} + column_meta_counts: dict[str, int] = {} models_with_grants = 0 for name, model in models.items(): @@ -111,7 +111,7 @@ def scan_project(registry: Any) -> PolicyInitScan: ) -def _has_select_grant(config: Dict[str, Any]) -> bool: +def _has_select_grant(config: dict[str, Any]) -> bool: """True when a model's resolved dbt ``config`` declares a non-empty ``grants.select``. Mirrors the engine's ``config.grants.select`` dotted lookup: the roles a model grants SELECT @@ -134,7 +134,7 @@ def _has_select_grant(config: Dict[str, Any]) -> bool: # --- YAML emitter (string-templated; comments cannot survive yaml.dump) ------ -def _header_lines() -> List[str]: +def _header_lines() -> list[str]: """The top comment block: ownership, the ``policy test`` pointer, and the two footguns. Deliberately never writes the literal permissive-default token: this scaffold only ever uses @@ -172,7 +172,7 @@ def _header_lines() -> List[str]: ] -def _defaults_lines() -> List[str]: +def _defaults_lines() -> list[str]: return [ "defaults:", " # Anything we cannot prove safe is treated as unsafe. Safe here because every ENABLED", @@ -185,7 +185,7 @@ def _defaults_lines() -> List[str]: ] -def _comment_out(yaml_lines: List[str]) -> List[str]: +def _comment_out(yaml_lines: list[str]) -> list[str]: """Comment out a block of 2-space-indented YAML lines, preserving relative indentation.""" return [f" # {line[2:]}" if line.startswith(" ") else f"# {line}" for line in yaml_lines] @@ -205,7 +205,7 @@ def _comment_out(yaml_lines: List[str]) -> List[str]: ] -def _provable_break_block_lines(enabled: bool) -> List[str]: +def _provable_break_block_lines(enabled: bool) -> list[str]: """The ``provable-break-block`` rule. Safe-by-construction: ``provable_test_break`` is a pure structural fact (always TRUE/FALSE, never UNKNOWN) evaluated at aggregate scope, so this block can only ever fire on a real, offline-verifiable breakage — never on a fail-safe @@ -223,7 +223,7 @@ def _provable_break_block_lines(enabled: bool) -> List[str]: ] + _comment_out(_PROVABLE_BREAK_BLOCK_YAML) -def _exposure_guard_lines(enabled: bool) -> List[str]: +def _exposure_guard_lines(enabled: bool) -> list[str]: """The ``exposure-guard`` rule — WARN (not block) when a change reaches an exposure. ``touches_exposure`` CAN be UNKNOWN on unresolved reach, but this is a non-blocking WARN @@ -240,7 +240,7 @@ def _exposure_guard_lines(enabled: bool) -> List[str]: ] + _comment_out(_EXPOSURE_GUARD_YAML) -def _meta_template_lines(row: MetaKeyCoverage, subject: str) -> List[str]: +def _meta_template_lines(row: MetaKeyCoverage, subject: str) -> list[str]: """A single COMMENTED, meta-keyed reach template for one discovered key. Prefixed by the REAL coverage from the scan and using a PRESENCE operator (``is_true``) — @@ -279,7 +279,7 @@ def _slug(key: str) -> str: return slug or "meta" -def _footer_lines() -> List[str]: +def _footer_lines() -> list[str]: return [ "", "# --- Next steps -------------------------------------------------------------", @@ -302,7 +302,7 @@ def emit_policy_yaml(scan: PolicyInitScan) -> str: Invariant: the string ``fail_open`` is never emitted (an open-when-unsure gate is not a gate). """ - lines: List[str] = [] + lines: list[str] = [] lines.extend(_header_lines()) lines.append("") lines.append("version: 1") @@ -310,7 +310,7 @@ def emit_policy_yaml(scan: PolicyInitScan) -> str: lines.extend(_defaults_lines()) lines.append("") - enabled_blocks: List[List[str]] = [] + enabled_blocks: list[list[str]] = [] if scan.tests_present: enabled_blocks.append(_provable_break_block_lines(enabled=True)) if scan.exposures_present: @@ -336,7 +336,7 @@ def emit_policy_yaml(scan: PolicyInitScan) -> str: lines.append("rules: []") # Disabled tool-owned rules, shown commented so the user knows why they were withheld. - disabled_blocks: List[List[str]] = [] + disabled_blocks: list[list[str]] = [] if not scan.tests_present: disabled_blocks.append(_provable_break_block_lines(enabled=False)) if not scan.exposures_present: @@ -360,11 +360,11 @@ def emit_policy_yaml(scan: PolicyInitScan) -> str: return text -def _meta_section_lines(scan: PolicyInitScan) -> List[str]: +def _meta_section_lines(scan: PolicyInitScan) -> list[str]: """Emit the commented meta-template section (empty when the scan found no meta keys).""" if not scan.model_meta_keys and not scan.column_meta_keys: return [] - lines: List[str] = [ + lines: list[str] = [ "", " # --- Commented meta templates (uncomment AFTER `policy test`) ---------", " # These are keyed to dbt `meta` the scan actually found. Each is prefixed with its real", @@ -376,7 +376,7 @@ def _meta_section_lines(scan: PolicyInitScan) -> List[str]: return lines -def _config_section_lines(scan: PolicyInitScan) -> List[str]: +def _config_section_lines(scan: PolicyInitScan) -> list[str]: """Emit the commented config-axis (PII over-grant) template, or nothing. Only offered when the scan found at least one model declaring ``config.grants.select`` — the @@ -411,8 +411,8 @@ def _config_section_lines(scan: PolicyInitScan) -> List[str]: ] -def _capped_templates(rows: List[MetaKeyCoverage], subject: str) -> List[str]: - lines: List[str] = [] +def _capped_templates(rows: list[MetaKeyCoverage], subject: str) -> list[str]: + lines: list[str] = [] for row in rows[:_META_TEMPLATE_CAP]: lines.append("") lines.extend(_meta_template_lines(row, subject)) @@ -432,7 +432,7 @@ def _capped_templates(rows: List[MetaKeyCoverage], subject: str) -> List[str]: def run_policy_init( manifest: str, catalog: str, - adapter: Optional[str], + adapter: str | None, output: str, force: bool, stdout: bool, diff --git a/parrant/lineage/provider.py b/parrant/lineage/provider.py index e4f8a5b..606b8eb 100644 --- a/parrant/lineage/provider.py +++ b/parrant/lineage/provider.py @@ -16,7 +16,7 @@ from __future__ import annotations -from typing import Any, Dict, List, Optional, Protocol, Set, runtime_checkable +from typing import Any, Protocol, runtime_checkable from parrant.models.schema import ( Column, @@ -58,13 +58,13 @@ def is_loaded(self) -> bool: """Whether :meth:`load` has completed and the graph is queryable.""" # --- model / graph access --------------------------------------------- - def get_models(self) -> Dict[str, Model]: + def get_models(self) -> dict[str, Model]: """All nodes (models, snapshots, seeds, sources) keyed by lowercased name.""" def get_model(self, model_name: str) -> Model: """One node by name (case-insensitive). Raise ``ModelNotFoundError`` if absent.""" - def get_manifest_downstream(self) -> Dict[str, Set[str]]: + def get_manifest_downstream(self) -> dict[str, set[str]]: """Model-level child map over the *whole* DAG, incl. nodes with no column info. Distinct from per-column edges: this is the reachability frontier the impact @@ -73,7 +73,7 @@ def get_manifest_downstream(self) -> Dict[str, Set[str]]: """ # --- column lineage (the core product) -------------------------------- - def get_column_lineage(self, model_name: str, column_name: str) -> List[ColumnLineage]: + def get_column_lineage(self, model_name: str, column_name: str) -> list[ColumnLineage]: """Per-column upstream edges for ``model.column`` (case-insensitive). Each :class:`ColumnLineage` carries ``source_columns`` (``model.col`` refs), @@ -83,7 +83,7 @@ def get_column_lineage(self, model_name: str, column_name: str) -> List[ColumnLi ``get_model(model).columns[column].lineage``; both must agree. """ - def get_column(self, model_name: str, column_name: str) -> Optional[Column]: + def get_column(self, model_name: str, column_name: str) -> Column | None: """Column truth (name, ``data_type``, description, lineage) or ``None`` if unknown. ``data_type`` may be ``None`` when the backend lacks catalog/compiler column truth @@ -91,7 +91,7 @@ def get_column(self, model_name: str, column_name: str) -> Optional[Column]: """ # --- row-set / predicate lineage (capability) ------------------------- - def get_filter_dependents(self, source_column: str) -> Set[str]: + def get_filter_dependents(self, source_column: str) -> set[str]: """Models that reference ``source_column`` ONLY in a predicate (WHERE/JOIN/HAVING/QUALIFY). Row-set dependents that value-lineage misses. Capability method: a backend that @@ -100,7 +100,7 @@ def get_filter_dependents(self, source_column: str) -> Set[str]: """ # --- provenance / quality signals ------------------------------------- - def get_dialect(self) -> Optional[str]: + def get_dialect(self) -> str | None: """The SQL dialect used to canonicalize expressions, or ``None`` if unknown. Consumed by the AST semantic-diff (:mod:`~parrant.lineage.changeset`) @@ -120,7 +120,7 @@ def is_catalog_backed(self, model_name: str) -> bool: A compiler-grade backend (Fusion) returns ``True`` for every built node. """ - def get_parse_failed_models(self) -> Set[str]: + def get_parse_failed_models(self) -> set[str]: """Nodes whose lineage the backend could not compute (had input, failed to derive). Feeds the ``partial`` confidence level. A backend with no notion of parse failure @@ -128,7 +128,7 @@ def get_parse_failed_models(self) -> Set[str]: not a deliberate choice not to analyze. """ - def get_opaque_models(self) -> Set[str]: + def get_opaque_models(self) -> set[str]: """Nodes the backend deliberately does NOT column-analyze (unparseable SQL). Semantic views chief among them; generally any node whose compiled SQL the backend @@ -140,7 +140,7 @@ def get_opaque_models(self) -> Set[str]: """ # --- raw SQL (leaky capability) --------------------------------------- - def get_compiled_sql(self, model_name: str) -> Optional[str]: + def get_compiled_sql(self, model_name: str) -> str | None: """The model's compiled SQL, or ``None`` when the backend has none. Leaky: raw SQL is a *production input*, exposed only because ``changeset`` uses a @@ -159,29 +159,29 @@ class ProjectMetadataProvider(Protocol): lineage provider can be paired with the *same* metadata provider unchanged. """ - def get_exposures(self) -> Dict[str, Exposure]: + def get_exposures(self) -> dict[str, Exposure]: """All exposures keyed by name.""" def get_exposure(self, exposure_name: str) -> Exposure: """One exposure by name. Raise if absent.""" - def get_column_tests(self, model: str, column: str) -> List[TestNode]: + def get_column_tests(self, model: str, column: str) -> list[TestNode]: """dbt tests targeting ``model.column`` (case-insensitive).""" - def get_tests_referencing(self, model: str, column: str) -> List[TestNode]: + def get_tests_referencing(self, model: str, column: str) -> list[TestNode]: """relationships tests whose *referenced* (parent) side is ``model.column``.""" - def get_model_tests(self, model: str) -> List[TestNode]: + def get_model_tests(self, model: str) -> list[TestNode]: """Every test that breaks if ``model`` is removed wholesale.""" - def get_test_unique_ids(self) -> Set[str]: + def get_test_unique_ids(self) -> set[str]: """All test ``unique_id``s present (verdict confirms a base test survived in head).""" def get_unattributable_test_count(self) -> int: """Tests that could not be attributed to a (model, column) — coverage honesty.""" # --- arbitrary dbt meta (metadata-agnostic access) -------------------- - def get_model_dbt_meta(self, model: str) -> Dict[str, Any]: + def get_model_dbt_meta(self, model: str) -> dict[str, Any]: """Arbitrary user-authored dbt ``meta`` on a model (case-insensitive), or ``{}``. Manifest-sourced and independent of the lineage engine, exactly like exposures and @@ -190,13 +190,13 @@ def get_model_dbt_meta(self, model: str) -> Dict[str, Any]: ``config.meta`` over top-level ``meta`` per dbt precedence. Absent meta ⇒ ``{}``. """ - def get_column_dbt_meta(self, model: str, column: str) -> Dict[str, Any]: + def get_column_dbt_meta(self, model: str, column: str) -> dict[str, Any]: """Arbitrary user-authored dbt ``meta`` on a column (case-insensitive), or ``{}``. Same contract as :meth:`get_model_dbt_meta`, scoped to one column. Absent ⇒ ``{}``. """ - def get_model_config(self, model: str) -> Dict[str, Any]: + def get_model_config(self, model: str) -> dict[str, Any]: """The node's resolved dbt ``config`` dict for a model (case-insensitive), or ``{}``. Manifest-sourced and metadata-agnostic, exactly like :meth:`get_model_dbt_meta`: diff --git a/parrant/lineage/semantic_diff.py b/parrant/lineage/semantic_diff.py index cd8cb44..a1703e2 100644 --- a/parrant/lineage/semantic_diff.py +++ b/parrant/lineage/semantic_diff.py @@ -76,7 +76,7 @@ def comment_free_token_signature(expr_sql: str, dialect: str | None = None) -> s """ try: tokens = tokenize(expr_sql, dialect=dialect) - except Exception: # noqa: BLE001 - fail-safe: any tokenize failure -> None (no lexical proof) + except Exception: return None return _TOKEN_RECORD_SEP.join( f"{token.token_type.name}{_TOKEN_FIELD_SEP}{token.text}" for token in tokens @@ -100,7 +100,7 @@ def canonicalize_expression(expr_sql: str, dialect: str | None = None) -> exp.Ex """ try: expression = parse_one(expr_sql, dialect=dialect) - except Exception: # noqa: BLE001 - fail-safe: any parse failure -> None (indeterminate) + except Exception: return None return normalize_identifiers(expression, dialect=dialect) diff --git a/parrant/lineage/service.py b/parrant/lineage/service.py index e73f69c..6b2fc38 100644 --- a/parrant/lineage/service.py +++ b/parrant/lineage/service.py @@ -1,9 +1,9 @@ -from pathlib import Path -from typing import Dict, List, Literal, Set, Optional, Any, Tuple, Union, TYPE_CHECKING -from dataclasses import dataclass, field -from collections import Counter import logging import re +from collections import Counter +from dataclasses import dataclass, field +from pathlib import Path +from typing import TYPE_CHECKING, Any, Literal, Optional from parrant.artifacts.exceptions import ModelNotFoundError from parrant.lineage.provider import LineageAndMetadataProvider @@ -27,13 +27,13 @@ # Higher rank == more severe. Used to keep the worst severity when the same # downstream node is reached by several changed columns. -_SEVERITY_RANK: Dict[str, int] = {"critical": 2, "low_impact": 1} +_SEVERITY_RANK: dict[str, int] = {"critical": 2, "low_impact": 1} # A downstream column's ``transformation_type`` → the plain-language *mechanism* by which # the change reaches it. This is the machine-readable twin of the markdown's mechanism # split (derived recompute / row-set filter / pass-through): it lets an agent or the # Impact Report envelope reason over *how* impact propagates, not just how many nodes. -_MECHANISM_LABELS: Dict[str, str] = { +_MECHANISM_LABELS: dict[str, str] = { "derived": "derived_recompute", "filter": "rowset_filter", "renamed": "renamed_passthrough", @@ -41,7 +41,7 @@ } -def _mechanism_label(transformation_type: Optional[str]) -> str: +def _mechanism_label(transformation_type: str | None) -> str: """Map a downstream column's ``transformation_type`` to its reach *mechanism* label. The single source of the recompute/filter/pass-through taxonomy predicates match on. @@ -53,8 +53,8 @@ def _mechanism_label(transformation_type: Optional[str]) -> str: def _reached_from_impact( - impact: Dict[str, Any], -) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]], List[Dict[str, Any]]]: + impact: dict[str, Any], +) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]]]: """Re-shape a single change's impact into reached NAMES + mechanism (no new traversal). ``get_column_impact`` already computes the reached models/columns/exposures; it only @@ -70,7 +70,7 @@ def _reached_from_impact( Pure re-shape of data already in ``impact``; deterministic ordering for stable reports. """ - reached_columns: List[Dict[str, Any]] = [ + reached_columns: list[dict[str, Any]] = [ { "model": column["model"], "column": column["column"], @@ -79,13 +79,13 @@ def _reached_from_impact( for column in impact.get("affected_columns", []) ] - model_mechanisms: Dict[str, Set[str]] = {} + model_mechanisms: dict[str, set[str]] = {} for column in impact.get("affected_columns", []): model_mechanisms.setdefault(column["model"], set()).add( _mechanism_label(column.get("transformation_type")) ) - reached_models: List[Dict[str, Any]] = [] + reached_models: list[dict[str, Any]] = [] for model in impact.get("affected_models", []): mechanisms = sorted(model_mechanisms.get(model["name"], set())) if mechanisms: @@ -95,21 +95,21 @@ def _reached_from_impact( else: reached_models.append({"name": model["name"], "mechanism": None}) - reached_exposures: List[Dict[str, Any]] = [ + reached_exposures: list[dict[str, Any]] = [ {"name": exposure["name"]} for exposure in impact.get("affected_exposures", []) ] return reached_models, reached_exposures, reached_columns -def _mechanism_breakdown(affected_columns: List[Dict[str, Any]]) -> Dict[str, int]: +def _mechanism_breakdown(affected_columns: list[dict[str, Any]]) -> dict[str, int]: """Count affected downstream columns by the mechanism that propagates the change. Pure aggregation over the ``transformation_type`` each affected column already carries — no new traversal. An unrecognized type is bucketed under its raw value so nothing is silently dropped. """ - breakdown: Dict[str, int] = {} + breakdown: dict[str, int] = {} for column in affected_columns: raw = column.get("transformation_type") or "unknown" label = _MECHANISM_LABELS.get(raw, raw) @@ -117,7 +117,7 @@ def _mechanism_breakdown(affected_columns: List[Dict[str, Any]]) -> Dict[str, in return breakdown -def _change_is_breaking(entry: Dict[str, Any]) -> bool: +def _change_is_breaking(entry: dict[str, Any]) -> bool: """Whether a ``by_change`` entry contributes its reach to the rebuild set (fail-closed). Only a *proven* additive change is safe to skip: ``kind == "added"`` with a semantic that @@ -134,7 +134,7 @@ def _change_is_breaking(entry: Dict[str, Any]) -> bool: return not proven_additive -def _unresolved_edges(model: Any) -> List[Dict[str, Any]]: +def _unresolved_edges(model: Any) -> list[dict[str, Any]]: """The model's unresolved-edge markers, or ``[]`` if none. Reads ``model.metadata["unresolved_edges"]`` — the complete, uncapped list the registry @@ -164,21 +164,21 @@ def _partial_edges_reason(registry: LineageAndMetadataProvider, name: str) -> st model = registry.get_model(name) except ModelNotFoundError: return "unresolved_edge" - reasons: List[str] = [ + reasons: list[str] = [ str(edge["reason"]) for edge in _unresolved_edges(model) if edge.get("reason") ] if not reasons: return "unresolved_edge" counts = Counter(reasons) - return sorted(counts.items(), key=lambda kv: (-kv[1], kv[0]))[0][0] + return min(counts.items(), key=lambda kv: (-kv[1], kv[0]))[0] def build_selection( - reachable: Set[str], - changed_models: Set[str], - by_change: List[Dict[str, Any]], - confidence: Dict[str, Any], -) -> Dict[str, Any]: + reachable: set[str], + changed_models: set[str], + by_change: list[dict[str, Any]], + confidence: dict[str, Any], +) -> dict[str, Any]: """Derive the policy-free minimal rebuild set from an already-computed changeset impact. Pure function over the diff facts (reachability, the per-change reach in ``by_change``, and @@ -203,7 +203,7 @@ def build_selection( """ # The universe every model is partitioned over: the strictly-downstream reachable set plus # the edited models themselves (which the DAG walk excludes but which always rebuild). - universe: Set[str] = set(reachable) | set(changed_models) + universe: set[str] = set(reachable) | set(changed_models) # The COMPLETE unanalyzable lists (uncapped in machine output). A model parrant could not # analyze is assumed affected — this is the whole reason the lists must be complete. Models @@ -211,7 +211,7 @@ def build_selection( # have columns but a phantom/unresolvable source edge, so they must ALWAYS rebuild regardless # of the widen branch below. This is a pure ADD (fail-safe) — a marker can only grow the # rebuild set, never shrink it or let a marker-carrying model land in ``skippable``. - unanalyzable: Set[str] = ( + unanalyzable: set[str] = ( set(confidence.get("no_column_info_models", [])) | set(confidence.get("parse_failed_models", [])) | set(confidence.get("partial_edges_models", [])) @@ -224,7 +224,7 @@ def build_selection( confidence.get("no_column_info_truncated") or confidence.get("parse_failed_truncated") ) - breaking_reached: Set[str] = set() + breaking_reached: set[str] = set() for entry in by_change: if not _change_is_breaking(entry): continue @@ -233,7 +233,7 @@ def build_selection( if name is not None: breaking_reached.add(name) - rebuild: Set[str] = ( + rebuild: set[str] = ( (set(changed_models) & universe) | (breaking_reached & universe) | (unanalyzable & universe) ) @@ -241,7 +241,7 @@ def build_selection( widened = level != "full" or truncated if widened: rebuild = set(universe) - skippable: List[str] = [] + skippable: list[str] = [] else: skippable = sorted(universe - rebuild) @@ -275,15 +275,15 @@ class _ReachPartition: ``parse_failed`` — partition ``reachable`` exactly. """ - resolved: Set[str] - no_column_info: Set[str] - parse_failed: Set[str] - catalog_backed: Set[str] - parsed: Set[str] - partial_edges: Set[str] = field(default_factory=set) + resolved: set[str] + no_column_info: set[str] + parse_failed: set[str] + catalog_backed: set[str] + parsed: set[str] + partial_edges: set[str] = field(default_factory=set) # Nodes we deliberately do NOT column-analyze (unparseable SQL, e.g. semantic views): # model-level reach preserved, column edges withheld. Pulled out of the resolved set. - opaque: Set[str] = field(default_factory=set) + opaque: set[str] = field(default_factory=set) # A ``SELECT *`` immediately followed by a column-set modifier (Snowflake EXCLUDE/RENAME/ @@ -294,7 +294,7 @@ class _ReachPartition: def _resolve_reason( registry: LineageAndMetadataProvider, name: str -) -> Tuple[Literal["no_column_info", "unresolved"], Optional[str]]: +) -> tuple[Literal["no_column_info", "unresolved"], str | None]: """(status, reason) for a reachable model that has no column info. Best-effort and advisory. A python model is surfaced as ``unresolved``/``python_model``; @@ -342,8 +342,8 @@ def _opaque_reason(registry: LineageAndMetadataProvider, name: str) -> str: def build_resolution( registry: LineageAndMetadataProvider, partition: _ReachPartition, - rebuild_models: Set[str], -) -> Tuple[Dict[str, Dict[str, Any]], Dict[str, Any]]: + rebuild_models: set[str], +) -> tuple[dict[str, dict[str, Any]], dict[str, Any]]: """Retain the reachable partition per model as a resolution status + advisory reason. Pure emission over already-computed facts (the confidence partition, the catalog-backed @@ -351,7 +351,7 @@ def build_resolution( Returns ``(per_model_map, resolution_summary_dict)``. Every reachable model appears exactly once in the map; the summary counts reconcile exactly with the confidence counts. """ - per_model: Dict[str, Dict[str, Any]] = {} + per_model: dict[str, dict[str, Any]] = {} for name in partition.catalog_backed: per_model[name] = ModelResolution(status="catalog_backed", reason=None).model_dump() for name in partition.parsed: @@ -407,7 +407,7 @@ def build_resolution( @dataclass class LineageSelector: model: str - column: Optional[str] + column: str | None upstream: bool downstream: bool @@ -439,14 +439,14 @@ def from_string(cls, selector: str) -> "LineageSelector": class LineageReferences: """Structured lineage references separating model mappings from special sets.""" - models: Dict[str, Dict[str, ColumnLineage]] = field(default_factory=dict) - exposures: Set[str] = field(default_factory=set) - sources: Set[str] = field(default_factory=set) - direct_refs: Set[str] = field(default_factory=set) + models: dict[str, dict[str, ColumnLineage]] = field(default_factory=dict) + exposures: set[str] = field(default_factory=set) + sources: set[str] = field(default_factory=set) + direct_refs: set[str] = field(default_factory=set) - def to_dict(self) -> Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]: + def to_dict(self) -> dict[str, dict[str, ColumnLineage] | set[str]]: """Convert to legacy dict format for backward compatibility.""" - result: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]] = {} + result: dict[str, dict[str, ColumnLineage] | set[str]] = {} result.update(self.models) if self.exposures: result["exposures"] = self.exposures @@ -457,9 +457,7 @@ def to_dict(self) -> Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]: return result @classmethod - def from_dict( - cls, data: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]] - ) -> "LineageReferences": + def from_dict(cls, data: dict[str, dict[str, ColumnLineage] | set[str]]) -> "LineageReferences": """Create from legacy dict format.""" refs = cls() for key, value in data.items(): @@ -477,7 +475,7 @@ def from_dict( class LineageService: """Service for handling lineage operations.""" - def __init__(self, catalog_path: Path, manifest_path: Path, adapter: Optional[str] = None): + def __init__(self, catalog_path: Path, manifest_path: Path, adapter: str | None = None): # Depend on the LineageProvider seam, not the concrete registry: the factory builds # and loads the SQLGlot-backed provider today, and is the single place a future # Fusion/warehouse backend would swap in. @@ -490,11 +488,11 @@ def get_coverage(self) -> Coverage: """Return coverage for the loaded artifacts.""" return self._coverage - def _dag_reachable_models(self, model_name: str) -> Set[str]: + def _dag_reachable_models(self, model_name: str) -> set[str]: """Transitive downstream models of model_name in the manifest DAG.""" downstream_map = self.registry.get_manifest_downstream() start = model_name.lower() - reachable: Set[str] = set() + reachable: set[str] = set() queue = [start] while queue: current = queue.pop() @@ -504,7 +502,7 @@ def _dag_reachable_models(self, model_name: str) -> Set[str]: queue.append(child) return reachable - def _partition_reachable(self, reachable: Set[str]) -> _ReachPartition: + def _partition_reachable(self, reachable: set[str]) -> _ReachPartition: """Split ``reachable`` by column-resolution outcome — the single source of truth. Both the confidence block and the per-model resolution status consume this, so the two @@ -516,11 +514,11 @@ def _partition_reachable(self, reachable: Set[str]) -> _ReachPartition: parse_failed_names = self.registry.get_parse_failed_models() opaque_names = self.registry.get_opaque_models() - resolved: Set[str] = set() - parse_failed: Set[str] = set() - no_column_info: Set[str] = set() - partial_edges: Set[str] = set() - opaque: Set[str] = set() + resolved: set[str] = set() + parse_failed: set[str] = set() + no_column_info: set[str] = set() + partial_edges: set[str] = set() + opaque: set[str] = set() for name in reachable: # Opaque takes priority over any columns a catalog entry might supply: we have no # column-level EDGES for these nodes (unparseable SQL), so they can never be treated @@ -556,7 +554,7 @@ def _partition_reachable(self, reachable: Set[str]) -> _ReachPartition: opaque=opaque, ) - def _impact_confidence(self, reachable: Set[str], resolved_models: int) -> Dict[str, Any]: + def _impact_confidence(self, reachable: set[str], resolved_models: int) -> dict[str, Any]: """Confidence block: "full" when every reachable model was analyzable, else "partial". The honest signal is the *coverage gap* — reachable downstream models we could @@ -588,9 +586,7 @@ def _impact_confidence(self, reachable: Set[str], resolved_models: int) -> Dict[ # through it, so we widen the rebuild rather than prove anything downstream skippable. The # widen at ``build_selection`` fires on ``partial``. level: Literal["full", "partial"] = ( - "full" - if not unanalyzable_reachable and not partial_edges and not opaque - else "partial" + "full" if not unanalyzable_reachable and not partial_edges and not opaque else "partial" ) # Machine surface carries the COMPLETE name lists (no cap) so a fail-closed # consumer can never miss a model we couldn't analyze/resolve; the display layer caps. @@ -612,7 +608,7 @@ def _impact_confidence(self, reachable: Set[str], resolved_models: int) -> Dict[ level=level, ).model_dump() - def get_model_info(self, selector: LineageSelector) -> Dict[str, Any]: + def get_model_info(self, selector: LineageSelector) -> dict[str, Any]: """Get model information based on selector.""" model = self.registry.get_model(selector.model) return { @@ -624,7 +620,7 @@ def get_model_info(self, selector: LineageSelector) -> Dict[str, Any]: "downstream": list(model.downstream) if selector.downstream else [], } - def get_column_info(self, selector: LineageSelector) -> Dict[str, Any]: + def get_column_info(self, selector: LineageSelector) -> dict[str, Any]: """Get column information and lineage based on selector.""" model = self.registry.get_model(selector.model) if not selector.column or selector.column not in model.columns: @@ -647,7 +643,7 @@ def get_column_info(self, selector: LineageSelector) -> Dict[str, Any]: ), } - def _split_qualified_name(self, qualified_name: str) -> Optional[tuple[str, str]]: + def _split_qualified_name(self, qualified_name: str) -> tuple[str, str] | None: """Split a fully qualified name into model and column parts. Returns None if invalid.""" if "." not in qualified_name: return None @@ -670,7 +666,7 @@ def _process_source_reference( def _merge_upstream_refs( self, target: LineageReferences, - source_dict: Dict[str, Union[Dict[str, ColumnLineage], Set[str]]], + source_dict: dict[str, dict[str, ColumnLineage] | set[str]], ) -> None: """Merge source refs dict into target LineageReferences.""" for key, value in source_dict.items(): @@ -691,7 +687,7 @@ def _process_model_reference( src_column: str, lineage: ColumnLineage, upstream_refs: LineageReferences, - visited: Set[str], + visited: set[str], ) -> None: """Process a model reference and add it to upstream_refs.""" try: @@ -713,8 +709,8 @@ def _process_model_reference( self._process_source_reference(f"{src_model}.{src_column}", upstream_refs) def _get_upstream_lineage( - self, model_name: str, column_name: str, visited: Optional[Set[str]] = None - ) -> Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]: + self, model_name: str, column_name: str, visited: set[str] | None = None + ) -> dict[str, dict[str, ColumnLineage] | set[str]]: """Recursively get all upstream column references.""" if visited is None: visited = set() @@ -740,7 +736,7 @@ def _get_upstream_lineage( column.lineage, key=lambda lineage: ( lineage.transformation_type, - sorted(lineage.source_columns)[0] if lineage.source_columns else "", + min(lineage.source_columns) if lineage.source_columns else "", ), ) for lineage in sorted_lineage: @@ -759,13 +755,13 @@ def _get_upstream_lineage( ) except Exception as e: - logger.warning(f"Failed to process lineage for {current_ref}: {str(e)}") + logger.warning(f"Failed to process lineage for {current_ref}: {e!s}") return upstream_refs.to_dict() def _get_immediate_downstream_lineage( self, model_name: str, column_name: str - ) -> Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]: + ) -> dict[str, dict[str, ColumnLineage] | set[str]]: """Get only immediate (non-recursive) downstream column references.""" column_name = strip_sql_comments(column_name).lower() current_ref = f"{model_name}.{column_name}" @@ -811,7 +807,7 @@ def _get_immediate_downstream_lineage( downstream_refs.models[other_name][col_name] = lineage except Exception as e: - logger.warning(f"Failed to process downstream model {other_name}: {str(e)}") + logger.warning(f"Failed to process downstream model {other_name}: {e!s}") if column_used_downstream and downstream_models_using_column: models_using_column = set(downstream_models_using_column) @@ -831,14 +827,14 @@ def _get_immediate_downstream_lineage( except Exception as e: logger.warning( - f"Failed to process immediate downstream lineage for {current_ref}: {str(e)}" + f"Failed to process immediate downstream lineage for {current_ref}: {e!s}" ) return downstream_refs.to_dict() def _get_downstream_lineage( - self, model_name: str, column_name: str, visited: Optional[Set[str]] = None - ) -> Dict[str, Union[Dict[str, ColumnLineage], Set[str]]]: + self, model_name: str, column_name: str, visited: set[str] | None = None + ) -> dict[str, dict[str, ColumnLineage] | set[str]]: """Get downstream column references following the model DAG, including exposures. Uses breadth-first traversal without shared mutable state to ensure determinism. @@ -898,7 +894,7 @@ def _get_downstream_lineage( key=lambda lineage: ( lineage.transformation_type, ( - sorted(lineage.source_columns)[0] + min(lineage.source_columns) if lineage.source_columns else "" ), @@ -925,19 +921,17 @@ def _get_downstream_lineage( except Exception as e: logger.warning( - f"Failed to process downstream model {other_name}: {str(e)}" + f"Failed to process downstream model {other_name}: {e!s}" ) if column_used_downstream: models_using_column.update(level_downstream_models) all_models_using_column.update(level_downstream_models) all_models_using_column.add(current_model) - models_using_column = set(sorted(models_using_column)) + models_using_column = set(models_using_column) except Exception as e: - logger.warning( - f"Failed to process downstream lineage for {current_ref}: {str(e)}" - ) + logger.warning(f"Failed to process downstream lineage for {current_ref}: {e!s}") next_level_nodes.sort() queue.extend(next_level_nodes) @@ -964,7 +958,7 @@ def _get_downstream_lineage( return downstream_refs.to_dict() - def get_column_impact(self, model_name: str, column_name: str) -> Dict[str, Any]: + def get_column_impact(self, model_name: str, column_name: str) -> dict[str, Any]: """Get impact analysis for a column - what would break if this column is modified. Returns: @@ -1154,9 +1148,7 @@ def get_column_impact(self, model_name: str, column_name: str) -> Dict[str, Any] raise @staticmethod - def _lookup_column_description( - registry: Any, model_name: str, column_name: str - ) -> Optional[str]: + def _lookup_column_description(registry: Any, model_name: str, column_name: str) -> str | None: """The dbt-authored description of a column, or None if it can't be resolved. Guarded so a stub service without a real registry simply yields no description @@ -1174,11 +1166,11 @@ def _lookup_column_description( def get_changeset_impact( self, - changes: List["ColumnChange"], + changes: list["ColumnChange"], base_service: Optional["LineageService"] = None, *, metabase: Optional["MetabaseReach"] = None, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Aggregate single-column impact across a changeset into one blast radius. ``self`` is the *head* service. Each change is fanned through @@ -1204,10 +1196,10 @@ def get_changeset_impact( # importing here keeps module load order simple and avoids any cycle. from parrant.lineage.changeset import ChangeKind - affected_models: Dict[str, Dict[str, Any]] = {} - affected_columns: Dict[Tuple[str, str], Dict[str, Any]] = {} - affected_exposures: Dict[str, Dict[str, Any]] = {} - by_change: List[Dict[str, Any]] = [] + affected_models: dict[str, dict[str, Any]] = {} + affected_columns: dict[tuple[str, str], dict[str, Any]] = {} + affected_exposures: dict[str, dict[str, Any]] = {} + by_change: list[dict[str, Any]] = [] unresolved = 0 for change in changes: @@ -1255,10 +1247,10 @@ def get_changeset_impact( # node of THIS change — the changed column itself plus every downstream column / # model the dbt reach already resolved. Appended, never re-walked. if metabase is not None: - columns_universe: Set[Tuple[str, str]] = {(change.model, change.column)} + columns_universe: set[tuple[str, str]] = {(change.model, change.column)} for affected in impact["affected_columns"]: columns_universe.add((affected["model"], affected["column"])) - models_universe: Set[str] = {change.model} | { + models_universe: set[str] = {change.model} | { model["name"] for model in impact["affected_models"] } for entry in metabase.reached_dashboards(columns_universe, models_universe): @@ -1318,12 +1310,12 @@ def get_changeset_impact( low_impact_count = len(deduped_columns) - critical_count - filter_count # Guarded so a stub service without a real registry omits confidence rather than erroring. - confidence: Optional[Dict[str, Any]] = None - selection: Optional[Dict[str, Any]] = None - resolution: Optional[Dict[str, Dict[str, Any]]] = None - resolution_summary: Optional[Dict[str, Any]] = None + confidence: dict[str, Any] | None = None + selection: dict[str, Any] | None = None + resolution: dict[str, dict[str, Any]] | None = None + resolution_summary: dict[str, Any] | None = None if getattr(self, "registry", None) is not None: - reachable: Set[str] = set() + reachable: set[str] = set() for change in changes: reachable |= self._dag_reachable_models(change.model) confidence = self._impact_confidence(reachable, len(affected_models)) diff --git a/parrant/lineage/sqlglot_provider.py b/parrant/lineage/sqlglot_provider.py index 940949a..d225cc7 100644 --- a/parrant/lineage/sqlglot_provider.py +++ b/parrant/lineage/sqlglot_provider.py @@ -18,8 +18,6 @@ from __future__ import annotations -from typing import List, Optional - from parrant.artifacts.exceptions import ModelNotFoundError from parrant.artifacts.registry import ModelRegistry from parrant.lineage.provider import LineageAndMetadataProvider @@ -35,7 +33,7 @@ class SqlglotLineageProvider(ModelRegistry): code or test that expects a registry keeps working. """ - def get_column_lineage(self, model_name: str, column_name: str) -> List[ColumnLineage]: + def get_column_lineage(self, model_name: str, column_name: str) -> list[ColumnLineage]: """Per-column upstream edges for ``model.column`` (see the interface). Convenience over ``get_model(model).columns[column].lineage``; empty list when the @@ -46,7 +44,7 @@ def get_column_lineage(self, model_name: str, column_name: str) -> List[ColumnLi return [] return list(column.lineage or []) - def get_column(self, model_name: str, column_name: str) -> Optional[Column]: + def get_column(self, model_name: str, column_name: str) -> Column | None: """Column truth for ``model.column`` (case-insensitive), or ``None`` if unknown.""" try: model = self.get_model(model_name) @@ -55,7 +53,7 @@ def get_column(self, model_name: str, column_name: str) -> Optional[Column]: columns = model.columns return columns.get(column_name) or columns.get(column_name.lower()) - def get_compiled_sql(self, model_name: str) -> Optional[str]: # type: ignore[override] + def get_compiled_sql(self, model_name: str) -> str | None: # type: ignore[override] """Compiled SQL for a model, or ``None`` when there is none. Softens ``ModelRegistry.get_compiled_sql`` (which raises ``ValueError`` / @@ -72,7 +70,7 @@ def get_compiled_sql(self, model_name: str) -> Optional[str]: # type: ignore[ov def build_sqlglot_provider( catalog_path: str, manifest_path: str, - adapter_override: Optional[str] = None, + adapter_override: str | None = None, ) -> LineageAndMetadataProvider: """Construct and :meth:`load` a SQLGlot-backed provider, ready to query. diff --git a/parrant/lineage/verdict.py b/parrant/lineage/verdict.py index d921e22..839d45e 100644 --- a/parrant/lineage/verdict.py +++ b/parrant/lineage/verdict.py @@ -13,10 +13,10 @@ from __future__ import annotations -from typing import Any, Dict, List, Optional, Set, Tuple +from typing import Any -from parrant.lineage.provider import LineageAndMetadataProvider, LineageProvider from parrant.lineage.changeset import ChangeKind, ColumnChange +from parrant.lineage.provider import LineageAndMetadataProvider, LineageProvider from parrant.models.schema import BreakFinding, OverrideVerb, TestNode # Change kinds that can orphan a test by making the column disappear. A rename is emitted by @@ -39,10 +39,10 @@ def _column_missing_in_head(head: LineageProvider, model: str, column: str) -> b def classify_provable_breaks( - changes: List[ColumnChange], + changes: list[ColumnChange], head_registry: LineageAndMetadataProvider, - base_registry: Optional[LineageAndMetadataProvider] = None, -) -> List[BreakFinding]: + base_registry: LineageAndMetadataProvider | None = None, +) -> list[BreakFinding]: """Return the dbt tests that a changeset provably breaks (BREAK-TEST). For each removed/renamed column we look up the tests that targeted it in the *base* @@ -63,10 +63,10 @@ def classify_provable_breaks( """ source = base_registry or head_registry head_test_ids = head_registry.get_test_unique_ids() - findings: List[BreakFinding] = [] + findings: list[BreakFinding] = [] # Dedup across the whole changeset: a relationships test whose child column AND # referenced parent key are both removed in one PR must count once, not twice. - seen: Set[str] = set() + seen: set[str] = set() def _emit(model: str, column: str, kind: str, test: TestNode, via_reference: bool) -> None: # Only a test that survives into head (still declared) can actually fail on build. @@ -118,7 +118,7 @@ def _emit(model: str, column: str, kind: str, test: TestNode, via_reference: boo return findings -def _has_meaning_shift(changes: Optional[List[ColumnChange]]) -> bool: +def _has_meaning_shift(changes: list[ColumnChange] | None) -> bool: """True when any NON-overridden change carries a proven-or-unprovable meaning shift. An ``EQUIVALENT`` edit is never emitted as a change, so a *set* ``semantic`` is always @@ -134,7 +134,7 @@ def _has_meaning_shift(changes: Optional[List[ColumnChange]]) -> bool: ) -def break_is_overridden(break_finding: BreakFinding, changes: Optional[List[ColumnChange]]) -> bool: +def break_is_overridden(break_finding: BreakFinding, changes: list[ColumnChange] | None) -> bool: """True when the change matching this provable break carries an ``allow-break`` override. Fail-safe: only the hard ``allow-break`` verb can demote a break — ``allow-change`` never @@ -155,7 +155,7 @@ def break_is_overridden(break_finding: BreakFinding, changes: Optional[List[Colu return False -def unexcused_break_count(breaks: List[BreakFinding], changes: Optional[List[ColumnChange]]) -> int: +def unexcused_break_count(breaks: list[BreakFinding], changes: list[ColumnChange] | None) -> int: """Number of provable breaks NOT excused by an ``allow-break`` override. This is the count the CI gate (``--fail-on tests``) must read: an acknowledged break is @@ -168,10 +168,10 @@ def unexcused_break_count(breaks: List[BreakFinding], changes: Optional[List[Col _REACHING_MECHANISMS = ("derived_recompute", "rowset_filter") -def _reaching_change_keys(by_change: Optional[List[Dict[str, Any]]]) -> Set[Tuple[str, str, str]]: +def _reaching_change_keys(by_change: list[dict[str, Any]] | None) -> set[tuple[str, str, str]]: """``(model, column, kind)`` keys of changes that reach a recompute/row-set column or an exposure — i.e. the changes that drive a blast-radius REVIEW.""" - keys: Set[Tuple[str, str, str]] = set() + keys: set[tuple[str, str, str]] = set() for entry in by_change or []: if not entry.get("resolved"): continue @@ -186,7 +186,7 @@ def _reaching_change_keys(by_change: Optional[List[Dict[str, Any]]]) -> Set[Tupl return keys -def _change_reaches(change: ColumnChange, reaching_keys: Set[Tuple[str, str, str]]) -> bool: +def _change_reaches(change: ColumnChange, reaching_keys: set[tuple[str, str, str]]) -> bool: """Whether a change contributes to the REVIEW tier: it meaning-shifts, or it reaches a recompute/row-set column or an exposure (per ``by_change``).""" if change.semantic is not None and change.semantic.is_breaking: @@ -195,7 +195,7 @@ def _change_reaches(change: ColumnChange, reaching_keys: Set[Tuple[str, str, str def _all_reaching_overridden( - changes: List[ColumnChange], by_change: Optional[List[Dict[str, Any]]] + changes: list[ColumnChange], by_change: list[dict[str, Any]] | None ) -> bool: """True only when EVERY change that drives a blast-radius review carries an override. @@ -213,10 +213,10 @@ def _all_reaching_overridden( def decide_verdict( - breaks: List[BreakFinding], - summary: Dict[str, Any], - changes: Optional[List[ColumnChange]] = None, - by_change: Optional[List[Dict[str, Any]]] = None, + breaks: list[BreakFinding], + summary: dict[str, Any], + changes: list[ColumnChange] | None = None, + by_change: list[dict[str, Any]] | None = None, ) -> str: """Collapse breaks + blast-radius summary + semantic axis into a single ruling. @@ -272,7 +272,7 @@ def decide_verdict( def _applied_record( change: ColumnChange, downgraded_from: str, downgraded_to: str -) -> Dict[str, Any]: +) -> dict[str, Any]: """A unified honored-override record. Same shape as the policy path so the report's ``overrides`` block and the ``overrides_applied`` count never diverge across the two gates.""" override = change.override @@ -303,7 +303,7 @@ def override_hint(change: ColumnChange, is_break: bool) -> str: return "matched a change that neither breaks nor reaches anything — safe to remove" -def ineffective_override_record(change: ColumnChange, is_break: bool) -> Dict[str, Any]: +def ineffective_override_record(change: ColumnChange, is_break: bool) -> dict[str, Any]: """A unified no-op override record (same base skeleton + a ``hint``).""" override = change.override assert override is not None @@ -315,16 +315,16 @@ def ineffective_override_record(change: ColumnChange, is_break: bool) -> Dict[st def applied_overrides( - changes: List[ColumnChange], - breaks: List[BreakFinding], - by_change: Optional[List[Dict[str, Any]]] = None, -) -> List[Dict[str, Any]]: + changes: list[ColumnChange], + breaks: list[BreakFinding], + by_change: list[dict[str, Any]] | None = None, +) -> list[dict[str, Any]]: """Honored-override records for the DEFAULT (no-policy) gate — one per override that actually lowered its change's contribution. Overrides that changed nothing are skipped here and surface via :func:`ineffective_overrides` instead.""" break_keys = {(b.change_model.lower(), b.change_column.lower()) for b in breaks} reaching = _reaching_change_keys(by_change) - records: List[Dict[str, Any]] = [] + records: list[dict[str, Any]] = [] for change in changes: if change.override is None: continue @@ -340,16 +340,16 @@ def applied_overrides( def ineffective_overrides( - changes: List[ColumnChange], - breaks: List[BreakFinding], - by_change: Optional[List[Dict[str, Any]]] = None, -) -> List[Dict[str, Any]]: + changes: list[ColumnChange], + breaks: list[BreakFinding], + by_change: list[dict[str, Any]] | None = None, +) -> list[dict[str, Any]]: """Override records that resolved to a REAL changed column but produced NO effect (the rename black-hole and friends). Distinct from stale overrides (no matching change at all), these must surface so the author isn't silently ignored.""" break_keys = {(b.change_model.lower(), b.change_column.lower()) for b in breaks} reaching = _reaching_change_keys(by_change) - records: List[Dict[str, Any]] = [] + records: list[dict[str, Any]] = [] for change in changes: if change.override is None: continue diff --git a/parrant/metabase/artifact.py b/parrant/metabase/artifact.py index bb3203e..cc10048 100644 --- a/parrant/metabase/artifact.py +++ b/parrant/metabase/artifact.py @@ -9,7 +9,6 @@ import json from pathlib import Path -from typing import Optional, Union from pydantic import ValidationError @@ -32,7 +31,7 @@ class MetabaseArtifactError(Exception): """ -def load_metabase_lineage(path: Optional[Union[str, Path]]) -> Optional[MetabaseLineage]: +def load_metabase_lineage(path: str | Path | None) -> MetabaseLineage | None: """Parse ``metabase_lineage.json``. Returns ``None`` when ``path`` is falsy or the file does not exist — the Metabase @@ -64,7 +63,7 @@ def load_metabase_lineage(path: Optional[Union[str, Path]]) -> Optional[Metabase raise MetabaseArtifactError(f"Invalid Metabase artifact {file_path}: {exc}") from exc -def dump_metabase_lineage(lineage: MetabaseLineage, path: Union[str, Path]) -> None: +def dump_metabase_lineage(lineage: MetabaseLineage, path: str | Path) -> None: """Write ``lineage`` to ``path`` as JSON (by-alias, so ``schema`` is emitted for the relation's ``schema_name`` field), pretty-printed and stable for diff-friendly snapshots.""" file_path = Path(path) diff --git a/parrant/metabase/cli.py b/parrant/metabase/cli.py index 0ab8522..1c0f229 100644 --- a/parrant/metabase/cli.py +++ b/parrant/metabase/cli.py @@ -11,7 +11,6 @@ import sys from importlib.metadata import PackageNotFoundError, version from pathlib import Path -from typing import Dict, Optional, Tuple import click @@ -28,7 +27,7 @@ def _extractor_version() -> str: return "0.0.0" -def _resolve_dialect(manifest: Optional[str], adapter: Optional[str]) -> Optional[str]: +def _resolve_dialect(manifest: str | None, adapter: str | None) -> str | None: if adapter: return adapter # ``--manifest`` is optional only when ``--adapter`` supplies the dialect directly; the @@ -39,7 +38,7 @@ def _resolve_dialect(manifest: Optional[str], adapter: Optional[str]) -> Optiona return reader.get_adapter() -def _load_dashboard_meta(path: Optional[str]) -> Dict: +def _load_dashboard_meta(path: str | None) -> dict: if not path: return {} return json.loads(Path(path).read_text(encoding="utf-8")) @@ -104,19 +103,19 @@ def _load_dashboard_meta(path: Optional[str]) -> Dict: ) def metabase_extract( metabase_url: str, - metabase_api_key: Optional[str], - metabase_username: Optional[str], - metabase_password: Optional[str], - database_ids: Tuple[int, ...], - manifest: Optional[str], - adapter: Optional[str], + metabase_api_key: str | None, + metabase_username: str | None, + metabase_password: str | None, + database_ids: tuple[int, ...], + manifest: str | None, + adapter: str | None, output: str, include_archived: bool, - dashboard_meta_file: Optional[str], - previous: Optional[str], + dashboard_meta_file: str | None, + previous: str | None, max_workers: int, timeout: int, - fail_under: Optional[float], + fail_under: float | None, ) -> None: """Snapshot Metabase card→column and card→dashboard lineage into an offline artifact.""" if not manifest and not adapter: @@ -176,6 +175,6 @@ def metabase_extract( sys.exit(1) -def main(args: Optional[list] = None, prog_name: Optional[str] = None) -> None: +def main(args: list | None = None, prog_name: str | None = None) -> None: """Entry point used by ``cli/main.py``'s dispatch branch.""" metabase_extract.main(args=args, prog_name=prog_name) diff --git a/parrant/metabase/client.py b/parrant/metabase/client.py index 44d92f6..7a41a48 100644 --- a/parrant/metabase/client.py +++ b/parrant/metabase/client.py @@ -16,7 +16,8 @@ import concurrent.futures import contextlib import time -from typing import Any, Callable, Dict, Iterator, List, Optional +from collections.abc import Callable, Iterator +from typing import Any try: # ``requests`` is a runtime dependency; import lazily so importing the type/schema import requests # modules never forces it (defensive — mirrors the offline guardrail). @@ -56,9 +57,9 @@ class MetabaseClient: def __init__( self, base_url: str, - api_key: Optional[str] = None, - username: Optional[str] = None, - password: Optional[str] = None, + api_key: str | None = None, + username: str | None = None, + password: str | None = None, session: Any = None, timeout: int = 30, max_retries: int = 5, @@ -78,7 +79,7 @@ def __init__( self.max_retries = max_retries self.page_size = page_size self._sleep = sleep - self._session_token: Optional[str] = None + self._session_token: str | None = None self._authenticated = False # --- auth ------------------------------------------------------------- @@ -118,7 +119,7 @@ def ensure_auth(self) -> None: ``_ensure_auth`` and duplicate the ``POST /api/session``.""" self._ensure_auth() - def _headers(self) -> Dict[str, str]: + def _headers(self) -> dict[str, str]: headers = {"Content-Type": "application/json"} if self._api_key: headers["x-api-key"] = self._api_key @@ -127,11 +128,11 @@ def _headers(self) -> Dict[str, str]: return headers # --- transport -------------------------------------------------------- - def _get(self, path: str, params: Optional[Dict[str, Any]] = None) -> Any: + def _get(self, path: str, params: dict[str, Any] | None = None) -> Any: """GET ``path`` with retry/backoff on 429/5xx; returns the parsed JSON body.""" self._ensure_auth() url = f"{self.base_url}{path}" - last_exc: Optional[Exception] = None + last_exc: Exception | None = None for attempt in range(self.max_retries + 1): try: resp = self._session.get( @@ -161,7 +162,7 @@ def _backoff(self, attempt: int, resp: Any = None) -> None: # still de-synchronizing concurrent extractors. self._sleep(delay + (attempt % 3) * 0.1) - def _paginate(self, path: str, extra_params: Optional[Dict[str, Any]] = None) -> Iterator[dict]: + def _paginate(self, path: str, extra_params: dict[str, Any] | None = None) -> Iterator[dict]: """Yield every item from a Metabase list endpoint. Metabase's bulk list endpoints (``/api/card``, ``/api/dashboard``, @@ -181,7 +182,7 @@ def _paginate(self, path: str, extra_params: Optional[Dict[str, Any]] = None) -> offset = 0 for _ in range(_MAX_PAGES): - params: Dict[str, Any] = dict(extra_params or {}) + params: dict[str, Any] = dict(extra_params or {}) params.update({"limit": self.page_size, "offset": offset}) page = self._get(path, params=params) items = page.get("data", []) if isinstance(page, dict) else (page or []) @@ -193,14 +194,14 @@ def _paginate(self, path: str, extra_params: Optional[Dict[str, Any]] = None) -> offset += len(items) # --- endpoints -------------------------------------------------------- - def list_cards(self, include_archived: bool = False) -> List[dict]: + def list_cards(self, include_archived: bool = False) -> list[dict]: """All cards (``GET /api/card``), MBQL pinned to the legacy (v4) serialization.""" params = dict(LEGACY_MBQL_PARAM) if include_archived: params["f"] = "archived" return list(self._paginate("/api/card", extra_params=params)) - def list_dashboards(self) -> List[dict]: + def list_dashboards(self) -> list[dict]: """All dashboards (``GET /api/dashboard``) — summary shells; use :meth:`get_dashboard` for each one's ``dashcards``.""" return list(self._paginate("/api/dashboard")) @@ -209,7 +210,7 @@ def get_dashboard(self, dashboard_id: int) -> dict: """One dashboard with its ``dashcards`` (``GET /api/dashboard/:id``).""" return self._get(f"/api/dashboard/{dashboard_id}") - def get_dashboards(self, dashboard_ids: List[int], max_workers: int = 8) -> Dict[int, dict]: + def get_dashboards(self, dashboard_ids: list[int], max_workers: int = 8) -> dict[int, dict]: """Fetch many dashboards concurrently, returning ``{id: detail}``. Auth is warmed once up front (:meth:`ensure_auth`) so worker threads never race on @@ -225,7 +226,7 @@ def get_dashboards(self, dashboard_ids: List[int], max_workers: int = 8) -> Dict return {did: self.get_dashboard(did) for did in dashboard_ids} workers = max(1, min(max_workers, len(dashboard_ids))) - results: Dict[int, dict] = {} + results: dict[int, dict] = {} executor = concurrent.futures.ThreadPoolExecutor(max_workers=workers) try: futures = {executor.submit(self.get_dashboard, did): did for did in dashboard_ids} @@ -239,11 +240,11 @@ def get_dashboards(self, dashboard_ids: List[int], max_workers: int = 8) -> Dict executor.shutdown(wait=True, cancel_futures=True) return results - def list_snippets(self) -> List[dict]: + def list_snippets(self) -> list[dict]: """All native-query snippets (``GET /api/native-query-snippet``).""" return list(self._paginate("/api/native-query-snippet")) - def server_version(self) -> Optional[str]: + def server_version(self) -> str | None: """Best-effort Metabase version tag (``GET /api/session/properties``) for provenance. Returns ``None`` if the endpoint is unavailable — version stamping is nice-to-have, diff --git a/parrant/metabase/extract.py b/parrant/metabase/extract.py index cb5bf1d..1cfe513 100644 --- a/parrant/metabase/extract.py +++ b/parrant/metabase/extract.py @@ -8,9 +8,10 @@ from __future__ import annotations +from collections.abc import Callable from dataclasses import dataclass, field from datetime import datetime, timezone -from typing import Any, Callable, Dict, List, Optional, Set +from typing import Any from parrant.metabase.client import MetabaseClient from parrant.metabase.resolvers import CardResolver, ResolvedCard @@ -33,32 +34,32 @@ class ExtractConfig: """Inputs to :func:`run_extract` that are not the network client itself.""" metabase_base_url: str - database_ids: List[int] + database_ids: list[int] extractor_version: str - dialect: Optional[str] = None + dialect: str | None = None include_archived: bool = False # Consumer-configurable dashboard meta mapping (spec Q8 /): the tool # never hardcodes an org's taxonomy. Shape: # {"by_collection": {: {...}}, "by_dashboard": {: {...}}} - dashboard_meta: Dict[str, Dict[Any, Dict[str, Any]]] = field(default_factory=dict) + dashboard_meta: dict[str, dict[Any, dict[str, Any]]] = field(default_factory=dict) # A previously-loaded snapshot for incremental reuse. When provided, dashboards whose # Metabase ``updated_at`` matches the previous snapshot are reused rather than refetched # (the N+1 detail fetch is the expensive part at 500+ dashboards). ``None`` = full extract. - previous: Optional["MetabaseLineage"] = None + previous: MetabaseLineage | None = None # Concurrency for the dashboard detail fan-out (``client.get_dashboards``). max_workers: int = 8 def build_dashboard_meta_resolver( - mapping: Dict[str, Dict[Any, Dict[str, Any]]], -) -> Callable[[dict], Dict[str, Any]]: + mapping: dict[str, dict[Any, dict[str, Any]]], +) -> Callable[[dict], dict[str, Any]]: """Turn the consumer mapping into ``dashboard_dict -> meta`` (per-dashboard overrides per-collection). Absent → ``{}``; the taxonomy is entirely the consumer's.""" by_collection = {str(k): v for k, v in (mapping.get("by_collection") or {}).items()} by_dashboard = {str(k): v for k, v in (mapping.get("by_dashboard") or {}).items()} - def resolve(dashboard: dict) -> Dict[str, Any]: - meta: Dict[str, Any] = {} + def resolve(dashboard: dict) -> dict[str, Any]: + meta: dict[str, Any] = {} collection_id = dashboard.get("collection_id") if collection_id is not None and str(collection_id) in by_collection: meta.update(by_collection[str(collection_id)]) @@ -70,12 +71,12 @@ def resolve(dashboard: dict) -> Dict[str, Any]: return resolve -def _dashcard_card_ids(dashboard: dict) -> List[int]: +def _dashcard_card_ids(dashboard: dict) -> list[int]: """Collect card ids from a dashboard's ``dashcards`` (or legacy ``ordered_cards``).""" entries = dashboard.get("dashcards") if entries is None: entries = dashboard.get("ordered_cards") or [] - ids: List[int] = [] + ids: list[int] = [] for entry in entries: card_id = entry.get("card_id") if not isinstance(card_id, int): @@ -86,7 +87,7 @@ def _dashcard_card_ids(dashboard: dict) -> List[int]: return ids -def _person(obj: Any) -> Optional[str]: +def _person(obj: Any) -> str | None: """A human identifier from a Metabase ``creator`` / ``last-edit-info`` object. Prefers email (actionable for ownership routing), falls back to a display name; ``None`` @@ -100,7 +101,7 @@ def _person(obj: Any) -> Optional[str]: return name or None -def _collection_name(obj: Any) -> Optional[str]: +def _collection_name(obj: Any) -> str | None: """The name of an embedded Metabase ``collection`` object, if present.""" if isinstance(obj, dict): name = obj.get("name") @@ -147,9 +148,9 @@ def run_extract(config: ExtractConfig, client: MetabaseClient) -> MetabaseLineag # malformed card with no ``database`` (db is None) keeps the old behavior of being resolved. resolver = CardResolver(meta, corpus, config.dialect) scoped_db_ids = set(config.database_ids) - cards: List[MetabaseCard] = [] - included_card_ids: Set[int] = set() - used_relations: Set[str] = set() + cards: list[MetabaseCard] = [] + included_card_ids: set[int] = set() + used_relations: set[str] = set() for raw_card in raw_cards: if not isinstance(raw_card.get("id"), int): continue @@ -180,7 +181,7 @@ def run_extract(config: ExtractConfig, client: MetabaseClient) -> MetabaseLineag prev_scope_matches = config.previous is not None and set( config.previous.provenance.database_ids ) == set(config.database_ids) - prev_by_id: Dict[int, MetabaseDashboard] = ( + prev_by_id: dict[int, MetabaseDashboard] = ( {d.dashboard_id: d for d in config.previous.dashboards} if config.previous is not None and prev_scope_matches else {} @@ -188,9 +189,9 @@ def run_extract(config: ExtractConfig, client: MetabaseClient) -> MetabaseLineag # Decide reuse vs fetch per shell; a shell is reusable only when both its and the previous # snapshot's ``updated_at`` are present and equal (a missing stamp forces a refetch). - shells_by_id: Dict[int, dict] = {} - fetch_ids: List[int] = [] - reused_ids: Set[int] = set() + shells_by_id: dict[int, dict] = {} + fetch_ids: list[int] = [] + reused_ids: set[int] = set() for shell in shells: shell_id = shell.get("id") if not isinstance(shell_id, int): @@ -210,7 +211,7 @@ def run_extract(config: ExtractConfig, client: MetabaseClient) -> MetabaseLineag details = client.get_dashboards(fetch_ids, max_workers=config.max_workers) if fetch_ids else {} - dashboards: List[MetabaseDashboard] = [] + dashboards: list[MetabaseDashboard] = [] for dashboard_id, shell in shells_by_id.items(): if dashboard_id in reused_ids: # Reused = unchanged since the previous snapshot, so carry its asset metadata @@ -254,7 +255,7 @@ def run_extract(config: ExtractConfig, client: MetabaseClient) -> MetabaseLineag dashboards.sort(key=lambda d: d.dashboard_id) # 4. Relations actually referenced (de-duplicated), + coverage + provenance. - relations: Dict[str, MetabaseRelation] = { + relations: dict[str, MetabaseRelation] = { key: rel for key, rel in meta.relations.items() if key in used_relations } coverage = _build_coverage(cards, dashboards, len(snippets)) @@ -276,7 +277,7 @@ def run_extract(config: ExtractConfig, client: MetabaseClient) -> MetabaseLineag def _build_coverage( - cards: List[MetabaseCard], dashboards: List[MetabaseDashboard], snippets_total: int + cards: list[MetabaseCard], dashboards: list[MetabaseDashboard], snippets_total: int ) -> MetabaseCoverage: column = sum(1 for c in cards if c.precision == "column") table_only = sum(1 for c in cards if c.precision == "table") diff --git a/parrant/metabase/join.py b/parrant/metabase/join.py index d5874c9..cf1f8f5 100644 --- a/parrant/metabase/join.py +++ b/parrant/metabase/join.py @@ -1,4 +1,4 @@ -""" — join the Metabase artifact's warehouse relations back to dbt models. +"""— join the Metabase artifact's warehouse relations back to dbt models. Offline, zero-credential by construction: this module imports ONLY the lineage provider protocol (for typing) and reads ``Model`` fields. It never imports the Metabase client, so @@ -15,7 +15,7 @@ from __future__ import annotations -from typing import Callable, Dict, Optional, Set +from collections.abc import Callable from parrant.lineage.provider import LineageProvider @@ -35,8 +35,8 @@ def normalize_relation(raw: str) -> str: def build_relation_index( provider: LineageProvider, - get_relation_name: Optional[Callable[[str], Optional[str]]] = None, -) -> Dict[str, str]: + get_relation_name: Callable[[str], str | None] | None = None, +) -> dict[str, str]: """Build ``normalized db.schema.table -> dbt model name`` (lowercased model keys). Prefer the manifest ``relation_name`` (via ``get_relation_name(model_name)``, when @@ -47,11 +47,11 @@ def build_relation_index( Pure and offline: reads only ``provider.get_models()`` plus the optional resolver. """ - full: Dict[str, str] = {} - schema_table_owners: Dict[str, Set[str]] = {} + full: dict[str, str] = {} + schema_table_owners: dict[str, set[str]] = {} for name, model in provider.get_models().items(): - keys: Set[str] = set() + keys: set[str] = set() if get_relation_name is not None: relation_name = get_relation_name(name) if relation_name: diff --git a/parrant/metabase/pmbql.py b/parrant/metabase/pmbql.py index 7a9f818..2f1c187 100644 --- a/parrant/metabase/pmbql.py +++ b/parrant/metabase/pmbql.py @@ -26,7 +26,7 @@ from __future__ import annotations -from typing import Any, Dict, List, Optional +from typing import Any # Ref heads whose ``[head, opts, target]`` pMBQL form reorders to legacy ``[head, target, opts]``. _REF_HEADS = {"field", "expression", "aggregation"} @@ -69,9 +69,9 @@ def normalize_dataset_query(query: dict) -> dict: return {"type": "query", "database": database, "query": _fold_stages(stages)} -def _fold_stages(stages: List[dict]) -> dict: +def _fold_stages(stages: list[dict]) -> dict: """Fold innermost-first ``stages`` into legacy nested ``source-query`` form.""" - folded: Optional[dict] = None + folded: dict | None = None for stage in stages: if not isinstance(stage, dict): continue @@ -84,7 +84,7 @@ def _fold_stages(stages: List[dict]) -> dict: def _normalize_stage(stage: dict) -> dict: """Convert one ``mbql.stage/mbql`` stage into a legacy query dict.""" - out: Dict[str, Any] = {} + out: dict[str, Any] = {} _set_source(out, stage.get("source-table"), stage.get("source-card")) for key in ("breakout", "aggregation", "fields", "order-by", "expressions"): @@ -106,7 +106,7 @@ def _normalize_stage(stage: dict) -> dict: def _normalize_join(join: dict) -> dict: """Convert a pMBQL join to the legacy ``{source-table, condition, alias}`` shape.""" - out: Dict[str, Any] = {} + out: dict[str, Any] = {} alias = join.get("alias") or join.get("join-alias") if alias: out["alias"] = alias @@ -131,7 +131,7 @@ def _normalize_join(join: dict) -> dict: return out -def _set_source(out: Dict[str, Any], source_table: Any, source_card: Any) -> None: +def _set_source(out: dict[str, Any], source_table: Any, source_card: Any) -> None: """Write the legacy ``source-table`` key (``card__`` for an upstream card).""" if isinstance(source_card, int): out["source-table"] = f"card__{source_card}" @@ -169,7 +169,7 @@ def _clean_opts(opts: dict) -> dict: return {key: value for key, value in opts.items() if key not in _OPTS_NOISE} -def _normalize_template_tags(tags: Any) -> Dict[str, Any]: +def _normalize_template_tags(tags: Any) -> dict[str, Any]: """Convert a pMBQL template-tags **list** into the legacy **dict** keyed by tag name. A dimension tag's ``dimension`` value is a pMBQL field ref → reorder it per the ref rule. @@ -179,7 +179,7 @@ def _normalize_template_tags(tags: Any) -> Dict[str, Any]: return tags if not isinstance(tags, list): return {} - out: Dict[str, Any] = {} + out: dict[str, Any] = {} for tag in tags: if not isinstance(tag, dict): continue diff --git a/parrant/metabase/reach.py b/parrant/metabase/reach.py index 5f32e74..917744f 100644 --- a/parrant/metabase/reach.py +++ b/parrant/metabase/reach.py @@ -1,4 +1,4 @@ -""" — the offline Metabase reach index: ``(dbt_model, column) -> cards -> dashboards``. +"""— the offline Metabase reach index: ``(dbt_model, column) -> cards -> dashboards``. Built once from a loaded :class:`MetabaseLineage` artifact + the relation join (:mod:`parrant.metabase.join`). Pure, offline, zero-credential — it imports ONLY @@ -15,8 +15,9 @@ from __future__ import annotations +from collections.abc import Iterable from datetime import datetime, timezone -from typing import Any, Dict, Iterable, List, Optional, Set, Tuple +from typing import Any from parrant.metabase.join import normalize_relation from parrant.models.schema import ( @@ -45,10 +46,10 @@ class MetabaseReach: def __init__( self, - column_cards: Dict[Tuple[str, str], Dict[int, str]], - model_cards: Dict[str, Set[int]], - dashboards_by_card: Dict[int, List[MetabaseDashboard]], - dashboard_by_name: Dict[str, MetabaseDashboard], + column_cards: dict[tuple[str, str], dict[int, str]], + model_cards: dict[str, set[int]], + dashboards_by_card: dict[int, list[MetabaseDashboard]], + dashboard_by_name: dict[str, MetabaseDashboard], ) -> None: self._column_cards = column_cards self._model_cards = model_cards @@ -56,7 +57,7 @@ def __init__( self._dashboard_by_name = dashboard_by_name @classmethod - def build(cls, lineage: MetabaseLineage, relation_index: Dict[str, str]) -> "MetabaseReach": + def build(cls, lineage: MetabaseLineage, relation_index: dict[str, str]) -> MetabaseReach: """Invert the artifact into the reach index using the dbt relation join. A card's column refs (column-precise) map ``(relation, column) -> (dbt_model, column)``; @@ -65,8 +66,8 @@ def build(cls, lineage: MetabaseLineage, relation_index: Dict[str, str]) -> "Met card never over-fires on a model-level match. Relations that don't join to any dbt model are dropped (honest: no guess). """ - column_cards: Dict[Tuple[str, str], Dict[int, str]] = {} - model_cards: Dict[str, Set[int]] = {} + column_cards: dict[tuple[str, str], dict[int, str]] = {} + model_cards: dict[str, set[int]] = {} for card in lineage.cards: for ref in card.columns: @@ -89,8 +90,8 @@ def build(cls, lineage: MetabaseLineage, relation_index: Dict[str, str]) -> "Met continue model_cards.setdefault(model, set()).add(card.card_id) - dashboards_by_card: Dict[int, List[MetabaseDashboard]] = {} - dashboard_by_name: Dict[str, MetabaseDashboard] = {} + dashboards_by_card: dict[int, list[MetabaseDashboard]] = {} + dashboard_by_name: dict[str, MetabaseDashboard] = {} for dashboard in lineage.dashboards: dashboard_by_name[dashboard_reach_name(dashboard.dashboard_id)] = dashboard for card_id in dashboard.card_ids: @@ -102,9 +103,9 @@ def build(cls, lineage: MetabaseLineage, relation_index: Dict[str, str]) -> "Met def reached_dashboards( self, - columns: Iterable[Tuple[str, str]], + columns: Iterable[tuple[str, str]], models: Iterable[str], - ) -> List[Dict[str, Any]]: + ) -> list[dict[str, Any]]: """Dashboards reached by any of ``columns`` (column-precise) or ``models`` (table grain). ``columns`` is the set of ``(model, column)`` the change touches — the changed column @@ -113,13 +114,13 @@ def reached_dashboards( dashboard, deterministically ordered by dashboard id, ready to append onto ``affected_exposures`` and each change's ``reached_exposures``. """ - hits: Dict[int, Dict[str, Any]] = {} + hits: dict[int, dict[str, Any]] = {} def _record( card_id: int, precision: str, - matched: Optional[Tuple[str, str]] = None, - role: Optional[str] = None, + matched: tuple[str, str] | None = None, + role: str | None = None, ) -> None: for dashboard in self._dashboards_by_card.get(card_id, ()): entry = hits.setdefault( @@ -148,7 +149,7 @@ def _record( for card_id in self._model_cards.get(model.lower(), ()): _record(card_id, "table") - entries: List[Dict[str, Any]] = [] + entries: list[dict[str, Any]] = [] for dashboard_id in sorted(hits): info = hits[dashboard_id] dashboard: MetabaseDashboard = info["dashboard"] @@ -175,7 +176,7 @@ def _record( ) return entries - def dashboard_meta(self, name: str) -> Optional[Dict[str, Any]]: + def dashboard_meta(self, name: str) -> dict[str, Any] | None: """Resolve a reached dashboard's meta for the policy ``reach.where`` clause. ``name`` is the synthetic ``metabase.dashboard.``. Returns the dashboard's ``meta`` @@ -186,7 +187,7 @@ def dashboard_meta(self, name: str) -> Optional[Dict[str, Any]]: dashboard = self._dashboard_by_name.get(name) if dashboard is None: return None - merged: Dict[str, Any] = {"source": "metabase", "name": dashboard.name} + merged: dict[str, Any] = {"source": "metabase", "name": dashboard.name} if dashboard.url is not None: merged["url"] = dashboard.url merged.update(dashboard.meta) @@ -194,10 +195,10 @@ def dashboard_meta(self, name: str) -> Optional[Dict[str, Any]]: def build_reach_confidence( - lineage: Optional[MetabaseLineage], - reached_entries: List[Dict[str, Any]], + lineage: MetabaseLineage | None, + reached_entries: list[dict[str, Any]], max_age_hours: float = 24.0, - now: Optional[datetime] = None, + now: datetime | None = None, ) -> MetabaseReachConfidence: """Summarise the honesty of the appended Metabase reach. @@ -209,7 +210,7 @@ def build_reach_confidence( return MetabaseReachConfidence(level="absent") generated_at = lineage.provenance.generated_at - age_hours: Optional[float] = None + age_hours: float | None = None parsed = _parse_iso8601(generated_at) if parsed is not None: reference = now or datetime.now(timezone.utc) @@ -235,7 +236,7 @@ def build_reach_confidence( ) -def _parse_iso8601(value: str) -> Optional[datetime]: +def _parse_iso8601(value: str) -> datetime | None: """Parse an ISO-8601 timestamp, tolerating a trailing ``Z``. ``None`` on failure.""" try: normalized = value.replace("Z", "+00:00") if value.endswith("Z") else value diff --git a/parrant/metabase/resolvers.py b/parrant/metabase/resolvers.py index 287e1af..6450756 100644 --- a/parrant/metabase/resolvers.py +++ b/parrant/metabase/resolvers.py @@ -18,12 +18,12 @@ import re from dataclasses import dataclass, field -from typing import Any, Dict, List, Literal, Optional, Set, Tuple +from typing import Any, Literal -from parrant.models.schema import MetabaseColumnRef -from parrant.parser.sql_parser import SQLColumnParser from parrant.metabase.pmbql import normalize_dataset_query from parrant.metabase.warehouse_meta import CardCorpus, WarehouseMeta +from parrant.models.schema import MetabaseColumnRef +from parrant.parser.sql_parser import SQLColumnParser # The roles a resolved column may carry (matches ``MetabaseColumnRef.role``). Role = Literal["field", "breakout", "aggregation", "filter", "join", "order", "native"] @@ -43,16 +43,16 @@ class ResolvedCard: """The unified resolver output the extractor turns into a ``MetabaseCard``.""" precision: str # "column" | "table" | "none" - columns: List[MetabaseColumnRef] = field(default_factory=list) - table_relations: List[str] = field(default_factory=list) - upstream_card_ids: List[int] = field(default_factory=list) - snippet_ids: List[int] = field(default_factory=list) - unresolved_reason: Optional[str] = None + columns: list[MetabaseColumnRef] = field(default_factory=list) + table_relations: list[str] = field(default_factory=list) + upstream_card_ids: list[int] = field(default_factory=list) + snippet_ids: list[int] = field(default_factory=list) + unresolved_reason: str | None = None -def _iter_field_refs(node: Any) -> List[list]: +def _iter_field_refs(node: Any) -> list[list]: """Recursively collect every MBQL ``["field", , opts]`` clause under ``node``.""" - found: List[list] = [] + found: list[list] = [] if isinstance(node, list): if node and node[0] == "field": found.append(node) @@ -72,15 +72,15 @@ def __init__( self, meta: WarehouseMeta, corpus: CardCorpus, - dialect: Optional[str], - parser: Optional[SQLColumnParser] = None, + dialect: str | None, + parser: SQLColumnParser | None = None, ) -> None: self.meta = meta self.corpus = corpus self.dialect = dialect self.parser = parser or SQLColumnParser(dialect) - self._cache: Dict[int, ResolvedCard] = {} - self._resolving: Set[int] = set() + self._cache: dict[int, ResolvedCard] = {} + self._resolving: set[int] = set() # --- entry point ------------------------------------------------------ def resolve_card(self, card: dict) -> ResolvedCard: @@ -151,13 +151,13 @@ def resolve_mbql(self, query: dict, card: dict) -> ResolvedCard: unresolved_reason=reason, ) - def _resolve_mbql_query(self, query: dict, acc: "_MbqlAccumulator") -> None: + def _resolve_mbql_query(self, query: dict, acc: _MbqlAccumulator) -> None: """Resolve one (possibly nested) MBQL query into ``acc``. ``name_map`` (column name → (relation_key, column)) lets field-by-name refs — common when the source is a card / nested query — resolve against the source's output. """ - name_map: Dict[str, Tuple[str, str]] = {} + name_map: dict[str, tuple[str, str]] = {} source = query.get("source-table") if isinstance(source, str) and source.startswith("card__"): try: @@ -180,7 +180,7 @@ def _resolve_mbql_query(self, query: dict, acc: "_MbqlAccumulator") -> None: self._resolve_mbql_query(nested, acc) # Each clause contributes fields with a distinct role. - clause_roles: List[Tuple[str, Role]] = [ + clause_roles: list[tuple[str, Role]] = [ ("fields", "field"), ("breakout", "breakout"), ("aggregation", "aggregation"), @@ -214,8 +214,8 @@ def _collect_clause( self, node: Any, role: Role, - acc: "_MbqlAccumulator", - name_map: Dict[str, Tuple[str, str]], + acc: _MbqlAccumulator, + name_map: dict[str, tuple[str, str]], ) -> None: if node is None: return @@ -255,11 +255,11 @@ def resolve_native(self, query: dict, card: dict) -> ResolvedCard: sql, tags, visited_cards=set(), visited_snippets=set() ) - columns: Dict[Tuple[str, str], MetabaseColumnRef] = {} + columns: dict[tuple[str, str], MetabaseColumnRef] = {} for ref in dim_columns: columns[(ref.relation, ref.column)] = ref - table_relations: Set[str] = set() - reason: Optional[str] = None + table_relations: set[str] = set() + reason: str | None = None try: result = self.parser.parse_column_lineage(expanded) @@ -329,9 +329,9 @@ def _map_source_column( self, source: str, role: Role, - synthetic: Dict[int, Set[str]], - columns: Dict[Tuple[str, str], MetabaseColumnRef], - table_relations: Set[str], + synthetic: dict[int, set[str]], + columns: dict[tuple[str, str], MetabaseColumnRef], + table_relations: set[str], ) -> None: """Map a parser ``table.column`` source to a relation, or degrade to table grain.""" if "." in source: @@ -354,8 +354,8 @@ def _map_source_column( relation=key, column=column.lower(), role=role, confidence="medium" ) - def _synthetic_relations(self, synthetic: Dict[int, Set[str]]) -> Set[str]: - out: Set[str] = set() + def _synthetic_relations(self, synthetic: dict[int, set[str]]) -> set[str]: + out: set[str] = set() for relations in synthetic.values(): out.update(relations) return out @@ -365,9 +365,9 @@ def _expand_tags( self, sql: str, tags: dict, - visited_cards: Set[int], - visited_snippets: Set[int], - ) -> Tuple[str, Set[int], Set[int], Dict[int, Set[str]], List[MetabaseColumnRef]]: + visited_cards: set[int], + visited_snippets: set[int], + ) -> tuple[str, set[int], set[int], dict[int, set[str]], list[MetabaseColumnRef]]: """Substitute template tags so the SQL parses, returning expansion side-channels. Returns ``(expanded_sql, upstream_card_ids, snippet_ids, synthetic, dim_columns)`` @@ -375,10 +375,10 @@ def _expand_tags( (so ``__card_`` tokens resolve to table-grain reach) and ``dim_columns`` are the precisely-recovered field-filter columns (spec Q4). """ - upstream_card_ids: Set[int] = set() - snippet_ids: Set[int] = set() - synthetic: Dict[int, Set[str]] = {} - dim_columns: List[MetabaseColumnRef] = [] + upstream_card_ids: set[int] = set() + snippet_ids: set[int] = set() + synthetic: dict[int, set[str]] = {} + dim_columns: list[MetabaseColumnRef] = [] # Unwrap Metabase optional blocks [[ ... ]] so inner SQL/tags survive. expanded = sql.replace("[[", " ").replace("]]", " ") @@ -401,12 +401,12 @@ def _expand_tags( ) # Card references {{#123}} → a table token; inherit the card's relations. - def _card_sub(match: "re.Match[str]") -> str: + def _card_sub(match: re.Match[str]) -> str: card_id = int(match.group(1)) upstream_card_ids.add(card_id) if card_id not in visited_cards: sub = self._resolve_card_id(card_id) - relations: Set[str] = set(sub.table_relations) + relations: set[str] = set(sub.table_relations) for ref in sub.columns: relations.add(ref.relation) synthetic[card_id] = relations @@ -416,7 +416,7 @@ def _card_sub(match: "re.Match[str]") -> str: # Snippet references {{snippet: name}} → inline the snippet content (one level of # transitive expansion, cycle-guarded). - def _snippet_sub(match: "re.Match[str]") -> str: + def _snippet_sub(match: re.Match[str]) -> str: name = match.group(1).strip() snippet = self.corpus.snippet_by_name(name) if snippet is None: @@ -441,7 +441,7 @@ def _snippet_sub(match: "re.Match[str]") -> str: expanded = _SNIPPET_TAG_RE.sub(_snippet_sub, expanded) # Remaining variables (text/number/date/dimension placeholders) → safe literals. - def _var_sub(match: "re.Match[str]") -> str: + def _var_sub(match: re.Match[str]) -> str: name = match.group(1).strip() tag = tags.get(name) or {} if tag.get("type") == "dimension": @@ -451,7 +451,7 @@ def _var_sub(match: "re.Match[str]") -> str: expanded = _VAR_TAG_RE.sub(_var_sub, expanded) return expanded, upstream_card_ids, snippet_ids, synthetic, dim_columns - def _extract_tables(self, sql: str) -> Set[str]: + def _extract_tables(self, sql: str) -> set[str]: """Cheap table extraction for the parse-failed degrade path (sqlglot table walk).""" try: from sqlglot import exp, parse_one @@ -459,7 +459,7 @@ def _extract_tables(self, sql: str) -> Set[str]: parsed = parse_one(sql, dialect=self.dialect) except Exception: return set() - names: Set[str] = set() + names: set[str] = set() for table in parsed.find_all(exp.Table): parts = [p for p in (table.catalog, table.db, table.name) if p] if parts: @@ -469,7 +469,7 @@ def _extract_tables(self, sql: str) -> Set[str]: @dataclass class _MbqlAccumulator: - relations: Set[str] = field(default_factory=set) - columns: Dict[Tuple[str, str], MetabaseColumnRef] = field(default_factory=dict) - upstream_card_ids: Set[int] = field(default_factory=set) + relations: set[str] = field(default_factory=set) + columns: dict[tuple[str, str], MetabaseColumnRef] = field(default_factory=dict) + upstream_card_ids: set[int] = field(default_factory=set) unknown_field: bool = False diff --git a/parrant/metabase/warehouse_meta.py b/parrant/metabase/warehouse_meta.py index 7f41af8..bec1f52 100644 --- a/parrant/metabase/warehouse_meta.py +++ b/parrant/metabase/warehouse_meta.py @@ -1,4 +1,4 @@ -""" support — in-memory warehouse metadata + the card/snippet corpus. +"""support — in-memory warehouse metadata + the card/snippet corpus. :class:`WarehouseMeta` turns Metabase's bulk ``GET /api/database/:id/metadata`` into fast lookups both resolvers need: Table/Field **id** → warehouse relation/column (for MBQL) and @@ -11,8 +11,6 @@ from __future__ import annotations -from typing import Dict, List, Optional, Tuple - from parrant.models.schema import MetabaseRelation @@ -43,18 +41,18 @@ class WarehouseMeta: """Resolved Metabase warehouse metadata across one or more databases.""" def __init__(self) -> None: - self.relations: Dict[str, MetabaseRelation] = {} + self.relations: dict[str, MetabaseRelation] = {} # Field id -> (relation_key, column_name) - self._field_by_id: Dict[int, Tuple[str, str]] = {} + self._field_by_id: dict[int, tuple[str, str]] = {} # Table id -> relation_key - self._table_by_id: Dict[int, str] = {} + self._table_by_id: dict[int, str] = {} # name lookups: "table", "schema.table", "db.schema.table" -> relation_key - self._by_name: Dict[str, str] = {} + self._by_name: dict[str, str] = {} # name collisions on the bare-table key: once ambiguous, never guess (drop it) self._ambiguous_names: set = set() @classmethod - def from_database_metadata(cls, metadatas: List[dict]) -> "WarehouseMeta": + def from_database_metadata(cls, metadatas: list[dict]) -> WarehouseMeta: """Build from a list of ``/api/database/:id/metadata`` response bodies.""" meta = cls() for metadata in metadatas: @@ -101,16 +99,16 @@ def _register_name(self, name: str, key: str) -> None: self._by_name[name] = key # --- id lookups (MBQL) ------------------------------------------------ - def field(self, field_id: int) -> Optional[Tuple[str, str]]: + def field(self, field_id: int) -> tuple[str, str] | None: """``field_id`` → ``(relation_key, column)`` or ``None`` if unknown.""" return self._field_by_id.get(field_id) - def table(self, table_id: int) -> Optional[str]: + def table(self, table_id: int) -> str | None: """``table_id`` → relation_key or ``None`` if unknown.""" return self._table_by_id.get(table_id) # --- name lookups (native SQL) ---------------------------------------- - def resolve_name(self, raw_name: str) -> Optional[str]: + def resolve_name(self, raw_name: str) -> str | None: """Resolve a SQL table reference to a relation_key. Tries the fully-qualified name first, then ``schema.table``, then the bare table @@ -130,29 +128,29 @@ def resolve_name(self, raw_name: str) -> Optional[str]: return key return None - def relation(self, key: str) -> Optional[MetabaseRelation]: + def relation(self, key: str) -> MetabaseRelation | None: return self.relations.get(key) class CardCorpus: """Every fetched card + snippet, indexed for transitive template-tag expansion.""" - def __init__(self, cards: List[dict], snippets: List[dict]) -> None: - self.cards_by_id: Dict[int, dict] = { + def __init__(self, cards: list[dict], snippets: list[dict]) -> None: + self.cards_by_id: dict[int, dict] = { c["id"]: c for c in cards if isinstance(c.get("id"), int) } - self.snippets_by_id: Dict[int, dict] = { + self.snippets_by_id: dict[int, dict] = { s["id"]: s for s in snippets if isinstance(s.get("id"), int) } - self.snippets_by_name: Dict[str, dict] = { + self.snippets_by_name: dict[str, dict] = { str(s.get("name", "")).lower(): s for s in snippets if s.get("name") } - def card(self, card_id: int) -> Optional[dict]: + def card(self, card_id: int) -> dict | None: return self.cards_by_id.get(card_id) - def snippet_by_id(self, snippet_id: int) -> Optional[dict]: + def snippet_by_id(self, snippet_id: int) -> dict | None: return self.snippets_by_id.get(snippet_id) - def snippet_by_name(self, name: str) -> Optional[dict]: + def snippet_by_name(self, name: str) -> dict | None: return self.snippets_by_name.get(name.lower()) diff --git a/parrant/models/schema.py b/parrant/models/schema.py index 70fa3a0..1a59107 100644 --- a/parrant/models/schema.py +++ b/parrant/models/schema.py @@ -1,7 +1,7 @@ from enum import Enum +from typing import Any, Literal, Optional -from pydantic import BaseModel, Field, ConfigDict, model_validator -from typing import List, Optional, Set, Dict, Literal, Any +from pydantic import BaseModel, ConfigDict, Field, model_validator class OverrideVerb(str, Enum): @@ -33,16 +33,16 @@ class OverrideDirective(BaseModel): verb: OverrideVerb # Lowercased target column. ``None`` => model scope OR an unresolved line-adjacency. - column: Optional[str] = None + column: str | None = None reason: str # ``model`` when no column arg and the pragma precedes the first SELECT; ``column`` otherwise. scope: Literal["column", "model"] # 1-indexed line within the scanned (compiled) head SQL — NOTE: compiled-relative. source_line: int # Set by the changeset builder once it knows which model this SQL belongs to. - model: Optional[str] = None + model: str | None = None - def to_record(self) -> Dict[str, Any]: + def to_record(self) -> dict[str, Any]: """The report dict skeleton for an override record (stale / applied / ineffective).""" return { "model": self.model, @@ -86,10 +86,10 @@ class SemanticDiff(BaseModel): class ColumnLineage(BaseModel): - source_columns: Set[str] + source_columns: set[str] transformation_type: Literal["direct", "renamed", "derived"] - sql_expression: Optional[str] = None - description: Optional[str] = None + sql_expression: str | None = None + description: str | None = None class UnresolvedColumnEdge(BaseModel): @@ -121,16 +121,16 @@ class UnresolvedColumnEdge(BaseModel): "pivot_output", "other", ] - detail: Optional[str] = None + detail: str | None = None class Column(BaseModel): name: str model_name: str - description: Optional[str] = None - data_type: Optional[str] = None - lineage: Optional[List[ColumnLineage]] = Field(default_factory=list) # type: ignore - metadata: Optional[Dict[str, Any]] = None + description: str | None = None + data_type: str | None = None + lineage: list[ColumnLineage] | None = Field(default_factory=list) # type: ignore + metadata: dict[str, Any] | None = None @property def full_name(self) -> str: @@ -140,13 +140,13 @@ def full_name(self) -> str: class Exposure(BaseModel): name: str type: str - url: Optional[str] = None - description: Optional[str] = None - owner: Optional[Dict[str, Any]] = None + url: str | None = None + description: str | None = None + owner: dict[str, Any] | None = None unique_id: str - depends_on_models: Set[str] = Field(default_factory=set) - resource_path: Optional[str] = None - metadata: Optional[Dict[str, Any]] = None + depends_on_models: set[str] = Field(default_factory=set) + resource_path: str | None = None + metadata: dict[str, Any] | None = None class TestNode(BaseModel): @@ -166,16 +166,16 @@ class TestNode(BaseModel): test_name: str # The model the test is attached to (lowercased). ``None`` when it cannot be # attributed honestly (e.g. a singular/custom test with no clear model). - target_model: Optional[str] = None + target_model: str | None = None # The column under test (lowercased). ``None`` for tests with no column (e.g. # a model-level test) — never guessed. - target_column: Optional[str] = None + target_column: str | None = None # For ``relationships`` tests: the referenced ("parent") side. ``referenced_model`` # comes from ``kwargs.to`` (a ``ref(...)``), ``referenced_column`` from ``kwargs.field``. - referenced_model: Optional[str] = None - referenced_column: Optional[str] = None + referenced_model: str | None = None + referenced_column: str | None = None # ``original_file_path`` — the ``file:line`` the reviewer would open to fix it. - resource_path: Optional[str] = None + resource_path: str | None = None class BreakFinding(BaseModel): @@ -199,7 +199,7 @@ class of impact objective enough to *block* a PR — see the verdict classifier. # through its *referenced* (parent) side rather than its own target column. test_name: str test_unique_id: str - resource_path: Optional[str] = None + resource_path: str | None = None via_reference: bool = False def code(self) -> str: @@ -209,7 +209,7 @@ def code(self) -> str: class ModelDependency(BaseModel): model_name: str - depends_on: Set[str] + depends_on: set[str] class Model(BaseModel): @@ -218,41 +218,41 @@ class Model(BaseModel): name: str schema_name: str = Field(alias="schema") # Handle base model shadow attribute `schema` database: str - columns: Dict[str, Column] = Field(default_factory=dict) - metadata: Optional[Dict[str, Any]] = None - unique_id: Optional[str] = None - upstream: Set[str] = Field(default_factory=set) - downstream: Set[str] = Field(default_factory=set) + columns: dict[str, Column] = Field(default_factory=dict) + metadata: dict[str, Any] | None = None + unique_id: str | None = None + upstream: set[str] = Field(default_factory=set) + downstream: set[str] = Field(default_factory=set) # Upstream columns this model uses only in predicates (WHERE / JOIN / HAVING / QUALIFY). - predicate_sources: Set[str] = Field(default_factory=set) + predicate_sources: set[str] = Field(default_factory=set) # Upstream column -> the predicate condition text it appears in. - predicate_lineage: Dict[str, str] = Field(default_factory=dict) - compiled_sql: Optional[str] = None - language: Optional[str] = None + predicate_lineage: dict[str, str] = Field(default_factory=dict) + compiled_sql: str | None = None + language: str | None = None resource_type: Literal["model", "source", "seed", "test", "exposure", "snapshot"] - resource_path: Optional[str] = None - source_identifier: Optional[str] = None - source_name: Optional[str] = None - description: Optional[str] = None - tags: List[str] = Field(default_factory=list) + resource_path: str | None = None + source_identifier: str | None = None + source_name: str | None = None + description: str | None = None + tags: list[str] = Field(default_factory=list) class SQLParseResult(BaseModel): - column_lineage: Dict[str, List[ColumnLineage]] - star_sources: Set[str] = Field(default_factory=set) + column_lineage: dict[str, list[ColumnLineage]] + star_sources: set[str] = Field(default_factory=set) # Upstream columns referenced only in predicates (WHERE / JOIN ON / HAVING / QUALIFY), # never projected. A change to one of these alters this model's row-set (and therefore # its aggregates), so it is a real — if indirect — downstream impact. - predicate_sources: Set[str] = Field(default_factory=set) + predicate_sources: set[str] = Field(default_factory=set) # Upstream column -> the predicate condition text it appears in (the "why" for the # row-set impact, e.g. ``status = 'flagged'``). - predicate_lineage: Dict[str, str] = Field(default_factory=dict) + predicate_lineage: dict[str, str] = Field(default_factory=dict) # Columns whose upstream source parrant could not genuinely resolve — a phantom flatten # alias, a quoted pivot literal, a ``select * rename`` output, or a ``select *`` off an # unresolvable relation. The parser emits these INSTEAD of fabricating a source column, so # the set is complete and uncapped (display layers may cap; the machine surface must not). # ``model`` is left empty here (the parser has no model name); the registry stamps it. - unresolved_edges: List[UnresolvedColumnEdge] = Field(default_factory=list) + unresolved_edges: list[UnresolvedColumnEdge] = Field(default_factory=list) class Coverage(BaseModel): @@ -262,15 +262,15 @@ class Coverage(BaseModel): parse_failed: int skipped_no_sql: int not_in_catalog_count: int - failed_models: List[str] = Field(default_factory=list) - skipped_models: List[str] = Field(default_factory=list) + failed_models: list[str] = Field(default_factory=list) + skipped_models: list[str] = Field(default_factory=list) # Nodes whose compiled SQL the parser cannot read and that we deliberately do NOT # analyze at the column level (semantic views chief among them; more generally any # unparseable node). These are NOT a coverage failure — we chose not to parse them, so # they are excluded from the ``complete`` denominator and never counted as ``parse_failed``. # Model-level reach through them is still preserved from the manifest dependency graph. opaque: int = 0 - opaque_models: List[str] = Field(default_factory=list) + opaque_models: list[str] = Field(default_factory=list) complete: bool @@ -294,8 +294,8 @@ class ImpactConfidence(BaseModel): # len(parse_failed_models) == parse_failed always hold in machine output. # Display layers (markdown) cap the rendered names and set the *_truncated flags # below; the integer counts above remain the source of truth for totals. - no_column_info_models: List[str] = Field(default_factory=list) - parse_failed_models: List[str] = Field(default_factory=list) + no_column_info_models: list[str] = Field(default_factory=list) + parse_failed_models: list[str] = Field(default_factory=list) # Reachable models that DO expose columns but carry at least one unresolved-edge marker # (a phantom flatten alias, a quoted pivot literal, a ``select * rename`` output, or a # ``select *`` off an unresolvable relation — see :class:`UnresolvedColumnEdge`). Unlike @@ -305,7 +305,7 @@ class ImpactConfidence(BaseModel): # ``partial`` (which widens the rebuild). COMPLETE, uncapped machine list — a fail-closed # consumer force-rebuilds every one, so len(partial_edges_models) == partial_edges always. partial_edges: int = 0 - partial_edges_models: List[str] = Field(default_factory=list) + partial_edges_models: list[str] = Field(default_factory=list) # Reachable nodes we deliberately do NOT analyze at the column level: their compiled SQL # is unparseable (semantic views chief among them; generally any node the parser cannot # read). Distinct from ``parse_failed`` (a node we *tried* and failed to derive lineage @@ -316,7 +316,7 @@ class ImpactConfidence(BaseModel): # column edges through them, so we widen the rebuild rather than prove anything skippable). # COMPLETE, uncapped machine list — len(opaque_models) == opaque always in machine output. opaque: int = 0 - opaque_models: List[str] = Field(default_factory=list) + opaque_models: list[str] = Field(default_factory=list) # Display-only truncation signals: False in machine output (lists are complete), # set True only by a display layer when it elided names from the rendered list. no_column_info_truncated: bool = False @@ -344,10 +344,10 @@ class Selection(BaseModel): # and exits green): branch on this, never on the selector string's emptiness. has_rebuild: bool = False # Sorted, deduplicated dbt node names that must be rebuilt. - rebuild_models: List[str] = Field(default_factory=list) + rebuild_models: list[str] = Field(default_factory=list) # Sorted reachable complement — reached only by an additive/passthrough change at full # confidence. Informational: the consumer decides whether to actually skip these. - skippable_models: List[str] = Field(default_factory=list) + skippable_models: list[str] = Field(default_factory=list) # Space-joined ``rebuild_models`` — a drop-in for ``dbt build --select $(...)``. # Empty string exactly when ``has_rebuild`` is False. rebuild_selector: str = "" @@ -384,7 +384,7 @@ class ModelResolution(BaseModel): # For a ``partial_edges`` status it is the marker construct (phantom_alias, unexpandable_star, # fabricated_column, star_rename, pivot_output, other). For an ``opaque`` status it names WHY # we did not parse it (``semantic_view`` or the general ``unparseable_sql``). - reason: Optional[str] = None + reason: str | None = None class ResolutionReasonCount(BaseModel): @@ -425,7 +425,7 @@ class ResolutionSummary(BaseModel): # rebuild volume is driven by unresolved models rather than proven changes. rebuild_forced_by_nonresolution: int = 0 # Coarse reasons ranked by frequency (advisory): the ranked resolution-gap backlog. - top_reasons: List[ResolutionReasonCount] = Field(default_factory=list) + top_reasons: list[ResolutionReasonCount] = Field(default_factory=list) # --------------------------------------------------------------------------- @@ -446,10 +446,10 @@ class MetabaseProvenance(BaseModel): generated_at: str # ISO-8601 UTC — stamps snapshot age metabase_base_url: str - metabase_version: Optional[str] = None - database_ids: List[int] = Field(default_factory=list) + metabase_version: str | None = None + database_ids: list[int] = Field(default_factory=list) extractor_version: str - dbt_adapter: Optional[str] = None # dialect the native resolver parsed with + dbt_adapter: str | None = None # dialect the native resolver parsed with class MetabaseCoverage(BaseModel): @@ -465,8 +465,8 @@ class MetabaseCoverage(BaseModel): cards_unresolved: int # no warehouse relation resolved at all dashboards_total: int snippets_total: int - unresolved_card_ids: List[int] = Field(default_factory=list) # capped honesty sample - table_only_card_ids: List[int] = Field(default_factory=list) + unresolved_card_ids: list[int] = Field(default_factory=list) # capped honesty sample + table_only_card_ids: list[int] = Field(default_factory=list) class MetabaseRelation(BaseModel): @@ -514,23 +514,23 @@ class MetabaseCard(BaseModel): name: str query_kind: Literal["mbql", "native"] precision: Literal["column", "table", "none"] - collection_id: Optional[int] = None + collection_id: int | None = None archived: bool = False - database_id: Optional[int] = None - columns: List[MetabaseColumnRef] = Field(default_factory=list) - table_relations: List[str] = Field(default_factory=list) # relation keys, table grain - upstream_card_ids: List[int] = Field(default_factory=list) # {{#id}} / card__ deps - snippet_ids: List[int] = Field(default_factory=list) - unresolved_reason: Optional[str] = None # select_star | parse_failed | unknown_table | ... + database_id: int | None = None + columns: list[MetabaseColumnRef] = Field(default_factory=list) + table_relations: list[str] = Field(default_factory=list) # relation keys, table grain + upstream_card_ids: list[int] = Field(default_factory=list) # {{#id}} / card__ deps + snippet_ids: list[int] = Field(default_factory=list) + unresolved_reason: str | None = None # select_star | parse_failed | unknown_table | ... # snapshot-time last-modified stamp (ISO-8601 UTC), used for incremental reuse - updated_at: Optional[str] = None + updated_at: str | None = None # Asset metadata — a deep link to the card and human/ownership context, so a reached card # in impact/explorer output is clickable and attributable. All best-effort (None if absent). - url: Optional[str] = None # {base}/question/{card_id} - description: Optional[str] = None - collection_name: Optional[str] = None - creator: Optional[str] = None # email (else display name) of the card's creator - last_edited_by: Optional[str] = None # email/name from Metabase's last-edit-info + url: str | None = None # {base}/question/{card_id} + description: str | None = None + collection_name: str | None = None + creator: str | None = None # email (else display name) of the card's creator + last_edited_by: str | None = None # email/name from Metabase's last-edit-info class MetabaseDashboard(BaseModel): @@ -546,17 +546,17 @@ class MetabaseDashboard(BaseModel): dashboard_id: int name: str - collection_id: Optional[int] = None - url: Optional[str] = None - card_ids: List[int] = Field(default_factory=list) - meta: Dict[str, Any] = Field(default_factory=dict) # tier/owner/... for policy reach.where + collection_id: int | None = None + url: str | None = None + card_ids: list[int] = Field(default_factory=list) + meta: dict[str, Any] = Field(default_factory=dict) # tier/owner/... for policy reach.where # snapshot-time last-modified stamp (ISO-8601 UTC), used for incremental reuse - updated_at: Optional[str] = None + updated_at: str | None = None # Asset metadata (best-effort; None if absent). ``url`` is already populated above. - description: Optional[str] = None - collection_name: Optional[str] = None - creator: Optional[str] = None # email (else display name) of the dashboard's creator - last_edited_by: Optional[str] = None + description: str | None = None + collection_name: str | None = None + creator: str | None = None # email (else display name) of the dashboard's creator + last_edited_by: str | None = None class MetabaseLineage(BaseModel): @@ -573,9 +573,9 @@ class MetabaseLineage(BaseModel): schema_version: int = 2 provenance: MetabaseProvenance coverage: MetabaseCoverage - relations: Dict[str, MetabaseRelation] = Field(default_factory=dict) - cards: List[MetabaseCard] = Field(default_factory=list) - dashboards: List[MetabaseDashboard] = Field(default_factory=list) + relations: dict[str, MetabaseRelation] = Field(default_factory=dict) + cards: list[MetabaseCard] = Field(default_factory=list) + dashboards: list[MetabaseDashboard] = Field(default_factory=list) # =========================================================================== @@ -692,7 +692,7 @@ class ChangeCondition(BaseModel): field: Literal["kind", "semantic", "breaking", "model", "column"] op: Operator - value: Optional[Any] = None + value: Any | None = None class MetaCondition(BaseModel): @@ -700,7 +700,7 @@ class MetaCondition(BaseModel): key: str op: Operator - value: Optional[Any] = None + value: Any | None = None class ReachCondition(BaseModel): @@ -712,7 +712,7 @@ class ReachCondition(BaseModel): """ kind: ReachKind - mechanism: Optional[List[Mechanism]] = None + mechanism: list[Mechanism] | None = None where: "Predicate" # ge=1: min_count 0 is a vacuously-true reach (matches with zero reached objects), which is # never a meaningful gate — reject it at config load rather than silently always-firing. @@ -736,23 +736,23 @@ class Predicate(BaseModel): ``config`` / ``reach`` / ``structural``. The one-of invariant is enforced by a model validator. """ - all_: Optional[List["Predicate"]] = Field(default=None, alias="all") - any_: Optional[List["Predicate"]] = Field(default=None, alias="any") + all_: list["Predicate"] | None = Field(default=None, alias="all") + any_: list["Predicate"] | None = Field(default=None, alias="any") not_: Optional["Predicate"] = Field(default=None, alias="not") - change: Optional[ChangeCondition] = None - meta: Optional[MetaCondition] = None + change: ChangeCondition | None = None + meta: MetaCondition | None = None # ``inferred_meta`` shares ``MetaCondition``'s shape (key/op/value) but resolves the key's # value by folding UPSTREAM lineage rather than reading only the node's own declared meta — # see ``inferred_meta`` in ``lineage/policy.py``. An unresolvable value is UNKNOWN (fail-safe). - inferred_meta: Optional[MetaCondition] = None + inferred_meta: MetaCondition | None = None # ``config`` shares ``MetaCondition``'s shape (key/op/value) but resolves the DOTTED key # against the subject model's resolved dbt ``node.config`` (``grants.select``, # ``materialized``, ``tags`` …) rather than user ``meta``. Model-grained. A missing dotted # path resolves to the EMPTY SET for set operators (present, not unknown) and to # UNKNOWN_MISSING for scalar operators — see ``config`` in ``lineage/policy.py``. - config: Optional[MetaCondition] = None - reach: Optional[ReachCondition] = None - structural: Optional[StructuralCondition] = None + config: MetaCondition | None = None + reach: ReachCondition | None = None + structural: StructuralCondition | None = None model_config = ConfigDict(populate_by_name=True) @@ -790,10 +790,10 @@ class Action(BaseModel): type: ActionKind include: Literal["reached", "subject", "both"] = "reached" - mechanism: Optional[List[Mechanism]] = None - channel: Optional[str] = None - target: Optional[str] = None - message: Optional[str] = None + mechanism: list[Mechanism] | None = None + channel: str | None = None + target: str | None = None + message: str | None = None # --- rule + policy --------------------------------------------------------- @@ -803,15 +803,15 @@ class Rule(BaseModel): """A single ``predicate -> actions`` rule.""" id: str - description: Optional[str] = None + description: str | None = None scope: Literal["change", "aggregate"] = "change" predicate: Predicate - action: List[Action] + action: list[Action] # Two independent fail-safe knobs: on_missing_meta governs an undecidable leaf # caused by a missing meta key / unresolved reach; on_error governs an operator/type # mismatch (a genuine evaluation error). Each falls back to the matching PolicyDefaults knob. - on_missing_meta: Optional[MissingMetaPolicy] = None - on_error: Optional[MissingMetaPolicy] = None + on_missing_meta: MissingMetaPolicy | None = None + on_error: MissingMetaPolicy | None = None @model_validator(mode="before") @classmethod @@ -836,8 +836,8 @@ class PolicyDefaults(BaseModel): on_missing_meta: MissingMetaPolicy = MissingMetaPolicy.FAIL_CLOSED on_error: MissingMetaPolicy = MissingMetaPolicy.FAIL_CLOSED - on_meaning_changed: Optional[GateDecision] = None - on_indeterminate: Optional[GateDecision] = None + on_meaning_changed: GateDecision | None = None + on_indeterminate: GateDecision | None = None class Policy(BaseModel): @@ -845,7 +845,7 @@ class Policy(BaseModel): version: int defaults: PolicyDefaults = Field(default_factory=PolicyDefaults) - rules: List[Rule] = Field(default_factory=list) + rules: list[Rule] = Field(default_factory=list) # --- outputs --------------------------------------------------------------- @@ -864,16 +864,16 @@ class RuleHit(BaseModel): rule_id: str decision: GateDecision - change_model: Optional[str] = None - change_column: Optional[str] = None - matched_reach: List[str] = Field(default_factory=list) - actions: List[ActionKind] = Field(default_factory=list) + change_model: str | None = None + change_column: str | None = None + matched_reach: list[str] = Field(default_factory=list) + actions: list[ActionKind] = Field(default_factory=list) # override cap (backward-compatible): when an override pragma capped this hit, # ``decision`` holds the effective/capped value (so ``_combine_decision`` needs no change) # while ``original_decision`` keeps the pre-cap value so the report can show the delta. overridden: bool = False - original_decision: Optional[GateDecision] = None - override_reason: Optional[str] = None + original_decision: GateDecision | None = None + override_reason: str | None = None # the backtest trust signal (backward-compatible, default False keeps serialization additive): # True when this hit fired via a *fail-safe UNKNOWN* resolution (a blocking rule under # ``fail_closed`` firing on an undecidable predicate) rather than a proven TRUE match. This is @@ -882,17 +882,17 @@ class RuleHit(BaseModel): # proven) nor perturbed by override caps. ``unknown_cause`` splits missing-vs-error for the # drill-down; the bool is the primary signal. fired_on_unknown: bool = False - unknown_cause: Optional[Literal["missing", "error"]] = None + unknown_cause: Literal["missing", "error"] | None = None class PolicyVerdict(BaseModel): """The engine's output: a gate decision + accumulated build/test sets + notifications.""" decision: GateDecision - hits: List[RuleHit] = Field(default_factory=list) - build_set: List[str] = Field(default_factory=list) - test_set: List[str] = Field(default_factory=list) - notifications: List[Notification] = Field(default_factory=list) + hits: list[RuleHit] = Field(default_factory=list) + build_set: list[str] = Field(default_factory=list) + test_set: list[str] = Field(default_factory=list) + notifications: list[Notification] = Field(default_factory=list) evaluated_rules: int = 0 fired_rules: int = 0 # Honesty counters for fail-safe explainability (additive; see policy.py §7). @@ -928,8 +928,8 @@ class MetabaseReachConfidence(BaseModel): warning"; it never fabricates dashboard reach. """ - snapshot_generated_at: Optional[str] = None - snapshot_age_hours: Optional[float] = None + snapshot_generated_at: str | None = None + snapshot_age_hours: float | None = None stale: bool = False # age > threshold OR artifact missing dashboards_reached: int = 0 cards_column_precise: int = 0 # reached via a column-precise card @@ -986,11 +986,11 @@ class BacktestPointResult(BaseModel): source: str # human label for the changeset source (e.g. "git-diff (^..)") total_changes: int unmapped_changes: int = 0 # .sql paths that mapped to no model in the HEAD registry - parse_failures: List[str] = Field(default_factory=list) + parse_failures: list[str] = Field(default_factory=list) decision: str # allow / warn / block blast_radius: int = 0 - fired: List[BacktestFiredHit] = Field(default_factory=list) - sample_reach: List[str] = Field(default_factory=list) + fired: list[BacktestFiredHit] = Field(default_factory=list) + sample_reach: list[str] = Field(default_factory=list) class BacktestReport(BaseModel): @@ -1003,18 +1003,18 @@ class BacktestReport(BaseModel): mode: Literal["git-diff", "changesets"] policy_source: str - base: Optional[str] = None - head: Optional[str] = None + base: str | None = None + head: str | None = None prs_replayed: int = 0 prs_would_block: int = 0 prs_would_warn: int = 0 prs_skipped: int = 0 # points whose impact could not be computed (parse failures / errors) avg_blast_radius: float = 0.0 - rule_stats: List[BacktestRuleStat] = Field(default_factory=list) - points: List[BacktestPointResult] = Field(default_factory=list) + rule_stats: list[BacktestRuleStat] = Field(default_factory=list) + points: list[BacktestPointResult] = Field(default_factory=list) fidelity_note: str = "" - warnings: List[str] = Field(default_factory=list) - baseline_delta: Optional[Dict[str, Any]] = None + warnings: list[str] = Field(default_factory=list) + baseline_delta: dict[str, Any] | None = None # =========================================================================== @@ -1071,8 +1071,8 @@ class PolicyInitScan(BaseModel): # config-axis (PII over-grant) template, so it is only offered when dbt grants actually exist. models_with_grants: int = 0 # Sorted by n_present desc, then key (stable, most-covered-first). - model_meta_keys: List[MetaKeyCoverage] = Field(default_factory=list) - column_meta_keys: List[MetaKeyCoverage] = Field(default_factory=list) + model_meta_keys: list[MetaKeyCoverage] = Field(default_factory=list) + column_meta_keys: list[MetaKeyCoverage] = Field(default_factory=list) @property def tests_present(self) -> bool: diff --git a/parrant/parser/sql_parser.py b/parrant/parser/sql_parser.py index 0d9c703..cbb450f 100644 --- a/parrant/parser/sql_parser.py +++ b/parrant/parser/sql_parser.py @@ -1,8 +1,11 @@ -import re import logging +import re +from collections.abc import Callable from dataclasses import dataclass, field -from sqlglot import parse_one, exp -from typing import Dict, List, Set, Optional, Any, Callable, Literal, Tuple, cast +from typing import Any, Literal, cast + +from sqlglot import exp, parse_one + from parrant.models.schema import ( ColumnLineage, OverrideDirective, @@ -11,12 +14,12 @@ UnresolvedColumnEdge, ) from parrant.parser.sql_parser_utils import ( - get_table_aliases, - get_lateral_flatten_aliases, - get_flatten_alias_nodes, - get_table_context, get_all_tables_from_select, get_final_selects, + get_flatten_alias_nodes, + get_lateral_flatten_aliases, + get_table_aliases, + get_table_context, split_qualified_name, strip_sql_comments, ) @@ -47,7 +50,7 @@ _TRAILING_IDENT_RE = re.compile(r'(?:"(?P[^"]+)"|(?P[A-Za-z_]\w*))\s*$') -def _extract_select_alias(line: str) -> Optional[str]: +def _extract_select_alias(line: str) -> str | None: """Best-effort pull of the projected column name from a SELECT-list line. Prefers the token after a case-insensitive `` as `` (the explicit alias); else the last @@ -70,7 +73,7 @@ def _extract_select_alias(line: str) -> Optional[str]: return None -def _adjacent_column(lines: List[str], pragma_idx: int) -> Optional[str]: +def _adjacent_column(lines: list[str], pragma_idx: int) -> str | None: """Scan forward from the pragma line to the next non-blank, non-comment line and extract its projected column via :func:`_extract_select_alias`. Best-effort (``None`` => stale).""" for j in range(pragma_idx + 1, len(lines)): @@ -81,7 +84,7 @@ def _adjacent_column(lines: List[str], pragma_idx: int) -> Optional[str]: return None -def parse_override_directives(sql: str) -> Tuple[List[OverrideDirective], List[str]]: +def parse_override_directives(sql: str) -> tuple[list[OverrideDirective], list[str]]: """Parse ``-- lineage:allow-(change|break) ...`` pragmas from raw head SQL. Returns ``(directives, warnings)``. A pragma is DROPPED (contributing a human warning @@ -94,14 +97,14 @@ def parse_override_directives(sql: str) -> Tuple[List[OverrideDirective], List[s (``column`` may be ``None`` when adjacency can't resolve => the caller marks it stale). """ lines = sql.splitlines() - first_select_idx: Optional[int] = None + first_select_idx: int | None = None for i, line in enumerate(lines): if _SELECT_WORD_RE.search(line): first_select_idx = i break - directives: List[OverrideDirective] = [] - warnings: List[str] = [] + directives: list[OverrideDirective] = [] + warnings: list[str] = [] for i, line in enumerate(lines): m = _OVERRIDE_LINE_RE.search(line) if not m: @@ -127,7 +130,7 @@ def parse_override_directives(sql: str) -> Tuple[List[OverrideDirective], List[s # Look for column= OUTSIDE the quoted reason span (so reason="see column=x" is safe). args_wo_reason = args[: reason_match.start()] + args[reason_match.end() :] column_match = _COLUMN_ARG_RE.search(args_wo_reason) - column: Optional[str] + column: str | None scope: Literal["column", "model"] if column_match: column = column_match.group("val").lower() @@ -154,22 +157,22 @@ def parse_override_directives(sql: str) -> Tuple[List[OverrideDirective], List[s class ParserContext: """Context object containing parser state and dependencies.""" - aliases: Dict[str, str] + aliases: dict[str, str] table_context: str - cte_sources: Dict[str, Dict[str, str]] - cte_to_model: Optional[Dict[str, str]] - cte_transformation_types: Dict[str, Dict[str, str]] = field(default_factory=dict) - cte_sql_expressions: Dict[str, Dict[str, Optional[str]]] = field(default_factory=dict) - cte_base_tables: Dict[str, Set[str]] = field(default_factory=dict) + cte_sources: dict[str, dict[str, str]] + cte_to_model: dict[str, str] | None + cte_transformation_types: dict[str, dict[str, str]] = field(default_factory=dict) + cte_sql_expressions: dict[str, dict[str, str | None]] = field(default_factory=dict) + cte_base_tables: dict[str, set[str]] = field(default_factory=dict) # Additional per-column sources contributed by non-left UNION branches of a CTE. # cte_sources holds a single primary source per column; these are merged in on top # so a CTE built from a UNION is not reduced to only its left-most branch. - cte_extra_sources: Dict[str, Dict[str, Set[str]]] = field(default_factory=dict) - column_definitions: Optional[Dict[str, Any]] = None + cte_extra_sources: dict[str, dict[str, set[str]]] = field(default_factory=dict) + column_definitions: dict[str, Any] | None = None class CTEHandler: - def extract_cte_model_mappings_from_parsed(self, parsed: Any) -> Dict[str, str]: + def extract_cte_model_mappings_from_parsed(self, parsed: Any) -> dict[str, str]: mappings = {} for cte in parsed.find_all(exp.CTE): cte_name = cte.alias @@ -183,9 +186,9 @@ def extract_cte_model_mappings_from_parsed(self, parsed: Any) -> Dict[str, str]: def trace_base_tables( self, table: str, - cte_to_model: Optional[Dict[str, str]], - cte_sources: Dict[str, Dict[str, str]], - star_sources: Set[str], + cte_to_model: dict[str, str] | None, + cte_sources: dict[str, dict[str, str]], + star_sources: set[str], ) -> None: if cte_to_model is None: if table not in cte_sources: @@ -209,21 +212,21 @@ def trace_base_tables( class StarExpressionHandler: def __init__(self) -> None: - self._cte_handler: Optional[CTEHandler] = None + self._cte_handler: CTEHandler | None = None def is_star_expression(self, expr: Any) -> bool: return isinstance(expr, exp.Star) or ( isinstance(expr, exp.Column) and getattr(expr, "is_star", False) ) - def get_star_source_table(self, expr: Any, aliases: Dict[str, str], table_context: str) -> str: + def get_star_source_table(self, expr: Any, aliases: dict[str, str], table_context: str) -> str: if isinstance(expr, exp.Column) and expr.table: star_table_alias = str(expr.table) return aliases.get(star_table_alias, star_table_alias) else: return table_context - def get_excluded_columns(self, star_expr: exp.Star) -> List[str]: + def get_excluded_columns(self, star_expr: exp.Star) -> list[str]: excluded = [] if hasattr(star_expr, "args") and "except" in star_expr.args: except_clause = star_expr.args["except"] @@ -240,7 +243,7 @@ def get_excluded_columns(self, star_expr: exp.Star) -> List[str]: def get_cte_transformation_info( self, context: ParserContext, cte_name: str, col_name: str - ) -> tuple[str, Optional[str]]: + ) -> tuple[str, str | None]: trans_type = context.cte_transformation_types.get(cte_name, {}).get(col_name, "direct") sql_expr = context.cte_sql_expressions.get(cte_name, {}).get(col_name) return trans_type, sql_expr @@ -248,11 +251,11 @@ def get_cte_transformation_info( def expand_from_join_tables( self, select: Any, - all_tables: List[str], - excluded_col_names: Set[str], + all_tables: list[str], + excluded_col_names: set[str], context: ParserContext, - columns: Dict[str, List[ColumnLineage]], - star_sources: Set[str], + columns: dict[str, list[ColumnLineage]], + star_sources: set[str], ) -> None: for join_table in all_tables: if join_table in context.cte_sources: @@ -283,10 +286,10 @@ def expand_from_join_tables( def expand_from_cte( self, source_table: str, - excluded_col_names: Set[str], + excluded_col_names: set[str], context: ParserContext, - columns: Dict[str, List[ColumnLineage]], - star_sources: Set[str], + columns: dict[str, list[ColumnLineage]], + star_sources: set[str], ) -> bool: if source_table in context.cte_sources: if len(context.cte_sources[source_table]) > 0: @@ -319,7 +322,7 @@ def expand_from_cte( class ExpressionAnalyzer: def __init__(self, parser: "SQLColumnParser") -> None: self.parser = parser - self._handlers: Dict[type, Callable[[Any, ParserContext, bool], List[ColumnLineage]]] = {} + self._handlers: dict[type, Callable[[Any, ParserContext, bool], list[ColumnLineage]]] = {} self._register_default_handlers() def _register_default_handlers(self) -> None: @@ -327,13 +330,13 @@ def _register_default_handlers(self) -> None: self.register_handler(exp.Column, self._handle_column) def register_handler( - self, expr_type: type, handler: Callable[[Any, ParserContext, bool], List[ColumnLineage]] + self, expr_type: type, handler: Callable[[Any, ParserContext, bool], list[ColumnLineage]] ) -> None: self._handlers[expr_type] = handler def analyze( self, expr: Any, context: ParserContext, is_aliased: bool = False - ) -> List[ColumnLineage]: + ) -> list[ColumnLineage]: expr_type = type(expr) if expr_type in self._handlers: return self._handlers[expr_type](expr, context, is_aliased) @@ -341,12 +344,12 @@ def analyze( def _handle_alias( self, expr: exp.Alias, context: ParserContext, is_aliased: bool - ) -> List[ColumnLineage]: + ) -> list[ColumnLineage]: return self.analyze(expr.this, context, is_aliased=True) def _handle_column( self, expr: exp.Column, context: ParserContext, is_aliased: bool - ) -> List[ColumnLineage]: + ) -> list[ColumnLineage]: col_name = ( str(expr.this).lower() if hasattr(expr, "this") and expr.this else str(expr).lower() ) @@ -359,7 +362,7 @@ def _handle_column( return self.parser._analyze_column_reference(expr, col_name, context, is_aliased) - def _default_handler(self, expr: Any, context: ParserContext) -> List[ColumnLineage]: + def _default_handler(self, expr: Any, context: ParserContext) -> list[ColumnLineage]: source_cols = self.parser._extract_source_columns(expr, context) normalized_source_cols = self.parser._normalize_source_columns(source_cols) return [ @@ -372,7 +375,7 @@ def _default_handler(self, expr: Any, context: ParserContext) -> List[ColumnLine class SQLColumnParser: - def __init__(self, dialect: Optional[str] = None): + def __init__(self, dialect: str | None = None): self.dialect = dialect self._cte_handler = CTEHandler() self._star_handler = StarExpressionHandler() @@ -383,10 +386,10 @@ def parse_column_lineage(self, sql: str) -> SQLParseResult: parsed = parse_one(sql, dialect=self.dialect) cte_to_model = self._cte_handler.extract_cte_model_mappings_from_parsed(parsed) - cte_transformation_types: Dict[str, Dict[str, str]] = {} - cte_sql_expressions: Dict[str, Dict[str, Optional[str]]] = {} - cte_base_tables: Dict[str, Set[str]] = {} - cte_extra_sources: Dict[str, Dict[str, Set[str]]] = {} + cte_transformation_types: dict[str, dict[str, str]] = {} + cte_sql_expressions: dict[str, dict[str, str | None]] = {} + cte_base_tables: dict[str, set[str]] = {} + cte_extra_sources: dict[str, dict[str, set[str]]] = {} aliases = get_table_aliases(parsed) for cte in parsed.find_all(exp.CTE): @@ -401,11 +404,11 @@ def parse_column_lineage(self, sql: str) -> SQLParseResult: cte_extra_sources, ) - columns: Dict[str, List[ColumnLineage]] = {} - star_sources: Set[str] = set() + columns: dict[str, list[ColumnLineage]] = {} + star_sources: set[str] = set() # Unresolved-edge markers collected during this parse (see UnresolvedColumnEdge). # ``model`` is left empty; the registry stamps the real node name. - markers: List[UnresolvedColumnEdge] = [] + markers: list[UnresolvedColumnEdge] = [] flatten_aliases = get_lateral_flatten_aliases(parsed) # flatten pseudo-alias -> the REAL upstream source columns of the expression it unnests # (e.g. `flatten(payload:items) f` -> {`raw_events.payload`}). Lets a downstream @@ -425,7 +428,7 @@ def parse_column_lineage(self, sql: str) -> SQLParseResult: final_selects = get_final_selects(parsed) if not final_selects: - selects_to_process: List[Any] = list(parsed.find_all(exp.Select)) + selects_to_process: list[Any] = list(parsed.find_all(exp.Select)) else: selects_to_process = list(final_selects) # `select * from `: expand the CTE's own SELECT(s). Using @@ -559,10 +562,10 @@ def parse_column_lineage(self, sql: str) -> SQLParseResult: ) @staticmethod - def _dedupe_markers(markers: List[UnresolvedColumnEdge]) -> List[UnresolvedColumnEdge]: + def _dedupe_markers(markers: list[UnresolvedColumnEdge]) -> list[UnresolvedColumnEdge]: """Order-stable de-duplication of markers (same column/reason/detail collapse to one).""" - seen: Set[Tuple[str, str, Optional[str]]] = set() - unique: List[UnresolvedColumnEdge] = [] + seen: set[tuple[str, str, str | None]] = set() + unique: list[UnresolvedColumnEdge] = [] for marker in markers: key = (marker.column, marker.reason, marker.detail) if key not in seen: @@ -571,7 +574,7 @@ def _dedupe_markers(markers: List[UnresolvedColumnEdge]) -> List[UnresolvedColum return unique def _emit_star_rename_markers( - self, star_expr: exp.Star, markers: List[UnresolvedColumnEdge] + self, star_expr: exp.Star, markers: list[UnresolvedColumnEdge] ) -> None: """Declare each ``select * rename (old as new)`` output as an unresolved edge. @@ -601,10 +604,10 @@ def _emit_star_rename_markers( def _declare_phantom_edges( self, - columns: Dict[str, List[ColumnLineage]], - flatten_aliases: Set[str], - flatten_alias_sources: Dict[str, Set[str]], - markers: List[UnresolvedColumnEdge], + columns: dict[str, list[ColumnLineage]], + flatten_aliases: set[str], + flatten_alias_sources: dict[str, set[str]], + markers: list[UnresolvedColumnEdge], ) -> None: """Resolve or declare fabricated source tokens on every resolved column. @@ -627,7 +630,7 @@ def _declare_phantom_edges( """ for out_col, lineages in columns.items(): for lineage in lineages: - kept: Set[str] = set() + kept: set[str] = set() for token in lineage.source_columns: reason = self._phantom_token_reason(token, flatten_aliases) if reason is None: @@ -646,7 +649,7 @@ def _declare_phantom_edges( lineage.source_columns = kept @staticmethod - def _resolve_flatten_token(token: str, flatten_alias_sources: Dict[str, Set[str]]) -> Set[str]: + def _resolve_flatten_token(token: str, flatten_alias_sources: dict[str, set[str]]) -> set[str]: """Return the real upstream columns a flatten-qualified token (``f.value``) derives from. ``f`` is looked up in the flatten-alias -> flattened-source map. An empty result means the @@ -658,13 +661,13 @@ def _resolve_flatten_token(token: str, flatten_alias_sources: Dict[str, Set[str] def _build_flatten_alias_sources( self, parsed: Any, - cte_to_model: Optional[Dict[str, str]], - cte_sources: Dict[str, Dict[str, str]], - cte_transformation_types: Dict[str, Dict[str, str]], - cte_sql_expressions: Dict[str, Dict[str, Optional[str]]], - cte_base_tables: Dict[str, Set[str]], - cte_extra_sources: Dict[str, Dict[str, Set[str]]], - ) -> Dict[str, Set[str]]: + cte_to_model: dict[str, str] | None, + cte_sources: dict[str, dict[str, str]], + cte_transformation_types: dict[str, dict[str, str]], + cte_sql_expressions: dict[str, dict[str, str | None]], + cte_base_tables: dict[str, set[str]], + cte_extra_sources: dict[str, dict[str, set[str]]], + ) -> dict[str, set[str]]: """Map each flatten pseudo-alias to the real upstream columns of the expression it unnests. For every ``flatten() alias`` in the query, resolve ````'s columns through the @@ -678,7 +681,7 @@ def _build_flatten_alias_sources( alias whose expression traces to nothing real (a literal array, an untraceable path) maps to the empty set, which keeps the honest ``phantom_alias`` marker downstream. """ - raw_sources: Dict[str, Set[str]] = {} + raw_sources: dict[str, set[str]] = {} for alias, flattened_expr, select in get_flatten_alias_nodes(parsed): if select is None: raw_sources.setdefault(alias, set()) @@ -688,7 +691,7 @@ def _build_flatten_alias_sources( # Populate column_definitions so forward-reference resolution follows such a derived # column to its REAL upstream source instead of fabricating a same-name column on the # base relation. - column_definitions: Dict[str, Any] = {} + column_definitions: dict[str, Any] = {} for projected in select.expressions: projected_name = strip_sql_comments(projected.alias_or_name).lower() column_definitions[projected_name] = projected @@ -708,7 +711,7 @@ def _build_flatten_alias_sources( # honest here — both are genuine flattened sources; dropping one would hide an edge). raw_sources.setdefault(alias, set()).update(resolved) - expanded: Dict[str, Set[str]] = {} + expanded: dict[str, set[str]] = {} for alias in raw_sources: expanded[alias] = self._expand_flatten_sources(alias, raw_sources, set()) return expanded @@ -716,14 +719,14 @@ def _build_flatten_alias_sources( def _expand_flatten_sources( self, alias: str, - raw_sources: Dict[str, Set[str]], - visiting: Set[str], - ) -> Set[str]: + raw_sources: dict[str, set[str]], + visiting: set[str], + ) -> set[str]: """Transitively resolve a flatten alias's sources, replacing nested-flatten qualifiers.""" if alias in visiting: return set() visiting.add(alias) - out: Set[str] = set() + out: set[str] = set() for token in raw_sources.get(alias, set()): table_part, _ = split_qualified_name(token) qualifier = table_part.strip().strip('"').lower() if table_part else "" @@ -736,8 +739,8 @@ def _expand_flatten_sources( @staticmethod def _phantom_token_reason( - token: str, flatten_aliases: Set[str] - ) -> Optional[Literal["pivot_output", "phantom_alias"]]: + token: str, flatten_aliases: set[str] + ) -> Literal["pivot_output", "phantom_alias"] | None: """Classify a source token as a fabricated edge, or ``None`` if it is genuine. Returns ``"pivot_output"`` for a quoted pivot literal, ``"phantom_alias"`` for a @@ -756,12 +759,12 @@ def _phantom_token_reason( def _extract_predicate_lineage( self, parsed: Any, - cte_to_model: Optional[Dict[str, str]], - cte_sources: Dict[str, Dict[str, str]], - cte_transformation_types: Dict[str, Dict[str, str]], - cte_sql_expressions: Dict[str, Dict[str, Optional[str]]], - cte_base_tables: Dict[str, Set[str]], - ) -> Dict[str, str]: + cte_to_model: dict[str, str] | None, + cte_sources: dict[str, dict[str, str]], + cte_transformation_types: dict[str, dict[str, str]], + cte_sql_expressions: dict[str, dict[str, str | None]], + cte_base_tables: dict[str, set[str]], + ) -> dict[str, str]: """Resolve upstream columns referenced only in predicate clauses, with the condition. Column-value lineage is built from the projected ``SELECT`` list, so a column a @@ -773,7 +776,7 @@ def _extract_predicate_lineage( a predicate on a CTE that wraps an upstream model resolves to that model's column. The returned map is ``upstream_column -> predicate condition text`` (the "why"). """ - conditions_by_source: Dict[str, Set[str]] = {} + conditions_by_source: dict[str, set[str]] = {} for select in parsed.find_all(exp.Select): context = ParserContext( @@ -787,7 +790,7 @@ def _extract_predicate_lineage( column_definitions={}, ) - conditions: List[Any] = [] + conditions: list[Any] = [] for key in ("where", "having", "qualify"): wrapper = select.args.get(key) if wrapper is not None: @@ -822,7 +825,7 @@ def _extract_predicate_lineage( for source, conditions in conditions_by_source.items() } - def _extract_cte_model_mappings(self, sql: str) -> Dict[str, str]: + def _extract_cte_model_mappings(self, sql: str) -> dict[str, str]: """Extract mappings from CTE names to model names (legacy method using regex).""" mappings = {} # Pattern to handle: @@ -839,7 +842,7 @@ def _extract_cte_model_mappings(self, sql: str) -> Dict[str, str]: return mappings - def _normalize_table_ref(self, column: str, aliases: Dict[str, str], table_context: str) -> str: + def _normalize_table_ref(self, column: str, aliases: dict[str, str], table_context: str) -> str: column = strip_sql_comments(column) table_part, col = split_qualified_name(column) if not table_part: @@ -850,13 +853,13 @@ def _normalize_table_ref(self, column: str, aliases: Dict[str, str], table_conte def _build_cte_sources( self, parsed: Any, - cte_to_model: Optional[Dict[str, str]], - cte_transformation_types: Dict[str, Dict[str, str]], - cte_sql_expressions: Dict[str, Dict[str, Optional[str]]], - cte_base_tables: Dict[str, Set[str]], - cte_extra_sources: Dict[str, Dict[str, Set[str]]], - ) -> Dict[str, Dict[str, str]]: - cte_sources: Dict[str, Dict[str, str]] = {} + cte_to_model: dict[str, str] | None, + cte_transformation_types: dict[str, dict[str, str]], + cte_sql_expressions: dict[str, dict[str, str | None]], + cte_base_tables: dict[str, set[str]], + cte_extra_sources: dict[str, dict[str, set[str]]], + ) -> dict[str, dict[str, str]]: + cte_sources: dict[str, dict[str, str]] = {} for cte in parsed.find_all(exp.CTE): cte_name = cte.alias @@ -942,7 +945,7 @@ def _resolve_star_from_table_in_cte( self, expr: Any, select: Any, - aliases: Dict[str, str], + aliases: dict[str, str], table_context: str, ) -> str: if isinstance(expr, exp.Column) and expr.table: @@ -975,7 +978,7 @@ def _copy_cte_columns_with_exclusions( self, from_table: str, cte_name: str, - excluded_col_names: Set[str], + excluded_col_names: set[str], context: ParserContext, ) -> None: if from_table in context.cte_sources: @@ -1019,8 +1022,8 @@ def _resolve_column_source( self, column: str, table: str, - cte_sources: Dict[str, Dict[str, str]], - cte_to_model: Optional[Dict[str, str]] = None, + cte_sources: dict[str, dict[str, str]], + cte_to_model: dict[str, str] | None = None, ) -> str: column = strip_sql_comments(column) table_part, col_name = split_qualified_name(column) @@ -1045,7 +1048,7 @@ def _resolve_column_source( return f"{table}.{col_name_lower}" return column - def _resolve_base_table(self, table: str, cte_to_model: Dict[str, str]) -> str: + def _resolve_base_table(self, table: str, cte_to_model: dict[str, str]) -> str: """Follow cte_to_model transitively until reaching a table that is not a CTE. A single cte_to_model lookup can land on another CTE alias (e.g. a chain of @@ -1055,7 +1058,7 @@ def _resolve_base_table(self, table: str, cte_to_model: Dict[str, str]) -> str: infinite loops on recursive/self-referential mappings. """ current = table - visited: Set[str] = set() + visited: set[str] = set() while current in cte_to_model and current not in visited: visited.add(current) next_table = cte_to_model[current] @@ -1069,7 +1072,7 @@ def _handle_forward_reference( expr: exp.Column, col_name: str, context: ParserContext, - ) -> Optional[List[ColumnLineage]]: + ) -> list[ColumnLineage] | None: is_qualified = bool(expr.table) if ( not is_qualified @@ -1098,11 +1101,11 @@ def _analyze_column_reference( col_name: str, context: ParserContext, is_aliased: bool, - ) -> List[ColumnLineage]: + ) -> list[ColumnLineage]: source_col = self._normalize_table_ref( strip_sql_comments(str(expr)), context.aliases, context.table_context ) - table_part, col = split_qualified_name(source_col) + table_part, _col = split_qualified_name(source_col) table = table_part if table_part else context.table_context resolved_source = self._resolve_column_source( source_col, table, context.cte_sources, context.cte_to_model @@ -1138,9 +1141,9 @@ def _analyze_column_reference( ] def _normalize_extra_cte_sources( - self, extras_for_table: Dict[str, Set[str]], col_name: str - ) -> Set[str]: - normalized: Set[str] = set() + self, extras_for_table: dict[str, set[str]], col_name: str + ) -> set[str]: + normalized: set[str] = set() for extra in extras_for_table.get(col_name, set()): extra_table, extra_col = split_qualified_name(extra) if extra_table: @@ -1149,7 +1152,7 @@ def _normalize_extra_cte_sources( normalized.add(extra_col.lower()) return normalized - def _normalize_source_columns(self, source_cols: Set[str]) -> Set[str]: + def _normalize_source_columns(self, source_cols: set[str]) -> set[str]: """Normalize source columns, ensuring all are cleaned of comments and lowercase.""" normalized = set() for s in source_cols: @@ -1166,8 +1169,8 @@ def _handle_forward_reference_in_extraction( col: exp.Column, col_name: str, context: ParserContext, - visited_forward_refs: Set[str], - ) -> Optional[Set[str]]: + visited_forward_refs: set[str], + ) -> set[str] | None: is_qualified = bool(col.table) if ( not is_qualified @@ -1191,8 +1194,8 @@ def _extract_source_columns( self, expr: Any, context: ParserContext, - visited_forward_refs: Optional[Set[str]] = None, - ) -> Set[str]: + visited_forward_refs: set[str] | None = None, + ) -> set[str]: if visited_forward_refs is None: visited_forward_refs = set() diff --git a/parrant/parser/sql_parser_utils.py b/parrant/parser/sql_parser_utils.py index 4c6b966..aae16b5 100644 --- a/parrant/parser/sql_parser_utils.py +++ b/parrant/parser/sql_parser_utils.py @@ -1,6 +1,7 @@ import re +from typing import Any + from sqlglot import exp -from typing import Dict, List, Optional, Any def strip_sql_comments(text: str) -> str: @@ -25,7 +26,7 @@ def strip_sql_comments(text: str) -> str: return text.strip() -def get_table_aliases(parsed: Any) -> Dict[str, str]: +def get_table_aliases(parsed: Any) -> dict[str, str]: aliases = {} for table in parsed.find_all((exp.Table, exp.From, exp.Join)): if table.alias: @@ -54,7 +55,7 @@ def get_lateral_flatten_aliases(parsed: Any) -> set: return aliases -def _enclosing_select(node: Any) -> Optional[Any]: +def _enclosing_select(node: Any) -> Any | None: """Walk up the parent chain to the SELECT that owns this node (``None`` if unattached).""" parent = node.parent while parent is not None and not isinstance(parent, exp.Select): @@ -62,7 +63,7 @@ def _enclosing_select(node: Any) -> Optional[Any]: return parent -def _flatten_input_expression(explode: Any) -> Optional[Any]: +def _flatten_input_expression(explode: Any) -> Any | None: """Return the expression being unnested by a ``flatten`` (an ``exp.Explode``). The flattened value is either passed positionally (``flatten(x)`` -> ``explode.this == x``) @@ -88,7 +89,7 @@ def _flatten_input_expression(explode: Any) -> Optional[Any]: return inner -def get_flatten_alias_nodes(parsed: Any) -> List[tuple]: +def get_flatten_alias_nodes(parsed: Any) -> list[tuple]: """Return ``(alias, flattened_expression, enclosing_select)`` for each ``flatten`` in the query. Covers the two shapes Snowflake ``flatten`` parses into: ``lateral flatten(...) a`` — an @@ -101,7 +102,7 @@ def get_flatten_alias_nodes(parsed: Any) -> List[tuple]: no alias, or whose inner is not an ``Explode`` (e.g. a ``lateral (subquery)``), is skipped — only genuine flatten table-functions are returned. """ - nodes: List[tuple] = [] + nodes: list[tuple] = [] for holder in list(parsed.find_all(exp.Lateral)) + list(parsed.find_all(exp.TableFromRows)): alias = holder.alias if not alias: @@ -131,7 +132,7 @@ def get_table_context(select: Any) -> str: return "" -def get_all_tables_from_select(select: Any) -> List[str]: +def get_all_tables_from_select(select: Any) -> list[str]: tables = [] from_clause = select.find(exp.From) if from_clause: @@ -142,15 +143,13 @@ def get_all_tables_from_select(select: Any) -> List[str]: for join in select.find_all(exp.Join): if hasattr(join, "this"): join_table = join.this - if isinstance(join_table, exp.Table): - tables.append(str(join_table.name).lower()) - elif hasattr(join_table, "name"): + if isinstance(join_table, exp.Table) or hasattr(join_table, "name"): tables.append(str(join_table.name).lower()) return tables -def get_final_select(parsed: Any) -> Optional[Any]: +def get_final_select(parsed: Any) -> Any | None: query = parsed while hasattr(query, "this") and query.this: query = query.this @@ -164,7 +163,7 @@ def get_final_select(parsed: Any) -> Optional[Any]: return None -def get_final_selects(parsed: Any) -> List[Any]: +def get_final_selects(parsed: Any) -> list[Any]: """Return every top-level branch SELECT to process. For a ``UNION`` / ``UNION ALL`` (including chained/nested unions), returns the @@ -183,7 +182,7 @@ def get_final_selects(parsed: Any) -> List[Any]: query = query.this if isinstance(query, exp.Union): - selects: List[Any] = [] + selects: list[Any] = [] for side in (query.this, query.expression): selects.extend(get_final_selects(side)) return selects diff --git a/scripts/run_tests.py b/scripts/run_tests.py index 5791f78..ae1e157 100644 --- a/scripts/run_tests.py +++ b/scripts/run_tests.py @@ -1,11 +1,10 @@ +import os import subprocess import sys -import os from pathlib import Path -from typing import Optional -def run_tests(test_type: Optional[str] = None) -> int: +def run_tests(test_type: str | None = None) -> int: """Run tests with pytest.""" project_root = Path(__file__).parent.parent os.environ["PYTHONPATH"] = os.pathsep.join( diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index ae9563f..b7212c9 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -1,13 +1,14 @@ """Conftest for e2e tests - reuses integration conftest.""" -import pytest -from pathlib import Path -from typing import Dict, Any import sys +from pathlib import Path +from typing import Any + +import pytest @pytest.fixture(scope="session") -def dbt_artifacts() -> Dict[str, Any]: +def dbt_artifacts() -> dict[str, Any]: sys.path.insert(0, str(Path(__file__).parent.parent.parent)) from tests.resources.dbt_test_project.setup import setup_dbt_project diff --git a/tests/e2e/test_lineage_api.py b/tests/e2e/test_lineage_api.py index bb91b74..ebe867a 100644 --- a/tests/e2e/test_lineage_api.py +++ b/tests/e2e/test_lineage_api.py @@ -1,13 +1,14 @@ """End-to-end tests for lineage API.""" -import pytest import subprocess import time -import requests -from pathlib import Path -from typing import Dict, Any, Set, Iterator, TypedDict, cast +from collections.abc import Iterator from contextlib import contextmanager +from pathlib import Path +from typing import Any, TypedDict, cast +import pytest +import requests EXPECTED_CRYPTO_PORTFOLIO_COLUMNS = 37 TEST_MODEL = "int_trade_flow" @@ -88,18 +89,18 @@ def _terminate_process(process: subprocess.Popen[bytes]) -> None: pass -def _get_lineage_response(port: int, model: str, column: str) -> Dict[str, Any]: +def _get_lineage_response(port: int, model: str, column: str) -> dict[str, Any]: """Get lineage response from the API endpoint.""" endpoint = f"http://127.0.0.1:{port}/api/lineage/{model}/{column}" try: response = requests.get(endpoint, timeout=30) response.raise_for_status() - return cast(Dict[str, Any], response.json()) + return cast(dict[str, Any], response.json()) except requests.exceptions.RequestException as e: raise AssertionError(f"Failed to get lineage from {endpoint}: {e}") -def _extract_column_set(data: Dict[str, Any]) -> Set[str]: +def _extract_column_set(data: dict[str, Any]) -> set[str]: """Extract all column identifiers from the API response.""" if not data or "nodes" not in data: return set() @@ -115,7 +116,7 @@ def _extract_column_set(data: Dict[str, Any]) -> Set[str]: return columns -def _count_columns_by_model(data: Dict[str, Any], model_name: str) -> int: +def _count_columns_by_model(data: dict[str, Any], model_name: str) -> int: """Count columns for a specific model in the response.""" if not data or "nodes" not in data: return 0 @@ -143,11 +144,11 @@ class ResultDict(TypedDict): restart: int total_columns: int target_model_columns: int - columns: Set[str] + columns: set[str] def test_lineage_api_determinism_across_restarts( - dbt_artifacts: Dict[str, Any], server_port: int + dbt_artifacts: dict[str, Any], server_port: int ) -> None: """Verify lineage API returns identical results across server restarts.""" catalog_path = Path(dbt_artifacts["catalog_path"]) @@ -201,11 +202,11 @@ def test_lineage_api_determinism_across_restarts( class RequestResultDict(TypedDict): request: int total_columns: int - columns: Set[str] + columns: set[str] def test_lineage_api_determinism_within_instance( - dbt_artifacts: Dict[str, Any], server_port: int + dbt_artifacts: dict[str, Any], server_port: int ) -> None: """Verify multiple requests to the same server instance return identical results.""" catalog_path = Path(dbt_artifacts["catalog_path"]) @@ -247,7 +248,7 @@ def test_lineage_api_determinism_within_instance( ) -def test_lineage_api_column_count(dbt_artifacts: Dict[str, Any], server_port: int) -> None: +def test_lineage_api_column_count(dbt_artifacts: dict[str, Any], server_port: int) -> None: """Verify the API returns the expected number of columns for the target model.""" catalog_path = Path(dbt_artifacts["catalog_path"]) manifest_path = Path(dbt_artifacts["manifest_path"]) @@ -262,7 +263,7 @@ def test_lineage_api_column_count(dbt_artifacts: Dict[str, Any], server_port: in ) -def test_snapshot_api_support(dbt_artifacts: Dict[str, Any], server_port: int) -> None: +def test_snapshot_api_support(dbt_artifacts: dict[str, Any], server_port: int) -> None: """Verify that snapshots are accessible through the API and have correct resource_type.""" catalog_path = Path(dbt_artifacts["catalog_path"]) manifest_path = Path(dbt_artifacts["manifest_path"]) @@ -328,7 +329,7 @@ def collect_resource_types(nodes: list) -> None: pytest.skip(f"Could not test snapshot lineage endpoint: {e}") -def test_home_page_renders(dbt_artifacts: Dict[str, Any], server_port: int) -> None: +def test_home_page_renders(dbt_artifacts: dict[str, Any], server_port: int) -> None: """The root page must render — it is the only template-rendering route, so API-only coverage misses it entirely (issue #139 shipped a server that answered every /api/* call but 500'd on the page users actually open).""" @@ -343,7 +344,7 @@ def test_home_page_renders(dbt_artifacts: Dict[str, Any], server_port: int) -> N assert "Parrant" in response.text -def test_coverage_endpoint(dbt_artifacts: Dict[str, Any], server_port: int) -> None: +def test_coverage_endpoint(dbt_artifacts: dict[str, Any], server_port: int) -> None: """Verify /api/coverage exposes the artifact coverage block to the explorer UI.""" catalog_path = Path(dbt_artifacts["catalog_path"]) manifest_path = Path(dbt_artifacts["manifest_path"]) @@ -365,7 +366,7 @@ def test_coverage_endpoint(dbt_artifacts: Dict[str, Any], server_port: int) -> N def test_lineage_response_includes_confidence( - dbt_artifacts: Dict[str, Any], server_port: int + dbt_artifacts: dict[str, Any], server_port: int ) -> None: """Verify /api/lineage carries the impact confidence block in impact_summary.""" catalog_path = Path(dbt_artifacts["catalog_path"]) @@ -384,7 +385,7 @@ def test_lineage_response_includes_confidence( def test_upstream_lineage_returns_full_chain( - dbt_artifacts: Dict[str, Any], server_port: int + dbt_artifacts: dict[str, Any], server_port: int ) -> None: """Verify API returns full upstream chain with proper edges.""" catalog_path = Path(dbt_artifacts["catalog_path"]) diff --git a/tests/e2e/test_unresolved_edges_e2e.py b/tests/e2e/test_unresolved_edges_e2e.py index 14a8bc5..88e91d6 100644 --- a/tests/e2e/test_unresolved_edges_e2e.py +++ b/tests/e2e/test_unresolved_edges_e2e.py @@ -30,7 +30,7 @@ import sys import tempfile from pathlib import Path -from typing import Any, Dict, List, Set +from typing import Any import pytest from click.testing import CliRunner @@ -39,7 +39,7 @@ _FIXTURE_DIR = Path(__file__).parents[1] / "fixtures" / "unresolved_edges" sys.path.insert(0, str(_FIXTURE_DIR)) -import _build # type: ignore[import-not-found] # noqa: E402 (path-injected fixture builder) +import _build # type: ignore[import-not-found] # Materialize the abstract manifest + catalog once into a tmp dir for the whole module. _TMP = tempfile.TemporaryDirectory() @@ -57,14 +57,14 @@ _CHANGED_UPSTREAM = "stg_a" -def _json_from_output(output: str) -> Dict[str, Any]: +def _json_from_output(output: str) -> dict[str, Any]: """Parse the JSON report out of the CLI stdout (skipping any leading log lines).""" start = output.index("{") payload, _ = json.JSONDecoder().raw_decode(output[start:]) return payload -def _run(command: Any, args: List[str]) -> Dict[str, Any]: +def _run(command: Any, args: list[str]) -> dict[str, Any]: result = CliRunner().invoke(command, args, catch_exceptions=False) assert result.exit_code == 0, f"exit={result.exit_code}\noutput={result.output}" return _json_from_output(result.output) @@ -81,7 +81,7 @@ def _qualifier(token: str) -> str: return token.split(".", 1)[0].strip().strip('"').lower() -def _legit_upstreams(upstream_block: Dict[str, Any]) -> Set[str]: +def _legit_upstreams(upstream_block: dict[str, Any]) -> set[str]: """The REAL upstreams parrant grouped this column's edges under (ground-truth qualifiers). ``upstream.models`` keys are always real dbt nodes; ``sources``/``direct_refs`` are real too. @@ -93,18 +93,20 @@ def _legit_upstreams(upstream_block: Dict[str, Any]) -> Set[str]: return legit -def _all_source_tokens(upstream_block: Dict[str, Any]) -> Set[str]: - tokens: Set[str] = set() +def _all_source_tokens(upstream_block: dict[str, Any]) -> set[str]: + tokens: set[str] = set() for source_cols in upstream_block.get("models", {}).values(): for edge in source_cols.values(): tokens |= set(edge.get("source_columns", [])) return tokens -def _phantom_tokens(upstream_block: Dict[str, Any]) -> Set[str]: +def _phantom_tokens(upstream_block: dict[str, Any]) -> set[str]: legit = _legit_upstreams(upstream_block) return { - token for token in _all_source_tokens(upstream_block) if _qualifier(token) not in legit | {""} + token + for token in _all_source_tokens(upstream_block) + if _qualifier(token) not in legit | {""} } @@ -147,7 +149,7 @@ def test_flatten_column_has_no_phantom_source_via_cli() -> None: # Propagation: a change reaching a phantom-bearing model degrades confidence and forces its rebuild. # --------------------------------------------------------------------------------------------- # @pytest.fixture -def impact_reaching_phantom(tmp_path: Path) -> Dict[str, Any]: +def impact_reaching_phantom(tmp_path: Path) -> dict[str, Any]: """Run ``parrant impact`` for a change on ``stg_a`` (an upstream of the marker-bearing ``int_c``). diff --git a/tests/fixtures/unresolved_edges/_build.py b/tests/fixtures/unresolved_edges/_build.py index 1063d0b..23430c8 100644 --- a/tests/fixtures/unresolved_edges/_build.py +++ b/tests/fixtures/unresolved_edges/_build.py @@ -33,7 +33,7 @@ import json from pathlib import Path -from typing import Any, Dict, List, Tuple +from typing import Any _PROJECT = "demo" _DB = "DB" @@ -132,9 +132,9 @@ def _manifest_node( name: str, schema: str, compiled_code: str, - depends_on_nodes: List[str], - columns: List[str], -) -> Tuple[str, Dict[str, Any]]: + depends_on_nodes: list[str], + columns: list[str], +) -> tuple[str, dict[str, Any]]: unique_id = f"model.{_PROJECT}.{name}" return unique_id, { "unique_id": unique_id, @@ -156,10 +156,10 @@ def _manifest_node( } -def build_manifest() -> Dict[str, Any]: +def build_manifest() -> dict[str, Any]: source_id = f"source.{_PROJECT}.src_x.raw_x" - nodes: Dict[str, Any] = {} + nodes: dict[str, Any] = {} for uid, node in [ _manifest_node( name="stg_a", @@ -228,7 +228,7 @@ def build_manifest() -> Dict[str, Any]: } -def _catalog_node(*, name: str, schema: str, columns: List[str]) -> Tuple[str, Dict[str, Any]]: +def _catalog_node(*, name: str, schema: str, columns: list[str]) -> tuple[str, dict[str, Any]]: unique_id = f"model.{_PROJECT}.{name}" return unique_id, { "metadata": {"name": name, "schema": schema, "database": _DB, "type": "BASE TABLE"}, @@ -236,10 +236,10 @@ def _catalog_node(*, name: str, schema: str, columns: List[str]) -> Tuple[str, D } -def build_catalog() -> Dict[str, Any]: +def build_catalog() -> dict[str, Any]: # Only the models that must be catalog-backed. `stg_d` is deliberately catalog-missing # (its columns are recovered from compiled SQL). - nodes: Dict[str, Any] = {} + nodes: dict[str, Any] = {} for uid, node in [ _catalog_node(name="stg_a", schema="STG", columns=["col_1", "col_2", "col_3"]), _catalog_node(name="stg_b", schema="STG", columns=["v"]), @@ -270,7 +270,7 @@ def build_catalog() -> Dict[str, Any]: } -def write_fixtures(target_dir: Path) -> Tuple[Path, Path]: +def write_fixtures(target_dir: Path) -> tuple[Path, Path]: """Write the manifest + catalog into ``target_dir`` and return their paths.""" target_dir = Path(target_dir) manifest_path = target_dir / "manifest.json" diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 9c3279f..12408cd 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -1,7 +1,8 @@ -import pytest -import os -from pathlib import Path import sys +from pathlib import Path + +import pytest + @pytest.fixture(scope="session") def dbt_artifacts(): @@ -9,6 +10,6 @@ def dbt_artifacts(): # Import here to avoid circular imports sys.path.insert(0, str(Path(__file__).parent.parent.parent)) from tests.resources.dbt_test_project.setup import setup_dbt_project - + project_dir = Path(__file__).parent.parent / "resources" / "dbt_test_project" - return setup_dbt_project(project_dir) \ No newline at end of file + return setup_dbt_project(project_dir) diff --git a/tests/integration/test_compiled_sql_fallback_integration.py b/tests/integration/test_compiled_sql_fallback_integration.py index 7111aca..c746bf6 100644 --- a/tests/integration/test_compiled_sql_fallback_integration.py +++ b/tests/integration/test_compiled_sql_fallback_integration.py @@ -73,8 +73,7 @@ def test_registry_extracts_lineage_without_embedded_compiled_code( models_with_lineage = [ name for name, model in models.items() - if model.language == "sql" - and any(col.lineage for col in model.columns.values()) + if model.language == "sql" and any(col.lineage for col in model.columns.values()) ] assert models_with_lineage, "no lineage extracted despite compiled files on disk" diff --git a/tests/integration/test_exposures_integration.py b/tests/integration/test_exposures_integration.py index bdb4916..846e236 100644 --- a/tests/integration/test_exposures_integration.py +++ b/tests/integration/test_exposures_integration.py @@ -1,7 +1,9 @@ +from pathlib import Path + import pytest -from parrant.lineage.service import LineageService, LineageSelector + from parrant.lineage.display.text import TextDisplay -from pathlib import Path +from parrant.lineage.service import LineageSelector, LineageService @pytest.fixture @@ -9,7 +11,7 @@ def lineage_service(dbt_artifacts): """Create a LineageService instance.""" return LineageService( catalog_path=Path(dbt_artifacts["catalog_path"]), - manifest_path=Path(dbt_artifacts["manifest_path"]) + manifest_path=Path(dbt_artifacts["manifest_path"]), ) @@ -59,134 +61,143 @@ def test_display_downstream_with_models_and_exposures(lineage_service, capsys): def test_transactions_lineage_includes_exposures(lineage_service): """Test that transactions-related lineage includes both exposures.""" - selector = LineageSelector(model="transactions", column="transaction_id", upstream=False, downstream=True) + selector = LineageSelector( + model="transactions", column="transaction_id", upstream=False, downstream=True + ) column_info = lineage_service.get_column_info(selector) downstream_lineage = column_info["downstream"] - - assert "exposures" in downstream_lineage, \ - "transactions.transaction_id downstream lineage should include exposures" - + + assert ( + "exposures" in downstream_lineage + ), "transactions.transaction_id downstream lineage should include exposures" + exposures = downstream_lineage.get("exposures", set()) assert isinstance(exposures, set), "exposures should be a set" - - expected_exposures = { - "transactions_dashboard", - "api_transactions_endpoint" - } - - assert exposures == expected_exposures, \ - f"Expected exposures {expected_exposures}, got {exposures}" - - selector_amount = LineageSelector(model="transactions", column="amount", upstream=False, downstream=True) + + expected_exposures = {"transactions_dashboard", "api_transactions_endpoint"} + + assert ( + exposures == expected_exposures + ), f"Expected exposures {expected_exposures}, got {exposures}" + + selector_amount = LineageSelector( + model="transactions", column="amount", upstream=False, downstream=True + ) column_info_amount = lineage_service.get_column_info(selector_amount) exposures_amount = column_info_amount["downstream"].get("exposures", set()) - assert exposures_amount == expected_exposures, \ - f"transactions.amount should also include both exposures, got {exposures_amount}" + assert ( + exposures_amount == expected_exposures + ), f"transactions.amount should also include both exposures, got {exposures_amount}" def test_stg_transactions_lineage_includes_exposures(lineage_service): """Test that stg_transactions lineage (which feeds into transactions) includes exposures.""" # stg_transactions -> transactions -> exposures - selector = LineageSelector(model="stg_transactions", column="transaction_id", upstream=False, downstream=True) + selector = LineageSelector( + model="stg_transactions", column="transaction_id", upstream=False, downstream=True + ) column_info = lineage_service.get_column_info(selector) downstream_lineage = column_info["downstream"] - - assert "transactions" in downstream_lineage, \ - "stg_transactions.transaction_id should flow to transactions model" - + + assert ( + "transactions" in downstream_lineage + ), "stg_transactions.transaction_id should flow to transactions model" + exposures = downstream_lineage.get("exposures", set()) - expected_exposures = { - "transactions_dashboard", - "api_transactions_endpoint" - } - - assert exposures == expected_exposures, \ - f"stg_transactions.transaction_id should flow through transactions to both exposures, got {exposures}" + expected_exposures = {"transactions_dashboard", "api_transactions_endpoint"} + + assert ( + exposures == expected_exposures + ), f"stg_transactions.transaction_id should flow through transactions to both exposures, got {exposures}" def test_account_tiering_lineage_includes_report(lineage_service): """Test that account_tiering-related lineage includes the report exposure.""" - selector = LineageSelector(model="accounts_tiering", column="account_id", upstream=False, downstream=True) + selector = LineageSelector( + model="accounts_tiering", column="account_id", upstream=False, downstream=True + ) column_info = lineage_service.get_column_info(selector) downstream_lineage = column_info["downstream"] - + exposures = downstream_lineage.get("exposures", set()) expected_exposures = {"account_tiering_report"} - - assert exposures == expected_exposures, \ - f"accounts_tiering.account_id should include account_tiering_report, got {exposures}" + + assert ( + exposures == expected_exposures + ), f"accounts_tiering.account_id should include account_tiering_report, got {exposures}" def test_int_monthly_account_metrics_lineage_includes_report(lineage_service): """Test that int_monthly_account_metrics lineage includes account_tiering_report.""" # int_monthly_account_metrics -> accounts_tiering -> account_tiering_report - selector = LineageSelector(model="int_monthly_account_metrics", column="account_id", upstream=False, downstream=True) + selector = LineageSelector( + model="int_monthly_account_metrics", column="account_id", upstream=False, downstream=True + ) column_info = lineage_service.get_column_info(selector) downstream_lineage = column_info["downstream"] - + exposures = downstream_lineage.get("exposures", set()) expected_exposures = {"account_tiering_report"} - - assert exposures == expected_exposures, \ - f"int_monthly_account_metrics.account_id should include account_tiering_report, got {exposures}" + + assert ( + exposures == expected_exposures + ), f"int_monthly_account_metrics.account_id should include account_tiering_report, got {exposures}" def test_unrelated_lineage_excludes_transactions_exposures(lineage_service): """Test that unrelated lineage does not include transactions exposures.""" # account_holder column in int_monthly_account_metrics flows to accounts_tiering, # but NOT to transactions, so should NOT include transactions exposures - selector = LineageSelector(model="int_monthly_account_metrics", column="account_holder", upstream=False, downstream=True) + selector = LineageSelector( + model="int_monthly_account_metrics", + column="account_holder", + upstream=False, + downstream=True, + ) column_info = lineage_service.get_column_info(selector) downstream_lineage = column_info["downstream"] - + exposures = downstream_lineage.get("exposures", set()) - - assert "transactions_dashboard" not in exposures, \ - "account_holder lineage should NOT include transactions_dashboard" - assert "api_transactions_endpoint" not in exposures, \ - "account_holder lineage should NOT include api_transactions_endpoint" - + + assert ( + "transactions_dashboard" not in exposures + ), "account_holder lineage should NOT include transactions_dashboard" + assert ( + "api_transactions_endpoint" not in exposures + ), "account_holder lineage should NOT include api_transactions_endpoint" + def test_exposures_not_in_upstream_lineage(lineage_service): """Test that exposures are not included in upstream lineage (only downstream).""" - selector = LineageSelector(model="transactions", column="transaction_id", upstream=True, downstream=False) + selector = LineageSelector( + model="transactions", column="transaction_id", upstream=True, downstream=False + ) column_info = lineage_service.get_column_info(selector) upstream_lineage = column_info["upstream"] - - assert "exposures" not in upstream_lineage, \ - "Upstream lineage should not include exposures" + + assert "exposures" not in upstream_lineage, "Upstream lineage should not include exposures" def test_exposure_metadata_in_lineage_explorer(lineage_service): """Test that exposure metadata is correctly included in lineage explorer.""" from parrant.lineage.display.html.explore import LineageExplorer - + lineage_explorer = LineageExplorer(host="127.0.0.1", port=8000) lineage_explorer.set_lineage_service(lineage_service) - + lineage_explorer._process_lineage_tree("transactions", "transaction_id") - + data_dict = lineage_explorer.data.model_dump() - exposure_nodes = [ - node for node in data_dict["nodes"] - if node.get("type") == "exposure" - ] - - expected_exposure_names = { - "transactions_dashboard", - "api_transactions_endpoint" - } - + exposure_nodes = [node for node in data_dict["nodes"] if node.get("type") == "exposure"] + + expected_exposure_names = {"transactions_dashboard", "api_transactions_endpoint"} + actual_exposure_names = {node["model"] for node in exposure_nodes} - assert actual_exposure_names == expected_exposure_names, \ - f"Expected exposure nodes {expected_exposure_names}, got {actual_exposure_names}" - + assert ( + actual_exposure_names == expected_exposure_names + ), f"Expected exposure nodes {expected_exposure_names}, got {actual_exposure_names}" + # Check that exposure edges are created - exposure_edges = [ - edge for edge in data_dict["edges"] - if edge.get("type") == "exposure" - ] - - assert len(exposure_edges) > 0, \ - "Expected exposure edges to be created in the graph" + exposure_edges = [edge for edge in data_dict["edges"] if edge.get("type") == "exposure"] + assert len(exposure_edges) > 0, "Expected exposure edges to be created in the graph" diff --git a/tests/integration/test_json_output_integration.py b/tests/integration/test_json_output_integration.py index 6e26969..9d1787d 100644 --- a/tests/integration/test_json_output_integration.py +++ b/tests/integration/test_json_output_integration.py @@ -7,7 +7,6 @@ import json -import pytest from click.testing import CliRunner from parrant.cli.main import cli diff --git a/tests/integration/test_lineage_explorer_integration.py b/tests/integration/test_lineage_explorer_integration.py index aad1895..d64f016 100644 --- a/tests/integration/test_lineage_explorer_integration.py +++ b/tests/integration/test_lineage_explorer_integration.py @@ -1,9 +1,11 @@ +from pathlib import Path + import pytest from fastapi.testclient import TestClient -from parrant.lineage.display.html.explore import LineageExplorer + from parrant.artifacts.registry import ModelRegistry +from parrant.lineage.display.html.explore import LineageExplorer from parrant.lineage.service import LineageService -from pathlib import Path @pytest.fixture diff --git a/tests/integration/test_manifest_integration.py b/tests/integration/test_manifest_integration.py index 46543fd..9e32b1a 100644 --- a/tests/integration/test_manifest_integration.py +++ b/tests/integration/test_manifest_integration.py @@ -1,5 +1,3 @@ -from pathlib import Path - from parrant.artifacts.manifest import ManifestReader diff --git a/tests/integration/test_predicate_impact_integration.py b/tests/integration/test_predicate_impact_integration.py index 71d65fd..d4116a5 100644 --- a/tests/integration/test_predicate_impact_integration.py +++ b/tests/integration/test_predicate_impact_integration.py @@ -63,8 +63,6 @@ def test_filter_only_consumer_surfaces_as_filter_severity(dbt_artifacts): # Changing a column it PROJECTS (account_id) is an ordinary value impact, not 'filter'. value_impact = service.get_column_impact("transactions", "account_id") - value_cols = [ - c for c in value_impact["affected_columns"] if c["model"] == _FILTER_MODEL - ] + value_cols = [c for c in value_impact["affected_columns"] if c["model"] == _FILTER_MODEL] assert value_cols, "flagged_transaction_metrics projects account_id" assert all(c["severity"] != "filter" for c in value_cols), value_cols diff --git a/tests/integration/test_registry_integration.py b/tests/integration/test_registry_integration.py index 0bba2f6..d359f0c 100644 --- a/tests/integration/test_registry_integration.py +++ b/tests/integration/test_registry_integration.py @@ -1,6 +1,7 @@ from pathlib import Path import pytest + from parrant.artifacts.registry import ModelRegistry @@ -131,7 +132,7 @@ def test_select_star_lineage(registry): stg_model = models["stg_transactions"] int_model = models["int_transactions_enriched"] - for col_name, _ in stg_model.columns.items(): + for col_name in stg_model.columns: assert ( col_name in int_model.columns ), f"Column {col_name} from stg_transactions should exist in int_transactions_enriched" @@ -213,7 +214,7 @@ def test_snapshot_support(registry): if not snapshots: pytest.skip("No snapshots found in registry. Ensure dbt snapshot has been run.") else: - snapshot_name = list(snapshots.keys())[0] + snapshot_name = next(iter(snapshots.keys())) snapshot = snapshots[snapshot_name] else: snapshot = models[snapshot_name] @@ -354,9 +355,10 @@ def test_exposures_as_downstream_dependencies(registry): def test_impact_analysis(registry, dbt_artifacts): """Test impact analysis for a column - what would break if the column is modified.""" - from parrant.lineage.service import LineageService from pathlib import Path + from parrant.lineage.service import LineageService + service = LineageService( Path(dbt_artifacts["catalog_path"]), Path(dbt_artifacts["manifest_path"]) ) diff --git a/tests/live/seed.py b/tests/live/seed.py index da1df71..1218503 100644 --- a/tests/live/seed.py +++ b/tests/live/seed.py @@ -20,7 +20,7 @@ import time from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional +from typing import Any import requests @@ -52,12 +52,12 @@ class Auth: """ base_url: str - api_key: Optional[str] = None - session_id: Optional[str] = None - username: Optional[str] = None - password: Optional[str] = None + api_key: str | None = None + session_id: str | None = None + username: str | None = None + password: str | None = None - def headers(self) -> Dict[str, str]: + def headers(self) -> dict[str, str]: headers = {"Content-Type": "application/json"} if self.api_key: headers["x-api-key"] = self.api_key @@ -65,7 +65,7 @@ def headers(self) -> Dict[str, str]: headers["X-Metabase-Session"] = self.session_id return headers - def cli_args(self) -> List[str]: + def cli_args(self) -> list[str]: """Credential flags for ``parrant metabase-extract``. Prefers the API key; else username+password. A bare session id is not a CLI-facing @@ -92,7 +92,7 @@ class SeededContent: native_card_id: int mbql_card_id: int dashboard_id: int - card_ids: List[int] = field(default_factory=list) + card_ids: list[int] = field(default_factory=list) def wait_for_health(base_url: str, timeout: float = 180.0, interval: float = 3.0) -> None: @@ -101,7 +101,7 @@ def wait_for_health(base_url: str, timeout: float = 180.0, interval: float = 3.0 Metabase takes 30-90s to boot, so the default timeout is generous. """ deadline = time.monotonic() + timeout - last_error: Optional[str] = None + last_error: str | None = None while time.monotonic() < deadline: try: resp = requests.get(f"{base_url}/api/health", timeout=10) @@ -116,9 +116,9 @@ def wait_for_health(base_url: str, timeout: float = 180.0, interval: float = 3.0 def authenticate( base_url: str, - api_key: Optional[str] = None, - username: Optional[str] = None, - password: Optional[str] = None, + api_key: str | None = None, + username: str | None = None, + password: str | None = None, site_name: str = "parrant-live", ) -> Auth: """Resolve credentials into an :class:`Auth`. @@ -143,7 +143,7 @@ def authenticate( return Auth(base_url=base_url, session_id=session_id, username=username, password=password) -def _setup_token(base_url: str) -> Optional[str]: +def _setup_token(base_url: str) -> str | None: resp = requests.get(f"{base_url}/api/session/properties", timeout=15) if resp.status_code != 200: return None @@ -213,7 +213,7 @@ def find_sample_database(auth: Auth) -> int: def table_and_field_ids( auth: Auth, database_id: int, table_name: str = SAMPLE_TABLE -) -> Dict[str, Any]: +) -> dict[str, Any]: """Return ``{"table_id": int, "field_ids": {NAME: id}}`` for one Sample-DB table. Uses the bulk ``GET /api/database/:id/metadata`` so the MBQL card can be built from real @@ -264,7 +264,7 @@ def create_mbql_card( return _create_card(auth, name, dataset_query) -def _create_card(auth: Auth, name: str, dataset_query: Dict[str, Any]) -> int: +def _create_card(auth: Auth, name: str, dataset_query: dict[str, Any]) -> int: body = _request( auth, "POST", @@ -282,7 +282,7 @@ def _create_card(auth: Auth, name: str, dataset_query: Dict[str, Any]) -> int: return card_id -def create_dashboard(auth: Auth, card_ids: List[int], name: str = "parrant-live dashboard") -> int: +def create_dashboard(auth: Auth, card_ids: list[int], name: str = "parrant-live dashboard") -> int: """Create a dashboard (``POST /api/dashboard``) and place each card on it (``PUT``).""" body = _request(auth, "POST", "/api/dashboard", json={"name": name}) dashboard_id = (body or {}).get("id") diff --git a/tests/resources/dbt_test_project/setup.py b/tests/resources/dbt_test_project/setup.py index 36b6e3a..c9e98c0 100644 --- a/tests/resources/dbt_test_project/setup.py +++ b/tests/resources/dbt_test_project/setup.py @@ -1,8 +1,9 @@ import os from pathlib import Path -from typing import Dict, Any -from dbt.cli.main import dbtRunner +from typing import Any + import duckdb +from dbt.cli.main import dbtRunner def setup_test_db(project_dir: Path) -> Path: @@ -59,7 +60,7 @@ def setup_test_db(project_dir: Path) -> Path: return db_path -def setup_dbt_project(project_dir: Path) -> Dict[str, Any]: +def setup_dbt_project(project_dir: Path) -> dict[str, Any]: """Setup dbt project and return paths to artifacts.""" dbt = dbtRunner() @@ -132,7 +133,7 @@ def setup_dbt_project(project_dir: Path) -> Dict[str, Any]: conn.close() except Exception as e: - error_msg += f"\nDatabase inspection error: {str(e)}" + error_msg += f"\nDatabase inspection error: {e!s}" raise Exception(error_msg) diff --git a/tests/unit/artifacts/test_adapter_mapping.py b/tests/unit/artifacts/test_adapter_mapping.py index cc83853..f044ac5 100644 --- a/tests/unit/artifacts/test_adapter_mapping.py +++ b/tests/unit/artifacts/test_adapter_mapping.py @@ -4,7 +4,7 @@ import pytest -import parrant.artifacts.adapter_mapping as adapter_mapping +from parrant.artifacts import adapter_mapping from parrant.artifacts.adapter_mapping import ( ADAPTER_TO_DIALECT, normalize_adapter, diff --git a/tests/unit/artifacts/test_catalog.py b/tests/unit/artifacts/test_catalog.py index ffd6ad2..0e40a03 100644 --- a/tests/unit/artifacts/test_catalog.py +++ b/tests/unit/artifacts/test_catalog.py @@ -2,6 +2,7 @@ from pathlib import Path import pytest + from parrant.artifacts.catalog import CatalogReader diff --git a/tests/unit/artifacts/test_manifest.py b/tests/unit/artifacts/test_manifest.py index a869af1..e0e6b2b 100644 --- a/tests/unit/artifacts/test_manifest.py +++ b/tests/unit/artifacts/test_manifest.py @@ -2,6 +2,7 @@ from pathlib import Path import pytest + from parrant.artifacts.manifest import ManifestReader diff --git a/tests/unit/artifacts/test_meta_index.py b/tests/unit/artifacts/test_meta_index.py index 29ab0d4..6f2d92b 100644 --- a/tests/unit/artifacts/test_meta_index.py +++ b/tests/unit/artifacts/test_meta_index.py @@ -221,7 +221,9 @@ def test_manifest_get_model_config_returns_node_config(tmp_path): assert config["grants"] == {"select": ["pii_reader", "analyst"]} assert config["tags"] == ["nightly", "finance"] # absent / unknown -> empty dict, never guessed - reader2 = _write_manifest(tmp_path, {"model.p.bare": {"name": "bare", "resource_type": "model"}}) + reader2 = _write_manifest( + tmp_path, {"model.p.bare": {"name": "bare", "resource_type": "model"}} + ) reader2.load() assert reader2.get_model_config("bare") == {} assert reader2.get_model_config("nope") == {} diff --git a/tests/unit/lineage/test_backtest.py b/tests/unit/lineage/test_backtest.py index 208d4c5..2fac45a 100644 --- a/tests/unit/lineage/test_backtest.py +++ b/tests/unit/lineage/test_backtest.py @@ -36,7 +36,6 @@ SemanticChangeKind, ) - # --- git enumeration -------------------------------------------------------- diff --git a/tests/unit/lineage/test_backtest_display.py b/tests/unit/lineage/test_backtest_display.py index cd4f8ff..f25ec7a 100644 --- a/tests/unit/lineage/test_backtest_display.py +++ b/tests/unit/lineage/test_backtest_display.py @@ -12,18 +12,18 @@ def _report(rule_stats=None, **kwargs): - base = dict( - mode="git-diff", - policy_source="p.yml", - base="HEAD~3", - head="HEAD", - prs_replayed=3, - prs_would_block=1, - prs_would_warn=1, - avg_blast_radius=2.5, - rule_stats=rule_stats or [], - fidelity_note="NOTE: block tiers not exercised.", - ) + base = { + "mode": "git-diff", + "policy_source": "p.yml", + "base": "HEAD~3", + "head": "HEAD", + "prs_replayed": 3, + "prs_would_block": 1, + "prs_would_warn": 1, + "avg_blast_radius": 2.5, + "rule_stats": rule_stats or [], + "fidelity_note": "NOTE: block tiers not exercised.", + } base.update(kwargs) return BacktestReport(**base) diff --git a/tests/unit/lineage/test_changeset.py b/tests/unit/lineage/test_changeset.py index e335724..eb62279 100644 --- a/tests/unit/lineage/test_changeset.py +++ b/tests/unit/lineage/test_changeset.py @@ -6,7 +6,6 @@ """ from dataclasses import dataclass, field -from typing import Dict, List, Optional, Set import pytest @@ -18,29 +17,28 @@ build_git_changeset, scope_changes_to_models, ) -from parrant.models.schema import SemanticChangeKind from parrant.lineage.display.markdown import render_changeset_markdown from parrant.lineage.service import LineageService - +from parrant.models.schema import SemanticChangeKind # --- stubs ----------------------------------------------------------------- @dataclass class _Col: - data_type: Optional[str] + data_type: str | None @dataclass class _Model: - columns: Dict[str, _Col] + columns: dict[str, _Col] @dataclass class _Lin: """Stand-in for ColumnLineage (the per-column derivation signature source).""" - source_columns: Set[str] + source_columns: set[str] transformation_type: str sql_expression: str @@ -49,8 +47,8 @@ class _Lin: class _LinCol: """A column that also carries parsed per-column lineage, enabling a precise diff.""" - data_type: Optional[str] = None - lineage: List[_Lin] = field(default_factory=list) + data_type: str | None = None + lineage: list[_Lin] = field(default_factory=list) class _FakeRegistry: @@ -58,9 +56,9 @@ class _FakeRegistry: def __init__( self, - models: Dict[str, _Model], - compiled: Optional[Dict[str, str]] = None, - catalog_backed: Optional[set] = None, + models: dict[str, _Model], + compiled: dict[str, str] | None = None, + catalog_backed: set | None = None, ): self._models = models self._compiled = compiled or {} @@ -68,7 +66,7 @@ def __init__( # Pass an explicit set to simulate catalog-missing (manifest-only) models. self._catalog_backed = catalog_backed if catalog_backed is not None else set(models) - def get_models(self) -> Dict[str, _Model]: + def get_models(self) -> dict[str, _Model]: return self._models def is_catalog_backed(self, model_name: str) -> bool: @@ -83,7 +81,7 @@ def get_compiled_sql(self, model_name: str) -> str: class _FakeService: """Stub exposing get_column_impact, used to drive get_changeset_impact.""" - def __init__(self, impacts: Dict[tuple, dict]): + def __init__(self, impacts: dict[tuple, dict]): self._impacts = impacts def get_column_impact(self, model: str, column: str) -> dict: @@ -953,15 +951,15 @@ def test_markdown_empty_changeset_warns_when_structural_skipped(): @dataclass class _PathModel: - columns: Dict[str, _Col] - resource_path: Optional[str] + columns: dict[str, _Col] + resource_path: str | None class _PathRegistry: - def __init__(self, models: Dict[str, _PathModel]): + def __init__(self, models: dict[str, _PathModel]): self._models = models - def get_models(self) -> Dict[str, _PathModel]: + def get_models(self) -> dict[str, _PathModel]: return self._models @@ -1007,11 +1005,11 @@ def test_scope_changes_empty_when_no_overlap(): # --- override resolution ------------------------------------------------- -from parrant.lineage.changeset import ( # noqa: E402 +from parrant.lineage.changeset import ( OverrideResolution, resolve_overrides, ) -from parrant.models.schema import OverrideVerb # noqa: E402 +from parrant.models.schema import OverrideVerb def _lc(model, column): diff --git a/tests/unit/lineage/test_confidence_completeness.py b/tests/unit/lineage/test_confidence_completeness.py index c421bac..eef753b 100644 --- a/tests/unit/lineage/test_confidence_completeness.py +++ b/tests/unit/lineage/test_confidence_completeness.py @@ -7,8 +7,6 @@ ``len(list) == count`` always, with the display-only ``*_truncated`` flags False. """ -from typing import Dict - from parrant.models.schema import Column, Model from tests.unit.test_lineage_provider import InMemoryProvider, _model, _service_on @@ -24,7 +22,7 @@ def _root_with_blind_downstream(n_blind: int, parse_failed_count: int = 0) -> In "root", {"id": Column(name="id", model_name="root", data_type="int")}, ) - models: Dict[str, Model] = {"root": root} + models: dict[str, Model] = {"root": root} blind_names = [f"d{i:03d}" for i in range(n_blind)] for name in blind_names: models[name] = _model(name, {}) diff --git a/tests/unit/lineage/test_config_axis.py b/tests/unit/lineage/test_config_axis.py index 1b30ad4..c5886bf 100644 --- a/tests/unit/lineage/test_config_axis.py +++ b/tests/unit/lineage/test_config_axis.py @@ -37,7 +37,6 @@ SemanticChangeKind, ) - # --- fakes ------------------------------------------------------------------ @@ -122,7 +121,9 @@ def test_grants_subset_of_allowlist_does_not_fire(): def test_grants_outside_allowlist_fires(): """grants {loader, reporter} ⊄ [loader, transformer] -> not_subset_of is TRUE -> block.""" - registry = FakeRegistry(model_config={"customers": {"grants": {"select": ["loader", "reporter"]}}}) + registry = FakeRegistry( + model_config={"customers": {"grants": {"select": ["loader", "reporter"]}}} + ) policy = _config_policy("grants.select", "not_subset_of", ["loader", "transformer"]) verdict = _engine(policy, registry).evaluate([_change()]) assert verdict.blocks() @@ -136,7 +137,9 @@ def test_missing_grants_is_empty_set_not_unknown(): block even under the default fail_closed posture (the missing path is a PROVEN empty set, not an UNKNOWN that would route to on_missing_meta).""" registry = FakeRegistry(model_config={"customers": {"materialized": "table"}}) # no grants key - policy = _config_policy("grants.select", "not_subset_of", ["loader", "transformer"]) # blocking, fail_closed + policy = _config_policy( + "grants.select", "not_subset_of", ["loader", "transformer"] + ) # blocking, fail_closed verdict = _engine(policy, registry).evaluate([_change()]) assert verdict.decision is GateDecision.ALLOW assert verdict.fired_rules == 0 @@ -144,7 +147,9 @@ def test_missing_grants_is_empty_set_not_unknown(): def test_intersects_and_subset_of_sanity(): - registry = FakeRegistry(model_config={"customers": {"grants": {"select": ["loader", "reporter"]}}}) + registry = FakeRegistry( + model_config={"customers": {"grants": {"select": ["loader", "reporter"]}}} + ) # intersects [reporter] -> shares reporter -> TRUE v_int = _engine(_config_policy("grants.select", "intersects", ["reporter"]), registry).evaluate( [_change()] @@ -152,14 +157,15 @@ def test_intersects_and_subset_of_sanity(): assert v_int.blocks() # subset_of [loader, transformer, reporter] -> {loader, reporter} ⊆ -> TRUE v_sub = _engine( - _config_policy("grants.select", "subset_of", ["loader", "transformer", "reporter"]), registry + _config_policy("grants.select", "subset_of", ["loader", "transformer", "reporter"]), + registry, ).evaluate([_change()]) assert v_sub.blocks() # missing path + intersects -> [] shares nothing -> FALSE (empty set, no fire) empty = FakeRegistry(model_config={"customers": {}}) - v_missing = _engine(_config_policy("grants.select", "intersects", ["reporter"]), empty).evaluate( - [_change()] - ) + v_missing = _engine( + _config_policy("grants.select", "intersects", ["reporter"]), empty + ).evaluate([_change()]) assert v_missing.decision is GateDecision.ALLOW @@ -188,7 +194,9 @@ def test_scalar_missing_is_unknown_and_routes_to_on_missing_meta(): Under fail_closed (default) a *blocking* rule fires on UNKNOWN; under skip it is dropped and counted in skipped_missing_meta. This proves the scalar side behaves like a missing meta key, NOT like the set-op empty set.""" - registry = FakeRegistry(model_config={"customers": {"grants": {"select": ["loader"]}}}) # no materialized + registry = FakeRegistry( + model_config={"customers": {"grants": {"select": ["loader"]}}} + ) # no materialized # fail_closed (default) + blocking -> fires on UNKNOWN closed = _engine(_config_policy("materialized", "eq", "incremental"), registry).evaluate( [_change()] @@ -213,9 +221,9 @@ def test_dotted_traversal_into_nested_config(): model_config={"customers": {"grants": {"select": ["loader"], "insert": ["transformer"]}}} ) # grants.insert is a distinct nested path - verdict = _engine(_config_policy("grants.insert", "intersects", ["transformer"]), registry).evaluate( - [_change()] - ) + verdict = _engine( + _config_policy("grants.insert", "intersects", ["transformer"]), registry + ).evaluate([_change()]) assert verdict.blocks() @@ -283,9 +291,7 @@ def test_non_pii_over_granted_does_not_block(): model_config={"customers": {"grants": {"select": ["analyst"]}}}, column_meta={("customers", "created_at"): {"pii": False}}, ) - verdict = _engine(_pii_grants_policy(), registry).evaluate( - [_change(column="created_at")] - ) + verdict = _engine(_pii_grants_policy(), registry).evaluate([_change(column="created_at")]) assert verdict.decision is GateDecision.ALLOW @@ -437,9 +443,9 @@ def test_superset_of_positive_and_missing(): assert v_pos.blocks() # missing path -> [] ⊇ [loader] is FALSE -> no fire empty = FakeRegistry(model_config={"customers": {}}) - v_missing = _engine( - _config_policy("grants.select", "superset_of", ["loader"]), empty - ).evaluate([_change()]) + v_missing = _engine(_config_policy("grants.select", "superset_of", ["loader"]), empty).evaluate( + [_change()] + ) assert v_missing.decision is GateDecision.ALLOW diff --git a/tests/unit/lineage/test_inferred_meta.py b/tests/unit/lineage/test_inferred_meta.py index af3af29..fd09f5f 100644 --- a/tests/unit/lineage/test_inferred_meta.py +++ b/tests/unit/lineage/test_inferred_meta.py @@ -31,7 +31,6 @@ SemanticChangeKind, ) - # --- fakes ------------------------------------------------------------------ @@ -235,7 +234,9 @@ def test_secret_or_propagates_but_own_false_halts(): prop_lookup = MetaIndex(prop).inferred_meta("mart", "token", "secret") assert prop_lookup.present is True assert prop_lookup.value is True - assert _engine(_inferred_policy(key="secret"), prop).evaluate([_change("mart", "token")]).blocks() + assert ( + _engine(_inferred_policy(key="secret"), prop).evaluate([_change("mart", "token")]).blocks() + ) # (b) downstream own secret:false wins over an upstream secret:true. halt = FakeRegistry( @@ -245,7 +246,9 @@ def test_secret_or_propagates_but_own_false_halts(): halt_lookup = MetaIndex(halt).inferred_meta("mart", "token", "secret") assert halt_lookup.present is True assert halt_lookup.value is False - halt_verdict = _engine(_inferred_policy(key="secret"), halt).evaluate([_change("mart", "token")]) + halt_verdict = _engine(_inferred_policy(key="secret"), halt).evaluate( + [_change("mart", "token")] + ) assert halt_verdict.decision is GateDecision.ALLOW diff --git a/tests/unit/lineage/test_manifest_seeded_universe.py b/tests/unit/lineage/test_manifest_seeded_universe.py index 49c37f7..dec73b4 100644 --- a/tests/unit/lineage/test_manifest_seeded_universe.py +++ b/tests/unit/lineage/test_manifest_seeded_universe.py @@ -36,10 +36,9 @@ import json from parrant.artifacts.registry import ModelRegistry -from parrant.lineage.service import LineageService from parrant.lineage.changeset import ChangeKind, ChangesetBuilder from parrant.lineage.display.markdown import _confidence_reason_words - +from parrant.lineage.service import LineageService # --- fixture helpers ------------------------------------------------------- diff --git a/tests/unit/lineage/test_markdown_confidence_cap.py b/tests/unit/lineage/test_markdown_confidence_cap.py index 86b66f4..f1b9bf8 100644 --- a/tests/unit/lineage/test_markdown_confidence_cap.py +++ b/tests/unit/lineage/test_markdown_confidence_cap.py @@ -5,12 +5,12 @@ completeness of the machine lists the JSON surface already emitted. """ -from typing import Any, Dict, List +from typing import Any from parrant.lineage.display.markdown import render_changeset_markdown -def _report_with_unanalyzable(names: List[str]) -> Dict[str, Any]: +def _report_with_unanalyzable(names: list[str]) -> dict[str, Any]: return { "summary": {"affected_models": 0, "affected_columns": 0}, "changeset": {"total_changes": 1, "by_kind": {"logic_changed": 1}}, diff --git a/tests/unit/lineage/test_metabase_markdown.py b/tests/unit/lineage/test_metabase_markdown.py index 127e38a..9b0c9e0 100644 --- a/tests/unit/lineage/test_metabase_markdown.py +++ b/tests/unit/lineage/test_metabase_markdown.py @@ -34,8 +34,12 @@ def test_column_precise_dashboard_names_the_affected_field(): "precision": "column", "via_cards": [128], "via_columns": [ - {"model": "dim_accounts", "column": "balance", "card_id": 128, - "role": "field"}, + { + "model": "dim_accounts", + "column": "balance", + "card_id": 128, + "role": "field", + }, ], "meta": {}, } diff --git a/tests/unit/lineage/test_opaque_models.py b/tests/unit/lineage/test_opaque_models.py index ae23252..02594c1 100644 --- a/tests/unit/lineage/test_opaque_models.py +++ b/tests/unit/lineage/test_opaque_models.py @@ -16,9 +16,8 @@ import json from parrant.artifacts.registry import ModelRegistry -from parrant.lineage.service import LineageService from parrant.lineage.changeset import ChangesetBuilder - +from parrant.lineage.service import LineageService # --- fixture helpers ------------------------------------------------------- @@ -77,9 +76,7 @@ def _catalog_nodes(): def _manifest_nodes(orders_sql): return { - "model.p.dim_orders": _manifest_node( - "dim_orders", "select 1 as order_id, 1 as amount" - ), + "model.p.dim_orders": _manifest_node("dim_orders", "select 1 as order_id, 1 as amount"), "model.p.orders": _manifest_node("orders", orders_sql, depends_on=["dim_orders"]), "model.p.revenue_semantic_view": _manifest_node( "revenue_semantic_view", diff --git a/tests/unit/lineage/test_policy_engine.py b/tests/unit/lineage/test_policy_engine.py index b583370..81f9816 100644 --- a/tests/unit/lineage/test_policy_engine.py +++ b/tests/unit/lineage/test_policy_engine.py @@ -1191,11 +1191,11 @@ def test_semantic_knobs_fold_into_user_rules_most_severe_wins(): # --- override caps ------------------------------------------------------- -from parrant.lineage.policy import ( # noqa: E402 +from parrant.lineage.policy import ( applied_policy_overrides, ineffective_policy_overrides, ) -from parrant.models.schema import OverrideDirective, OverrideVerb # noqa: E402 +from parrant.models.schema import OverrideDirective, OverrideVerb def _ov(verb, column, reason="ack"): @@ -1272,18 +1272,16 @@ def test_applied_policy_overrides_shape_matches_default_gate(): records = applied_policy_overrides(verdict, [change]) assert len(records) == 1 r = records[0] - assert set( - [ - "model", - "column", - "verb", - "reason", - "downgraded_from", - "downgraded_to", - "source_line", - "scope", - ] - ) <= set(r.keys()) + assert { + "model", + "column", + "verb", + "reason", + "downgraded_from", + "downgraded_to", + "source_line", + "scope", + } <= set(r.keys()) assert r["verb"] == "allow-change" assert r["downgraded_from"] == "block" assert r["downgraded_to"] == "allow" diff --git a/tests/unit/lineage/test_policy_init.py b/tests/unit/lineage/test_policy_init.py index 746ced8..80eabeb 100644 --- a/tests/unit/lineage/test_policy_init.py +++ b/tests/unit/lineage/test_policy_init.py @@ -12,7 +12,6 @@ from parrant.lineage.policy_init import _flatten_meta_keys, _has_select_grant, emit_policy_yaml from parrant.models.schema import MetaKeyCoverage, PolicyInitScan - # --- pure scan helpers ------------------------------------------------------- diff --git a/tests/unit/lineage/test_resolution.py b/tests/unit/lineage/test_resolution.py index 047661f..829ce08 100644 --- a/tests/unit/lineage/test_resolution.py +++ b/tests/unit/lineage/test_resolution.py @@ -8,8 +8,6 @@ statuses reconcile exactly with the confidence counts and the rebuild/skippable sets. """ -from typing import Dict - from parrant.lineage.changeset import ChangesetBuilder from parrant.lineage.service import build_resolution from parrant.models.schema import Column, ColumnLineage, Model @@ -29,7 +27,7 @@ def _resolution_provider() -> InMemoryProvider: missing = _model("missing", {}) py_model = _model("py_model", {}, language="python") broke = _model("broke", {}) - models: Dict[str, Model] = { + models: dict[str, Model] = { "catalog_backed": catalog_model, "parsed": parsed_model, "star_cte": star_cte, diff --git a/tests/unit/lineage/test_selection.py b/tests/unit/lineage/test_selection.py index 3541c07..319bf6e 100644 --- a/tests/unit/lineage/test_selection.py +++ b/tests/unit/lineage/test_selection.py @@ -8,7 +8,7 @@ sentinel, and determinism). """ -from typing import Any, Dict, List, Optional, Set +from typing import Any from parrant.lineage.changeset import ChangesetBuilder from parrant.lineage.service import build_selection @@ -18,12 +18,12 @@ def _confidence( *, level: str = "full", - no_column_info_models: Optional[List[str]] = None, - parse_failed_models: Optional[List[str]] = None, - partial_edges_models: Optional[List[str]] = None, + no_column_info_models: list[str] | None = None, + parse_failed_models: list[str] | None = None, + partial_edges_models: list[str] | None = None, no_column_info_truncated: bool = False, parse_failed_truncated: bool = False, -) -> Dict[str, Any]: +) -> dict[str, Any]: return { "level": level, "no_column_info_models": no_column_info_models or [], @@ -37,11 +37,11 @@ def _confidence( def _change( *, kind: str, - semantic: Optional[str], - reached: Optional[List[str]] = None, + semantic: str | None, + reached: list[str] | None = None, resolved: bool = True, -) -> Dict[str, Any]: - entry: Dict[str, Any] = {"kind": kind, "semantic": semantic, "resolved": resolved} +) -> dict[str, Any]: + entry: dict[str, Any] = {"kind": kind, "semantic": semantic, "resolved": resolved} if reached is not None: entry["reached_models"] = [ {"name": name, "mechanism": "direct_passthrough"} for name in reached @@ -49,7 +49,7 @@ def _change( return entry -def _assert_partition(selection: Dict[str, Any], universe: Set[str]) -> None: +def _assert_partition(selection: dict[str, Any], universe: set[str]) -> None: """Every model in the universe has exactly one disposition — the headline invariant.""" rebuild = set(selection["rebuild_models"]) skippable = set(selection["skippable_models"]) @@ -200,7 +200,7 @@ def test_unresolved_change_still_rebuilds_its_own_model() -> None: # An unresolved change contributes no reach, but its edited model is in changed_models and so # is always rebuilt — the diff never silently drops a model it could not fan out. changed = {"orphan"} - reachable: Set[str] = set() + reachable: set[str] = set() by_change = [_change(kind="removed", semantic=None, resolved=False)] selection = build_selection(reachable, changed, by_change, _confidence()) diff --git a/tests/unit/lineage/test_star_passthrough_catalog_missing.py b/tests/unit/lineage/test_star_passthrough_catalog_missing.py index 5362ea6..54255f5 100644 --- a/tests/unit/lineage/test_star_passthrough_catalog_missing.py +++ b/tests/unit/lineage/test_star_passthrough_catalog_missing.py @@ -22,9 +22,9 @@ import json -from parrant.lineage.sqlglot_provider import build_sqlglot_provider -from parrant.lineage.policy import MetaIndex, evaluate_policy, parse_policy from parrant.lineage.changeset import ChangeKind, ColumnChange +from parrant.lineage.policy import MetaIndex, evaluate_policy, parse_policy +from parrant.lineage.sqlglot_provider import build_sqlglot_provider from parrant.models.schema import GateDecision, SemanticChangeKind DB, SCH = "analytics", "main" @@ -122,7 +122,10 @@ def _build_provider(tmp_path): "exposures": {}, } # Deferred build: only the changed mart is in the fresh catalog; upstreams are not rebuilt. - catalog = {"metadata": {}, "nodes": {f"model.analytics.{MART}": _catalog_node(MART, _MART_COLS)}} + catalog = { + "metadata": {}, + "nodes": {f"model.analytics.{MART}": _catalog_node(MART, _MART_COLS)}, + } m = tmp_path / "manifest.json" c = tmp_path / "catalog.json" diff --git a/tests/unit/lineage/test_unresolved_edges_propagation.py b/tests/unit/lineage/test_unresolved_edges_propagation.py index 3ff94cb..0718ccd 100644 --- a/tests/unit/lineage/test_unresolved_edges_propagation.py +++ b/tests/unit/lineage/test_unresolved_edges_propagation.py @@ -16,14 +16,12 @@ tokens, so empty sources must not be read as "clean" — the marker is the signal. """ -from typing import Optional, Set - from parrant.lineage.service import build_resolution, build_selection -from parrant.models.schema import Column, ColumnLineage, Model +from parrant.models.schema import Column, ColumnLineage from tests.unit.test_lineage_provider import InMemoryProvider, _model, _service_on -def _col(name: str, *, sources: Optional[Set[str]] = None) -> Column: +def _col(name: str, *, sources: set[str] | None = None) -> Column: lineage = ( [ColumnLineage(source_columns=sources, transformation_type="direct")] if sources is not None diff --git a/tests/unit/lineage/test_verdict.py b/tests/unit/lineage/test_verdict.py index 3d0e618..1affd1a 100644 --- a/tests/unit/lineage/test_verdict.py +++ b/tests/unit/lineage/test_verdict.py @@ -6,7 +6,6 @@ """ from dataclasses import dataclass, field -from typing import Dict, List, Tuple from parrant.lineage.changeset import ChangeKind, ColumnChange from parrant.lineage.verdict import classify_provable_breaks, decide_verdict @@ -15,31 +14,31 @@ @dataclass class _Model: - columns: Dict[str, object] + columns: dict[str, object] @dataclass class _FakeRegistry: """Stand-in for ModelRegistry covering only the classifier's call surface.""" - models: Dict[str, _Model] = field(default_factory=dict) - column_tests: Dict[Tuple[str, str], List[TestNode]] = field(default_factory=dict) - referenced_tests: Dict[Tuple[str, str], List[TestNode]] = field(default_factory=dict) + models: dict[str, _Model] = field(default_factory=dict) + column_tests: dict[tuple[str, str], list[TestNode]] = field(default_factory=dict) + referenced_tests: dict[tuple[str, str], list[TestNode]] = field(default_factory=dict) # Test unique_ids present in THIS manifest. On a head registry this is what lets the # classifier confirm a base test survived the change (see the rename-with-yml-update case). test_ids: set = field(default_factory=set) - model_tests: Dict[str, List[TestNode]] = field(default_factory=dict) + model_tests: dict[str, list[TestNode]] = field(default_factory=dict) - def get_models(self) -> Dict[str, _Model]: + def get_models(self) -> dict[str, _Model]: return self.models - def get_column_tests(self, model: str, column: str) -> List[TestNode]: + def get_column_tests(self, model: str, column: str) -> list[TestNode]: return list(self.column_tests.get((model.lower(), column.lower()), [])) - def get_tests_referencing(self, model: str, column: str) -> List[TestNode]: + def get_tests_referencing(self, model: str, column: str) -> list[TestNode]: return list(self.referenced_tests.get((model.lower(), column.lower()), [])) - def get_model_tests(self, model: str) -> List[TestNode]: + def get_model_tests(self, model: str) -> list[TestNode]: return list(self.model_tests.get(model.lower(), [])) def get_test_unique_ids(self) -> set: @@ -278,12 +277,12 @@ def test_verdict_omitting_changes_is_backward_compatible(): # --- overrides: decide_verdict + applied/ineffective override records ---- -from parrant.models.schema import OverrideDirective, OverrideVerb # noqa: E402 -from parrant.lineage.verdict import ( # noqa: E402 +from parrant.lineage.verdict import ( applied_overrides, ineffective_overrides, unexcused_break_count, ) +from parrant.models.schema import OverrideDirective, OverrideVerb def _directive(verb, column=None, scope="column", reason="because"): diff --git a/tests/unit/metabase/_fixtures.py b/tests/unit/metabase/_fixtures.py index c338d34..075bff4 100644 --- a/tests/unit/metabase/_fixtures.py +++ b/tests/unit/metabase/_fixtures.py @@ -11,24 +11,24 @@ import re import threading from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any RECORDED_PATH = Path(__file__).parents[2] / "resources" / "metabase" / "recorded.json" -def load_recorded() -> Dict[str, Any]: +def load_recorded() -> dict[str, Any]: return json.loads(RECORDED_PATH.read_text(encoding="utf-8")) def build_recorded( *, - cards: Optional[List[dict]] = None, - dashboards: Optional[List[dict]] = None, - dashboard_details: Optional[Dict[str, dict]] = None, - database_metadata: Optional[Dict[str, dict]] = None, - snippets: Optional[List[dict]] = None, - session_properties: Optional[dict] = None, -) -> Dict[str, Any]: + cards: list[dict] | None = None, + dashboards: list[dict] | None = None, + dashboard_details: dict[str, dict] | None = None, + database_metadata: dict[str, dict] | None = None, + snippets: list[dict] | None = None, + session_properties: dict | None = None, +) -> dict[str, Any]: """Assemble a full recorded-payload dict from inline parts, defaulting the rest. Every key :class:`FakeSession` may index is filled so a purpose-built corpus (a @@ -46,7 +46,7 @@ def build_recorded( class FakeResponse: - def __init__(self, body: Any, status_code: int = 200, headers: Optional[dict] = None): + def __init__(self, body: Any, status_code: int = 200, headers: dict | None = None): self._body = body self.status_code = status_code self.headers = headers or {} @@ -77,10 +77,10 @@ class FakeSession: a lock so the recorded ``(url, params)`` list never races or drops an entry. """ - def __init__(self, recorded: Dict[str, Any], fail_first: int = 0): + def __init__(self, recorded: dict[str, Any], fail_first: int = 0): self.recorded = recorded - self.get_calls: List[tuple] = [] - self.post_calls: List[tuple] = [] + self.get_calls: list[tuple] = [] + self.post_calls: list[tuple] = [] self._remaining_failures = fail_first self._lock = threading.Lock() diff --git a/tests/unit/metabase/test_incremental.py b/tests/unit/metabase/test_incremental.py index 61a690e..f494df1 100644 --- a/tests/unit/metabase/test_incremental.py +++ b/tests/unit/metabase/test_incremental.py @@ -9,8 +9,6 @@ from __future__ import annotations -from typing import Dict, List, Optional - from parrant.metabase.client import MetabaseClient from parrant.metabase.extract import ExtractConfig, run_extract from parrant.models.schema import ( @@ -65,9 +63,9 @@ class SpyClient(MetabaseClient): def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) - self.fetched_dashboard_ids: List[int] = [] + self.fetched_dashboard_ids: list[int] = [] - def get_dashboards(self, dashboard_ids: List[int], max_workers: int = 8) -> Dict[int, dict]: + def get_dashboards(self, dashboard_ids: list[int], max_workers: int = 8) -> dict[int, dict]: self.fetched_dashboard_ids.extend(dashboard_ids) return super().get_dashboards(dashboard_ids, max_workers=max_workers) @@ -81,22 +79,22 @@ def _spy_client(session: FakeSession) -> SpyClient: ) -def _config(previous: Optional[MetabaseLineage], **overrides) -> ExtractConfig: - kwargs = dict( - metabase_base_url="https://metabase.example.com", - database_ids=[2], - extractor_version="9.9.9", - dialect="snowflake", - previous=previous, - ) +def _config(previous: MetabaseLineage | None, **overrides) -> ExtractConfig: + kwargs = { + "metabase_base_url": "https://metabase.example.com", + "database_ids": [2], + "extractor_version": "9.9.9", + "dialect": "snowflake", + "previous": previous, + } kwargs.update(overrides) return ExtractConfig(**kwargs) # type: ignore[arg-type] def _prev_snapshot( - dashboards: List[MetabaseDashboard], + dashboards: list[MetabaseDashboard], schema_version: int = 2, - database_ids: Optional[List[int]] = None, + database_ids: list[int] | None = None, ) -> MetabaseLineage: return MetabaseLineage( schema_version=schema_version, diff --git a/tests/unit/metabase/test_join_reach.py b/tests/unit/metabase/test_join_reach.py index 9e4af81..eb3e791 100644 --- a/tests/unit/metabase/test_join_reach.py +++ b/tests/unit/metabase/test_join_reach.py @@ -1,4 +1,4 @@ -""" unit tests — the offline relation join and the ``(model,column) -> card -> dashboard`` +"""unit tests — the offline relation join and the ``(model,column) -> card -> dashboard`` reach index. Pure (no dbt build), so they run under ``test-unit``.""" from __future__ import annotations diff --git a/tests/unit/metabase/test_pmbql.py b/tests/unit/metabase/test_pmbql.py index 8aeb10a..cd6be8f 100644 --- a/tests/unit/metabase/test_pmbql.py +++ b/tests/unit/metabase/test_pmbql.py @@ -4,7 +4,7 @@ import json from pathlib import Path -from typing import Any, Dict +from typing import Any from parrant.metabase.pmbql import is_pmbql, normalize_dataset_query from parrant.metabase.resolvers import CardResolver @@ -23,11 +23,11 @@ F_PEOPLE_NAME = 48 -def load_pmbql() -> Dict[str, Any]: +def load_pmbql() -> dict[str, Any]: return json.loads(_PMBQL_PATH.read_text(encoding="utf-8")) -def _cards_by_id() -> Dict[int, dict]: +def _cards_by_id() -> dict[int, dict]: return {c["id"]: c for c in load_pmbql()["cards"]} diff --git a/tests/unit/metabase/test_resolvers.py b/tests/unit/metabase/test_resolvers.py index 28a40e6..0c50cc0 100644 --- a/tests/unit/metabase/test_resolvers.py +++ b/tests/unit/metabase/test_resolvers.py @@ -1,4 +1,4 @@ -""" — the two resolvers + warehouse-meta normalization, against the recorded fixture.""" +"""— the two resolvers + warehouse-meta normalization, against the recorded fixture.""" from __future__ import annotations diff --git a/tests/unit/parser/test_sql_parser.py b/tests/unit/parser/test_sql_parser.py index 2541496..f493afe 100644 --- a/tests/unit/parser/test_sql_parser.py +++ b/tests/unit/parser/test_sql_parser.py @@ -795,11 +795,12 @@ def test_complex_query_structure(): def test_table_names_normalized_from_sql() -> None: """Test that table names extracted from SQL are normalized to lowercase.""" + from sqlglot import exp, parse_one + from parrant.parser.sql_parser_utils import ( - get_table_context, get_all_tables_from_select, + get_table_context, ) - from sqlglot import parse_one, exp # Test with uppercase table names (Snowflake style) sql = "SELECT * FROM RAW_ORDERS_TABLE" @@ -974,7 +975,7 @@ def test_comments_in_join_condition() -> None: assert "order_id" in lineage # All source columns should be clean - for col_name, lineage_list in lineage.items(): + for lineage_list in lineage.values(): for lineage_item in lineage_list: for src in lineage_item.source_columns: assert "/*" not in src and "*/" not in src @@ -1053,7 +1054,7 @@ def test_comments_in_complex_query() -> None: assert "order_id" in lineage # Verify all source columns are clean - for col_name, lineage_list in lineage.items(): + for lineage_list in lineage.values(): for lineage_item in lineage_list: for src in lineage_item.source_columns: assert "/*" not in src and "*/" not in src diff --git a/tests/unit/registry/test_coverage.py b/tests/unit/registry/test_coverage.py index dc49003..30a0fcb 100644 --- a/tests/unit/registry/test_coverage.py +++ b/tests/unit/registry/test_coverage.py @@ -10,8 +10,8 @@ import pytest -from parrant.artifacts.registry import ModelRegistry from parrant.artifacts.exceptions import RegistryNotLoadedError +from parrant.artifacts.registry import ModelRegistry def _write(tmp_path, catalog_data, manifest_data): diff --git a/tests/unit/registry/test_registry.py b/tests/unit/registry/test_registry.py index e0edd14..83ad652 100644 --- a/tests/unit/registry/test_registry.py +++ b/tests/unit/registry/test_registry.py @@ -1,7 +1,10 @@ import json + import pytest -from parrant.artifacts.registry import ModelRegistry + from parrant.artifacts.exceptions import ModelNotFoundError, RegistryNotLoadedError +from parrant.artifacts.registry import ModelRegistry + @pytest.fixture def sample_catalog(tmp_path): @@ -18,9 +21,9 @@ def sample_catalog(tmp_path): "name": "id", "description": "Primary key", "data_type": "integer", - "model_name": "customers" + "model_name": "customers", } - } + }, }, "model.jaffle_shop.orders": { "unique_id": "model.jaffle_shop.orders", @@ -29,8 +32,12 @@ def sample_catalog(tmp_path): "database": "raw", "columns": { "id": {"name": "id", "data_type": "integer", "model_name": "orders"}, - "customer_id": {"name": "customer_id", "data_type": "integer", "model_name": "orders"} - } + "customer_id": { + "name": "customer_id", + "data_type": "integer", + "model_name": "orders", + }, + }, }, "model.jaffle_shop.order_items": { "unique_id": "model.jaffle_shop.order_items", @@ -38,19 +45,28 @@ def sample_catalog(tmp_path): "schema": "jaffle_shop", "database": "raw", "columns": { - "order_id": {"name": "order_id", "data_type": "integer", "model_name": "order_items"}, - "quantity": {"name": "quantity", "data_type": "integer", "model_name": "order_items"} - } - } + "order_id": { + "name": "order_id", + "data_type": "integer", + "model_name": "order_items", + }, + "quantity": { + "name": "quantity", + "data_type": "integer", + "model_name": "order_items", + }, + }, + }, } } - + catalog_path = tmp_path / "catalog.json" with open(catalog_path, "w") as f: json.dump(catalog_data, f) - + return catalog_path + @pytest.fixture def sample_manifest(tmp_path): """Create a sample manifest file for testing.""" @@ -59,78 +75,72 @@ def sample_manifest(tmp_path): "model.jaffle_shop.customers": { "name": "customers", "resource_type": "model", - "depends_on": {"nodes": []} + "depends_on": {"nodes": []}, }, "model.jaffle_shop.orders": { "name": "orders", "resource_type": "model", - "depends_on": { - "nodes": ["model.jaffle_shop.customers"] - } + "depends_on": {"nodes": ["model.jaffle_shop.customers"]}, }, "model.jaffle_shop.order_items": { "name": "order_items", "resource_type": "model", - "depends_on": { - "nodes": ["model.jaffle_shop.orders"] - } - } + "depends_on": {"nodes": ["model.jaffle_shop.orders"]}, + }, } } - + manifest_path = tmp_path / "manifest.json" with open(manifest_path, "w") as f: json.dump(manifest_data, f) - + return manifest_path + def test_registry_basics(sample_catalog, sample_manifest): """Test registry initialization, loading, and error handling.""" registry = ModelRegistry(sample_catalog, sample_manifest) assert not registry.is_loaded - + with pytest.raises(RegistryNotLoadedError): registry.get_models() - + registry.load() assert registry.is_loaded - + models = registry.get_models() assert len(models) > 0 - + model = registry.get_model("customers") assert model is not None assert model.name == "customers" - + with pytest.raises(ModelNotFoundError): registry.get_model("foobar_model") + def test_model_dependencies(sample_catalog, sample_manifest): """Test model dependency relationships.""" registry = ModelRegistry(sample_catalog, sample_manifest) registry.load() - + models = registry.get_models() - + dependency_tests = { - "customers": { - "upstream": set(), - "downstream": {"orders"} - }, - "orders": { - "upstream": {"customers"}, - "downstream": {"order_items"} - }, - "order_items": { - "upstream": {"orders"}, - "downstream": set() - } + "customers": {"upstream": set(), "downstream": {"orders"}}, + "orders": {"upstream": {"customers"}, "downstream": {"order_items"}}, + "order_items": {"upstream": {"orders"}, "downstream": set()}, } - + for model_name, expected in dependency_tests.items(): model = models[model_name] - assert model.upstream == expected["upstream"], f"Expected {model_name}.upstream to be {expected['upstream']}, got {model.upstream}" - assert model.downstream == expected["downstream"], f"Expected {model_name}.downstream to be {expected['downstream']}, got {model.downstream}" + assert ( + model.upstream == expected["upstream"] + ), f"Expected {model_name}.upstream to be {expected['upstream']}, got {model.upstream}" + assert ( + model.downstream == expected["downstream"] + ), f"Expected {model_name}.downstream to be {expected['downstream']}, got {model.downstream}" + def test_adapter_override(tmp_path): """Test that adapter override takes precedence over manifest adapter.""" @@ -142,38 +152,40 @@ def test_adapter_override(tmp_path): "name": "test_model", "schema": "test", "database": "test_db", - "columns": {} + "columns": {}, } } } - + manifest_data = { - "metadata": { - "adapter_type": "sqlserver" # Should be normalized to "tsql" - }, + "metadata": {"adapter_type": "sqlserver"}, # Should be normalized to "tsql" "nodes": { "model.test.test_model": { "name": "test_model", "resource_type": "model", - "depends_on": {"nodes": []} + "depends_on": {"nodes": []}, } - } + }, } - + catalog_path = tmp_path / "catalog.json" manifest_path = tmp_path / "manifest.json" - + with open(catalog_path, "w") as f: json.dump(catalog_data, f) with open(manifest_path, "w") as f: json.dump(manifest_data, f) - + # Test without override - should normalize sqlserver to tsql registry_no_override = ModelRegistry(str(catalog_path), str(manifest_path)) registry_no_override.load() assert registry_no_override._dialect == "tsql", "Expected sqlserver to be normalized to tsql" - + # Test with override - should use the override - registry_with_override = ModelRegistry(str(catalog_path), str(manifest_path), adapter_override="bigquery") + registry_with_override = ModelRegistry( + str(catalog_path), str(manifest_path), adapter_override="bigquery" + ) registry_with_override.load() - assert registry_with_override._dialect == "bigquery", "Expected adapter override to take precedence" + assert ( + registry_with_override._dialect == "bigquery" + ), "Expected adapter override to take precedence" diff --git a/tests/unit/registry/test_unresolved_edges.py b/tests/unit/registry/test_unresolved_edges.py index 59eaa47..7c5de38 100644 --- a/tests/unit/registry/test_unresolved_edges.py +++ b/tests/unit/registry/test_unresolved_edges.py @@ -17,7 +17,7 @@ _FIXTURE_DIR = Path(__file__).parents[2] / "fixtures" / "unresolved_edges" sys.path.insert(0, str(_FIXTURE_DIR)) -import _build # type: ignore[import-not-found] # noqa: E402 (path-injected fixture builder) +import _build # type: ignore[import-not-found] # Materialize the manifest + catalog once into a tmp dir for the whole module. _TMP = tempfile.TemporaryDirectory() @@ -146,9 +146,7 @@ def test_fabricated_detection_keeps_legit_passthrough_failsafe(): assert "stg_b.v" in sources # kept — column exists upstream assert not any("fab_" in token for token in sources) # fabricated tokens dropped - fabricated_details = { - m["detail"] for m in _markers(model) if m["column"] == "g_owner" - } + fabricated_details = {m["detail"] for m in _markers(model) if m["column"] == "g_owner"} assert fabricated_details == { "stg_a.fab_v1", "stg_a.fab_v2", diff --git a/tests/unit/test_lineage_provider.py b/tests/unit/test_lineage_provider.py index 83203af..637c47a 100644 --- a/tests/unit/test_lineage_provider.py +++ b/tests/unit/test_lineage_provider.py @@ -18,7 +18,7 @@ from __future__ import annotations -from typing import Any, Dict, List, Optional, Set +from typing import Any from parrant.artifacts.exceptions import ModelNotFoundError from parrant.lineage.changeset import ChangeKind, ChangesetBuilder @@ -49,16 +49,16 @@ class InMemoryProvider: def __init__( self, - models: Dict[str, Model], + models: dict[str, Model], *, - exposures: Optional[Dict[str, Exposure]] = None, - dialect: Optional[str] = None, - downstream: Optional[Dict[str, Set[str]]] = None, - compiled: Optional[Dict[str, str]] = None, - column_tests: Optional[Dict[tuple, List[TestNode]]] = None, - catalog_backed: Optional[Set[str]] = None, - parse_failed: Optional[Set[str]] = None, - opaque: Optional[Set[str]] = None, + exposures: dict[str, Exposure] | None = None, + dialect: str | None = None, + downstream: dict[str, set[str]] | None = None, + compiled: dict[str, str] | None = None, + column_tests: dict[tuple, list[TestNode]] | None = None, + catalog_backed: set[str] | None = None, + parse_failed: set[str] | None = None, + opaque: set[str] | None = None, ) -> None: self._models = {name.lower(): model for name, model in models.items()} self._exposures = exposures or {} @@ -84,7 +84,7 @@ def is_loaded(self) -> bool: return self._loaded # --- model / graph access --------------------------------------------- - def get_models(self) -> Dict[str, Model]: + def get_models(self) -> dict[str, Model]: return self._models def get_model(self, model_name: str) -> Model: @@ -93,15 +93,15 @@ def get_model(self, model_name: str) -> Model: raise ModelNotFoundError(f"Model '{model_name}' not found") return model - def get_manifest_downstream(self) -> Dict[str, Set[str]]: + def get_manifest_downstream(self) -> dict[str, set[str]]: return self._downstream # --- column lineage ---------------------------------------------------- - def get_column_lineage(self, model_name: str, column_name: str) -> List[ColumnLineage]: + def get_column_lineage(self, model_name: str, column_name: str) -> list[ColumnLineage]: column = self.get_column(model_name, column_name) return list(column.lineage or []) if column is not None else [] - def get_column(self, model_name: str, column_name: str) -> Optional[Column]: + def get_column(self, model_name: str, column_name: str) -> Column | None: try: model = self.get_model(model_name) except ModelNotFoundError: @@ -109,10 +109,10 @@ def get_column(self, model_name: str, column_name: str) -> Optional[Column]: return model.columns.get(column_name) or model.columns.get(column_name.lower()) # --- capabilities ------------------------------------------------------ - def get_filter_dependents(self, source_column: str) -> Set[str]: + def get_filter_dependents(self, source_column: str) -> set[str]: return set() - def get_dialect(self) -> Optional[str]: + def get_dialect(self) -> str | None: return self._dialect def get_coverage(self) -> Coverage: @@ -131,17 +131,17 @@ def get_coverage(self) -> Coverage: def is_catalog_backed(self, model_name: str) -> bool: return model_name.lower() in self._catalog_backed - def get_parse_failed_models(self) -> Set[str]: + def get_parse_failed_models(self) -> set[str]: return set(self._parse_failed) - def get_opaque_models(self) -> Set[str]: + def get_opaque_models(self) -> set[str]: return set(self._opaque) - def get_compiled_sql(self, model_name: str) -> Optional[str]: + def get_compiled_sql(self, model_name: str) -> str | None: return self._compiled.get(model_name.lower()) # --- metadata ---------------------------------------------------------- - def get_exposures(self) -> Dict[str, Exposure]: + def get_exposures(self) -> dict[str, Exposure]: return self._exposures def get_exposure(self, exposure_name: str) -> Exposure: @@ -150,39 +150,39 @@ def get_exposure(self, exposure_name: str) -> Exposure: raise ValueError(f"Exposure '{exposure_name}' not found") return exposure - def get_column_tests(self, model: str, column: str) -> List[TestNode]: + def get_column_tests(self, model: str, column: str) -> list[TestNode]: return list(self._column_tests.get((model.lower(), column.lower()), [])) - def get_tests_referencing(self, model: str, column: str) -> List[TestNode]: + def get_tests_referencing(self, model: str, column: str) -> list[TestNode]: return [] - def get_model_tests(self, model: str) -> List[TestNode]: - found: List[TestNode] = [] + def get_model_tests(self, model: str) -> list[TestNode]: + found: list[TestNode] = [] for (test_model, _column), tests in self._column_tests.items(): if test_model == model.lower(): found.extend(tests) return found - def get_test_unique_ids(self) -> Set[str]: + def get_test_unique_ids(self) -> set[str]: return {t.unique_id for tests in self._column_tests.values() for t in tests} def get_unattributable_test_count(self) -> int: return 0 - def get_model_dbt_meta(self, model: str) -> Dict[str, Any]: + def get_model_dbt_meta(self, model: str) -> dict[str, Any]: return {} - def get_column_dbt_meta(self, model: str, column: str) -> Dict[str, Any]: + def get_column_dbt_meta(self, model: str, column: str) -> dict[str, Any]: return {} - def get_model_config(self, model: str) -> Dict[str, Any]: + def get_model_config(self, model: str) -> dict[str, Any]: return {} # --- fixtures -------------------------------------------------------------- -def _model(name: str, columns: Dict[str, Column], **kwargs) -> Model: +def _model(name: str, columns: dict[str, Column], **kwargs) -> Model: return Model( name=name, schema="main", database="main", resource_type="model", columns=columns, **kwargs )