1use crate::{
10 ast::{
11 AlterRoutineStmt, ColumnDef, ColumnType, CreateFunction, FunctionBody, FunctionParamMode,
12 FunctionReturns, RoutineColumnTypeReference,
13 },
14 type_resolution::canonical_routine_type_name,
15 SQLError,
16};
17
18pub trait RoutineTypeCatalog {
19 fn try_describe_table(&self, reference: &str) -> Result<Option<Vec<ColumnDef>>, String>;
20 fn resolve_catalog_column_type(&self, name: &str) -> Option<ColumnType>;
21 fn resolve_catalog_column_type_name(&self, name: &str) -> Result<ColumnType, SQLError>;
22 fn resolve_catalog_domain_type_by_oid(&self, oid: u32) -> Option<ColumnType>;
23}
24
25pub fn resolve_routine_type_references(
26 catalog: &dyn RoutineTypeCatalog,
27 def: &mut CreateFunction,
28) -> Result<(), SQLError> {
29 for parameter in &mut def.params {
30 parameter.type_name = resolve_routine_type_name_with_reference(
31 catalog,
32 ¶meter.type_name,
33 ROUTINE_PARAMETER_PSEUDO_TYPES,
34 parameter.type_reference.as_ref(),
35 )?;
36 parameter.type_reference = None;
37 }
38 match &mut def.returns {
39 FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
40 *type_name = resolve_routine_type_name_with_reference(
41 catalog,
42 type_name,
43 ROUTINE_RESULT_PSEUDO_TYPES,
44 def.return_type_reference.as_ref(),
45 )?;
46 }
47 FunctionReturns::None | FunctionReturns::Table => {}
48 }
49 def.return_type_reference = None;
50 Ok(())
51}
52
53pub fn resolve_alter_routine_identity_types(
54 catalog: &dyn RoutineTypeCatalog,
55 stmt: &AlterRoutineStmt,
56) -> Result<Option<Vec<String>>, SQLError> {
57 resolve_routine_identity_types(
58 catalog,
59 stmt.arg_types.as_deref(),
60 &stmt.arg_type_references,
61 "ALTER routine",
62 )
63}
64
65pub fn resolve_routine_identity_types(
66 catalog: &dyn RoutineTypeCatalog,
67 types: Option<&[String]>,
68 references: &[Option<RoutineColumnTypeReference>],
69 context: &str,
70) -> Result<Option<Vec<String>>, SQLError> {
71 let Some(types) = types else {
72 if !references.is_empty() {
73 return Err(SQLError::Internal(format!(
74 "{context} omitted its identity types but retained type references"
75 )));
76 }
77 return Ok(None);
78 };
79 if !references.is_empty() && references.len() != types.len() {
80 return Err(SQLError::Internal(format!(
81 "{context} has {} identity types but {} type references",
82 types.len(),
83 references.len()
84 )));
85 }
86 types
87 .iter()
88 .enumerate()
89 .map(|(index, type_name)| {
90 resolve_routine_type_name_with_reference(
91 catalog,
92 type_name,
93 ROUTINE_PARAMETER_PSEUDO_TYPES,
94 references.get(index).and_then(Option::as_ref),
95 )
96 .map(|resolved| canonical_routine_type_name(&resolved))
97 })
98 .collect::<Result<Vec<_>, _>>()
99 .map(Some)
100}
101
102const POLYMORPHIC_PSEUDO_TYPES: &[&str] = &[
103 "anyelement",
104 "anyarray",
105 "anynonarray",
106 "anyenum",
107 "anyrange",
108 "anymultirange",
109 "anycompatible",
110 "anycompatiblearray",
111 "anycompatiblenonarray",
112 "anycompatiblerange",
113 "anycompatiblemultirange",
114];
115
116const ROUTINE_PARAMETER_PSEUDO_TYPES: &[&str] = &[
117 "record",
118 "refcursor",
119 "cstring",
120 "any",
121 "void",
122 "trigger",
123 "internal",
124 "event_trigger",
125 "anyelement",
126 "anyarray",
127 "anynonarray",
128 "anyenum",
129 "anyrange",
130 "anymultirange",
131 "anycompatible",
132 "anycompatiblearray",
133 "anycompatiblenonarray",
134 "anycompatiblerange",
135 "anycompatiblemultirange",
136];
137
138const ROUTINE_RESULT_PSEUDO_TYPES: &[&str] = &[
139 "record",
140 "refcursor",
141 "cstring",
142 "any",
143 "void",
144 "trigger",
145 "internal",
146 "event_trigger",
147 "anyelement",
148 "anyarray",
149 "anynonarray",
150 "anyenum",
151 "anyrange",
152 "anymultirange",
153 "anycompatible",
154 "anycompatiblearray",
155 "anycompatiblenonarray",
156 "anycompatiblerange",
157 "anycompatiblemultirange",
158];
159
160fn resolve_routine_type_name_with_reference(
161 catalog: &dyn RoutineTypeCatalog,
162 type_name: &str,
163 allowed_pseudo_types: &[&str],
164 structured_reference: Option<&RoutineColumnTypeReference>,
165) -> Result<String, SQLError> {
166 let mut base = type_name.trim();
167 let mut array_dimensions = 0usize;
168 while let Some(element) = base.strip_suffix("[]") {
169 base = element.trim_end();
170 array_dimensions += 1;
171 }
172 let resolved = if base
173 .get(base.len().saturating_sub("%type".len())..)
174 .is_some_and(|suffix| suffix.eq_ignore_ascii_case("%type"))
175 {
176 let reference = structured_reference.ok_or_else(|| {
177 SQLError::Internal(format!(
178 "routine type reference `{type_name}` is missing structured relation-column identity"
179 ))
180 })?;
181 let table = reference.relation_reference();
182 let columns = catalog
183 .try_describe_table(&table)
184 .map_err(|error| {
185 SQLError::Internal(format!(
186 "resolve routine type reference `{type_name}`: {error}"
187 ))
188 })?
189 .ok_or_else(|| SQLError::UnknownTable(table.clone()))?;
190 columns
191 .into_iter()
192 .find(|definition| definition.name == reference.column)
193 .map(|definition| definition.ty)
194 .ok_or_else(|| SQLError::UnknownColumn(reference.type_reference()))?
195 } else {
196 let canonical = canonical_routine_type_name(base);
197 if allowed_pseudo_types.contains(&canonical.as_str()) {
198 if array_dimensions != 0 {
199 return Err(SQLError::Routine {
200 sqlstate: "42704".into(),
201 message: format!("type `{type_name}` does not exist"),
202 });
203 }
204 return Ok(canonical);
205 }
206 catalog.resolve_catalog_column_type_name(base)?
207 };
208 let mut resolved = resolved;
209 for _ in 0..array_dimensions {
210 resolved = ColumnType::Array(Box::new(resolved));
211 }
212 Ok(resolved.sql_name())
213}
214
215pub fn resolve_plpgsql_datum_types(
216 catalog: &dyn RoutineTypeCatalog,
217 function: &mut crate::plpgsql::PLpgSQLFunction,
218) -> Result<(), SQLError> {
219 for datum in &mut function.datums {
220 let crate::plpgsql::PLpgSQLDatum::Var(variable) = datum else {
221 continue;
222 };
223 if variable.type_reference.is_none() {
224 if let Some(ty) = variable
225 .type_oid
226 .and_then(|oid| catalog.resolve_catalog_domain_type_by_oid(oid))
227 {
228 variable.type_name = ty.sql_name();
229 continue;
230 }
231 }
232 variable.type_name = resolve_routine_type_name_with_reference(
233 catalog,
234 &variable.type_name,
235 &[
236 "record",
237 "refcursor",
238 "anyelement",
239 "anyarray",
240 "anynonarray",
241 "anyenum",
242 "anyrange",
243 "anymultirange",
244 "anycompatible",
245 "anycompatiblearray",
246 "anycompatiblenonarray",
247 "anycompatiblerange",
248 "anycompatiblemultirange",
249 ],
250 variable.type_reference.as_ref(),
251 )?;
252 variable.type_reference = None;
253 }
254 Ok(())
255}
256
257pub(super) fn validate_routine_declaration(
258 catalog: &dyn RoutineTypeCatalog,
259 def: &CreateFunction,
260) -> Result<(), SQLError> {
261 validate_variadic_declaration(catalog, def)?;
262 let inputs = validate_routine_input_types(def)?;
263 if matches!(def.body, FunctionBody::Statements(_)) && inputs.any {
264 return Err(routine_definition_error(
265 "SQL function with unquoted function body cannot have polymorphic arguments",
266 ));
267 }
268 validate_routine_output_types(def, &inputs)
269}
270
271pub(super) fn routine_parameter_regrole_constants(
272 catalog: &dyn RoutineTypeCatalog,
273 def: &CreateFunction,
274) -> crate::catalog::regrole_dependencies::StoredRegroleConstants {
275 let mut constants = crate::catalog::regrole_dependencies::StoredRegroleConstants::default();
276 for parameter in &def.params {
277 let Some(default) = parameter.default.as_ref() else {
278 continue;
279 };
280 let target = catalog
281 .resolve_catalog_column_type(¶meter.type_name)
282 .or_else(|| ColumnType::from_sql_name(¶meter.type_name).ok());
283 constants.collect_expression(default, target.as_ref());
284 }
285 constants
286}
287
288fn validate_variadic_declaration(
289 catalog: &dyn RoutineTypeCatalog,
290 def: &CreateFunction,
291) -> Result<(), SQLError> {
292 let variadic_positions = def
293 .params
294 .iter()
295 .enumerate()
296 .filter_map(|(index, parameter)| {
297 (parameter.mode == FunctionParamMode::Variadic).then_some(index)
298 })
299 .collect::<Vec<_>>();
300 if variadic_positions.len() > 1 {
301 return Err(routine_definition_error(
302 "VARIADIC parameter must be the last parameter",
303 ));
304 }
305 if let Some(&variadic_index) = variadic_positions.first() {
306 let parameter = &def.params[variadic_index];
307 if !routine_declaration_is_array(catalog, ¶meter.type_name) {
308 return Err(routine_definition_error(
309 "VARIADIC parameter must be an array",
310 ));
311 }
312 let has_later_input = def.params[variadic_index + 1..].iter().any(|parameter| {
313 matches!(
314 parameter.mode,
315 FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
316 )
317 });
318 if has_later_input || def.is_procedure && variadic_index + 1 != def.params.len() {
319 return Err(routine_definition_error(
320 "VARIADIC parameter must be the last parameter",
321 ));
322 }
323 }
324 Ok(())
325}
326
327#[derive(Default)]
328struct PolymorphicInputs {
329 simple: bool,
330 compatible: bool,
331 any: bool,
332}
333
334fn validate_routine_input_types(def: &CreateFunction) -> Result<PolymorphicInputs, SQLError> {
335 let mut inputs = PolymorphicInputs::default();
336 for parameter in &def.params {
337 let type_name = canonical_routine_type_name(¶meter.type_name);
338 let is_input = matches!(
339 parameter.mode,
340 FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
341 );
342 if let Some(family) = polymorphic_family(&type_name) {
343 inputs.any |= is_input;
344 if is_input {
345 match family {
346 RoutinePolymorphicFamily::Simple => inputs.simple = true,
347 RoutinePolymorphicFamily::Compatible => inputs.compatible = true,
348 }
349 }
350 continue;
351 }
352 if ROUTINE_PARAMETER_PSEUDO_TYPES.contains(&type_name.as_str()) {
353 let supported = match type_name.as_str() {
354 "record" => !is_input || def.language == "plpgsql",
355 "refcursor" => true,
356 _ => false,
357 };
358 if !supported {
359 return Err(routine_definition_error(format!(
360 "{} routines cannot have arguments of type {type_name}",
361 def.language
362 )));
363 }
364 }
365 }
366 Ok(inputs)
367}
368
369fn validate_routine_output_types(
370 def: &CreateFunction,
371 inputs: &PolymorphicInputs,
372) -> Result<(), SQLError> {
373 let mut output_types = def
374 .output_params()
375 .into_iter()
376 .map(|parameter| parameter.type_name.as_str())
377 .collect::<Vec<_>>();
378 if let FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } =
379 &def.returns
380 {
381 output_types.push(type_name);
382 }
383 for output_type in output_types {
384 let type_name = canonical_routine_type_name(output_type);
385 match polymorphic_family(&type_name) {
386 Some(RoutinePolymorphicFamily::Simple) if !inputs.simple => {
387 return Err(routine_definition_error(format!(
388 "cannot determine result data type: a result of type {type_name} requires at least one simple polymorphic input"
389 )));
390 }
391 Some(RoutinePolymorphicFamily::Compatible) if !inputs.compatible => {
392 return Err(routine_definition_error(format!(
393 "cannot determine result data type: a result of type {type_name} requires at least one compatible polymorphic input"
394 )));
395 }
396 None if ROUTINE_RESULT_PSEUDO_TYPES.contains(&type_name.as_str())
397 && !matches!(type_name.as_str(), "record" | "refcursor" | "void")
398 && !(type_name == "trigger"
399 && def.language == "plpgsql"
400 && !def.is_procedure
401 && def.params.is_empty()
402 && matches!(def.returns, FunctionReturns::Scalar { .. })) =>
403 {
404 return Err(routine_definition_error(format!(
405 "{} routines cannot return type {type_name}",
406 def.language
407 )));
408 }
409 Some(_) | None => {}
410 }
411 }
412 Ok(())
413}
414
415#[derive(Debug, Clone, Copy, PartialEq, Eq)]
416enum RoutinePolymorphicFamily {
417 Simple,
418 Compatible,
419}
420
421fn polymorphic_family(type_name: &str) -> Option<RoutinePolymorphicFamily> {
422 if !POLYMORPHIC_PSEUDO_TYPES.contains(&type_name) {
423 return None;
424 }
425 Some(if type_name.starts_with("anycompatible") {
426 RoutinePolymorphicFamily::Compatible
427 } else {
428 RoutinePolymorphicFamily::Simple
429 })
430}
431
432fn routine_declaration_is_array(catalog: &dyn RoutineTypeCatalog, type_name: &str) -> bool {
433 let canonical = canonical_routine_type_name(type_name);
434 canonical.ends_with("[]")
435 || matches!(
436 canonical.as_str(),
437 "anyarray" | "anycompatiblearray" | "int2vector" | "oidvector"
438 )
439 || catalog
440 .resolve_catalog_column_type(&canonical)
441 .is_some_and(|ty| routine_column_type_is_array(&ty))
442}
443
444fn routine_column_type_is_array(ty: &ColumnType) -> bool {
445 match ty {
446 ColumnType::Array(_) | ColumnType::AnyArray => true,
447 ColumnType::Domain { base, .. } => routine_column_type_is_array(base),
448 _ => false,
449 }
450}
451
452pub(super) fn routine_definition_error(message: impl Into<String>) -> SQLError {
453 SQLError::Routine {
454 sqlstate: "42P13".into(),
455 message: message.into(),
456 }
457}