1use 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#[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.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
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 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
371pub(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}