Skip to main content

uqa_sql/catalog/node_tree/
expressions.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Bind catalog expression nodes to declared column and domain-value types.
8
9mod coercion;
10mod constructs;
11mod operators;
12
13use super::{values, Field, Node};
14use crate::ast::{BinaryOp, Expr, FunctionBinding, FunctionCallSyntax};
15use crate::catalog::type_metadata::{pg_type_collation_oid, pg_type_modifier, pg_type_oid};
16use crate::type_resolution::{
17    binary_operator_catalog_entry, binary_operator_types, common_context_expression_type,
18    FunctionTypeResolver,
19};
20use crate::{ColumnType, RowSchema, SQLError};
21
22pub struct RoutineIdentity {
23    pub oid: i64,
24    pub argument_types: Vec<ColumnType>,
25    pub result_type: ColumnType,
26}
27
28pub trait ExpressionRoutines {
29    fn resolve(
30        &self,
31        name: &str,
32        binding: Option<&FunctionBinding>,
33        argument_types: &[Option<ColumnType>],
34    ) -> Result<RoutineIdentity, SQLError>;
35}
36
37pub struct ExpressionContext<'a> {
38    pub schema: &'a RowSchema,
39    pub domain_value: Option<&'a ColumnType>,
40    pub types: Option<&'a dyn FunctionTypeResolver>,
41    pub routines: &'a dyn ExpressionRoutines,
42}
43
44struct TypedNode {
45    node: Node,
46    ty: ColumnType,
47}
48
49impl ExpressionContext<'_> {
50    pub fn check(&self, expression: &Expr) -> Result<Node, SQLError> {
51        self.encode(expression, Some(&ColumnType::Boolean))
52            .map(|value| value.node)
53    }
54
55    fn expression_type(&self, expression: &Expr) -> Result<Option<ColumnType>, SQLError> {
56        let plan = crate::plan::ExpressionPlan::lower(expression.clone());
57        common_context_expression_type(&plan.scalar, self.schema, &[], self.types)
58    }
59
60    fn encode(
61        &self,
62        expression: &Expr,
63        expected: Option<&ColumnType>,
64    ) -> Result<TypedNode, SQLError> {
65        let value = match expression {
66            Expr::Column(column) => self.column(column, None)?,
67            Expr::QualifiedColumn { qualifier, column } => self.column(column, Some(qualifier))?,
68            Expr::Literal(value) => self.literal(expression, value, expected)?,
69            Expr::TypedLiteral { value, ty } => {
70                let ty = self.resolve_type(ty)?;
71                TypedNode {
72                    node: values::constant(value, &ty)?,
73                    ty,
74                }
75            }
76            Expr::Binary { op, lhs, rhs } => self.binary(*op, lhs, rhs)?,
77            Expr::UnaryMinus(argument) => self.unary_minus(argument)?,
78            Expr::Array(elements) => self.array(elements, expected)?,
79            Expr::InList {
80                expr,
81                list,
82                negated,
83            } => self.in_list(expr, list, *negated)?,
84            Expr::Case {
85                base,
86                when,
87                else_branch,
88            } => self.case(expression, base.as_deref(), when, else_branch.as_deref())?,
89            Expr::And(items) => self.boolean("and", items)?,
90            Expr::Or(items) => self.boolean("or", items)?,
91            Expr::Not(item) => self.boolean("not", std::slice::from_ref(item.as_ref()))?,
92            Expr::IsNull { expr, negated } => {
93                let arg = self.encode(expr, None)?;
94                TypedNode {
95                    node: Node::new(
96                        "NULLTEST",
97                        [
98                            ("arg", arg.node.into()),
99                            ("nulltesttype", i64::from(*negated).into()),
100                            ("argisrow", false.into()),
101                            ("location", (-1).into()),
102                        ],
103                    ),
104                    ty: ColumnType::Boolean,
105                }
106            }
107            Expr::Between { expr, low, high } => self.boolean(
108                "and",
109                &[
110                    Expr::Binary {
111                        op: BinaryOp::GreaterEqual,
112                        lhs: expr.clone(),
113                        rhs: low.clone(),
114                    },
115                    Expr::Binary {
116                        op: BinaryOp::LessEqual,
117                        lhs: expr.clone(),
118                        rhs: high.clone(),
119                    },
120                ],
121            )?,
122            Expr::Func {
123                order_syntax,
124                name,
125                binding,
126                args,
127                distinct: false,
128                order_by,
129                filter: None,
130            } if order_by.is_empty() => {
131                if let Some(value) = self.construct(expression, name, binding.as_ref(), args)? {
132                    value
133                } else {
134                    self.function(name, binding.as_ref(), args, *order_syntax)?
135                }
136            }
137            Expr::Cast { expr, ty, .. } => {
138                let ty = self.resolve_type(ty)?;
139                let unknown = self.expression_type(expr)?.is_none();
140                let mut input_type = &ty;
141                while let ColumnType::Domain { base, .. } = input_type {
142                    input_type = base;
143                }
144                let input_type = input_type.without_type_modifiers();
145                let inner = self.encode(expr, unknown.then_some(&input_type))?;
146                Self::coerce(inner, &ty, 1)?
147            }
148            _ => {
149                return Err(SQLError::Unsupported(
150                    "catalog expression node encoding for this expression".into(),
151                ))
152            }
153        };
154        if let Some(ty) = expected {
155            Self::coerce(value, ty, 2)
156        } else {
157            Ok(value)
158        }
159    }
160
161    fn literal(
162        &self,
163        expression: &Expr,
164        value: &uqa_core::Value,
165        expected: Option<&ColumnType>,
166    ) -> Result<TypedNode, SQLError> {
167        let source = self.expression_type(expression)?;
168        let ty = source
169            .as_ref()
170            .or(expected)
171            .cloned()
172            .unwrap_or(ColumnType::Text);
173        let value = crate::type_resolution::coerce_common_context_value(
174            value.clone(),
175            source.as_ref(),
176            Some(&ty),
177        )?;
178        Ok(TypedNode {
179            node: values::constant(&value, &ty)?,
180            ty,
181        })
182    }
183
184    fn resolve_type(&self, name: &str) -> Result<ColumnType, SQLError> {
185        if let Some(resolver) = self.types {
186            if let Some(ty) = resolver.resolve_type_name(name)? {
187                return Ok(ty);
188            }
189        }
190        ColumnType::from_sql_name(name)
191    }
192
193    fn column(&self, name: &str, qualifier: Option<&str>) -> Result<TypedNode, SQLError> {
194        if let Some(ty) = self.domain_value {
195            if name != "value" || qualifier.is_some() {
196                return Err(SQLError::UnknownColumn(name.into()));
197            }
198            return Ok(TypedNode {
199                node: Node::new(
200                    "COERCETODOMAINVALUE",
201                    [
202                        ("typeId", pg_type_oid(ty).into()),
203                        ("typeMod", pg_type_modifier(ty).into()),
204                        ("collation", pg_type_collation_oid(ty).into()),
205                        ("location", (-1).into()),
206                    ],
207                ),
208                ty: ty.clone(),
209            });
210        }
211        let position = qualifier
212            .map_or_else(
213                || self.schema.unqualified_position(name),
214                |qualifier| self.schema.qualified_position(qualifier, name),
215            )
216            .ok_or_else(|| SQLError::UnknownColumn(name.into()))?;
217        let ty = self
218            .schema
219            .column_type(position)
220            .ok_or_else(|| SQLError::Internal("catalog column has no declared type".into()))?;
221        let ordinal = i64::try_from(position + 1)
222            .map_err(|_| SQLError::Internal("column ordinal overflow".into()))?;
223        Ok(TypedNode {
224            node: Node::new(
225                "VAR",
226                [
227                    ("varno", 1.into()),
228                    ("varattno", ordinal.into()),
229                    ("vartype", pg_type_oid(ty).into()),
230                    ("vartypmod", pg_type_modifier(ty).into()),
231                    ("varcollid", pg_type_collation_oid(ty).into()),
232                    ("varnullingrels", Field::List(vec![Field::Atom("b".into())])),
233                    ("varlevelsup", 0.into()),
234                    ("varreturningtype", 0.into()),
235                    ("varnosyn", 1.into()),
236                    ("varattnosyn", ordinal.into()),
237                    ("location", (-1).into()),
238                ],
239            ),
240            ty: ty.clone(),
241        })
242    }
243
244    fn binary(&self, op: BinaryOp, lhs: &Expr, rhs: &Expr) -> Result<TypedNode, SQLError> {
245        let types = [self.expression_type(lhs)?, self.expression_type(rhs)?];
246        let [left, right, result] =
247            binary_operator_types(op, types[0].as_ref(), types[1].as_ref())?;
248        let identity = binary_operator_catalog_entry(op, [&left, &right])?;
249        let arguments = [
250            self.encode(lhs, Some(&left))?,
251            self.encode(rhs, Some(&right))?,
252        ];
253        Ok(operator_node(
254            identity.oid,
255            identity.function_oid,
256            arguments,
257            result,
258        ))
259    }
260
261    fn boolean(&self, operator: &str, args: &[Expr]) -> Result<TypedNode, SQLError> {
262        let args = args
263            .iter()
264            .map(|arg| {
265                self.encode(arg, Some(&ColumnType::Boolean))
266                    .map(|value| value.node.into())
267            })
268            .collect::<Result<_, _>>()?;
269        Ok(TypedNode {
270            node: Node::new(
271                "BOOLEXPR",
272                [
273                    ("boolop", Field::Atom(operator.into())),
274                    ("args", Field::List(args)),
275                    ("location", (-1).into()),
276                ],
277            ),
278            ty: ColumnType::Boolean,
279        })
280    }
281
282    fn function(
283        &self,
284        name: &str,
285        binding: Option<&FunctionBinding>,
286        arguments: &[Expr],
287        syntax: FunctionCallSyntax,
288    ) -> Result<TypedNode, SQLError> {
289        let types = arguments
290            .iter()
291            .map(|arg| self.expression_type(arg))
292            .collect::<Result<Vec<_>, _>>()?;
293        let routine = self.routines.resolve(name, binding, &types)?;
294        if routine.argument_types.len() != arguments.len() {
295            return Err(SQLError::Internal(
296                "catalog routine arity differs from bound arguments".into(),
297            ));
298        }
299        let arguments = arguments
300            .iter()
301            .zip(&routine.argument_types)
302            .map(|(arg, ty)| self.encode(arg, Some(ty)))
303            .collect::<Result<Vec<_>, _>>()?;
304        let collation = arguments
305            .iter()
306            .map(|argument| pg_type_collation_oid(&argument.ty))
307            .find(|oid| *oid != 0)
308            .unwrap_or(0);
309        Ok(TypedNode {
310            node: Node::new(
311                "FUNCEXPR",
312                [
313                    ("funcid", routine.oid.into()),
314                    ("funcresulttype", pg_type_oid(&routine.result_type).into()),
315                    ("funcretset", false.into()),
316                    ("funcvariadic", false.into()),
317                    (
318                        "funcformat",
319                        if syntax == FunctionCallSyntax::Extract {
320                            3.into()
321                        } else {
322                            0.into()
323                        },
324                    ),
325                    (
326                        "funccollid",
327                        pg_type_collation_oid(&routine.result_type).into(),
328                    ),
329                    ("inputcollid", collation.into()),
330                    (
331                        "args",
332                        Field::List(arguments.into_iter().map(|arg| arg.node.into()).collect()),
333                    ),
334                    ("location", (-1).into()),
335                ],
336            ),
337            ty: routine.result_type,
338        })
339    }
340}
341
342fn operator_node(
343    oid: i64,
344    function_oid: i64,
345    arguments: impl IntoIterator<Item = TypedNode>,
346    result: ColumnType,
347) -> TypedNode {
348    let arguments: Vec<_> = arguments.into_iter().collect();
349    let input_collation = arguments
350        .iter()
351        .map(|argument| pg_type_collation_oid(&argument.ty))
352        .find(|oid| *oid != 0)
353        .unwrap_or(0);
354    TypedNode {
355        node: Node::new(
356            "OPEXPR",
357            [
358                ("opno", oid.into()),
359                ("opfuncid", function_oid.into()),
360                ("opresulttype", pg_type_oid(&result).into()),
361                ("opretset", false.into()),
362                ("opcollid", pg_type_collation_oid(&result).into()),
363                ("inputcollid", input_collation.into()),
364                (
365                    "args",
366                    Field::List(arguments.into_iter().map(|arg| arg.node.into()).collect()),
367                ),
368                ("location", (-1).into()),
369            ],
370        ),
371        ty: result,
372    }
373}