Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 0 additions & 4 deletions metadata-ingestion/docs/dev_guides/lineage_urn_casing.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
23 changes: 1 addition & 22 deletions metadata-ingestion/src/datahub/ingestion/run/pipeline_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.
Expand Down
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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 (
Comment thread
aviraj-gour marked this conversation as resolved.
SchemaResolver as _SchemaResolver,
match_columns_to_schema,
Expand Down
67 changes: 2 additions & 65 deletions metadata-ingestion/src/datahub/sql_parsing/_models.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
import functools
from typing import Any, Optional, Tuple

import sqlglot
from pydantic import BaseModel


Expand Down Expand Up @@ -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,
)
21 changes: 14 additions & 7 deletions metadata-ingestion/src/datahub/sql_parsing/sqlglot_lineage.py
Original file line number Diff line number Diff line change
Expand Up @@ -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],
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
)

Expand Down Expand Up @@ -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,
)
Expand Down Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
)
Expand All @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
25 changes: 13 additions & 12 deletions metadata-ingestion/tests/unit/sql_parsing/test_schemaresolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
Loading
Loading