1mod 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}