1use crate::ast::{ColumnDef, ColumnType, Expr, TableCheck};
10use crate::SQLError;
11use crate::ScalarExpr;
12use uqa_core::Value;
13
14pub fn bind_parent_check_columns(parent: &str, expr: &mut Expr) -> Result<(), SQLError> {
15 let relation =
16 uqa_core::RelationIdentity::from_legacy_name(parent).map_err(SQLError::Internal)?;
17 crate::schema::generated::bind_schema_column_references(expr, parent);
18 crate::schema::generated::bind_schema_column_references(expr, &relation.name);
19 Ok(())
20}
21
22pub(super) fn remove_identity_casts(
24 expression: &mut Expr,
25 columns: &[ColumnDef],
26) -> Result<(), SQLError> {
27 crate::catalog::stored_ast::visit_stored_expression(expression, &mut |node| {
28 while let Expr::Cast { expr, ty, .. } = node {
29 let Ok(target) = ColumnType::from_sql_name(ty) else {
30 break;
31 };
32 let source = match expr.as_ref() {
33 Expr::Column(name) => columns
34 .iter()
35 .find(|column| column.name == *name)
36 .map(|column| column.ty.clone()),
37 Expr::TypedLiteral { ty, .. } => ColumnType::from_sql_name(ty).ok(),
38 Expr::Literal(_) => crate::scalar_type(
39 &crate::plan::ExpressionPlan::lower(*expr.clone()).scalar,
40 &crate::RowSchema::default(),
41 &[],
42 )?,
43 _ => None,
44 };
45 if source.as_ref() != Some(&target) {
46 break;
47 }
48 *node = *expr.clone();
49 }
50 Ok(())
51 })
52}
53
54pub fn same_check_expression(
55 left: &Expr,
56 right: &Expr,
57 columns: &[ColumnDef],
58) -> Result<bool, SQLError> {
59 fn canonical(expression: &Expr, columns: &[ColumnDef]) -> Result<ScalarExpr, SQLError> {
60 let mut expression = expression.clone();
61 remove_identity_casts(&mut expression, columns)?;
62 let mut scalar = crate::plan::ExpressionPlan::lower(expression).scalar;
63 let mut failure = None;
64 crate::plan::rewrite_scalar_expression(&mut scalar, &mut |node| {
65 if let ScalarExpr::TypedLiteral {
66 value: Value::Int(value),
67 ty,
68 parameter_index: None,
69 ..
70 } = node
71 {
72 if matches!(ColumnType::from_sql_name(ty), Ok(ColumnType::Integer))
73 && i32::try_from(*value).is_ok()
74 {
75 *node = ScalarExpr::Literal(Value::Int(*value));
76 }
77 return;
78 }
79 let ScalarExpr::Cast { expr, ty, implicit } = node else {
80 return;
81 };
82 *implicit = false;
84 let Ok(target) = ColumnType::from_sql_name(ty) else {
85 return;
86 };
87 if let ScalarExpr::Column(name) = expr.as_ref() {
88 if columns
89 .iter()
90 .any(|column| column.name == *name && column.ty == target)
91 {
92 *node = *expr.clone();
93 }
94 } else if let ScalarExpr::Literal(value @ Value::Str(_)) = expr.as_ref() {
95 if target == ColumnType::Integer {
97 match crate::expr::cast_value(value, ty) {
98 Ok(value) => *node = ScalarExpr::Literal(value),
99 Err(error) => failure = Some(error),
100 }
101 }
102 }
103 });
104 if let Some(error) = failure {
105 return Err(error);
106 }
107 Ok(scalar)
108 }
109 Ok(canonical(left, columns)? == canonical(right, columns)?)
110}
111
112pub fn duplicate_check(table: &str, name: &str) -> SQLError {
113 error(
114 "42710",
115 format!("constraint \"{name}\" for relation \"{table}\" already exists"),
116 )
117}
118
119fn error(sqlstate: &str, message: String) -> SQLError {
120 SQLError::Routine {
121 sqlstate: sqlstate.into(),
122 message,
123 }
124}
125
126pub fn validate_check_merge(
128 table: &str,
129 existing: &TableCheck,
130 incoming: &TableCheck,
131 columns: &[ColumnDef],
132) -> Result<(), SQLError> {
133 let name = incoming.name.as_deref().unwrap_or("<unnamed>");
134 if !same_check_expression(&existing.expr, &incoming.expr, columns)? {
135 return Err(duplicate_check(table, name));
136 }
137 let conflict = if existing.no_inherit {
138 Some("non-inherited")
139 } else if incoming.no_inherit {
140 Some("inherited")
141 } else if incoming.validated && existing.enforced && !existing.validated {
142 Some("NOT VALID")
143 } else if (!incoming.is_local && incoming.enforced && !existing.enforced)
144 || (incoming.is_local && !incoming.enforced && existing.enforced)
145 {
146 Some("NOT ENFORCED")
147 } else {
148 None
149 };
150 if let Some(conflict) = conflict {
151 return Err(error("42P17", format!("constraint \"{name}\" conflicts with {conflict} constraint on relation \"{table}\"")));
152 }
153 Ok(())
154}
155
156pub fn merge_inherited_check(
158 inherited: &mut Vec<TableCheck>,
159 check: TableCheck,
160 columns: &[ColumnDef],
161) -> Result<(), SQLError> {
162 let Some(existing) = inherited
163 .iter_mut()
164 .find(|existing| existing.name.is_some() && existing.name == check.name)
165 else {
166 inherited.push(check);
167 return Ok(());
168 };
169 if !same_check_expression(&existing.expr, &check.expr, columns)? {
170 return Err(error(
171 "42710",
172 format!(
173 "check constraint name \"{}\" appears multiple times but with different expressions",
174 check.name.as_deref().unwrap_or("<unnamed>")
175 ),
176 ));
177 }
178 existing.enforced |= check.enforced;
179 existing.validated = existing.enforced;
180 Ok(())
181}
182
183#[cfg(test)]
184mod tests {
185 use super::*;
186
187 #[test]
188 fn inherited_checks_compare_cooked_int4_constants_without_erasing_other_types() {
189 let untyped = Expr::Literal(Value::Int(0));
190 for ty in ["integer", "int4"] {
191 let cooked = Expr::TypedLiteral {
192 value: Value::Int(0),
193 ty: ty.into(),
194 };
195 assert!(same_check_expression(&untyped, &cooked, &[]).unwrap());
196 assert!(same_check_expression(&cooked, &untyped, &[]).unwrap());
197 let mut cast = Expr::Cast {
198 implicit: false,
199 expr: Box::new(cooked.clone()),
200 ty: "integer".into(),
201 };
202 assert!(same_check_expression(&cast, &untyped, &[]).unwrap());
203 remove_identity_casts(&mut cast, &[]).unwrap();
204 assert_eq!(cast, cooked);
205 }
206 for ty in ["smallint", "bigint", "oid"] {
207 let cooked = Expr::TypedLiteral {
208 value: Value::Int(0),
209 ty: ty.into(),
210 };
211 assert!(!same_check_expression(&untyped, &cooked, &[]).unwrap());
212 }
213 }
214 #[test]
215 fn inherited_check_equality_ignores_coercion_display_origin() {
216 let implicit = Expr::Cast {
217 implicit: true,
218 expr: Box::new(Expr::Column("value".into())),
219 ty: "bigint".into(),
220 };
221 let mut explicit = implicit.clone();
222 let Expr::Cast {
223 implicit: origin, ..
224 } = &mut explicit
225 else {
226 unreachable!()
227 };
228 *origin = false;
229 assert!(same_check_expression(&implicit, &explicit, &[]).unwrap());
230 }
231}