Skip to content
Open
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
68 changes: 57 additions & 11 deletions sqlmesh/core/table_diff.py
Original file line number Diff line number Diff line change
Expand Up @@ -256,7 +256,7 @@ def __init__(
self.target_alias = target_alias

cols: t.List[str] = ensure_list(skip_columns)
self.skip_columns = {
self._requested_skip_columns = {
normalize_identifiers(
exp.parse_identifier(col),
dialect=self.model_dialect or self.dialect,
Expand All @@ -275,30 +275,76 @@ def source_schema(self) -> t.Dict[str, exp.DataType]:
def target_schema(self) -> t.Dict[str, exp.DataType]:
return self.adapter.columns(self.target_table)

@cached_property
def skip_columns(self) -> t.Set[str]:
"""The names of the columns to skip, as reported by the engine for either table.

Names that don't exist in a table are ignored for that table.
"""
skipped = set()
for name in self._requested_skip_columns:
for schema in (self.source_schema, self.target_schema):
if (resolved := self._find_column_name(name, schema)) is not None:
skipped.add(resolved)
return skipped

def _find_column_name(self, name: str, schema: t.Dict[str, exp.DataType]) -> t.Optional[str]:
"""Maps a user-supplied column name to the column name reported by the engine.

User-supplied names are normalized using the dialect, which may change their casing
compared to the casing the engine reports. An exact match always wins, otherwise a
case-insensitive match is used as long as it is unambiguous.

Returns None if there is no match and raises if the match is ambiguous.
"""
if name in schema:
return name

matches = [c for c in schema if c.casefold() == name.casefold()]
if len(matches) > 1:
raise SQLMeshError(
f"Column '{name}' is ambiguous, it matches multiple columns: {', '.join(matches)}"
)
return matches[0] if matches else None

def _resolve_column_name(self, name: str, schema: t.Dict[str, exp.DataType]) -> str:
resolved = self._find_column_name(name, schema)
if resolved is None:
raise SQLMeshError(
f"Column '{name}' does not exist. Available columns: {', '.join(schema)}"
)
return resolved

@cached_property
def key_columns(self) -> t.Tuple[t.List[exp.Column], t.List[exp.Column], t.List[str]]:
dialect = self.model_dialect or self.dialect

# If the columns to join on are explicitly specified, then just return them
if isinstance(self._on, (list, tuple)):
identifiers = [normalize_identifiers(c, dialect=dialect) for c in self._on]
s_index = [exp.column(c, "s") for c in identifiers]
t_index = [exp.column(c, "t") for c in identifiers]
return s_index, t_index, [i.name for i in identifiers]
names = [normalize_identifiers(c, dialect=dialect).name for c in self._on]
s_names = [self._resolve_column_name(n, self.source_schema) for n in names]
t_names = [self._resolve_column_name(n, self.target_schema) for n in names]
s_index = [exp.column(c, "s") for c in s_names]
t_index = [exp.column(c, "t") for c in t_names]
# The source and target spellings of a column can differ, so keep both
return s_index, t_index, list(dict.fromkeys(s_names + t_names))

# Otherwise, we need to parse them out of the supplied "on" condition
index_cols = []
s_index = []
t_index = []

normalize_identifiers(self._on, dialect=dialect)
for col in self._on.find_all(exp.Column):
# Work on a copy so the caller's expression isn't modified
on = normalize_identifiers(self._on.copy(), dialect=dialect)
for col in on.find_all(exp.Column):
table = col.table.lower()
if table in ("s", "t"):
schema = self.source_schema if table == "s" else self.target_schema
col.set("this", exp.to_identifier(self._resolve_column_name(col.name, schema)))
(s_index if table == "s" else t_index).append(col)
index_cols.append(col.name)
if col.table.lower() == "s":
s_index.append(col)
elif col.table.lower() == "t":
t_index.append(col)

# Like the list form above, index_cols can contain both source and target spellings
index_cols = list(dict.fromkeys(index_cols))
s_index = list(dict.fromkeys(s_index))
t_index = list(dict.fromkeys(t_index))
Expand Down
166 changes: 166 additions & 0 deletions tests/core/test_table_diff.py
Original file line number Diff line number Diff line change
Expand Up @@ -1246,3 +1246,169 @@ def test_data_diff_nulls_in_some_grain_columns():
"null value",
"null value modified",
]


def _create_uppercase_tables() -> t.Any:
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()

columns_to_types = {
"KEY1": exp.DataType.build("int"),
"KEY2": exp.DataType.build("int"),
"VALUE": exp.DataType.build("varchar"),
"OTHER": exp.DataType.build("varchar"),
}
engine_adapter.create_table("src", columns_to_types)
engine_adapter.create_table("target", columns_to_types)

src_df = pd.DataFrame(
[(1, 1, "a", "x"), (2, 2, "b", "y"), (3, 3, "src only", "z")],
columns=list(columns_to_types),
)
target_df = pd.DataFrame(
[(1, 1, "a", "x"), (2, 2, "b modified", "y2"), (4, 4, "target only", "z")],
columns=list(columns_to_types),
)
engine_adapter.insert_append("src", src_df)
engine_adapter.insert_append("target", target_df)
return engine_adapter


@pytest.mark.parametrize(
"on",
[
["KEY1"],
["key1"],
["KEY1", "KEY2"],
["key1", "KEY2"],
exp.condition("s.KEY1 = t.KEY1 AND s.KEY2 = t.KEY2"),
],
)
def test_data_diff_non_lowercase_key_columns(on):
engine_adapter = _create_uppercase_tables()

diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=on).row_diff()

assert diff.join_count == 2
assert diff.s_only_count == 1
assert diff.t_only_count == 1
assert diff.full_match_count == 1
assert diff.partial_match_count == 1
assert diff.s_sample["VALUE"].tolist() == ["src only"]
assert diff.t_sample["VALUE"].tolist() == ["target only"]


def test_data_diff_non_lowercase_skip_columns():
engine_adapter = _create_uppercase_tables()

diff = TableDiff(
adapter=engine_adapter,
source="src",
target="target",
on=["KEY1", "KEY2"],
skip_columns=["OTHER"],
).row_diff()

assert "OTHER" not in diff.s_sample.columns
assert "OTHER" not in diff.t_sample.columns
assert diff.join_count == 2
assert diff.partial_match_count == 1

# Skipping a non-existent column is still a no-op
diff = TableDiff(
adapter=engine_adapter,
source="src",
target="target",
on=["KEY1"],
skip_columns=["other", "does_not_exist"],
).row_diff()
assert "OTHER" not in diff.s_sample.columns


def test_data_diff_key_column_does_not_exist():
engine_adapter = _create_uppercase_tables()

with pytest.raises(SQLMeshError, match="missing_key"):
TableDiff(
adapter=engine_adapter, source="src", target="target", on=["missing_key"]
).row_diff()


def test_data_diff_key_column_exact_match_preferred():
engine_adapter = _create_uppercase_tables()
table_diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=["KEY1"])

schema = {
"key1": exp.DataType.build("int"),
"KEY1": exp.DataType.build("int"),
"Key1": exp.DataType.build("int"),
}
assert table_diff._resolve_column_name("key1", schema) == "key1"
assert table_diff._resolve_column_name("KEY1", schema) == "KEY1"
with pytest.raises(SQLMeshError, match="ambiguous"):
table_diff._resolve_column_name("kEy1", schema)


def test_data_diff_key_columns_different_casing_between_tables():
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()

engine_adapter.create_table(
"src", {"KEY1": exp.DataType.build("int"), "VALUE": exp.DataType.build("varchar")}
)
engine_adapter.create_table(
"target", {"key1": exp.DataType.build("int"), "VALUE": exp.DataType.build("varchar")}
)
engine_adapter.insert_append(
"src", pd.DataFrame([(1, "a"), (2, "b")], columns=["KEY1", "VALUE"])
)
engine_adapter.insert_append(
"target", pd.DataFrame([(1, "a"), (3, "c")], columns=["key1", "VALUE"])
)

for on in (["KEY1"], exp.condition("s.KEY1 = t.key1")):
diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=on).row_diff()
assert diff.join_count == 1
assert diff.s_only_count == 1
assert diff.t_only_count == 1


