use super::{
compile_expr, compile_from_node, compile_returning_clause, range_var_name, Expr, NodeEnum,
Result, SQLError,
};
#[expect(
clippy::too_many_lines,
reason = "ordered PostgreSQL lowering preserves syntax and error precedence"
)]
pub(super) fn compile_merge(stmt: &pg_query::protobuf::MergeStmt) -> Result<crate::ast::MergeStmt> {
use crate::ast::{MergeStmt, MergeWhen};
use pg_query::protobuf::{CmdType, MergeMatchKind};
let relation = stmt
.relation
.as_ref()
.ok_or_else(|| SQLError::Internal("MERGE without target".into()))?;
let target = range_var_name(relation);
let include_descendants = relation.inh;
let target_alias = relation
.alias
.as_ref()
.map(|a| a.aliasname.clone())
.filter(|s| !s.is_empty());
let target_qualifier = target_alias
.clone()
.unwrap_or_else(|| relation.relname.clone());
let source_node = stmt
.source_relation
.as_deref()
.ok_or_else(|| SQLError::Internal("MERGE without USING".into()))?;
let source = compile_from_node(source_node)?;
let join_condition_node = stmt
.join_condition
.as_deref()
.ok_or_else(|| SQLError::Internal("MERGE without ON".into()))?;
let join_condition = compile_expr(join_condition_node)?;
let mut when_clauses: Vec<MergeWhen> = Vec::with_capacity(stmt.merge_when_clauses.len());
let mut final_matched = false;
let mut final_not_matched_by_source = false;
let mut final_not_matched_by_target = false;
for clause in &stmt.merge_when_clauses {
let Some(NodeEnum::MergeWhenClause(w)) = clause.node.as_ref() else {
return Err(SQLError::Internal(
"MERGE contains a malformed WHEN clause".into(),
));
};
let condition = w
.condition
.as_deref()
.map(|c| compile_expr(c))
.transpose()?;
let match_kind = w.match_kind();
let final_clause = match match_kind {
MergeMatchKind::MergeWhenMatched => &mut final_matched,
MergeMatchKind::MergeWhenNotMatchedBySource => &mut final_not_matched_by_source,
MergeMatchKind::MergeWhenNotMatchedByTarget => &mut final_not_matched_by_target,
MergeMatchKind::Undefined => {
return Err(SQLError::Internal(
"MERGE WHEN clause has no match kind".into(),
));
}
};
if *final_clause {
return Err(SQLError::Routine {
sqlstate: "42601".into(),
message: "unreachable WHEN clause specified after unconditional WHEN clause".into(),
});
}
if condition.is_none() {
*final_clause = true;
}
let cmd = w.command_type();
match cmd {
CmdType::CmdUpdate => {
if matches!(match_kind, MergeMatchKind::MergeWhenNotMatchedByTarget) {
return Err(SQLError::Internal(
"MERGE UPDATE is not valid for WHEN NOT MATCHED BY TARGET".into(),
));
}
let mut assignments: Vec<(String, Expr)> = Vec::new();
for tgt in &w.target_list {
let Some(NodeEnum::ResTarget(rt)) = tgt.node.as_ref() else {
return Err(SQLError::Internal(
"MERGE UPDATE contains a malformed assignment".into(),
));
};
let val = rt
.val
.as_ref()
.ok_or_else(|| SQLError::Internal("MERGE UPDATE without value".into()))?;
assignments.push((rt.name.clone(), compile_expr(val)?));
}
when_clauses.push(match match_kind {
MergeMatchKind::MergeWhenMatched => MergeWhen::UpdateMatched {
condition,
assignments,
},
MergeMatchKind::MergeWhenNotMatchedBySource => {
MergeWhen::UpdateNotMatchedBySource {
condition,
assignments,
}
}
MergeMatchKind::MergeWhenNotMatchedByTarget | MergeMatchKind::Undefined => {
unreachable!()
}
});
}
CmdType::CmdDelete => {
if matches!(match_kind, MergeMatchKind::MergeWhenNotMatchedByTarget) {
return Err(SQLError::Internal(
"MERGE DELETE is not valid for WHEN NOT MATCHED BY TARGET".into(),
));
}
when_clauses.push(match match_kind {
MergeMatchKind::MergeWhenMatched => MergeWhen::DeleteMatched { condition },
MergeMatchKind::MergeWhenNotMatchedBySource => {
MergeWhen::DeleteNotMatchedBySource { condition }
}
MergeMatchKind::MergeWhenNotMatchedByTarget | MergeMatchKind::Undefined => {
unreachable!()
}
});
}
CmdType::CmdInsert => {
if !matches!(match_kind, MergeMatchKind::MergeWhenNotMatchedByTarget) {
return Err(SQLError::Internal(
"MERGE INSERT is only valid for WHEN NOT MATCHED BY TARGET".into(),
));
}
let mut columns: Vec<String> = Vec::with_capacity(w.target_list.len());
for tgt in &w.target_list {
let Some(NodeEnum::ResTarget(rt)) = tgt.node.as_ref() else {
return Err(SQLError::Internal(
"MERGE INSERT contains a malformed target column".into(),
));
};
columns.push(rt.name.clone());
}
let values: Vec<Expr> = w
.values
.iter()
.map(compile_expr)
.collect::<Result<Vec<_>>>()?;
when_clauses.push(MergeWhen::InsertNotMatched {
condition,
columns,
values,
});
}
CmdType::CmdNothing => {
when_clauses.push(match match_kind {
MergeMatchKind::MergeWhenMatched => MergeWhen::NothingMatched { condition },
MergeMatchKind::MergeWhenNotMatchedBySource => {
MergeWhen::NothingNotMatchedBySource { condition }
}
MergeMatchKind::MergeWhenNotMatchedByTarget => {
MergeWhen::NothingNotMatched { condition }
}
MergeMatchKind::Undefined => unreachable!(),
});
}
other => {
return Err(SQLError::Unsupported(format!(
"MERGE WHEN command {other:?}"
)));
}
}
}
let (returning, returning_aliases) = compile_returning_clause(stmt.returning_clause.as_ref())?;
Ok(MergeStmt {
target,
target_qualifier,
target_alias,
include_descendants,
source,
join_condition,
when_clauses,
returning,
returning_aliases,
})
}