Skip to main content

uqa_sql/type_resolution/
common.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7use crate::ast::ColumnType;
8use crate::{SQLError, SQLParam};
9use uqa_core::Value;
10
11use crate::{scalar_call_arguments, RowSchema, ScalarExpr};
12
13use super::{scalar_type_inner, FunctionTypeResolver};
14
15/// Decoded function-call argument names, effective overload types, and whether the call used explicit `VARIADIC` syntax.
16#[doc(hidden)]
17pub type FunctionCallArgumentSignature = (Vec<Option<String>>, Vec<Option<ColumnType>>, bool);
18
19/// Build the PostgreSQL-compatible overload signature for one physical function call using the shared common-context typing rule.
20#[doc(hidden)]
21pub fn function_call_argument_signature(
22    arguments: &[ScalarExpr],
23    schema: &RowSchema,
24    params: &[SQLParam],
25    resolver: Option<&dyn FunctionTypeResolver>,
26) -> Result<FunctionCallArgumentSignature, SQLError> {
27    let call_arguments = scalar_call_arguments(arguments)?;
28    let explicit_variadic = call_arguments
29        .iter()
30        .any(|argument| argument.explicit_variadic);
31    let mut argument_names = Vec::with_capacity(call_arguments.len());
32    let mut argument_types = Vec::with_capacity(call_arguments.len());
33    for argument in call_arguments {
34        argument_names.push(argument.name.map(str::to_string));
35        let argument_type =
36            common_context_expression_type(argument.value, schema, params, resolver)?;
37        argument_types.push(effective_overload_argument_type_with_params(
38            argument.value,
39            argument_type,
40            params,
41        ));
42    }
43    Ok((argument_names, argument_types, explicit_variadic))
44}
45
46pub(super) fn local_routine_name(name: &str) -> String {
47    let lower = name.to_ascii_lowercase();
48    lower
49        .strip_prefix("pg_catalog.")
50        .unwrap_or(&lower)
51        .to_string()
52}
53
54pub(super) fn numeric_type() -> ColumnType {
55    ColumnType::Numeric {
56        precision: None,
57        scale: None,
58    }
59}
60
61pub(super) fn base_type(mut ty: &ColumnType) -> &ColumnType {
62    while let ColumnType::Domain { base, .. } = ty {
63        ty = base;
64    }
65    ty.without_temporal_modifiers()
66}
67
68pub fn values_column_types(
69    rows: &[Vec<ScalarExpr>],
70    params: &[SQLParam],
71) -> Result<Vec<Option<ColumnType>>, SQLError> {
72    let width = rows.first().map_or(0, Vec::len);
73    let empty = RowSchema::default();
74    let mut types = vec![None; width];
75    for row in rows {
76        if row.len() != width {
77            return Err(SQLError::TypeMismatch(
78                "VALUES lists must all be the same length".into(),
79            ));
80        }
81        for (position, expression) in row.iter().enumerate() {
82            types[position] = merge_optional_types(
83                types[position].take(),
84                common_context_expression_type(expression, &empty, params, None)?,
85            )?;
86        }
87    }
88    Ok(types
89        .into_iter()
90        .map(|ty| ty.or(Some(ColumnType::Text)))
91        .collect())
92}
93
94/// Resolve an expression participating in `PostgreSQL`'s common-type selection. Bare string and NULL literals retain the parser's `unknown` type until the surrounding VALUES, set operation, CASE, or array context selects a concrete type.
95pub fn common_context_expression_type(
96    expression: &ScalarExpr,
97    schema: &RowSchema,
98    params: &[SQLParam],
99    resolver: Option<&dyn FunctionTypeResolver>,
100) -> Result<Option<ColumnType>, SQLError> {
101    if matches!(expression, ScalarExpr::Literal(Value::Str(_) | Value::Null)) {
102        return Ok(None);
103    }
104    scalar_type_inner(expression, schema, params, resolver)
105}
106
107/// Preserve parser-level `unknown` identity for fixed built-in overload selection.
108#[doc(hidden)]
109pub fn effective_overload_argument_type(
110    expression: &ScalarExpr,
111    resolved: Option<ColumnType>,
112) -> Option<ColumnType> {
113    if matches!(expression, ScalarExpr::Literal(Value::Str(_) | Value::Null))
114        || matches!(expression, ScalarExpr::Param(_))
115            && matches!(resolved.as_ref(), Some(ColumnType::Text))
116    {
117        None
118    } else {
119        resolved
120    }
121}
122
123/// Preserve an explicitly typed scalar parameter while retaining the legacy `unknown` treatment of untyped text-valued [`SQLParam::Scalar`] parameters.
124#[doc(hidden)]
125pub fn effective_overload_argument_type_with_params(
126    expression: &ScalarExpr,
127    resolved: Option<ColumnType>,
128    params: &[SQLParam],
129) -> Option<ColumnType> {
130    if let ScalarExpr::Param(index) = expression {
131        if index
132            .checked_sub(1)
133            .and_then(|index| params.get(index))
134            .is_some_and(|parameter| parameter.declared_scalar_type().is_some())
135        {
136            return resolved;
137        }
138    }
139    effective_overload_argument_type(expression, resolved)
140}
141
142pub(super) fn parameter_type(parameter: &SQLParam) -> Option<ColumnType> {
143    match parameter {
144        SQLParam::Scalar(value) => value_type(value),
145        SQLParam::TypedScalar { ty, .. } => Some(ty.clone()),
146        SQLParam::Vector(values) => u32::try_from(values.len()).ok().map(ColumnType::Vector),
147        SQLParam::Tensor(values) => values
148            .first()
149            .and_then(|values| u32::try_from(values.len()).ok())
150            .map(ColumnType::Tensor),
151    }
152}
153
154pub(super) fn value_type(value: &Value) -> Option<ColumnType> {
155    match value {
156        Value::Null | Value::Map(_) => None,
157        Value::Void => Some(ColumnType::Void),
158        Value::Row(_) | Value::Record(_) => Some(ColumnType::Record),
159        Value::Bool(_) => Some(ColumnType::Boolean),
160        Value::Int(value) if i32::try_from(*value).is_ok() => Some(ColumnType::Integer),
161        Value::Int(_) => Some(ColumnType::BigInteger),
162        Value::Float(_) => Some(ColumnType::DoublePrecision),
163        Value::Decimal(_) => Some(numeric_type()),
164        Value::Str(_) => Some(ColumnType::Text),
165        Value::FixedChar(value) => u32::try_from(value.chars().count())
166            .ok()
167            .map(ColumnType::Character),
168        Value::Bytes(_) => Some(ColumnType::Bytea),
169        Value::Temporal(value) => Some(match value {
170            uqa_core::TemporalValue::Date { .. } => ColumnType::Date,
171            uqa_core::TemporalValue::Time { .. } => ColumnType::Time,
172            uqa_core::TemporalValue::TimeTz { .. } => ColumnType::TimeTz,
173            uqa_core::TemporalValue::Timestamp { .. } => ColumnType::Timestamp,
174            uqa_core::TemporalValue::TimestampTz { .. } => ColumnType::TimestampTz,
175            uqa_core::TemporalValue::Interval { .. } => ColumnType::Interval,
176        }),
177        Value::Json(_) => Some(ColumnType::Json),
178        Value::JsonB(_) => Some(ColumnType::JsonB),
179        Value::Array(array) => {
180            let mut element = None;
181            merge_array_element_types(array.elements(), &mut element)?;
182            element.map(|element| ColumnType::Array(Box::new(element)))
183        }
184        Value::List(values) => {
185            let mut element = None;
186            for value in values {
187                element = merge_optional_types(element, value_type(value)).ok()?;
188            }
189            element.map(|element| ColumnType::Array(Box::new(element)))
190        }
191    }
192}
193
194fn merge_array_element_types(values: &[Value], element: &mut Option<ColumnType>) -> Option<()> {
195    for value in values {
196        if let Value::List(nested) = value {
197            merge_array_element_types(nested, element)?;
198        } else {
199            *element = merge_optional_types(element.take(), value_type(value)).ok()?;
200        }
201    }
202    Some(())
203}
204
205pub(super) fn merge_optional_types(
206    left: Option<ColumnType>,
207    right: Option<ColumnType>,
208) -> Result<Option<ColumnType>, SQLError> {
209    match (left, right) {
210        (None, other) | (other, None) => Ok(other),
211        (Some(left), Some(right)) => common_type(&left, &right).map(Some),
212    }
213}
214
215pub fn common_type(left: &ColumnType, right: &ColumnType) -> Result<ColumnType, SQLError> {
216    if left == right {
217        return Ok(left.clone());
218    }
219    if left != left.without_temporal_modifiers() || right != right.without_temporal_modifiers() {
220        return common_type(
221            left.without_temporal_modifiers(),
222            right.without_temporal_modifiers(),
223        );
224    }
225    if matches!(left, ColumnType::Domain { .. }) || matches!(right, ColumnType::Domain { .. }) {
226        return common_type(base_type(left), base_type(right));
227    }
228    if let Some(numeric) = common_numeric_type(left, right) {
229        return Ok(numeric);
230    }
231    if matches!(left, ColumnType::Oid) && is_integral_type(right)
232        || matches!(right, ColumnType::Oid) && is_integral_type(left)
233    {
234        return Ok(ColumnType::Oid);
235    }
236    if left.is_character_string() && right.is_character_string() {
237        return Ok(match left {
238            ColumnType::Bpchar | ColumnType::Character(_) => ColumnType::Bpchar,
239            ColumnType::Varchar(_) => ColumnType::Varchar(None),
240            ColumnType::Name => ColumnType::Name,
241            _ => ColumnType::Text,
242        });
243    }
244    match (left, right) {
245        (ColumnType::Date, ColumnType::Timestamp) | (ColumnType::Timestamp, ColumnType::Date) => {
246            Ok(ColumnType::Timestamp)
247        }
248        (ColumnType::Date | ColumnType::Timestamp, ColumnType::TimestampTz)
249        | (ColumnType::TimestampTz, ColumnType::Date | ColumnType::Timestamp) => {
250            Ok(ColumnType::TimestampTz)
251        }
252        (ColumnType::Array(left), ColumnType::Array(right)) => {
253            common_type(left, right).map(|element| ColumnType::Array(Box::new(element)))
254        }
255        _ => Err(SQLError::TypeMismatch(format!(
256            "types {} and {} cannot be matched",
257            left.sql_name(),
258            right.sql_name()
259        ))),
260    }
261}
262
263pub(super) fn case_output_type(
264    expression: &ScalarExpr,
265    common: &ColumnType,
266    schema: &RowSchema,
267    params: &[SQLParam],
268    resolver: Option<&dyn FunctionTypeResolver>,
269) -> Result<ColumnType, SQLError> {
270    let ScalarExpr::Case {
271        base,
272        when,
273        else_branch,
274    } = expression
275    else {
276        return Ok(common.clone());
277    };
278    let mut output = None;
279    let mut include = |expression: Option<&ScalarExpr>| -> Result<(), SQLError> {
280        let ty = expression
281            .map(|expression| common_context_expression_type(expression, schema, params, resolver))
282            .transpose()?
283            .flatten();
284        let ty = ty
285            .filter(|ty| ty.regtype_name() == common.regtype_name())
286            .unwrap_or_else(|| common.without_type_modifiers());
287        output = merge_optional_types(output.take(), Some(ty))?;
288        Ok(())
289    };
290    for (condition, value) in when {
291        match constant_case_condition(base.as_deref(), condition, schema, params, resolver) {
292            Some(Value::Bool(false) | Value::Null) => {}
293            Some(Value::Bool(true)) => {
294                include(Some(value))?;
295                return Ok(output.unwrap_or_else(|| common.without_type_modifiers()));
296            }
297            _ => include(Some(value))?,
298        }
299    }
300    include(else_branch.as_deref())?;
301    Ok(output.unwrap_or_else(|| common.without_type_modifiers()))
302}
303
304fn constant_case_condition(
305    base: Option<&ScalarExpr>,
306    condition: &ScalarExpr,
307    schema: &RowSchema,
308    params: &[SQLParam],
309    resolver: Option<&dyn FunctionTypeResolver>,
310) -> Option<Value> {
311    let Some(base) = base else {
312        return constant_value(condition);
313    };
314    let left = constant_value(base)?;
315    let right = constant_value(condition)?;
316    let left_type = common_context_expression_type(base, schema, params, resolver).ok()?;
317    let right_type = common_context_expression_type(condition, schema, params, resolver).ok()?;
318    let operand_type = match (left_type, right_type) {
319        (Some(left), Some(right)) => super::equality_operand_type(&left, &right).ok()?,
320        (Some(known), None) | (None, Some(known)) => known.without_type_modifiers(),
321        (None, None) => ColumnType::Text,
322    };
323    let ty = operand_type.sql_name();
324    let left = crate::expr::cast_value(&left, &ty).ok()?;
325    let right = crate::expr::cast_value(&right, &ty).ok()?;
326    crate::expr::eval_binary_values(crate::ast::BinaryOp::Equal, &left, &right).ok()
327}
328
329fn constant_value(expression: &ScalarExpr) -> Option<Value> {
330    match expression {
331        ScalarExpr::Literal(value) => Some(value.clone()),
332        ScalarExpr::Cast { expr, ty } => crate::expr::cast_value(&constant_value(expr)?, ty).ok(),
333        ScalarExpr::Binary { op, lhs, rhs } => {
334            crate::expr::eval_binary_values(*op, &constant_value(lhs)?, &constant_value(rhs)?).ok()
335        }
336        _ => None,
337    }
338}
339
340fn is_integral_type(ty: &ColumnType) -> bool {
341    matches!(
342        base_type(ty),
343        ColumnType::SmallInteger | ColumnType::Integer | ColumnType::BigInteger
344    )
345}
346
347pub(super) fn common_numeric_type(left: &ColumnType, right: &ColumnType) -> Option<ColumnType> {
348    let rank = numeric_rank(left)?.max(numeric_rank(right)?);
349    Some(match rank {
350        0 => ColumnType::SmallInteger,
351        1 => ColumnType::Integer,
352        2 => ColumnType::BigInteger,
353        3 => numeric_type(),
354        4 => ColumnType::Real,
355        _ => ColumnType::DoublePrecision,
356    })
357}
358
359pub(super) fn numeric_rank(ty: &ColumnType) -> Option<u8> {
360    match ty {
361        ColumnType::SmallInteger => Some(0),
362        ColumnType::Integer => Some(1),
363        ColumnType::BigInteger => Some(2),
364        ColumnType::Numeric { .. } => Some(3),
365        ColumnType::Real => Some(4),
366        ColumnType::DoublePrecision => Some(5),
367        _ => None,
368    }
369}
370
371/// Array dimensions belong to values; `PostgreSQL` operator signatures identify an array by its scalar element type, including an element domain's identity.
372pub(super) fn same_operator_type(left: &ColumnType, right: &ColumnType) -> bool {
373    fn element(mut ty: &ColumnType) -> &ColumnType {
374        while let ColumnType::Array(inner) = ty {
375            ty = inner;
376        }
377        ty
378    }
379    let left = base_type(left);
380    let right = base_type(right);
381    match (left, right) {
382        (ColumnType::Array(left), ColumnType::Array(right)) => {
383            element(left).without_type_modifiers() == element(right).without_type_modifiers()
384        }
385        _ => left.without_type_modifiers() == right.without_type_modifiers(),
386    }
387}