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
11 changes: 9 additions & 2 deletions sqlmesh/core/engine_adapter/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -2302,7 +2309,7 @@ def merge(
]
),
),
)
),
)
for source_query in source_queries:
with source_query as query:
Expand Down
48 changes: 48 additions & 0 deletions tests/core/engine_adapter/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down