uqa_sql/routines/
body_parameters.rs1use super::{
10 body_validation::routine_parameter_values,
11 compilation::{lower_sql_routine_statement, RoutineCompilationContext, SQLRoutineLowering},
12 routine_local_name,
13};
14use crate::{
15 ast::{CreateFunction, Expr, FunctionBody, FunctionParam, FunctionParamMode, Statement},
16 binding::{bind_routine_parameter_references, RoutineParameterScope},
17 catalog::stored_ast::{visit_stored_statement_expressions, visit_stored_statement_projections},
18 plan::UnifiedPlan,
19 SQLError, SQLParam, ScalarExpr,
20};
21use std::cell::Cell;
22use std::collections::BTreeSet;
23
24#[must_use]
26pub fn sql_body_parameters(def: &CreateFunction) -> Vec<&FunctionParam> {
27 def.params
28 .iter()
29 .filter(|parameter| is_sql_body_parameter(parameter))
30 .collect()
31}
32
33#[must_use]
35pub const fn is_sql_body_parameter(parameter: &FunctionParam) -> bool {
36 matches!(
37 parameter.mode,
38 FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
39 )
40}
41
42#[must_use]
44pub fn sql_body_parameter_names(def: &CreateFunction) -> Vec<String> {
45 match def.body {
46 FunctionBody::Statements(_) => def
47 .sql_body_parameters()
48 .into_iter()
49 .map(|parameter| parameter.name.to_string())
50 .collect(),
51 FunctionBody::Source(_) => sql_body_parameters(def)
52 .iter()
53 .map(|parameter| parameter.name.clone())
54 .collect(),
55 }
56}
57
58pub fn sql_body_parameter_scope(
60 def: &CreateFunction,
61 params: &[SQLParam],
62) -> Result<RoutineParameterScope, SQLError> {
63 Ok(RoutineParameterScope::new(
64 &routine_local_name(&def.name)?,
65 sql_body_parameter_names(def),
66 params
67 .iter()
68 .map(|param| param.declared_scalar_type().cloned())
69 .collect(),
70 ))
71}
72
73pub fn record_sql_standard_body_parameters(
75 context: &RoutineCompilationContext<'_>,
76 def: &mut CreateFunction,
77) -> Result<bool, SQLError> {
78 let FunctionBody::Statements(statements) = &def.body else {
79 return Ok(false);
80 };
81 let params = routine_parameter_values(context.types, def);
82 let recording = ParameterRecording {
83 context,
84 scope: sql_body_parameter_scope(def, ¶ms)?,
85 params,
86 sites: ParameterSites {
87 function: routine_local_name(&def.name)?,
88 names: sql_body_parameter_names(def),
89 },
90 };
91 let mut recorded = statements.clone();
92 let mut changed = false;
93 for statement in &mut recorded {
94 changed |= recording.record(statement)?;
95 }
96 if changed {
97 def.body = FunctionBody::Statements(recorded);
98 }
99 Ok(changed)
100}
101
102struct ParameterSites {
104 function: String,
106 names: Vec<String>,
107}
108
109impl ParameterSites {
110 fn position(&self, expression: &Expr) -> Option<usize> {
112 let name = match expression {
113 Expr::Column(name) => name,
114 Expr::QualifiedColumn { qualifier, column } if *qualifier == self.function => column,
115 _ => return None,
116 };
117 self.names
118 .iter()
119 .position(|candidate| !candidate.is_empty() && candidate == name)
120 }
121
122 fn collect(&self, statement: &mut Statement) -> Result<Vec<usize>, SQLError> {
124 let mut positions = Vec::new();
125 visit_stored_statement_expressions(statement, &mut |expression| {
126 positions.extend(self.position(expression));
127 Ok(())
128 })?;
129 Ok(positions)
130 }
131
132 fn replace(&self, statement: &mut Statement, chosen: &BTreeSet<usize>) -> Result<(), SQLError> {
134 let ordinal = Cell::new(0);
135 visit_stored_statement_projections(
136 statement,
137 &mut |projection| {
138 let name = match &projection.expr {
139 Expr::Column(name) | Expr::QualifiedColumn { column: name, .. } => name,
140 _ => return Ok(()),
141 };
142 if projection.alias.is_none()
143 && self.position(&projection.expr).is_some()
144 && chosen.contains(&ordinal.get())
145 {
146 projection.alias = Some(name.clone());
147 }
148 Ok(())
149 },
150 &mut |expression| {
151 if let Some(position) = self.position(expression) {
152 if chosen.contains(&ordinal.get()) {
153 *expression = Expr::Param(position + 1);
154 }
155 ordinal.set(ordinal.get() + 1);
156 }
157 Ok(())
158 },
159 )
160 }
161}
162
163struct ParameterRecording<'a> {
165 context: &'a RoutineCompilationContext<'a>,
166 scope: RoutineParameterScope,
167 params: Vec<SQLParam>,
168 sites: ParameterSites,
169}
170
171impl ParameterRecording<'_> {
172 fn record(&self, statement: &mut Statement) -> Result<bool, SQLError> {
174 let candidates = self.sites.collect(statement)?;
175 if candidates.is_empty() {
176 return Ok(false);
177 }
178 let original = self.resolved_counts(statement.clone())?;
179 let mut chosen = BTreeSet::new();
180 for (ordinal, position) in candidates.iter().enumerate() {
181 if original[*position] == 0 {
182 continue;
183 }
184 let mut probe = statement.clone();
185 self.sites.replace(&mut probe, &BTreeSet::from([ordinal]))?;
186 if self
187 .resolved_counts(probe)
188 .is_ok_and(|counts| counts[*position] < original[*position])
189 {
190 chosen.insert(ordinal);
191 }
192 }
193 if chosen.is_empty() {
194 return Ok(false);
195 }
196 self.sites.replace(statement, &chosen)?;
197 Ok(true)
198 }
199
200 fn resolved_counts(&self, statement: Statement) -> Result<Vec<usize>, SQLError> {
202 let mut plan = lower_sql_routine_statement(
203 self.context,
204 statement,
205 SQLRoutineLowering {
206 bind_catalog_dependencies: true,
207 persisted_definition: false,
208 preserve_target_expressions: false,
209 },
210 )?;
211 let before = parameter_counts(&mut plan, self.params.len());
212 let snapshot = self.context.catalog.binding_snapshot()?;
213 bind_routine_parameter_references(
214 self.context.routines,
215 &mut plan,
216 &self.params,
217 &snapshot.context(),
218 &self.scope,
219 )?;
220 Ok(parameter_counts(&mut plan, self.params.len())
221 .into_iter()
222 .zip(before)
223 .map(|(after, before)| after.saturating_sub(before))
224 .collect())
225 }
226}
227
228fn parameter_counts(plan: &mut UnifiedPlan, count: usize) -> Vec<usize> {
230 let mut counts = vec![0; count];
231 plan.rewrite_scalar_expressions(&mut |expression| {
232 if let ScalarExpr::Param(index) = expression {
233 if let Some(slot) = index.checked_sub(1).and_then(|slot| counts.get_mut(slot)) {
234 *slot += 1;
235 }
236 }
237 });
238 counts
239}