Skip to main content

uqa_sql/plpgsql/
binding.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Variable binding across expressions, queries, and statements.
8
9use super::{Expr, FromClause, MergeWhen, Projection, Result, SelectStmt, Statement, Value, CTE};
10
11/// Resolves routine variables while a compiled expression / statement
12/// is being specialized for one execution.
13pub trait VariableResolver {
14    /// Current value of an unqualified name. `Ok(None)` leaves the
15    /// column reference for the engine to resolve.
16    fn resolve_name(&mut self, name: &str) -> Result<Option<Value>>;
17    /// Current value of `qualifier.column` (record field access).
18    fn resolve_qualified(&mut self, qualifier: &str, column: &str) -> Result<Option<Value>>;
19    /// Value of a positional `$n` reference (function arguments).
20    fn resolve_param(&mut self, index: usize) -> Result<Option<Value>>;
21}
22
23/// Rewrite an expression, substituting resolvable variable references
24/// with literals. References the resolver declines stay untouched.
25pub fn bind_expr(expr: &Expr, r: &mut dyn VariableResolver) -> Result<Expr> {
26    Ok(match expr {
27        Expr::Column(name) => match r.resolve_name(name)? {
28            Some(value) => Expr::Literal(value),
29            None => expr.clone(),
30        },
31        Expr::QualifiedColumn {
32            qualifier, column, ..
33        } => match r.resolve_qualified(qualifier, column)? {
34            Some(value) => Expr::Literal(value),
35            None => expr.clone(),
36        },
37        Expr::Param(index) => match r.resolve_param(*index)? {
38            Some(value) => Expr::Literal(value),
39            None => expr.clone(),
40        },
41        Expr::Default | Expr::Literal(_) | Expr::Star | Expr::QualifiedStar(_) => expr.clone(),
42        Expr::Func {
43            name,
44            binding,
45            args,
46            distinct,
47            order_by,
48            filter,
49        } => Expr::Func {
50            name: name.clone(),
51            binding: binding.clone(),
52            args: bind_exprs(args, r)?,
53            distinct: *distinct,
54            order_by: bind_order_by(order_by, r)?,
55            filter: match filter {
56                Some(f) => Some(Box::new(bind_expr(f, r)?)),
57                None => None,
58            },
59        },
60        Expr::Array(items) => Expr::Array(bind_exprs(items, r)?),
61        Expr::Row(items) => Expr::Row(bind_exprs(items, r)?),
62        Expr::Binary { op, lhs, rhs } => Expr::Binary {
63            op: *op,
64            lhs: Box::new(bind_expr(lhs, r)?),
65            rhs: Box::new(bind_expr(rhs, r)?),
66        },
67        Expr::UnaryMinus(inner) => Expr::UnaryMinus(Box::new(bind_expr(inner, r)?)),
68        Expr::Not(inner) => Expr::Not(Box::new(bind_expr(inner, r)?)),
69        Expr::And(items) => Expr::And(bind_exprs(items, r)?),
70        Expr::Or(items) => Expr::Or(bind_exprs(items, r)?),
71        Expr::IsNull { expr, negated } => Expr::IsNull {
72            expr: Box::new(bind_expr(expr, r)?),
73            negated: *negated,
74        },
75        Expr::Between { expr, low, high } => Expr::Between {
76            expr: Box::new(bind_expr(expr, r)?),
77            low: Box::new(bind_expr(low, r)?),
78            high: Box::new(bind_expr(high, r)?),
79        },
80        Expr::InList {
81            expr,
82            list,
83            negated,
84        } => Expr::InList {
85            expr: Box::new(bind_expr(expr, r)?),
86            list: bind_exprs(list, r)?,
87            negated: *negated,
88        },
89        Expr::WindowCall { name, args, spec } => Expr::WindowCall {
90            name: name.clone(),
91            args: bind_exprs(args, r)?,
92            spec: crate::ast::WindowSpec {
93                partition_by: bind_exprs(&spec.partition_by, r)?,
94                order_by: bind_order_by(&spec.order_by, r)?,
95                frame: spec.frame.clone(),
96            },
97        },
98        Expr::Case {
99            base,
100            when,
101            else_branch,
102        } => Expr::Case {
103            base: match base {
104                Some(b) => Some(Box::new(bind_expr(b, r)?)),
105                None => None,
106            },
107            when: when
108                .iter()
109                .map(|(c, v)| Ok((bind_expr(c, r)?, bind_expr(v, r)?)))
110                .collect::<Result<Vec<_>>>()?,
111            else_branch: match else_branch {
112                Some(e) => Some(Box::new(bind_expr(e, r)?)),
113                None => None,
114            },
115        },
116        Expr::Cast { expr, ty } => Expr::Cast {
117            expr: Box::new(bind_expr(expr, r)?),
118            ty: ty.clone(),
119        },
120        Expr::ScalarSubquery(body) => Expr::ScalarSubquery(Box::new(bind_select(body, r)?)),
121        Expr::Exists { body, negated } => Expr::Exists {
122            body: Box::new(bind_select(body, r)?),
123            negated: *negated,
124        },
125        Expr::InSubquery {
126            expr,
127            body,
128            negated,
129        } => Expr::InSubquery {
130            expr: Box::new(bind_expr(expr, r)?),
131            body: Box::new(bind_select(body, r)?),
132            negated: *negated,
133        },
134    })
135}
136
137pub(super) fn bind_exprs(exprs: &[Expr], r: &mut dyn VariableResolver) -> Result<Vec<Expr>> {
138    exprs.iter().map(|e| bind_expr(e, r)).collect()
139}
140
141pub(super) fn bind_opt_expr(
142    expr: Option<&Expr>,
143    r: &mut dyn VariableResolver,
144) -> Result<Option<Expr>> {
145    match expr {
146        Some(e) => Ok(Some(bind_expr(e, r)?)),
147        None => Ok(None),
148    }
149}
150
151pub(super) fn bind_order_by(
152    items: &[crate::ast::OrderBy],
153    r: &mut dyn VariableResolver,
154) -> Result<Vec<crate::ast::OrderBy>> {
155    items
156        .iter()
157        .map(|o| {
158            Ok(crate::ast::OrderBy {
159                expr: bind_expr(&o.expr, r)?,
160                descending: o.descending,
161                nulls: o.nulls,
162            })
163        })
164        .collect()
165}
166
167pub(super) fn bind_projections(
168    items: &[Projection],
169    r: &mut dyn VariableResolver,
170) -> Result<Vec<Projection>> {
171    items
172        .iter()
173        .map(|p| {
174            Ok(Projection {
175                expr: bind_expr(&p.expr, r)?,
176                alias: p.alias.clone(),
177            })
178        })
179        .collect()
180}
181
182pub(super) fn bind_assignments(
183    items: &[(String, Expr)],
184    r: &mut dyn VariableResolver,
185) -> Result<Vec<(String, Expr)>> {
186    items
187        .iter()
188        .map(|(name, e)| Ok((name.clone(), bind_expr(e, r)?)))
189        .collect()
190}
191
192pub(super) fn bind_ctes(items: &[CTE], r: &mut dyn VariableResolver) -> Result<Vec<CTE>> {
193    items
194        .iter()
195        .map(|cte| {
196            Ok(CTE {
197                name: cte.name.clone(),
198                columns: cte.columns.clone(),
199                recursive: cte.recursive,
200                query: Box::new(bind_select(&cte.query, r)?),
201            })
202        })
203        .collect()
204}
205
206pub(super) fn bind_rows(
207    rows: &[Vec<Expr>],
208    r: &mut dyn VariableResolver,
209) -> Result<Vec<Vec<Expr>>> {
210    rows.iter().map(|row| bind_exprs(row, r)).collect()
211}
212
213/// Rewrite a `SELECT` body, substituting resolvable variables.
214pub fn bind_select(stmt: &SelectStmt, r: &mut dyn VariableResolver) -> Result<SelectStmt> {
215    Ok(SelectStmt {
216        projections: bind_projections(&stmt.projections, r)?,
217        values: bind_rows(&stmt.values, r)?,
218        from: match stmt.from.as_ref() {
219            Some(f) => Some(bind_from(f, r)?),
220            None => None,
221        },
222        r#where: bind_opt_expr(stmt.r#where.as_ref(), r)?,
223        group_by: bind_exprs(&stmt.group_by, r)?,
224        grouping_sets: stmt
225            .grouping_sets
226            .iter()
227            .map(|set| bind_exprs(set, r))
228            .collect::<Result<Vec<_>>>()?,
229        having: bind_opt_expr(stmt.having.as_ref(), r)?,
230        order_by: bind_order_by(&stmt.order_by, r)?,
231        limit: bind_opt_expr(stmt.limit.as_ref(), r)?,
232        offset: bind_opt_expr(stmt.offset.as_ref(), r)?,
233        with: bind_ctes(&stmt.with, r)?,
234        set_op: match stmt.set_op.as_ref() {
235            Some(op) => Some(Box::new(crate::ast::SetOp {
236                kind: op.kind,
237                all: op.all,
238                left: op
239                    .left
240                    .as_ref()
241                    .map(|left| bind_select(left, r).map(Box::new))
242                    .transpose()?,
243                right: bind_select(&op.right, r)?,
244                combined_order_by: bind_order_by(&op.combined_order_by, r)?,
245                combined_limit: bind_opt_expr(op.combined_limit.as_ref(), r)?,
246                combined_offset: bind_opt_expr(op.combined_offset.as_ref(), r)?,
247            })),
248            None => None,
249        },
250        distinct: stmt.distinct,
251        distinct_on: bind_exprs(&stmt.distinct_on, r)?,
252        locking: stmt.locking.clone(),
253    })
254}
255
256pub(super) fn bind_from(from: &FromClause, r: &mut dyn VariableResolver) -> Result<FromClause> {
257    Ok(match from {
258        FromClause::Table { .. } => from.clone(),
259        FromClause::Join {
260            left,
261            right,
262            kind,
263            on,
264            using,
265            natural,
266            lateral,
267        } => FromClause::Join {
268            left: Box::new(bind_from(left, r)?),
269            right: Box::new(bind_from(right, r)?),
270            kind: *kind,
271            on: bind_opt_expr(on.as_ref(), r)?,
272            using: using.clone(),
273            natural: *natural,
274            lateral: *lateral,
275        },
276        FromClause::Values {
277            rows,
278            alias,
279            column_aliases,
280        } => FromClause::Values {
281            rows: bind_rows(rows, r)?,
282            alias: alias.clone(),
283            column_aliases: column_aliases.clone(),
284        },
285        FromClause::Function {
286            name,
287            output_name,
288            relation,
289            args,
290            alias,
291            column_aliases,
292            column_types,
293        } => FromClause::Function {
294            name: name.clone(),
295            output_name: output_name.clone(),
296            relation: relation.clone(),
297            args: bind_exprs(args, r)?,
298            alias: alias.clone(),
299            column_aliases: column_aliases.clone(),
300            column_types: column_types.clone(),
301        },
302        FromClause::Subquery {
303            body,
304            alias,
305            column_aliases,
306        } => FromClause::Subquery {
307            body: Box::new(bind_select(body, r)?),
308            alias: alias.clone(),
309            column_aliases: column_aliases.clone(),
310        },
311    })
312}
313
314/// Rewrite a full statement, substituting resolvable variables in
315/// every expression position. Statements without expression payloads
316/// pass through unchanged.
317pub fn bind_statement(stmt: &Statement, r: &mut dyn VariableResolver) -> Result<Statement> {
318    Ok(match stmt {
319        Statement::Select(body) => Statement::Select(Box::new(bind_select(body, r)?)),
320        Statement::Insert(insert) => {
321            let mut out = insert.clone();
322            out.with = bind_ctes(&insert.with, r)?;
323            out.rows = bind_rows(&insert.rows, r)?;
324            out.select_source = match insert.select_source.as_ref() {
325                Some(body) => Some(Box::new(bind_select(body, r)?)),
326                None => None,
327            };
328            out.on_conflict = match insert.on_conflict.as_ref() {
329                Some(oc) => Some(crate::ast::OnConflict {
330                    conflict_columns: oc.conflict_columns.clone(),
331                    action: match &oc.action {
332                        crate::ast::OnConflictAction::Nothing => {
333                            crate::ast::OnConflictAction::Nothing
334                        }
335                        crate::ast::OnConflictAction::Update {
336                            assignments,
337                            r#where,
338                        } => crate::ast::OnConflictAction::Update {
339                            assignments: bind_assignments(assignments, r)?,
340                            r#where: bind_opt_expr(r#where.as_ref(), r)?,
341                        },
342                    },
343                }),
344                None => None,
345            };
346            out.returning = bind_projections(&insert.returning, r)?;
347            Statement::Insert(out)
348        }
349        Statement::Update(update) => {
350            let mut out = update.clone();
351            out.assignments = bind_assignments(&update.assignments, r)?;
352            out.r#where = bind_opt_expr(update.r#where.as_ref(), r)?;
353            out.with = bind_ctes(&update.with, r)?;
354            out.from = match update.from.as_ref() {
355                Some(f) => Some(bind_from(f, r)?),
356                None => None,
357            };
358            out.returning = bind_projections(&update.returning, r)?;
359            Statement::Update(out)
360        }
361        Statement::Delete(delete) => {
362            let mut out = delete.clone();
363            out.r#where = bind_opt_expr(delete.r#where.as_ref(), r)?;
364            out.with = bind_ctes(&delete.with, r)?;
365            out.using = match delete.using.as_ref() {
366                Some(f) => Some(bind_from(f, r)?),
367                None => None,
368            };
369            out.returning = bind_projections(&delete.returning, r)?;
370            Statement::Delete(out)
371        }
372        Statement::Values { rows } => Statement::Values {
373            rows: bind_rows(rows, r)?,
374        },
375        Statement::CreateTableAs {
376            name,
377            if_not_exists,
378            body,
379        } => Statement::CreateTableAs {
380            name: name.clone(),
381            if_not_exists: *if_not_exists,
382            body: Box::new(bind_select(body, r)?),
383        },
384        Statement::Explain {
385            analyze,
386            verbose,
387            format,
388            body,
389        } => Statement::Explain {
390            analyze: *analyze,
391            verbose: *verbose,
392            format: format.clone(),
393            body: Box::new(bind_statement(body, r)?),
394        },
395        Statement::Merge(merge) => {
396            let mut out = merge.clone();
397            out.source = bind_from(&merge.source, r)?;
398            out.join_condition = bind_expr(&merge.join_condition, r)?;
399            out.when_clauses = merge
400                .when_clauses
401                .iter()
402                .map(|w| bind_merge_when(w, r))
403                .collect::<Result<Vec<_>>>()?;
404            out.returning = bind_projections(&merge.returning, r)?;
405            Statement::Merge(out)
406        }
407        Statement::Call { name, args } => Statement::Call {
408            name: name.clone(),
409            args: bind_exprs(args, r)?,
410        },
411        other => other.clone(),
412    })
413}
414
415pub(super) fn bind_merge_when(when: &MergeWhen, r: &mut dyn VariableResolver) -> Result<MergeWhen> {
416    Ok(match when {
417        MergeWhen::UpdateMatched {
418            condition,
419            assignments,
420        } => MergeWhen::UpdateMatched {
421            condition: bind_opt_expr(condition.as_ref(), r)?,
422            assignments: bind_assignments(assignments, r)?,
423        },
424        MergeWhen::DeleteMatched { condition } => MergeWhen::DeleteMatched {
425            condition: bind_opt_expr(condition.as_ref(), r)?,
426        },
427        MergeWhen::InsertNotMatched {
428            condition,
429            columns,
430            values,
431        } => MergeWhen::InsertNotMatched {
432            condition: bind_opt_expr(condition.as_ref(), r)?,
433            columns: columns.clone(),
434            values: bind_exprs(values, r)?,
435        },
436        MergeWhen::NothingMatched { condition } => MergeWhen::NothingMatched {
437            condition: bind_opt_expr(condition.as_ref(), r)?,
438        },
439        MergeWhen::NothingNotMatched { condition } => MergeWhen::NothingNotMatched {
440            condition: bind_opt_expr(condition.as_ref(), r)?,
441        },
442    })
443}