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/.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 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/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 d6846e8..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,12 +190,20 @@ 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]: + """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/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 925f511..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,19 +164,21 @@ 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) - 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 @@ -201,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 @@ -209,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", [])) @@ -222,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 @@ -231,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) ) @@ -239,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) @@ -273,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/ @@ -292,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``; @@ -340,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 @@ -349,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: @@ -405,7 +407,7 @@ def build_resolution( @dataclass class LineageSelector: model: str - column: Optional[str] + column: str | None upstream: bool downstream: bool @@ -437,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 @@ -455,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(): @@ -475,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. @@ -488,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() @@ -502,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 @@ -514,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 @@ -554,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 @@ -586,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. @@ -610,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 { @@ -622,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: @@ -645,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 @@ -668,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(): @@ -689,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: @@ -711,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() @@ -738,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: @@ -757,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}" @@ -809,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) @@ -829,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. @@ -896,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 "" ), @@ -923,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) @@ -962,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: @@ -1152,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 @@ -1172,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 @@ -1202,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: @@ -1253,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): @@ -1316,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 4cdb923..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 "" @@ -735,7 +738,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] + ) -> 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 @@ -754,12 +759,12 @@ def _phantom_token_reason(token: str, flatten_aliases: Set[str]) -> Optional[str 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 @@ -771,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( @@ -785,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: @@ -820,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: @@ -837,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: @@ -848,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 @@ -940,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: @@ -973,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: @@ -1017,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) @@ -1043,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 @@ -1053,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] @@ -1067,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 @@ -1096,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 @@ -1136,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: @@ -1147,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: @@ -1164,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 @@ -1189,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/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 5e8b577..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,13 +70,33 @@ 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" [tool.mypy] python_version = "3.10" -files = ["dbt_column_lineage", "tests"] +files = ["parrant", "tests"] show_error_codes = true pretty = true explicit_package_bases = true 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 1cfc35a..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 # 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 760f990..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 @@ -268,7 +269,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 +349,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/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 d92c41c..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 # 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 2582a39..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,36 +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]: 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 )