1use super::{
10 apply_attribute_change, type_contains_composite, AttributeChange, CompositeTypeCatalog,
11};
12use crate::{
13 ast::{
14 ColumnDef, ColumnType, Expr, PartitionBound, PartitionRangeDatum, PartitionSpec, Statement,
15 TableCheck,
16 },
17 plan::{QueryPlan, UnifiedPlan},
18 type_resolution::FunctionTypeResolver,
19 SQLError, ScalarExpr,
20};
21use uqa_core::Value;
22
23pub struct CompositeConstantChange<'a> {
24 pub target: u32,
25 pub change: &'a AttributeChange,
26 pub catalog: &'a dyn CompositeTypeCatalog,
27 pub types: &'a dyn FunctionTypeResolver,
28 pub rename: Option<&'a crate::binding::composite_rename::CompositeFieldRename<'a>>,
29}
30
31impl CompositeConstantChange<'_> {
32 pub fn value(&self, value: &mut Value, ty: &ColumnType) -> Result<bool, SQLError> {
33 if matches!(value, Value::Null) || !type_contains_composite(ty, self.target, self.catalog)?
34 {
35 return Ok(false);
36 }
37 *value = apply_attribute_change(
38 std::mem::replace(value, Value::Null),
39 ty,
40 self.target,
41 self.change,
42 self.catalog,
43 )?;
44 Ok(true)
45 }
46
47 fn typed_value(
48 &self,
49 value: &mut Value,
50 name: &str,
51 source: &mut Option<Box<super::CompositeConstantSource>>,
52 ) -> Result<bool, SQLError> {
53 let Some(ty) = self.types.resolve_type_name(name)? else {
54 return if crate::ast::UserTypeIdentity::parse(name).is_some() {
55 Err(SQLError::Internal(format!(
56 "stored constant type {name} disappeared"
57 )))
58 } else {
59 Ok(false)
60 };
61 };
62 if !matches!(value, Value::Null) && type_contains_composite(&ty, self.target, self.catalog)?
63 {
64 if source.is_none() {
65 *source = Some(Box::new(super::CompositeConstantSource::capture(
66 value,
67 &ty,
68 Some(self.catalog),
69 self.types.enum_labels(),
70 )?));
71 }
72 source
73 .as_mut()
74 .expect("captured composite source")
75 .retain_enum_oids(self.types.enum_labels())?;
76 if let AttributeChange::Type { name, to, .. } = self.change {
77 *value = source
78 .as_ref()
79 .expect("captured composite source")
80 .project_type_change(&ty, self.target, name, to, self.catalog)?;
81 return Ok(true);
82 }
83 }
84 self.value(value, &ty)
85 }
86
87 fn syntax_node(&self, node: &mut Expr) -> Result<bool, SQLError> {
88 match node {
89 Expr::CompositeRow { binding, .. } => {
90 super::constructor::retain_argument_types(binding, self.types, Some(self.catalog))
91 }
92 Expr::TypedLiteral {
93 value,
94 ty,
95 composite_source,
96 } => self.typed_value(value, ty, composite_source),
97 _ => Ok(false),
98 }
99 }
100
101 fn scalar_node(&self, node: &mut ScalarExpr) -> Result<bool, SQLError> {
102 match node {
103 ScalarExpr::CompositeRow { binding, .. } => {
104 super::constructor::retain_argument_types(binding, self.types, Some(self.catalog))
105 }
106 ScalarExpr::TypedLiteral {
107 value,
108 ty,
109 composite_source,
110 ..
111 } => self.typed_value(value, ty, composite_source),
112 _ => Ok(false),
113 }
114 }
115
116 pub fn expression(&self, expression: &mut Expr) -> Result<bool, SQLError> {
117 self.expression_in_schema(expression, &crate::RowSchema::default())
118 }
119
120 pub fn expression_in_schema(
121 &self,
122 expression: &mut Expr,
123 schema: &crate::RowSchema,
124 ) -> Result<bool, SQLError> {
125 let mut changed = self
126 .rename
127 .map_or(Ok(false), |rename| rename.expression(expression, schema))?;
128 crate::catalog::stored_ast::visit_stored_expression(expression, &mut |node| {
129 changed |= self.syntax_node(node)?;
130 Ok(())
131 })?;
132 Ok(changed)
133 }
134
135 pub fn statement(&self, statement: &mut Statement) -> Result<bool, SQLError> {
136 self.statement_in_scope(statement, None, None)
137 }
138
139 pub fn statement_in_scope(
140 &self,
141 statement: &mut Statement,
142 definition: Option<&crate::ast::CreateFunction>,
143 outer: Option<&crate::RowSchema>,
144 ) -> Result<bool, SQLError> {
145 let mut changed = self.rename.map_or(Ok(false), |rename| {
146 rename.statement(statement, definition, outer)
147 })?;
148 crate::catalog::stored_ast::visit_stored_statement_expressions(statement, &mut |node| {
149 changed |= self.syntax_node(node)?;
150 Ok(())
151 })?;
152 Ok(changed)
153 }
154
155 pub fn query(&self, query: &mut QueryPlan) -> Result<bool, SQLError> {
156 let changed = self
157 .rename
158 .map_or(Ok(false), |rename| rename.query(query))?;
159 self.plan_nodes(|visit| query.rewrite_scalar_expressions(visit))
160 .map(|values| values || changed)
161 }
162
163 pub fn plan(&self, plan: &mut UnifiedPlan) -> Result<bool, SQLError> {
164 let changed = self.rename.is_some_and(|rename| rename.bound_plan(plan));
165 self.plan_nodes(|visit| plan.rewrite_scalar_expressions(visit))
166 .map(|values| values || changed)
167 }
168
169 pub fn expression_plan(
170 &self,
171 plan: &mut crate::plan::ExpressionPlan,
172 ) -> Result<bool, SQLError> {
173 let mut changed = self
174 .plan_nodes(|visit| crate::plan::rewrite_scalar_expression(&mut plan.scalar, visit))?;
175 for query in &mut plan.subqueries {
176 changed |= self.query(query)?;
177 }
178 Ok(changed)
179 }
180
181 fn plan_nodes(
182 &self,
183 visit: impl FnOnce(&mut dyn FnMut(&mut ScalarExpr)),
184 ) -> Result<bool, SQLError> {
185 let mut changed = false;
186 let mut failure = None;
187 visit(&mut |node| {
188 if failure.is_some() {
189 return;
190 }
191 match self.scalar_node(node) {
192 Ok(value) => changed |= value,
193 Err(error) => failure = Some(error),
194 }
195 });
196 failure.map_or(Ok(changed), Err)
197 }
198
199 pub fn columns(&self, columns: &mut [ColumnDef]) -> Result<bool, SQLError> {
200 let schema = crate::RowSchema::with_types(
201 columns.iter().map(|column| column.name.clone()).collect(),
202 columns
203 .iter()
204 .map(|column| Some(column.ty.clone()))
205 .collect(),
206 );
207 let mut changed = false;
208 for column in columns {
209 for expression in column.default.iter_mut().chain(column.check.iter_mut()) {
210 changed |= self.expression_in_schema(expression, &schema)?;
211 }
212 if let Some(generated) = &mut column.generated {
213 changed |= self.expression_in_schema(&mut generated.expression, &schema)?;
214 }
215 }
216 Ok(changed)
217 }
218
219 pub fn checks(
220 &self,
221 checks: &mut [TableCheck],
222 columns: &[ColumnDef],
223 ) -> Result<bool, SQLError> {
224 let schema = crate::RowSchema::with_types(
225 columns.iter().map(|column| column.name.clone()).collect(),
226 columns
227 .iter()
228 .map(|column| Some(column.ty.clone()))
229 .collect(),
230 );
231 let mut changed = false;
232 for check in checks {
233 changed |= self.expression_in_schema(&mut check.expr, &schema)?;
234 if let Some(partition) = &mut check.partition_constraint {
235 changed |= self.bound(&mut partition.bound, &partition.spec, columns)?;
236 changed |= self.spec(&mut partition.spec)?;
237 }
238 }
239 Ok(changed)
240 }
241
242 pub fn spec(&self, spec: &mut PartitionSpec) -> Result<bool, SQLError> {
243 let mut changed = false;
244 for key in &mut spec.keys {
245 changed |= self.expression(key)?;
246 }
247 Ok(changed)
248 }
249
250 pub fn bound(
251 &self,
252 bound: &mut PartitionBound,
253 spec: &PartitionSpec,
254 columns: &[ColumnDef],
255 ) -> Result<bool, SQLError> {
256 let mut changed = false;
257 for (position, key) in spec.keys.iter().enumerate() {
258 let ty = crate::semantics::partition::partition_key_type(self.types, key, columns)?;
259 match bound {
260 PartitionBound::List(values) => {
261 for expression in values {
262 changed |= self.bound_value(expression, &ty)?;
263 }
264 }
265 PartitionBound::Range { lower, upper } => {
266 for point in [lower.get_mut(position), upper.get_mut(position)]
267 .into_iter()
268 .flatten()
269 {
270 if let PartitionRangeDatum::Value(expression) = point {
271 changed |= self.bound_value(expression, &ty)?;
272 }
273 }
274 }
275 PartitionBound::Default | PartitionBound::Hash { .. } => {}
276 }
277 }
278 Ok(changed)
279 }
280
281 fn bound_value(&self, expression: &mut Expr, ty: &ColumnType) -> Result<bool, SQLError> {
282 match expression {
283 Expr::Literal(value) | Expr::TypedLiteral { value, .. } => self.value(value, ty),
284 _ => Err(SQLError::Internal(
285 "partition bound datum was not evaluated".into(),
286 )),
287 }
288 }
289}
290
291#[cfg(test)]
292mod tests;