diff --git a/metadata-ingestion/docs/dev_guides/lineage_urn_casing.md b/metadata-ingestion/docs/dev_guides/lineage_urn_casing.md index 646700b4750b..28170bc1af12 100644 --- a/metadata-ingestion/docs/dev_guides/lineage_urn_casing.md +++ b/metadata-ingestion/docs/dev_guides/lineage_urn_casing.md @@ -234,10 +234,6 @@ ingest-time only: existing metadata is updated only when its source is re-ingest (`ingest_data_platform_instance_aspect`), and the instance name must also match the casing that was emitted. If the log shows `Loaded 0 URNs`, drop `platform_instance` and read the whole platform / env instead. -- **Requires the SQL-parser dependency (`sqlglot`).** Every intended BI/dashboard connector already - bundles it, so the target use case needs no extra install. If you enable the flag on a source that - doesn't, the feature reports a clear failure (`install acryl-datahub[sql-parser]`) and emits lineage - unchanged. - **Only reconciles full-aspect (UPSERT) lineage, not PATCH.** A lineage aspect emitted as a patch (e.g. `dataJobInputOutput` via `DatasetPatchBuilder.add_upstream_lineage` / `DataJobPatchBuilder`, used by some dbt / Airflow / Spark paths) is emitted unchanged and counted under diff --git a/metadata-ingestion/src/datahub/ingestion/run/pipeline_config.py b/metadata-ingestion/src/datahub/ingestion/run/pipeline_config.py index 6f66ba1c10fe..546089f57378 100644 --- a/metadata-ingestion/src/datahub/ingestion/run/pipeline_config.py +++ b/metadata-ingestion/src/datahub/ingestion/run/pipeline_config.py @@ -103,10 +103,7 @@ class AutoResolveLineageUrnsConfig(ConfigModel): enabled: bool = Field( default=False, description="Whether to reconcile the casing of upstream warehouse URN " - "references in lineage against the casing stored in DataHub. Requires the " - "SQL-parser dependency (`sqlglot`) — install `acryl-datahub[sql-parser]` or a " - "connector extra that bundles it. Every intended BI/dashboard connector already " - "does, so the target use case needs no extra install.", + "references in lineage against the casing stored in DataHub.", ) upstream_platforms: List[UpstreamPlatformCasing] = Field( default_factory=list, @@ -149,24 +146,6 @@ def _require_platforms_or_resolve_all_when_enabled( ) return self - @model_validator(mode="after") - def _require_sql_parser_when_enabled(self) -> "AutoResolveLineageUrnsConfig": - # Fail fast at config parse (only when enabled) if the SQL parser is missing, - # rather than deep in the processor at run time. Resolution reuses the - # SchemaResolver, which depends on sqlglot; sqlglot is not in the ingestion core, - # so a source whose extra doesn't bundle it would otherwise fail mid-run. - if self.enabled: - try: - import sqlglot # noqa: F401 - except ImportError as e: - raise ValueError( - "auto_resolve_lineage_urns is enabled but the SQL parser it relies " - "on is not installed. Install it with " - "`pip install acryl-datahub[sql-parser]` (or a connector extra that " - "bundles it), or set enabled: false." - ) from e - return self - class FlagsConfig(ConfigModel): """Experimental flags for the ingestion pipeline. diff --git a/metadata-ingestion/src/datahub/ingestion/workunit_processors/auto_resolve_lineage_urns.py b/metadata-ingestion/src/datahub/ingestion/workunit_processors/auto_resolve_lineage_urns.py index a5052fe64013..3a6b146cc917 100644 --- a/metadata-ingestion/src/datahub/ingestion/workunit_processors/auto_resolve_lineage_urns.py +++ b/metadata-ingestion/src/datahub/ingestion/workunit_processors/auto_resolve_lineage_urns.py @@ -1,11 +1,9 @@ -# NOTE: `from __future__ import annotations` keeps the schema_resolver type hints -# (imported only under TYPE_CHECKING) as strings, so importing this module does not -# pull in sqlglot. This module is imported eagerly on every source's -# get_workunit_processors() path, so module load must stay sqlglot-free (guarded by -# test_module_import_does_not_pull_sqlglot). The sqlglot-heavy schema_resolver imports -# are therefore deferred to a single chokepoint in __init__, which runs only after -# should_enable() confirms the feature is on and a graph exists — off the module-load -# path, but honest about the dependency (see __init__). +# `from __future__ import annotations` is load-bearing: SchemaInfo is TYPE_CHECKING-only +# and appears in the _Resolution dataclass field and _schema_of's return annotation, +# both of which would otherwise be evaluated at class-body execution. This module is +# imported eagerly on every source's get_workunit_processors() path, so the +# schema_resolver imports are deferred to a single chokepoint in __init__, which runs +# only after should_enable() confirms the feature is on and a graph exists. from __future__ import annotations import logging @@ -236,13 +234,9 @@ def __init__(self, ctx: WorkunitProcessorContext) -> None: # Preloaded URN indexes per platform, one per configured entry that loaded. A cache # only: a miss is asked of DataHub. self._alias_resolvers: Dict[str, List[UrnAliasResolver]] = {} - # Resolve the sqlglot-backed schema_resolver callables once, here — a single - # honest chokepoint rather than imports buried in two leaf methods. Deferred into - # __init__ (not module level) so importing this module stays sqlglot-free - # (guarded by test_module_import_does_not_pull_sqlglot); __init__ runs only after - # should_enable() confirmed the feature is on. The sqlglot dependency itself is - # validated up front by AutoResolveLineageUrnsConfig (fail-fast at config parse - # when enabled), so these imports are guaranteed to succeed here. + # A single chokepoint rather than imports buried in two leaf methods. Deferred + # into __init__ so sources that never enable the feature don't import + # schema_resolver at all. from datahub.sql_parsing.schema_resolver import ( SchemaResolver as _SchemaResolver, match_columns_to_schema, diff --git a/metadata-ingestion/src/datahub/sql_parsing/_models.py b/metadata-ingestion/src/datahub/sql_parsing/_models.py index ca7b38f33c63..d9d4945f24a3 100644 --- a/metadata-ingestion/src/datahub/sql_parsing/_models.py +++ b/metadata-ingestion/src/datahub/sql_parsing/_models.py @@ -1,7 +1,6 @@ import functools from typing import Any, Optional, Tuple -import sqlglot from pydantic import BaseModel @@ -61,76 +60,14 @@ def __eq__(self, other: object) -> bool: return False return self.identity == other.identity - def as_sqlglot_table(self) -> sqlglot.exp.Table: - return sqlglot.exp.Table( - catalog=( - sqlglot.exp.Identifier(this=self.database) if self.database else None - ), - db=sqlglot.exp.Identifier(this=self.db_schema) if self.db_schema else None, - this=sqlglot.exp.Identifier(this=self.table), - ) - def qualified( self, - dialect: sqlglot.Dialect, default_db: Optional[str] = None, default_schema: Optional[str] = None, ) -> "_TableName": - database = self.database or default_db - db_schema = self.db_schema or default_schema - return _TableName( - database=database, - db_schema=db_schema, + database=self.database or default_db, + db_schema=self.db_schema or default_schema, table=self.table, parts=self.parts, ) - - @classmethod - def from_sqlglot_table( - cls, - table: sqlglot.exp.Table, - default_db: Optional[str] = None, - default_schema: Optional[str] = None, - ) -> "_TableName": - # Handle Snowflake semantic views: SEMANTIC_VIEW(table_name ...) - # In this case, table.this is a SemanticView expression, and we need to - # extract the actual table from within it. - if isinstance(table.this, sqlglot.exp.SemanticView): - # The SemanticView.this contains the actual table reference - inner_table = table.this.this - if isinstance(inner_table, sqlglot.exp.Table): - # Recursively extract from the inner table - return cls.from_sqlglot_table(inner_table, default_db, default_schema) - elif isinstance(inner_table, sqlglot.exp.Identifier): - # Simple table name - return cls( - database=table.catalog or default_db, - db_schema=table.db or default_schema, - table=inner_table.name, - parts=None, - ) - - if isinstance(table.this, sqlglot.exp.Dot): - # Multi-part tables (>3 parts) have extra parts in a Dot expression. - # Dot is left-associative (a.b.c = Dot(Dot(a,b),c)), so collect right-side - # identifiers while walking left, then reverse. - parts = [] - exp = table.this - while isinstance(exp, sqlglot.exp.Dot): - parts.append(exp.expression.name) - exp = exp.this - parts.append(exp.name) - parts.reverse() - table_name = ".".join(parts) - else: - table_name = table.this.name - - parts_tuple = tuple(p.name for p in table.parts) if table.parts else None - - return cls( - database=table.catalog or default_db, - db_schema=table.db or default_schema, - table=table_name, - parts=parts_tuple, - ) diff --git a/metadata-ingestion/src/datahub/sql_parsing/sqlglot_lineage.py b/metadata-ingestion/src/datahub/sql_parsing/sqlglot_lineage.py index 4ff4e1665bc6..8dbcbe14254f 100644 --- a/metadata-ingestion/src/datahub/sql_parsing/sqlglot_lineage.py +++ b/metadata-ingestion/src/datahub/sql_parsing/sqlglot_lineage.py @@ -138,6 +138,16 @@ def _restore_mssql_temp_table_prefix( return table_name +def _table_name_as_sqlglot_table(table: _TableName) -> sqlglot.exp.Table: + return sqlglot.exp.Table( + catalog=( + sqlglot.exp.Identifier(this=table.database) if table.database else None + ), + db=sqlglot.exp.Identifier(this=table.db_schema) if table.db_schema else None, + this=sqlglot.exp.Identifier(this=table.table), + ) + + def _table_name_from_sqlglot_table( table: sqlglot.exp.Table, dialect: Optional[sqlglot.Dialect], @@ -146,8 +156,7 @@ def _table_name_from_sqlglot_table( ) -> _TableName: """Create a _TableName from a sqlglot Table, handling MSSQL temp table prefixes. - This is a dialect-aware wrapper around _TableName.from_sqlglot_table that - restores MSSQL temp table prefixes (# or ##) that SQLGlot strips during parsing. + Restores MSSQL temp table prefixes (# or ##) that SQLGlot strips during parsing. Args: table: The SQLGlot Table expression @@ -200,8 +209,7 @@ def _table_name_from_sqlglot_table( # Handle Dot expressions (more than 3-part names). # Dot is left-associative (a.b.c = Dot(Dot(a,b),c)), so collect right-side - # identifiers while walking left, then reverse. Mirror of the traversal in - # `_TableName.from_sqlglot_table`; kept in sync intentionally. + # identifiers while walking left, then reverse. if isinstance(table.this, sqlglot.exp.Dot): all_parts_exp: List[sqlglot.exp.Expression] = [] exp: sqlglot.exp.Expression = table.this @@ -949,7 +957,7 @@ def _prepare_query_columns( normalized_table_schema[col_normalized] = col_type or "UNKNOWN" sqlglot_db_schema.add_table( - table.as_sqlglot_table(), + _table_name_as_sqlglot_table(table), column_mapping=normalized_table_schema, ) @@ -1410,7 +1418,6 @@ def _get_direct_raw_col_upstreams( and dialect is not None ): table_ref = table_ref.qualified( - dialect=dialect, default_db=default_db, default_schema=default_schema, ) @@ -2151,7 +2158,7 @@ def _sqlglot_lineage_inner( # For select statements, qualification will be a no-op. For other statements, this # is where the qualification actually happens. qualified_table = table.qualified( - dialect=dialect, default_db=default_db, default_schema=default_schema + default_db=default_db, default_schema=default_schema ) urn, schema_info = schema_resolver.resolve_table(qualified_table) diff --git a/metadata-ingestion/tests/unit/dremio/test_dremio_schema_resolver.py b/metadata-ingestion/tests/unit/dremio/test_dremio_schema_resolver.py index 52448d529f5b..3a1623b82122 100644 --- a/metadata-ingestion/tests/unit/dremio/test_dremio_schema_resolver.py +++ b/metadata-ingestion/tests/unit/dremio/test_dremio_schema_resolver.py @@ -3,6 +3,7 @@ from datahub.ingestion.source.dremio.dremio_source import DremioSchemaResolver from datahub.sql_parsing._models import _TableName +from datahub.sql_parsing.sqlglot_lineage import _table_name_from_sqlglot_table class TestDremioSchemaResolver: @@ -193,7 +194,7 @@ def _parse_table_from_sql(self, sql: str) -> _TableName: """Helper to parse a table reference from SQL using SQLGlot.""" parsed = sqlglot.parse_one(sql, dialect="dremio") for table in parsed.find_all(sqlglot.exp.Table): - return _TableName.from_sqlglot_table(table) + return _table_name_from_sqlglot_table(table, None) raise ValueError(f"No table found in SQL: {sql}") @pytest.mark.parametrize( @@ -262,7 +263,6 @@ def test_sqlglot_with_default_db(self, resolver): table = self._parse_table_from_sql(sql) qualified_table = table.qualified( - dialect=sqlglot.Dialect.get_or_raise("dremio"), default_db="dremio", default_schema=None, ) @@ -287,7 +287,7 @@ def test_sqlglot_complex_query_with_joins(self, resolver): parsed_tables = [] for table in parsed.find_all(sqlglot.exp.Table): - parsed_tables.append(_TableName.from_sqlglot_table(table)) + parsed_tables.append(_table_name_from_sqlglot_table(table, None)) assert len(parsed_tables) == 2 @@ -315,7 +315,7 @@ def test_multi_part_table_name_with_parts(self): tables = list(parsed.find_all(sqlglot.exp.Table)) table = tables[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) urn = resolver.get_urn_for_table(table_name) # Verify the full hierarchy is preserved @@ -350,7 +350,7 @@ def test_unquoted_multi_part_table_names(self, resolver): sql = "SELECT * FROM MySpace.folder1.folder2.folder3.table" parsed = sqlglot.parse_one(sql, dialect="dremio") tables = list(parsed.find_all(sqlglot.exp.Table)) - table_name = _TableName.from_sqlglot_table(tables[0]) + table_name = _table_name_from_sqlglot_table(tables[0], None) # Verify SQLGlot correctly populates parts for unquoted identifiers assert table_name.parts is not None diff --git a/metadata-ingestion/tests/unit/sql_parsing/test_schemaresolver.py b/metadata-ingestion/tests/unit/sql_parsing/test_schemaresolver.py index 5299ae9f768b..a9203dc60c14 100644 --- a/metadata-ingestion/tests/unit/sql_parsing/test_schemaresolver.py +++ b/metadata-ingestion/tests/unit/sql_parsing/test_schemaresolver.py @@ -13,6 +13,7 @@ _TableName, match_columns_to_schema, ) +from datahub.sql_parsing.sqlglot_lineage import _table_name_from_sqlglot_table def create_default_schema_resolver(urn: str) -> SchemaResolver: @@ -310,7 +311,7 @@ def test_multi_part_table_name_5_parts(self): tables = list(parsed.find_all(sqlglot.exp.Table)) table = tables[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) assert table_name.parts is not None assert len(table_name.parts) == 5 @@ -329,7 +330,7 @@ def test_multi_part_table_name_4_parts(self): tables = list(parsed.find_all(sqlglot.exp.Table)) table = tables[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) assert table_name.parts is not None assert len(table_name.parts) == 4 @@ -342,7 +343,7 @@ def test_parts_are_strings_not_objects(self): tables = list(parsed.find_all(sqlglot.exp.Table)) table = tables[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) assert table_name.parts is not None for part in table_name.parts: @@ -355,7 +356,7 @@ def test_hashability_with_parts(self): tables = list(parsed.find_all(sqlglot.exp.Table)) table = tables[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) # Should be able to hash it h = hash(table_name) @@ -402,14 +403,14 @@ def test_parts_in_set_operations(self): parsed2 = parse_one(sql2) parsed3 = parse_one(sql3) - tn1 = _TableName.from_sqlglot_table( - list(parsed1.find_all(sqlglot.exp.Table))[0] + tn1 = _table_name_from_sqlglot_table( + list(parsed1.find_all(sqlglot.exp.Table))[0], None ) - tn2 = _TableName.from_sqlglot_table( - list(parsed2.find_all(sqlglot.exp.Table))[0] + tn2 = _table_name_from_sqlglot_table( + list(parsed2.find_all(sqlglot.exp.Table))[0], None ) - tn3 = _TableName.from_sqlglot_table( - list(parsed3.find_all(sqlglot.exp.Table))[0] + tn3 = _table_name_from_sqlglot_table( + list(parsed3.find_all(sqlglot.exp.Table))[0], None ) # Set should deduplicate tn1 and tn3 @@ -497,7 +498,7 @@ def test_deep_hierarchy_6_plus_parts(self): parsed = parse_one(sql) table = list(parsed.find_all(Table))[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) assert table_name.parts is not None assert len(table_name.parts) == 7 @@ -509,7 +510,7 @@ def test_special_characters_in_parts(self): parsed = parse_one(sql) table = list(parsed.find_all(Table))[0] - table_name = _TableName.from_sqlglot_table(table) + table_name = _table_name_from_sqlglot_table(table, None) assert table_name.parts is not None assert "my-source" in table_name.parts diff --git a/metadata-ingestion/tests/unit/sql_parsing/test_table_name.py b/metadata-ingestion/tests/unit/sql_parsing/test_table_name.py index 9bd733bd2dc2..64b06c6769b4 100644 --- a/metadata-ingestion/tests/unit/sql_parsing/test_table_name.py +++ b/metadata-ingestion/tests/unit/sql_parsing/test_table_name.py @@ -95,14 +95,14 @@ def test_none_dialect_returns_unchanged(self): class TestTableNameFromSqlglotTable: - """Tests for _TableName.from_sqlglot_table() method (basic functionality).""" + """Basic table name extraction.""" def test_basic_table_extraction(self): """Basic table name extraction should work.""" table = sqlglot.exp.Table( this=sqlglot.exp.Identifier(this="my_table"), ) - result = _TableName.from_sqlglot_table(table) + result = _table_name_from_sqlglot_table(table, None) assert result.table == "my_table" assert result.database is None assert result.db_schema is None @@ -114,7 +114,7 @@ def test_qualified_table_extraction(self): db=sqlglot.exp.Identifier(this="my_schema"), this=sqlglot.exp.Identifier(this="my_table"), ) - result = _TableName.from_sqlglot_table(table) + result = _table_name_from_sqlglot_table(table, None) assert result.table == "my_table" assert result.database == "my_db" assert result.db_schema == "my_schema" @@ -124,8 +124,8 @@ def test_default_db_and_schema(self): table = sqlglot.exp.Table( this=sqlglot.exp.Identifier(this="my_table"), ) - result = _TableName.from_sqlglot_table( - table, default_db="default_db", default_schema="default_schema" + result = _table_name_from_sqlglot_table( + table, None, default_db="default_db", default_schema="default_schema" ) assert result.table == "my_table" assert result.database == "default_db" @@ -138,8 +138,8 @@ def test_explicit_overrides_default(self): db=sqlglot.exp.Identifier(this="explicit_schema"), this=sqlglot.exp.Identifier(this="my_table"), ) - result = _TableName.from_sqlglot_table( - table, default_db="default_db", default_schema="default_schema" + result = _table_name_from_sqlglot_table( + table, None, default_db="default_db", default_schema="default_schema" ) assert result.database == "explicit_db" assert result.db_schema == "explicit_schema" @@ -883,22 +883,9 @@ def test_qualified_adds_defaults(self): """qualified() should add default db/schema if not present.""" table = _TableName(table="my_table") qualified = table.qualified( - dialect=get_dialect("mssql"), default_db="default_db", default_schema="default_schema", ) assert qualified.database == "default_db" assert qualified.db_schema == "default_schema" assert qualified.table == "my_table" - - def test_qualified_preserves_temp_prefix(self): - """qualified() should preserve # prefix on temp tables.""" - table = _TableName(table="#temptable") - qualified = table.qualified( - dialect=get_dialect("mssql"), - default_db="mydb", - default_schema="dbo", - ) - assert qualified.table == "#temptable" - assert qualified.database == "mydb" - assert qualified.db_schema == "dbo" diff --git a/metadata-ingestion/tests/unit/workunit_processors/test_auto_resolve_lineage_urns.py b/metadata-ingestion/tests/unit/workunit_processors/test_auto_resolve_lineage_urns.py index 083758749fe9..839be0d2afb0 100644 --- a/metadata-ingestion/tests/unit/workunit_processors/test_auto_resolve_lineage_urns.py +++ b/metadata-ingestion/tests/unit/workunit_processors/test_auto_resolve_lineage_urns.py @@ -1095,14 +1095,17 @@ def test_workunit_level_counters_track_lineage_and_modified(): def test_module_import_does_not_pull_sqlglot(): - # Importing this module (e.g. via the workunit_processors package) must not drag - # in sqlglot, or connectors that don't declare it would break. The invariant rests - # on deferred imports + `from __future__ import annotations`; assert it in a fresh - # interpreter, since this test session may already have sqlglot loaded. + # The resolver chain must not drag in sqlglot, or connectors that don't declare it + # would break. Assert in a fresh interpreter, since this test session may already + # have sqlglot loaded. code = ( "import sys; " "import datahub.ingestion.workunit_processors.auto_resolve_lineage_urns; " - "assert 'sqlglot' not in sys.modules, 'sqlglot imported at module load'" + "assert 'sqlglot' not in sys.modules, 'sqlglot imported at module load'; " + "import datahub.sql_parsing.schema_resolver; " + "assert 'sqlglot' not in sys.modules, 'schema_resolver pulled in sqlglot'; " + "import datahub.sql_parsing.schema_resolver_provider; " + "assert 'sqlglot' not in sys.modules, 'schema_resolver_provider pulled in sqlglot'" ) result = subprocess.run( [sys.executable, "-c", code], capture_output=True, text=True @@ -1292,25 +1295,6 @@ def test_disabled_under_bare_mock_ctx(): assert AutoResolveLineageUrnsProcessor.should_enable(mock.MagicMock()) is False -def test_config_requires_sql_parser_only_when_enabled(monkeypatch): - # sqlglot is not in the ingestion core. Enabling the feature without it must fail - # fast at config parse (only when enabled), with an actionable message — not deep in - # the processor at run time. Simulate the missing dependency by nulling the module. - monkeypatch.setitem(sys.modules, "sqlglot", None) - - # Disabled: no requirement, config validates fine. - AutoResolveLineageUrnsConfig(enabled=False) - - # Enabled: the SQL parser is required, so config validation fails. - with pytest.raises(pydantic.ValidationError, match="sql-parser"): - AutoResolveLineageUrnsConfig( - enabled=True, - upstream_platforms=[ - UpstreamPlatformCasing(platform="snowflake", env="PROD") - ], - ) - - # --- identity from a shared index, columns from our own load -----------------------