uqa_execution/type_resolution/
common.rs1use 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#[doc(hidden)]
17pub type FunctionCallArgumentSignature = (Vec<Option<String>>, Vec<Option<ColumnType>>, bool);
18
19#[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
94pub 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#[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#[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}