Skip to content

Commit 8162c7a

Browse files
committed
fix(table_diff): resolve key and skip column names against engine-reported casing
User-supplied `on` and `skip_columns` names are normalized with the connection dialect, which lowercases them on engines like BigQuery and DuckDB, while `adapter.columns()` reports the stored casing. Exact lookups then failed with a KeyError for keys and silently ignored skip columns. Resolve each name against the source and target schemas: an exact match wins, otherwise a unique case-insensitive match is used. Ambiguous names and missing key columns raise a clear SQLMeshError. Also stop mutating the caller's `on` expression. Fixes #6067 Signed-off-by: Akshay Chame <akshaychame2@gmail.com>
1 parent 271b855 commit 8162c7a

2 files changed

Lines changed: 223 additions & 11 deletions

File tree

‎sqlmesh/core/table_diff.py‎

Lines changed: 57 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -256,7 +256,7 @@ def __init__(
256256
self.target_alias = target_alias
257257

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

278+
@cached_property
279+
def skip_columns(self) -> t.Set[str]:
280+
"""The names of the columns to skip, as reported by the engine for either table.
281+
282+
Names that don't exist in a table are ignored for that table.
283+
"""
284+
skipped = set()
285+
for name in self._requested_skip_columns:
286+
for schema in (self.source_schema, self.target_schema):
287+
if (resolved := self._find_column_name(name, schema)) is not None:
288+
skipped.add(resolved)
289+
return skipped
290+
291+
def _find_column_name(self, name: str, schema: t.Dict[str, exp.DataType]) -> t.Optional[str]:
292+
"""Maps a user-supplied column name to the column name reported by the engine.
293+
294+
User-supplied names are normalized using the dialect, which may change their casing
295+
compared to the casing the engine reports. An exact match always wins, otherwise a
296+
case-insensitive match is used as long as it is unambiguous.
297+
298+
Returns None if there is no match and raises if the match is ambiguous.
299+
"""
300+
if name in schema:
301+
return name
302+
303+
matches = [c for c in schema if c.casefold() == name.casefold()]
304+
if len(matches) > 1:
305+
raise SQLMeshError(
306+
f"Column '{name}' is ambiguous, it matches multiple columns: {', '.join(matches)}"
307+
)
308+
return matches[0] if matches else None
309+
310+
def _resolve_column_name(self, name: str, schema: t.Dict[str, exp.DataType]) -> str:
311+
resolved = self._find_column_name(name, schema)
312+
if resolved is None:
313+
raise SQLMeshError(
314+
f"Column '{name}' does not exist. Available columns: {', '.join(schema)}"
315+
)
316+
return resolved
317+
278318
@cached_property
279319
def key_columns(self) -> t.Tuple[t.List[exp.Column], t.List[exp.Column], t.List[str]]:
280320
dialect = self.model_dialect or self.dialect
281321

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

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

294-
normalize_identifiers(self._on, dialect=dialect)
295-
for col in self._on.find_all(exp.Column):
337+
# Work on a copy so the caller's expression isn't modified
338+
on = normalize_identifiers(self._on.copy(), dialect=dialect)
339+
for col in on.find_all(exp.Column):
340+
table = col.table.lower()
341+
if table in ("s", "t"):
342+
schema = self.source_schema if table == "s" else self.target_schema
343+
col.set("this", exp.to_identifier(self._resolve_column_name(col.name, schema)))
344+
(s_index if table == "s" else t_index).append(col)
296345
index_cols.append(col.name)
297-
if col.table.lower() == "s":
298-
s_index.append(col)
299-
elif col.table.lower() == "t":
300-
t_index.append(col)
301346

347+
# Like the list form above, index_cols can contain both source and target spellings
302348
index_cols = list(dict.fromkeys(index_cols))
303349
s_index = list(dict.fromkeys(s_index))
304350
t_index = list(dict.fromkeys(t_index))

‎tests/core/test_table_diff.py‎

Lines changed: 166 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1246,3 +1246,169 @@ def test_data_diff_nulls_in_some_grain_columns():
12461246
"null value",
12471247
"null value modified",
12481248
]
1249+
1250+
1251+
def _create_uppercase_tables() -> t.Any:
1252+
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()
1253+
1254+
columns_to_types = {
1255+
"KEY1": exp.DataType.build("int"),
1256+
"KEY2": exp.DataType.build("int"),
1257+
"VALUE": exp.DataType.build("varchar"),
1258+
"OTHER": exp.DataType.build("varchar"),
1259+
}
1260+
engine_adapter.create_table("src", columns_to_types)
1261+
engine_adapter.create_table("target", columns_to_types)
1262+
1263+
src_df = pd.DataFrame(
1264+
[(1, 1, "a", "x"), (2, 2, "b", "y"), (3, 3, "src only", "z")],
1265+
columns=list(columns_to_types),
1266+
)
1267+
target_df = pd.DataFrame(
1268+
[(1, 1, "a", "x"), (2, 2, "b modified", "y2"), (4, 4, "target only", "z")],
1269+
columns=list(columns_to_types),
1270+
)
1271+
engine_adapter.insert_append("src", src_df)
1272+
engine_adapter.insert_append("target", target_df)
1273+
return engine_adapter
1274+
1275+
1276+
@pytest.mark.parametrize(
1277+
"on",
1278+
[
1279+
["KEY1"],
1280+
["key1"],
1281+
["KEY1", "KEY2"],
1282+
["key1", "KEY2"],
1283+
exp.condition("s.KEY1 = t.KEY1 AND s.KEY2 = t.KEY2"),
1284+
],
1285+
)
1286+
def test_data_diff_non_lowercase_key_columns(on):
1287+
engine_adapter = _create_uppercase_tables()
1288+
1289+
diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=on).row_diff()
1290+
1291+
assert diff.join_count == 2
1292+
assert diff.s_only_count == 1
1293+
assert diff.t_only_count == 1
1294+
assert diff.full_match_count == 1
1295+
assert diff.partial_match_count == 1
1296+
assert diff.s_sample["VALUE"].tolist() == ["src only"]
1297+
assert diff.t_sample["VALUE"].tolist() == ["target only"]
1298+
1299+
1300+
def test_data_diff_non_lowercase_skip_columns():
1301+
engine_adapter = _create_uppercase_tables()
1302+
1303+
diff = TableDiff(
1304+
adapter=engine_adapter,
1305+
source="src",
1306+
target="target",
1307+
on=["KEY1", "KEY2"],
1308+
skip_columns=["OTHER"],
1309+
).row_diff()
1310+
1311+
assert "OTHER" not in diff.s_sample.columns
1312+
assert "OTHER" not in diff.t_sample.columns
1313+
assert diff.join_count == 2
1314+
assert diff.partial_match_count == 1
1315+
1316+
# Skipping a non-existent column is still a no-op
1317+
diff = TableDiff(
1318+
adapter=engine_adapter,
1319+
source="src",
1320+
target="target",
1321+
on=["KEY1"],
1322+
skip_columns=["other", "does_not_exist"],
1323+
).row_diff()
1324+
assert "OTHER" not in diff.s_sample.columns
1325+
1326+
1327+
def test_data_diff_key_column_does_not_exist():
1328+
engine_adapter = _create_uppercase_tables()
1329+
1330+
with pytest.raises(SQLMeshError, match="missing_key"):
1331+
TableDiff(
1332+
adapter=engine_adapter, source="src", target="target", on=["missing_key"]
1333+
).row_diff()
1334+
1335+
1336+
def test_data_diff_key_column_exact_match_preferred():
1337+
engine_adapter = _create_uppercase_tables()
1338+
table_diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=["KEY1"])
1339+
1340+
schema = {
1341+
"key1": exp.DataType.build("int"),
1342+
"KEY1": exp.DataType.build("int"),
1343+
"Key1": exp.DataType.build("int"),
1344+
}
1345+
assert table_diff._resolve_column_name("key1", schema) == "key1"
1346+
assert table_diff._resolve_column_name("KEY1", schema) == "KEY1"
1347+
with pytest.raises(SQLMeshError, match="ambiguous"):
1348+
table_diff._resolve_column_name("kEy1", schema)
1349+
1350+
1351+
def test_data_diff_key_columns_different_casing_between_tables():
1352+
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()
1353+
1354+
engine_adapter.create_table(
1355+
"src", {"KEY1": exp.DataType.build("int"), "VALUE": exp.DataType.build("varchar")}
1356+
)
1357+
engine_adapter.create_table(
1358+
"target", {"key1": exp.DataType.build("int"), "VALUE": exp.DataType.build("varchar")}
1359+
)
1360+
engine_adapter.insert_append(
1361+
"src", pd.DataFrame([(1, "a"), (2, "b")], columns=["KEY1", "VALUE"])
1362+
)
1363+
engine_adapter.insert_append(
1364+
"target", pd.DataFrame([(1, "a"), (3, "c")], columns=["key1", "VALUE"])
1365+
)
1366+
1367+
for on in (["KEY1"], exp.condition("s.KEY1 = t.key1")):
1368+
diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=on).row_diff()
1369+
assert diff.join_count == 1
1370+
assert diff.s_only_count == 1
1371+
assert diff.t_only_count == 1
1372+
1373+
1374+
def test_data_diff_on_expression_not_mutated():
1375+
engine_adapter = _create_uppercase_tables()
1376+
on = exp.condition("s.KEY1 = t.KEY1")
1377+
expected_sql = on.sql()
1378+
1379+
TableDiff(adapter=engine_adapter, source="src", target="target", on=on).row_diff()
1380+
1381+
assert on.sql() == expected_sql
1382+
1383+
1384+
def test_data_diff_skip_columns_resolution():
1385+
engine_adapter = _create_uppercase_tables()
1386+
engine_adapter.create_table(
1387+
"extra", {"KEY1": exp.DataType.build("int"), "EXTRA": exp.DataType.build("int")}
1388+
)
1389+
1390+
# A column that only exists in one of the tables is skipped in that table
1391+
table_diff = TableDiff(
1392+
adapter=engine_adapter,
1393+
source="src",
1394+
target="extra",
1395+
on=["KEY1"],
1396+
skip_columns=["other", "extra"],
1397+
)
1398+
assert table_diff.skip_columns == {"OTHER", "EXTRA"}
1399+
1400+
# Ambiguous names are reported instead of being silently ignored
1401+
ambiguous = {"OTHER": exp.DataType.build("int"), "Other": exp.DataType.build("int")}
1402+
with pytest.raises(SQLMeshError, match="ambiguous"):
1403+
table_diff._find_column_name("other", ambiguous)
1404+
1405+
1406+
def test_data_diff_generated_sql_uses_resolved_column_names(mocker: MockerFixture):
1407+
engine_adapter = _create_uppercase_tables()
1408+
spy_execute = mocker.spy(engine_adapter, "_execute")
1409+
1410+
TableDiff(adapter=engine_adapter, source="src", target="target", on=["key1", "key2"]).row_diff()
1411+
1412+
executed = [str(call.args[0]) for call in spy_execute.call_args_list]
1413+
assert any('"s"."KEY1"' in sql and '"t"."KEY2"' in sql for sql in executed)
1414+
assert not any('"s"."key1"' in sql for sql in executed)

0 commit comments

Comments
 (0)