1use 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}
25pub 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
37pub 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 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 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
182pub 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
251fn 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}