uqa_sql/routines/
invocation.rs1use crate::{
9 assignment::routines::RoutineValueContext,
10 ast::{
11 ColumnType, CreateFunction, FunctionParamMode, FunctionReturns, RoutineInvocationBinding,
12 },
13 expr::value_type_name,
14 routines::declaration::RoutineTypeCatalog,
15 type_resolution::canonical_routine_type_name,
16 SQLError,
17};
18use uqa_core::Value;
19pub fn output_column_names(def: &CreateFunction) -> Vec<String> {
20 def.output_params()
21 .iter()
22 .enumerate()
23 .map(|(idx, p)| {
24 if p.name.is_empty() {
25 format!("column{}", idx + 1)
26 } else {
27 p.name.clone()
28 }
29 })
30 .collect()
31}
32
33pub fn call_signature(name: &str, args: &[(Option<String>, Value)]) -> String {
34 let types = args
35 .iter()
36 .map(|(arg_name, value)| match arg_name {
37 Some(arg_name) => format!("{arg_name} => {}", value_type_name(value)),
38 None => value_type_name(value).to_string(),
39 })
40 .collect::<Vec<_>>()
41 .join(", ");
42 format!("{name}({types})")
43}
44
45pub fn routine_resolution_error(
46 kind: &str,
47 name: &str,
48 args: &[(Option<String>, Value)],
49 suffix: &str,
50) -> SQLError {
51 SQLError::Routine {
52 sqlstate: if suffix == "is not unique" {
53 "42725".into()
54 } else {
55 "42883".into()
56 },
57 message: format!("{kind} {} {suffix}", call_signature(name, args)),
58 }
59}
60
61pub fn runtime_argument_types(
62 args: &[(Option<String>, Value)],
63) -> Result<Vec<Option<ColumnType>>, SQLError> {
64 args.iter()
65 .map(|(_, value)| {
66 if matches!(value, Value::Null) {
67 Ok(None)
68 } else {
69 ColumnType::from_sql_name(value_type_name(value)).map(Some)
70 }
71 })
72 .collect()
73}
74
75pub fn specialized_definition(
76 definition: &CreateFunction,
77 invocation: &RoutineInvocationBinding,
78) -> Result<Option<CreateFunction>, SQLError> {
79 if invocation.parameter_types.len() != definition.params.len() {
80 return Err(SQLError::Internal(format!(
81 "routine `{}` has {} concrete parameter types for {} parameters",
82 definition.name,
83 invocation.parameter_types.len(),
84 definition.params.len()
85 )));
86 }
87 let parameters_match = definition
88 .params
89 .iter()
90 .zip(&invocation.parameter_types)
91 .all(|(parameter, type_name)| parameter.type_name == *type_name);
92 let return_type_matches = match (&invocation.return_type, &definition.returns) {
93 (Some(concrete), FunctionReturns::Scalar { type_name })
94 | (Some(concrete), FunctionReturns::SetOf { type_name }) => concrete == type_name,
95 (None, _) | (Some(_), FunctionReturns::None | FunctionReturns::Table) => true,
96 };
97 if parameters_match && return_type_matches {
98 return Ok(None);
99 }
100 let mut specialized = definition.clone();
101 for (parameter, type_name) in specialized
102 .params
103 .iter_mut()
104 .zip(&invocation.parameter_types)
105 {
106 parameter.type_name.clone_from(type_name);
107 }
108 if let Some(return_type) = &invocation.return_type {
109 match &mut specialized.returns {
110 FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
111 type_name.clone_from(return_type);
112 }
113 FunctionReturns::None | FunctionReturns::Table => {}
114 }
115 }
116 Ok(Some(specialized))
117}
118
119pub fn validate_anonymous_record_column_types(
120 source_types: &[Option<crate::ast::ColumnType>],
121 target_types: &[String],
122) -> Result<(), SQLError> {
123 if source_types.len() != target_types.len() {
124 return Err(anonymous_record_shape_error());
125 }
126 for (source, target) in source_types.iter().zip(target_types) {
127 let Some(source) = source else {
128 continue;
129 };
130 let source = crate::type_resolution::canonical_column_type_name(source);
131 let target = canonical_routine_type_name(target);
132 if !crate::type_resolution::routine_type_accepts_implicit_cast(&source, &target) {
133 return Err(anonymous_record_shape_error());
134 }
135 }
136 Ok(())
137}
138
139pub fn runtime_record_column_type(value: &Value) -> Option<crate::ast::ColumnType> {
140 if matches!(value, Value::Null) {
141 return None;
142 }
143 crate::ast::ColumnType::from_sql_name(crate::expr::value_type_name(value)).ok()
144}
145
146pub fn coerce_anonymous_record_value(
147 context: &dyn RoutineValueContext,
148 value: &Value,
149 type_name: &str,
150) -> Result<Value, SQLError> {
151 let target = context
152 .catalog_column_type(type_name)
153 .or_else(|| crate::ast::ColumnType::from_sql_name(type_name).ok());
154 let Some(target) = target else {
155 return crate::assignment::routines::coerce_routine_value(context, value, type_name);
156 };
157 crate::assignment::conversion::convert_value_to_column_type_with_context(
158 context,
159 value.clone(),
160 &target,
161 )
162 .map_err(|error| match error {
163 SQLError::TypeMismatch(message) if message.starts_with("value too long for type ") => {
164 SQLError::Routine {
165 sqlstate: "22001".into(),
166 message,
167 }
168 }
169 other => other,
170 })
171}
172
173pub fn anonymous_record_shape_error() -> SQLError {
174 SQLError::Routine {
175 sqlstate: "42P13".into(),
176 message: "return type mismatch in function declared to return record".into(),
177 }
178}
179pub fn call_output_schema(
180 catalog: &dyn RoutineTypeCatalog,
181 definition: &crate::ast::CreateFunction,
182 parameter_types: &[String],
183) -> Result<Option<crate::RowSchema>, SQLError> {
184 let output_indices = definition
185 .params
186 .iter()
187 .enumerate()
188 .filter_map(|(index, parameter)| {
189 matches!(
190 parameter.mode,
191 FunctionParamMode::Out | FunctionParamMode::InOut | FunctionParamMode::Table
192 )
193 .then_some(index)
194 })
195 .collect::<Vec<_>>();
196 if output_indices.is_empty() {
197 return Ok(None);
198 }
199 let columns = output_column_names(definition);
200 let column_types = output_indices
201 .into_iter()
202 .map(|index| {
203 catalog
204 .resolve_catalog_column_type(¶meter_types[index])
205 .or_else(|| crate::ast::ColumnType::from_sql_name(¶meter_types[index]).ok())
206 .map(Some)
207 .ok_or_else(|| {
208 SQLError::TypeMismatch(format!("unknown type `{}`", parameter_types[index]))
209 })
210 })
211 .collect::<Result<Vec<_>, SQLError>>()?;
212 Ok(Some(crate::RowSchema::with_types(columns, column_types)))
213}