Skip to main content

uqa_sql/routines/
body_parameters.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! The parameters a SQL routine's body names, by name or by position, as `PostgreSQL`'s parser resolves them in each statement of the body.
8
9use 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/// The parameters a SQL body can name: the routine's input parameters, which `get_func_input_arg_names` lists.
25#[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/// Whether the body can name `parameter`: an input parameter. A procedure's output parameter takes a placeholder in `CALL` but is not a parameter of its body.
34#[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/// The names a SQL body gives its input parameters, in order. A body given as a string names each by its own name; a SQL-standard body names the input at position `n` by the parameter declared at position `n` among all parameters, as `interpret_AS_clause` names them, so an input that follows an output parameter goes by that output parameter's name.
43#[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
58/// The scope of the parameters a SQL body names, typed as `params` types them.
59pub 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
73/// Record in the statements of a SQL-standard body each name that resolves to a parameter, as the positional parameter it names, so that a later change to the relations the body reads cannot take the name: `PostgreSQL` stores the analyzed statements, in which a parameter reference stays a parameter. Returns whether a statement changed.
74pub 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, &params)?,
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
102/// The names in a body that could refer to the routine's parameters.
103struct ParameterSites {
104    /// The routine's unqualified name, which qualifies its parameter names.
105    function: String,
106    names: Vec<String>,
107}
108
109impl ParameterSites {
110    /// The position of the parameter `expression` could refer to: a parameter name, alone or qualified by the routine's name.
111    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    /// The parameter positions of the names in `statement` that could refer to parameters, in the order the stored statement visitor reaches them.
123    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    /// Replace the names whose ordinals `chosen` holds with the parameters they name. A select list item that is such a name keeps its output name as an alias, so that the statement's output columns keep their names.
133    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
163/// What finding the parameter references of a SQL-standard body needs: the body's names are resolved by the same analysis that compiles the body.
164struct ParameterRecording<'a> {
165    context: &'a RoutineCompilationContext<'a>,
166    scope: RoutineParameterScope,
167    params: Vec<SQLParam>,
168    sites: ParameterSites,
169}
170
171impl ParameterRecording<'_> {
172    /// Record the parameter references of one statement. A name is a parameter reference exactly when writing the parameter in its place removes one of the references analysis resolves to that parameter by name; a name that resolves to a column leaves them as they were, or makes the statement invalid when the parameter does not type as the column does.
173    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    /// The number of references to each parameter that analysis of `statement` resolves by name.
201    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
228/// The number of references to each of `count` parameters in `plan`.
229fn 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}