diff --git a/sqlmesh/core/engine_adapter/base.py b/sqlmesh/core/engine_adapter/base.py index 930cdf7cd4..fb873c6ca9 100644 --- a/sqlmesh/core/engine_adapter/base.py +++ b/sqlmesh/core/engine_adapter/base.py @@ -2288,7 +2288,14 @@ def merge( else: match_expressions = when_matched.copy().expressions - match_expressions.append( + # Engines like Databricks require WHEN NOT MATCHED [BY TARGET] clauses to come before + # any WHEN NOT MATCHED BY SOURCE clause, so the insert goes in front of those + insert_index = next( + (i for i, when in enumerate(match_expressions) if when.args.get("source")), + len(match_expressions), + ) + match_expressions.insert( + insert_index, exp.When( matched=False, source=False, @@ -2302,7 +2309,7 @@ def merge( ] ), ), - ) + ), ) for source_query in source_queries: with source_query as query: diff --git a/tests/core/engine_adapter/test_base.py b/tests/core/engine_adapter/test_base.py index 1971ba3bbc..07f6b9a2a8 100644 --- a/tests/core/engine_adapter/test_base.py +++ b/tests/core/engine_adapter/test_base.py @@ -1346,6 +1346,54 @@ def test_merge_when_matched_multiple(make_mocked_engine_adapter: t.Callable, ass ) +def test_merge_when_not_matched_by_source(make_mocked_engine_adapter: t.Callable, assert_exp_eq): + adapter = make_mocked_engine_adapter(EngineAdapter) + + adapter.merge( + target_table="target", + source_table=t.cast(exp.Select, parse_one('SELECT "ID", val FROM source')), + target_columns_to_types={ + "ID": exp.DataType.build("int"), + "val": exp.DataType.build("int"), + }, + unique_key=[exp.to_identifier("ID", quoted=True)], + when_matched=exp.Whens( + expressions=[ + exp.When( + matched=True, + source=False, + then=exp.Update( + expressions=[ + exp.column("val", "__MERGE_TARGET__").eq( + exp.column("val", "__MERGE_SOURCE__") + ), + ], + ), + ), + exp.When(matched=False, source=True, then=exp.Delete()), + ] + ), + ) + + # the generated WHEN NOT MATCHED clause has to come before WHEN NOT MATCHED BY SOURCE + assert_exp_eq( + adapter.cursor.execute.call_args[0][0], + """ +MERGE INTO "target" AS "__MERGE_TARGET__" USING ( + SELECT + "ID", + "val" + FROM "source" +) AS "__MERGE_SOURCE__" + ON "__MERGE_TARGET__"."ID" = "__MERGE_SOURCE__"."ID" + WHEN MATCHED THEN UPDATE SET "__MERGE_TARGET__"."val" = "__MERGE_SOURCE__"."val" + WHEN NOT MATCHED THEN INSERT ("ID", "val") + VALUES ("__MERGE_SOURCE__"."ID", "__MERGE_SOURCE__"."val") + WHEN NOT MATCHED BY SOURCE THEN DELETE +""", + ) + + def test_merge_filter(make_mocked_engine_adapter: t.Callable, assert_exp_eq): adapter = make_mocked_engine_adapter(EngineAdapter)