def test_data_diff_on_expression_not_mutated():
engine_adapter = _create_uppercase_tables()
on = exp.condition("s.KEY1 = t.KEY1")
expected_sql = on.sql()

TableDiff(adapter=engine_adapter, source="src", target="target", on=on).row_diff()

assert on.sql() == expected_sql


def test_data_diff_skip_columns_resolution():
engine_adapter = _create_uppercase_tables()
engine_adapter.create_table(
"extra", {"KEY1": exp.DataType.build("int"), "EXTRA": exp.DataType.build("int")}
)

# A column that only exists in one of the tables is skipped in that table
table_diff = TableDiff(
adapter=engine_adapter,
source="src",
target="extra",
on=["KEY1"],
skip_columns=["other", "extra"],
)
assert table_diff.skip_columns == {"OTHER", "EXTRA"}

# Ambiguous names are reported instead of being silently ignored
ambiguous = {"OTHER": exp.DataType.build("int"), "Other": exp.DataType.build("int")}
with pytest.raises(SQLMeshError, match="ambiguous"):
table_diff._find_column_name("other", ambiguous)


def test_data_diff_generated_sql_uses_resolved_column_names(mocker: MockerFixture):
engine_adapter = _create_uppercase_tables()
spy_execute = mocker.spy(engine_adapter, "_execute")

TableDiff(adapter=engine_adapter, source="src", target="target", on=["key1", "key2"]).row_diff()

executed = [str(call.args[0]) for call in spy_execute.call_args_list]
assert any('"s"."KEY1"' in sql and '"t"."KEY2"' in sql for sql in executed)
assert not any('"s"."key1"' in sql for sql in executed)