From 8162c7ad9b8a2cf8fccbb8227d3076b3d8a63581 Mon Sep 17 00:00:00 2001 From: Akshay Chame Date: Thu, 1 Oct 2026 12:13:01 +0530 Subject: [PATCH] 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 --- sqlmesh/core/table_diff.py | 68 +++++++++++--- tests/core/test_table_diff.py | 166 ++++++++++++++++++++++++++++++++++ 2 files changed, 223 insertions(+), 11 deletions(-) diff --git a/sqlmesh/core/table_diff.py b/sqlmesh/core/table_diff.py index 97cb0c19ba..30619e0d3a 100644 --- a/sqlmesh/core/table_diff.py +++ b/sqlmesh/core/table_diff.py @@ -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, @@ -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)) diff --git a/tests/core/test_table_diff.py b/tests/core/test_table_diff.py index c2e293e4c2..a2a49431c3 100644 --- a/tests/core/test_table_diff.py +++ b/tests/core/test_table_diff.py @@ -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)