Skip to main content

uqa_sql/semantics/
conflict.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Unique-index inference, target validation, and SQL predicate implication.
8use crate::{
9    ast::{BinaryOp, ColumnDef, TableConstraintSet},
10    binding::snapshot::BindingSnapshot,
11    catalog::index::EnforcedKey,
12    plan::AggregateClassifier,
13    plan::{ConflictPlan, ExpressionPlan, InsertPlan},
14    routines::RoutineResolution,
15    RowSchema, SQLError, SQLParam, ScalarExpr as Expr,
16};
17use std::collections::BTreeSet;
18use uqa_core::Value;
19
20pub trait ConflictCatalog {
21    fn try_describe_table(&self, table: &str) -> Result<Option<Vec<ColumnDef>>, String>;
22    fn enforced_keys(&self, table: &str) -> Result<Vec<EnforcedKey>, String>;
23    fn try_declared_table_constraints(&self, table: &str) -> Result<TableConstraintSet, String>;
24}
25/// Capture the stored-expression catalog scope only when an inference expression requires binding.
26pub trait InferenceBindingScope {
27    fn binding_scope(&self) -> Result<BindingSnapshot, SQLError>;
28}
29#[derive(Clone, Copy)]
30pub struct InferenceContext<'a> {
31    pub catalog: &'a dyn ConflictCatalog,
32    pub aggregates: &'a dyn AggregateClassifier,
33    pub routines: &'a dyn RoutineResolution,
34    pub binding: &'a dyn InferenceBindingScope,
35}
36
37/// Analyze the inference clause in the INSERT target's scope before uniqueness arbitration, including commands that produce no input rows.
38pub fn prepare_inference_predicate<'a>(
39    context: InferenceContext<'_>,
40    statement: &'a InsertPlan,
41    params: &[SQLParam],
42) -> Result<std::borrow::Cow<'a, InsertPlan>, SQLError> {
43    if statement
44        .on_conflict
45        .as_ref()
46        .is_none_or(|conflict| conflict.predicate.is_none() && conflict.expressions.is_empty())
47    {
48        return Ok(std::borrow::Cow::Borrowed(statement));
49    }
50    let columns = context
51        .catalog
52        .try_describe_table(&statement.table)
53        .map_err(SQLError::Internal)?
54        .ok_or_else(|| SQLError::UnknownTable(statement.table.clone()))?;
55    let schema = RowSchema::with_qualified_types(
56        &statement.target_qualifier,
57        columns.iter().map(|column| column.name.clone()).collect(),
58        columns
59            .iter()
60            .map(|column| Some(column.ty.clone()))
61            .collect(),
62    );
63    let mut statement = statement.clone();
64    if let Some(conflict) = &mut statement.on_conflict {
65        for expression in conflict
66            .expressions
67            .iter_mut()
68            .chain(conflict.predicate.iter_mut().map(Box::as_mut))
69        {
70            prepare_inference_expression(
71                context,
72                expression,
73                &statement.target_qualifier,
74                &schema,
75                &columns,
76                params,
77            )?;
78        }
79        // A rewritten view expression can resolve to a simple base-table attribute.
80        let expressions = std::mem::take(&mut conflict.expressions);
81        for expression in expressions {
82            if let Expr::Column(column) = expression {
83                conflict.conflict_columns.push(column);
84            } else {
85                conflict.expressions.push(expression);
86            }
87        }
88    }
89    Ok(std::borrow::Cow::Owned(statement))
90}
91
92fn prepare_inference_expression(
93    context: InferenceContext<'_>,
94    expression: &mut Expr,
95    qualifier: &str,
96    schema: &RowSchema,
97    columns: &[crate::ast::ColumnDef],
98    params: &[SQLParam],
99) -> Result<(), SQLError> {
100    let mut has_subquery = false;
101    expression.visit(&mut |part| {
102        has_subquery |= matches!(
103            part,
104            Expr::ScalarSubquery(_) | Expr::Exists { .. } | Expr::InSubquery { .. }
105        );
106    });
107    if has_subquery {
108        return Err(predicate_error(
109            "0A000",
110            "cannot use subquery in index inference",
111        ));
112    }
113    if crate::semantics::aggregates::contains_aggregate(context.aggregates, expression) {
114        return Err(predicate_error(
115            "42803",
116            "aggregate functions are not allowed in index inference",
117        ));
118    }
119    if crate::semantics::windows::expr_has_window(expression) {
120        return Err(predicate_error(
121            "42P20",
122            "window functions are not allowed in index inference",
123        ));
124    }
125    let mut plan = ExpressionPlan {
126        scalar: expression.clone(),
127        subqueries: Vec::new(),
128    };
129    let binding = context.binding.binding_scope()?;
130    crate::binding::bind_expression_plan_routines_for_storage(
131        context.routines,
132        &mut plan,
133        params,
134        &binding.context(),
135        schema,
136    )?;
137    crate::plan::rewrite_scalar_expression(&mut plan.scalar, &mut |expression| {
138        if let Expr::QualifiedColumn {
139            qualifier: source,
140            column,
141        } = expression
142        {
143            if source == qualifier {
144                *expression = Expr::Column(column.clone());
145            }
146        }
147    });
148    crate::plan::rewrite_scalar_expression(&mut plan.scalar, &mut |expression| {
149        if let Expr::Cast { expr, ty, .. } = expression {
150            if let Expr::Column(name) = expr.as_ref() {
151                if columns.iter().any(|column| {
152                    column.name == *name
153                        && crate::ast::ColumnType::from_sql_name(ty).ok().as_ref()
154                            == Some(&column.ty)
155                }) {
156                    *expression = Expr::Column(name.clone());
157                }
158            }
159        }
160    });
161    *expression = inference_identity(plan.scalar);
162    Ok(())
163}
164
165fn inference_identity(mut expression: Expr) -> Expr {
166    // PostgreSQL expression equality ignores coercion display form, but retains the conversion itself and its destination type.
167    crate::plan::rewrite_scalar_expression(&mut expression, &mut |node| {
168        if let Expr::Cast { implicit, .. } = node {
169            *implicit = false;
170        }
171    });
172    expression
173}
174
175fn predicate_error(sqlstate: &str, message: &str) -> SQLError {
176    SQLError::Routine {
177        sqlstate: sqlstate.into(),
178        message: message.into(),
179    }
180}
181
182/// Match the inference expressions analyzed by [`prepare_inference_predicate`] against the stored unique keys.
183pub fn conflict_key_indices(
184    catalog: &dyn ConflictCatalog,
185    table: &str,
186    keys: &[EnforcedKey],
187    conflict: &ConflictPlan,
188) -> Result<Vec<usize>, SQLError> {
189    if let Some(name) = &conflict.constraint {
190        return constraint_target_index(catalog, table, keys, name);
191    }
192    if conflict.conflict_columns.is_empty() && conflict.expressions.is_empty() {
193        return Ok((0..keys.len()).collect());
194    }
195    validate_conflict_columns(catalog, table, &conflict.conflict_columns)?;
196    let target = conflict
197        .conflict_columns
198        .iter()
199        .map(String::as_str)
200        .collect::<BTreeSet<_>>();
201    let indexes = keys
202        .iter()
203        .enumerate()
204        .filter_map(|(index, key)| {
205            (key.columns
206                .iter()
207                .map(String::as_str)
208                .collect::<BTreeSet<_>>()
209                == target
210                && {
211                    let expressions = key
212                        .keys
213                        .iter()
214                        .filter_map(|key| match key {
215                            crate::ast::IndexKey::Expression(expr) => Some(inference_identity(
216                                ExpressionPlan::lower((**expr).clone()).scalar,
217                            )),
218                            crate::ast::IndexKey::Column(_) => None,
219                        })
220                        .collect::<Vec<_>>();
221                    expressions
222                        .iter()
223                        .all(|expr| conflict.expressions.contains(expr))
224                        && conflict
225                            .expressions
226                            .iter()
227                            .all(|expr| expressions.contains(expr))
228                }
229                && key.predicate.as_deref().is_none_or(|required| {
230                    conflict.predicate.as_deref().is_some_and(|given| {
231                        implies(
232                            given,
233                            &inference_identity(ExpressionPlan::lower(required.clone()).scalar),
234                        )
235                    })
236                }))
237            .then_some(index)
238        })
239        .collect::<Vec<_>>();
240    if indexes.is_empty() {
241        return Err(SQLError::Routine {
242            sqlstate: "42P10".into(),
243            message:
244                "there is no unique or exclusion constraint matching the ON CONFLICT specification"
245                    .into(),
246        });
247    }
248    Ok(indexes)
249}
250
251/// Prove that every row for which `given` is true also makes `required` true. SQL NULL is never treated as false in a proof that would admit an invalid arbiter.
252fn implies(given: &Expr, required: &Expr) -> bool {
253    if given == required || matches!(required, Expr::Literal(Value::Bool(true))) {
254        return true;
255    }
256    if matches!(given, Expr::Literal(Value::Bool(false) | Value::Null)) {
257        return true;
258    }
259    if let Expr::And(parts) = required {
260        return parts.iter().all(|part| implies(given, part));
261    }
262    if let Expr::Or(parts) = given {
263        return parts.iter().all(|part| implies(part, required));
264    }
265    if let Expr::And(parts) = given {
266        if parts.iter().any(|part| implies(part, required)) {
267            return true;
268        }
269    }
270    if let Expr::Or(parts) = required {
271        return parts.iter().any(|part| implies(given, part));
272    }
273    if let Some((left, given_op, given_value)) = comparison(given) {
274        if let Some((right, required_op, required_value)) = comparison(required) {
275            return left == right
276                && comparison_implies(given_op, given_value, required_op, required_value);
277        }
278        if let Expr::IsNull {
279            expr,
280            negated: true,
281        } = required
282        {
283            return left == expr.as_ref() && !matches!(given_value, Value::Null);
284        }
285    }
286    false
287}
288
289fn comparison(expr: &Expr) -> Option<(&Expr, BinaryOp, &Value)> {
290    let Expr::Binary { op, lhs, rhs } = expr else {
291        return None;
292    };
293    if !matches!(
294        op,
295        BinaryOp::Equal
296            | BinaryOp::NotEqual
297            | BinaryOp::Less
298            | BinaryOp::LessEqual
299            | BinaryOp::Greater
300            | BinaryOp::GreaterEqual
301    ) {
302        return None;
303    }
304    if let Expr::Literal(value) = rhs.as_ref() {
305        return Some((lhs, *op, value));
306    }
307    if let Expr::Literal(value) = lhs.as_ref() {
308        let reversed = match op {
309            BinaryOp::Less => BinaryOp::Greater,
310            BinaryOp::LessEqual => BinaryOp::GreaterEqual,
311            BinaryOp::Greater => BinaryOp::Less,
312            BinaryOp::GreaterEqual => BinaryOp::LessEqual,
313            other => *other,
314        };
315        return Some((rhs, reversed, value));
316    }
317    None
318}
319
320fn comparison_implies(given: BinaryOp, left: &Value, required: BinaryOp, right: &Value) -> bool {
321    use BinaryOp::{Equal, Greater, GreaterEqual, Less, LessEqual, NotEqual};
322    if matches!(left, Value::Null) || matches!(right, Value::Null) {
323        return false;
324    }
325    if std::mem::discriminant(left) != std::mem::discriminant(right) {
326        return false;
327    }
328    let order = left.cmp(right);
329    match (given, required) {
330        (Equal, Equal) | (NotEqual, NotEqual) => order.is_eq(),
331        (Equal, NotEqual) => !order.is_eq(),
332        (Equal | GreaterEqual, Greater) | (GreaterEqual, NotEqual) => order.is_gt(),
333        (Equal | LessEqual, Less) | (LessEqual, NotEqual) => order.is_lt(),
334        (Equal | Greater | GreaterEqual, GreaterEqual) | (Greater, Greater | NotEqual) => {
335            order.is_ge()
336        }
337        (Equal | Less | LessEqual, LessEqual) | (Less, Less | NotEqual) => order.is_le(),
338        _ => false,
339    }
340}
341
342pub fn validate_conflict_target(
343    catalog: &dyn ConflictCatalog,
344    table: &str,
345    conflict: &ConflictPlan,
346) -> Result<(), SQLError> {
347    let keys = catalog
348        .enforced_keys(table)
349        .map_err(|error| SQLError::Internal(format!("conflict target keys: {error}")))?;
350    conflict_key_indices(catalog, table, &keys, conflict).map(|_| ())
351}
352
353fn constraint_target_index(
354    catalog: &dyn ConflictCatalog,
355    table: &str,
356    keys: &[EnforcedKey],
357    name: &str,
358) -> Result<Vec<usize>, SQLError> {
359    if let Some(index) = keys
360        .iter()
361        .position(|key| key.constraint_owned && key.name.as_deref() == Some(name))
362    {
363        return Ok(vec![index]);
364    }
365    let snapshot = catalog
366        .try_declared_table_constraints(table)
367        .map_err(SQLError::Internal)?;
368    let columns = catalog
369        .try_describe_table(table)
370        .map_err(SQLError::Internal)?
371        .ok_or_else(|| SQLError::UnknownTable(table.into()))?;
372    let exists = snapshot
373        .checks
374        .iter()
375        .any(|check| check.name.as_deref() == Some(name))
376        || snapshot
377            .foreign_keys
378            .iter()
379            .any(|key| key.name.as_deref() == Some(name))
380        || columns.iter().any(|column| {
381            column.not_null_name.as_deref() == Some(name)
382                || column.check_name.as_deref() == Some(name)
383                || column
384                    .references
385                    .as_ref()
386                    .is_some_and(|reference| reference.name.as_deref() == Some(name))
387        });
388    if exists {
389        return Err(SQLError::Routine {
390            sqlstate: "42809".into(),
391            message: "constraint in ON CONFLICT clause has no associated index".into(),
392        });
393    }
394    Err(SQLError::Routine {
395        sqlstate: "42704".into(),
396        message: format!("constraint \"{name}\" for table \"{table}\" does not exist"),
397    })
398}
399
400fn validate_conflict_columns(
401    catalog: &dyn ConflictCatalog,
402    table: &str,
403    names: &[String],
404) -> Result<(), SQLError> {
405    let columns = catalog
406        .try_describe_table(table)
407        .map_err(SQLError::Internal)?
408        .ok_or_else(|| SQLError::UnknownTable(table.into()))?;
409    for name in names {
410        if !columns.iter().any(|column| column.name == *name)
411            && !matches!(
412                name.as_str(),
413                "ctid" | "tableoid" | "xmin" | "xmax" | "cmin" | "cmax"
414            )
415        {
416            return Err(SQLError::UnknownColumn(name.clone()));
417        }
418    }
419    Ok(())
420}
421
422#[cfg(test)]
423mod tests {
424    use super::*;
425
426    #[test]
427    fn inference_implication_ignores_cast_origin_but_preserves_cast_types() {
428        let predicate = |implicit, ty: &str, bound| {
429            inference_identity(Expr::Binary {
430                op: BinaryOp::Greater,
431                lhs: Box::new(Expr::Cast {
432                    implicit,
433                    expr: Box::new(Expr::Column("value".into())),
434                    ty: ty.into(),
435                }),
436                rhs: Box::new(Expr::Literal(Value::Int(bound))),
437            })
438        };
439        for implicit in [false, true] {
440            let required = predicate(implicit, "bigint", 0);
441            let equivalent = predicate(!implicit, "bigint", 0);
442            assert_eq!(required, equivalent);
443            assert!(implies(&equivalent, &required));
444            assert!(implies(&predicate(!implicit, "bigint", 1), &required));
445            assert!(!implies(&predicate(!implicit, "bigint", -1), &required));
446            assert!(!implies(&predicate(!implicit, "integer", 0), &required));
447        }
448    }
449}