Skip to content

Commit 3ce6bfa

Browse files
committed
fix(table_diff): support key columns stored with non-normalized casing
fixed tests Signed-off-by: Anant <75747269+Anant-gif@users.noreply.github.com>
1 parent e30fe61 commit 3ce6bfa

2 files changed

Lines changed: 168 additions & 4 deletions

File tree

‎sqlmesh/core/table_diff.py‎

Lines changed: 15 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -282,9 +282,11 @@ def key_columns(self) -> t.Tuple[t.List[exp.Column], t.List[exp.Column], t.List[
282282
# If the columns to join on are explicitly specified, then just return them
283283
if isinstance(self._on, (list, tuple)):
284284
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]
285+
s_names = [self._resolve_column_name(c.name, self.source_schema) for c in identifiers]
286+
t_names = [self._resolve_column_name(c.name, self.target_schema) for c in identifiers]
287+
s_index = [exp.column(c, "s") for c in s_names]
288+
t_index = [exp.column(c, "t") for c in t_names]
289+
return s_index, t_index, s_names
288290

289291
# Otherwise, we need to parse them out of the supplied "on" condition
290292
index_cols = []
@@ -293,18 +295,27 @@ def key_columns(self) -> t.Tuple[t.List[exp.Column], t.List[exp.Column], t.List[
293295

294296
normalize_identifiers(self._on, dialect=dialect)
295297
for col in self._on.find_all(exp.Column):
296-
index_cols.append(col.name)
297298
if col.table.lower() == "s":
299+
col = exp.column(self._resolve_column_name(col.name, self.source_schema), col.table)
298300
s_index.append(col)
299301
elif col.table.lower() == "t":
302+
col = exp.column(self._resolve_column_name(col.name, self.target_schema), col.table)
300303
t_index.append(col)
304+
index_cols.append(col.name)
301305

302306
index_cols = list(dict.fromkeys(index_cols))
303307
s_index = list(dict.fromkeys(s_index))
304308
t_index = list(dict.fromkeys(t_index))
305309

306310
return s_index, t_index, index_cols
307311

312+
def _resolve_column_name(self, name: str, schema: t.Dict[str, exp.DataType]) -> str:
313+
if name in schema:
314+
return name
315+
316+
lowercase_name = name.lower()
317+
return next((c for c in schema if c.lower() == lowercase_name), name)
318+
308319
@property
309320
def source_key_expression(self) -> exp.Expr:
310321
s_index, _, _ = self.key_columns

‎tests/core/test_table_diff.py‎

Lines changed: 153 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1197,6 +1197,159 @@ def test_data_diff_sample_limit():
11971197
assert len(diff.joined_sample) == 3
11981198

11991199

1200+
def test_data_diff_non_lowercase_key_columns():
1201+
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()
1202+
1203+
columns_to_types = {
1204+
"KEY1": exp.DataType.build("int"),
1205+
"Key2": exp.DataType.build("varchar"),
1206+
"VALUE": exp.DataType.build("varchar"),
1207+
}
1208+
1209+
engine_adapter.create_table("src", columns_to_types)
1210+
engine_adapter.create_table("target", columns_to_types)
1211+
1212+
src_records = [
1213+
(1, "a", "value"),
1214+
(2, "b", "source"),
1215+
(3, "c", "source only"),
1216+
]
1217+
1218+
target_records = [
1219+
(1, "a", "value"),
1220+
(2, "b", "target"),
1221+
(4, "d", "target only"),
1222+
]
1223+
1224+
src_df = pd.DataFrame(data=src_records, columns=columns_to_types.keys())
1225+
target_df = pd.DataFrame(data=target_records, columns=columns_to_types.keys())
1226+
1227+
engine_adapter.insert_append("src", src_df)
1228+
engine_adapter.insert_append("target", target_df)
1229+
1230+
# casing of the supplied key should not matter
1231+
for on in (["KEY1", "Key2"], ["key1", "KEY2"]):
1232+
table_diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=on)
1233+
1234+
_, _, col_names = table_diff.key_columns
1235+
assert col_names == ["KEY1", "Key2"]
1236+
1237+
diff = table_diff.row_diff()
1238+
1239+
assert diff.join_count == 2
1240+
assert diff.full_match_count == 1
1241+
assert diff.partial_match_count == 1
1242+
assert diff.s_only_count == 1
1243+
assert diff.t_only_count == 1
1244+
1245+
table_diff = TableDiff(adapter=engine_adapter, source="src", target="target", on=["KEY1"])
1246+
1247+
_, _, col_names = table_diff.key_columns
1248+
assert col_names == ["KEY1"]
1249+
1250+
diff = table_diff.row_diff()
1251+
1252+
assert diff.join_count == 2
1253+
assert diff.full_match_count == 1
1254+
assert diff.partial_match_count == 1
1255+
assert diff.s_only_count == 1
1256+
assert diff.t_only_count == 1
1257+
1258+
1259+
def test_data_diff_key_columns_with_differing_case_between_source_and_target():
1260+
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()
1261+
1262+
source_columns_to_types = {
1263+
"KEY1": exp.DataType.build("int"),
1264+
"Key2": exp.DataType.build("varchar"),
1265+
"value": exp.DataType.build("varchar"),
1266+
}
1267+
target_columns_to_types = {
1268+
"key1": exp.DataType.build("int"),
1269+
"KEY2": exp.DataType.build("varchar"),
1270+
"value": exp.DataType.build("varchar"),
1271+
}
1272+
1273+
engine_adapter.create_table("src", source_columns_to_types)
1274+
engine_adapter.create_table("target", target_columns_to_types)
1275+
1276+
engine_adapter.insert_append(
1277+
"src",
1278+
pd.DataFrame(
1279+
data=[(1, "a", "value"), (2, "b", "source")],
1280+
columns=source_columns_to_types.keys(),
1281+
),
1282+
)
1283+
engine_adapter.insert_append(
1284+
"target",
1285+
pd.DataFrame(
1286+
data=[(1, "a", "value"), (2, "b", "target")],
1287+
columns=target_columns_to_types.keys(),
1288+
),
1289+
)
1290+
1291+
table_diff = TableDiff(
1292+
adapter=engine_adapter, source="src", target="target", on=["key1", "KEY2"]
1293+
)
1294+
1295+
s_index, t_index, col_names = table_diff.key_columns
1296+
assert [c.sql() for c in s_index] == ["s.KEY1", "s.Key2"]
1297+
assert [c.sql() for c in t_index] == ["t.key1", "t.KEY2"]
1298+
assert col_names == ["KEY1", "Key2"]
1299+
1300+
diff = table_diff.row_diff()
1301+
1302+
assert diff.join_count == 2
1303+
assert diff.full_match_count == 1
1304+
assert diff.partial_match_count == 1
1305+
assert diff.s_only_count == 0
1306+
assert diff.t_only_count == 0
1307+
1308+
# the key columns are excluded from the per column match stats
1309+
assert diff.column_stats.index.tolist() == ["value"]
1310+
1311+
1312+
def test_data_diff_non_lowercase_key_columns_in_on_condition():
1313+
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()
1314+
1315+
columns_to_types = {
1316+
"KEY1": exp.DataType.build("int"),
1317+
"Key2": exp.DataType.build("varchar"),
1318+
"VALUE": exp.DataType.build("varchar"),
1319+
}
1320+
1321+
engine_adapter.create_table("src", columns_to_types)
1322+
engine_adapter.create_table("target", columns_to_types)
1323+
1324+
src_df = pd.DataFrame(
1325+
data=[(1, "a", "value"), (2, "b", "source")], columns=columns_to_types.keys()
1326+
)
1327+
target_df = pd.DataFrame(
1328+
data=[(1, "a", "value"), (2, "b", "target")], columns=columns_to_types.keys()
1329+
)
1330+
1331+
engine_adapter.insert_append("src", src_df)
1332+
engine_adapter.insert_append("target", target_df)
1333+
1334+
table_diff = TableDiff(
1335+
adapter=engine_adapter,
1336+
source="src",
1337+
target="target",
1338+
on=exp.condition('s."KEY1" = t."KEY1" AND s."Key2" = t."Key2"'),
1339+
)
1340+
1341+
_, _, col_names = table_diff.key_columns
1342+
assert col_names == ["KEY1", "Key2"]
1343+
1344+
diff = table_diff.row_diff()
1345+
1346+
assert diff.join_count == 2
1347+
assert diff.full_match_count == 1
1348+
assert diff.partial_match_count == 1
1349+
assert diff.s_only_count == 0
1350+
assert diff.t_only_count == 0
1351+
1352+
12001353
def test_data_diff_nulls_in_some_grain_columns():
12011354
engine_adapter = DuckDBConnectionConfig().create_engine_adapter()
12021355

0 commit comments

Comments
 (0)