Skip to main content

uqa_execution/type_resolution/
common.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7use uqa_core::Value;
8use uqa_sql::ast::ColumnType;
9use uqa_sql::{SQLError, SQLParam};
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
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 matches!(left, ColumnType::Domain { .. }) || matches!(right, ColumnType::Domain { .. }) {
220        return common_type(base_type(left), base_type(right));
221    }
222    if let Some(numeric) = common_numeric_type(left, right) {
223        return Ok(numeric);
224    }
225    if matches!(left, ColumnType::Oid) && is_integral_type(right)
226        || matches!(right, ColumnType::Oid) && is_integral_type(left)
227    {
228        return Ok(ColumnType::Oid);
229    }
230    if left.is_character_string() && right.is_character_string() {
231        return Ok(match left {
232            ColumnType::Bpchar | ColumnType::Character(_) => ColumnType::Bpchar,
233            ColumnType::Varchar(_) => ColumnType::Varchar(None),
234            ColumnType::Name => ColumnType::Name,
235            _ => ColumnType::Text,
236        });
237    }
238    match (left, right) {
239        (ColumnType::Date, ColumnType::Timestamp) | (ColumnType::Timestamp, ColumnType::Date) => {
240            Ok(ColumnType::Timestamp)
241        }
242        (ColumnType::Date | ColumnType::Timestamp, ColumnType::TimestampTz)
243        | (ColumnType::TimestampTz, ColumnType::Date | ColumnType::Timestamp) => {
244            Ok(ColumnType::TimestampTz)
245        }
246        (ColumnType::Array(left), ColumnType::Array(right)) => {
247            common_type(left, right).map(|element| ColumnType::Array(Box::new(element)))
248        }
249        _ => Err(SQLError::TypeMismatch(format!(
250            "types {} and {} cannot be matched",
251            left.sql_name(),
252            right.sql_name()
253        ))),
254    }
255}
256
257fn is_integral_type(ty: &ColumnType) -> bool {
258    matches!(
259        base_type(ty),
260        ColumnType::SmallInteger | ColumnType::Integer | ColumnType::BigInteger
261    )
262}
263
264pub(super) fn common_numeric_type(left: &ColumnType, right: &ColumnType) -> Option<ColumnType> {
265    let rank = numeric_rank(left)?.max(numeric_rank(right)?);
266    Some(match rank {
267        0 => ColumnType::SmallInteger,
268        1 => ColumnType::Integer,
269        2 => ColumnType::BigInteger,
270        3 => numeric_type(),
271        4 => ColumnType::Real,
272        _ => ColumnType::DoublePrecision,
273    })
274}
275
276pub(super) fn numeric_rank(ty: &ColumnType) -> Option<u8> {
277    match ty {
278        ColumnType::SmallInteger => Some(0),
279        ColumnType::Integer => Some(1),
280        ColumnType::BigInteger => Some(2),
281        ColumnType::Numeric { .. } => Some(3),
282        ColumnType::Real => Some(4),
283        ColumnType::DoublePrecision => Some(5),
284        _ => None,
285    }
286}