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_user_type_by_oid(&self, oid: u32) -> Option<ColumnType>;
23 fn require_type_usage(&self, ty: &ColumnType) -> Result<(), SQLError>;
25 fn format_type(&self, ty: &ColumnType) -> Result<String, SQLError>;
27 fn format_type_oid(&self, oid: u32) -> Result<String, SQLError> {
29 match self.resolve_catalog_user_type_by_oid(oid) {
30 Some(ty) => self.format_type(&ty),
31 None => Ok(crate::catalog::type_metadata::catalog_type_name(i64::from(oid)).into()),
32 }
33 }
34 fn composite_types(&self) -> Option<&dyn crate::expr::composites::CompositeTypeCatalog> {
36 None
37 }
38}
39
40#[must_use]
42pub fn result_type_name(def: &CreateFunction) -> &str {
43 let outputs = def.output_params();
44 if def.is_procedure {
45 return if outputs.is_empty() { "void" } else { "record" };
46 }
47 match &def.returns {
48 FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => type_name,
49 FunctionReturns::Table | FunctionReturns::None => match outputs.as_slice() {
50 [output] => &output.type_name,
51 [] => "void",
52 _ => "record",
53 },
54 }
55}
56
57pub fn resolve_routine_type_references(
59 context: &super::compilation::RoutineCompilationContext<'_>,
60 def: &mut CreateFunction,
61) -> Result<(), SQLError> {
62 let catalog = context.types;
63 let mut have_defaults = false;
64 let mut after_variadic = false;
65 for index in 0..def.params.len() {
66 let (previous, rest) = def.params.split_at_mut(index);
67 let parameter = &mut rest[0];
68 let written = parameter.written_type.take();
69 parameter.type_name = resolve_used_routine_type(
70 catalog,
71 ¶meter.type_name,
72 ROUTINE_PARAMETER_PSEUDO_TYPES,
73 parameter.type_reference.as_ref(),
74 MissingRoutineType::Parameter(written.as_deref()),
75 )?;
76 parameter.type_reference = None;
77 let input = !matches!(
78 parameter.mode,
79 FunctionParamMode::Out | FunctionParamMode::Table
80 );
81 let output = !matches!(
82 parameter.mode,
83 FunctionParamMode::In | FunctionParamMode::Variadic
84 );
85 if input && after_variadic {
86 return Err(routine_definition_error(
87 "VARIADIC parameter must be the last input parameter",
88 ));
89 }
90 if output && def.is_procedure && after_variadic {
91 return Err(routine_definition_error(
92 "VARIADIC parameter must be the last parameter",
93 ));
94 }
95 if parameter.mode == FunctionParamMode::Variadic {
96 after_variadic = true;
97 if !variadic_type_is_array(catalog, ¶meter.type_name) {
98 return Err(routine_definition_error(
99 "VARIADIC parameter must be an array",
100 ));
101 }
102 }
103 if !parameter.name.is_empty()
104 && previous.iter().any(|earlier| {
105 earlier.name == parameter.name
106 && parameter_names_conflict(parameter.mode, earlier.mode)
107 })
108 {
109 return Err(routine_definition_error(format!(
110 "parameter name \"{}\" used more than once",
111 parameter.name
112 )));
113 }
114 if let Some(default) = &mut parameter.default {
115 if !input {
116 return Err(routine_definition_error(
117 "only input parameters can have default values",
118 ));
119 }
120 parameter.default_type =
121 super::defaults::analyze_parameter_default(context, default, ¶meter.type_name)?;
122 have_defaults = true;
123 } else if input && have_defaults {
124 return Err(routine_definition_error(
125 "input parameters after one with a default value must also have defaults",
126 ));
127 } else if def.is_procedure && have_defaults {
128 return Err(routine_definition_error(
129 "procedure OUT parameters cannot appear after one with a default value",
130 ));
131 }
132 }
133 let written = def.return_written_type.take();
134 match &mut def.returns {
135 FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
136 *type_name = resolve_used_routine_type(
137 catalog,
138 type_name,
139 ROUTINE_RESULT_PSEUDO_TYPES,
140 def.return_type_reference.as_ref(),
141 MissingRoutineType::Result(written.as_deref()),
142 )?;
143 }
144 FunctionReturns::None | FunctionReturns::Table => {}
145 }
146 def.return_type_reference = None;
147 Ok(())
148}
149
150fn parameter_names_conflict(current: FunctionParamMode, earlier: FunctionParamMode) -> bool {
152 let pure_input = |mode| matches!(mode, FunctionParamMode::In | FunctionParamMode::Variadic);
153 let pure_output = |mode| matches!(mode, FunctionParamMode::Out | FunctionParamMode::Table);
154 !(pure_input(current) && pure_output(earlier) || pure_input(earlier) && pure_output(current))
155}
156
157fn variadic_type_is_array(catalog: &dyn RoutineTypeCatalog, type_name: &str) -> bool {
159 canonical_routine_type_name(type_name) == "any"
160 || routine_declaration_is_array(catalog, type_name)
161}
162
163#[derive(Clone, Copy)]
165enum MissingRoutineType<'a> {
166 Parameter(Option<&'a str>),
167 Result(Option<&'a str>),
168}
169
170impl MissingRoutineType<'_> {
171 fn error(self, type_name: &str) -> SQLError {
173 let (written, quoted) = match self {
174 Self::Parameter(written) => (written, false),
175 Self::Result(written) => (written, true),
176 };
177 let name = written.map_or_else(|| type_name.to_string(), str::to_string);
178 SQLError::Routine {
179 sqlstate: "42704".into(),
180 message: if quoted {
181 format!("type \"{name}\" does not exist")
182 } else {
183 format!("type {name} does not exist")
184 },
185 }
186 }
187}
188
189pub fn resolve_alter_routine_identity_types(
190 catalog: &dyn RoutineTypeCatalog,
191 stmt: &AlterRoutineStmt,
192) -> Result<Option<Vec<String>>, SQLError> {
193 resolve_routine_identity_types(
194 catalog,
195 stmt.arg_types.as_deref(),
196 &stmt.arg_type_references,
197 "ALTER routine",
198 )
199}
200
201pub fn resolve_routine_identity_types(
202 catalog: &dyn RoutineTypeCatalog,
203 types: Option<&[String]>,
204 references: &[Option<RoutineColumnTypeReference>],
205 context: &str,
206) -> Result<Option<Vec<String>>, SQLError> {
207 let Some(types) = types else {
208 if !references.is_empty() {
209 return Err(SQLError::Internal(format!(
210 "{context} omitted its identity types but retained type references"
211 )));
212 }
213 return Ok(None);
214 };
215 if !references.is_empty() && references.len() != types.len() {
216 return Err(SQLError::Internal(format!(
217 "{context} has {} identity types but {} type references",
218 types.len(),
219 references.len()
220 )));
221 }
222 types
223 .iter()
224 .enumerate()
225 .map(|(index, type_name)| {
226 resolve_routine_type_name_with_reference(
227 catalog,
228 type_name,
229 ROUTINE_PARAMETER_PSEUDO_TYPES,
230 references.get(index).and_then(Option::as_ref),
231 )
232 .map(|resolved| canonical_routine_type_name(&resolved))
233 })
234 .collect::<Result<Vec<_>, _>>()
235 .map(Some)
236}
237
238const POLYMORPHIC_PSEUDO_TYPES: &[&str] = &[
239 "anyelement",
240 "anyarray",
241 "anynonarray",
242 "anyenum",
243 "anyrange",
244 "anymultirange",
245 "anycompatible",
246 "anycompatiblearray",
247 "anycompatiblenonarray",
248 "anycompatiblerange",
249 "anycompatiblemultirange",
250];
251
252const ROUTINE_PARAMETER_PSEUDO_TYPES: &[&str] = &[
253 "record",
254 "refcursor",
255 "cstring",
256 "any",
257 "void",
258 "trigger",
259 "internal",
260 "event_trigger",
261 "anyelement",
262 "anyarray",
263 "anynonarray",
264 "anyenum",
265 "anyrange",
266 "anymultirange",
267 "anycompatible",
268 "anycompatiblearray",
269 "anycompatiblenonarray",
270 "anycompatiblerange",
271 "anycompatiblemultirange",
272];
273
274const ROUTINE_RESULT_PSEUDO_TYPES: &[&str] = &[
275 "record",
276 "refcursor",
277 "cstring",
278 "any",
279 "void",
280 "trigger",
281 "internal",
282 "event_trigger",
283 "anyelement",
284 "anyarray",
285 "anynonarray",
286 "anyenum",
287 "anyrange",
288 "anymultirange",
289 "anycompatible",
290 "anycompatiblearray",
291 "anycompatiblenonarray",
292 "anycompatiblerange",
293 "anycompatiblemultirange",
294];
295
296enum DeclaredRoutineType {
298 Pseudo(String),
299 Catalog(ColumnType),
300}
301
302impl DeclaredRoutineType {
303 fn into_catalog_name(self) -> String {
305 match self {
306 Self::Pseudo(name) => name,
307 Self::Catalog(ty) => ty.catalog_name(),
308 }
309 }
310}
311
312fn resolve_used_routine_type(
314 catalog: &dyn RoutineTypeCatalog,
315 type_name: &str,
316 allowed_pseudo_types: &[&str],
317 structured_reference: Option<&RoutineColumnTypeReference>,
318 missing: MissingRoutineType<'_>,
319) -> Result<String, SQLError> {
320 let declared = resolve_declared_routine_type(
321 catalog,
322 type_name,
323 allowed_pseudo_types,
324 structured_reference,
325 Some(missing),
326 )?;
327 if let DeclaredRoutineType::Catalog(ty) = &declared {
328 catalog.require_type_usage(ty)?;
329 }
330 Ok(declared.into_catalog_name())
331}
332
333fn resolve_routine_type_name_with_reference(
334 catalog: &dyn RoutineTypeCatalog,
335 type_name: &str,
336 allowed_pseudo_types: &[&str],
337 structured_reference: Option<&RoutineColumnTypeReference>,
338) -> Result<String, SQLError> {
339 resolve_declared_routine_type(
340 catalog,
341 type_name,
342 allowed_pseudo_types,
343 structured_reference,
344 None,
345 )
346 .map(DeclaredRoutineType::into_catalog_name)
347}
348
349fn resolve_declared_routine_type(
350 catalog: &dyn RoutineTypeCatalog,
351 type_name: &str,
352 allowed_pseudo_types: &[&str],
353 structured_reference: Option<&RoutineColumnTypeReference>,
354 missing: Option<MissingRoutineType<'_>>,
355) -> Result<DeclaredRoutineType, SQLError> {
356 let mut base = type_name.trim();
357 let mut array_dimensions = 0usize;
358 while let Some(element) = base.strip_suffix("[]") {
359 base = element.trim_end();
360 array_dimensions += 1;
361 }
362 let resolved = if base
363 .get(base.len().saturating_sub("%type".len())..)
364 .is_some_and(|suffix| suffix.eq_ignore_ascii_case("%type"))
365 {
366 let reference = structured_reference.ok_or_else(|| {
367 SQLError::Internal(format!(
368 "routine type reference `{type_name}` is missing structured relation-column identity"
369 ))
370 })?;
371 let table = reference.relation_reference();
372 let columns = catalog
373 .try_describe_table(&table)
374 .map_err(|error| {
375 SQLError::Internal(format!(
376 "resolve routine type reference `{type_name}`: {error}"
377 ))
378 })?
379 .ok_or_else(|| SQLError::UnknownTable(table.clone()))?;
380 columns
381 .into_iter()
382 .find(|definition| definition.name == reference.column)
383 .map(|definition| definition.ty)
384 .ok_or_else(|| SQLError::UnknownColumn(reference.type_reference()))?
385 } else {
386 let canonical = canonical_routine_type_name(base);
387 if allowed_pseudo_types.contains(&canonical.as_str()) {
388 if array_dimensions != 0 {
389 return Err(SQLError::Routine {
390 sqlstate: "42704".into(),
391 message: format!("type `{type_name}` does not exist"),
392 });
393 }
394 return Ok(DeclaredRoutineType::Pseudo(canonical));
395 }
396 match missing {
397 Some(missing) if catalog.resolve_catalog_column_type(base).is_none() => {
398 return Err(missing.error(type_name));
399 }
400 _ => catalog.resolve_catalog_column_type_name(base)?,
401 }
402 };
403 let mut resolved = resolved;
404 for _ in 0..array_dimensions {
405 resolved = ColumnType::Array(Box::new(resolved));
406 }
407 Ok(DeclaredRoutineType::Catalog(resolved))
408}
409
410pub fn resolve_plpgsql_datum_types(
411 catalog: &dyn RoutineTypeCatalog,
412 function: &mut crate::plpgsql::PLpgSQLFunction,
413) -> Result<(), SQLError> {
414 for datum in &mut function.datums {
415 let crate::plpgsql::PLpgSQLDatum::Var(variable) = datum else {
416 continue;
417 };
418 if variable.type_reference.is_none() {
419 if let Some(ty) = variable
420 .type_oid
421 .and_then(|oid| catalog.resolve_catalog_user_type_by_oid(oid))
422 {
423 variable.type_name = ty.catalog_name();
425 continue;
426 }
427 }
428 variable.type_name = resolve_routine_type_name_with_reference(
429 catalog,
430 &variable.type_name,
431 &[
432 "record",
433 "refcursor",
434 "anyelement",
435 "anyarray",
436 "anynonarray",
437 "anyenum",
438 "anyrange",
439 "anymultirange",
440 "anycompatible",
441 "anycompatiblearray",
442 "anycompatiblenonarray",
443 "anycompatiblerange",
444 "anycompatiblemultirange",
445 ],
446 variable.type_reference.as_ref(),
447 )?;
448 variable.type_reference = None;
449 }
450 Ok(())
451}
452
453pub(super) fn validate_routine_declaration(def: &CreateFunction) -> Result<(), SQLError> {
454 let inputs = validate_routine_input_types(def)?;
455 if matches!(def.body, FunctionBody::Statements(_)) && inputs.any {
456 return Err(routine_definition_error(
457 "SQL function with unquoted function body cannot have polymorphic arguments",
458 ));
459 }
460 validate_routine_output_types(def, &inputs)
461}
462
463pub(super) fn routine_parameter_regrole_constants(
464 catalog: &dyn RoutineTypeCatalog,
465 def: &CreateFunction,
466) -> crate::catalog::regrole_dependencies::StoredRegroleConstants {
467 let mut constants = crate::catalog::regrole_dependencies::StoredRegroleConstants::default();
468 for parameter in &def.params {
469 let Some(default) = parameter.default.as_ref() else {
470 continue;
471 };
472 let target = catalog
473 .resolve_catalog_column_type(¶meter.type_name)
474 .or_else(|| ColumnType::from_sql_name(¶meter.type_name).ok());
475 constants.collect_expression(default, target.as_ref());
476 }
477 constants
478}
479
480#[derive(Default)]
481struct PolymorphicInputs {
482 simple: bool,
483 compatible: bool,
484 any: bool,
485}
486
487fn validate_routine_input_types(def: &CreateFunction) -> Result<PolymorphicInputs, SQLError> {
488 let mut inputs = PolymorphicInputs::default();
489 for parameter in &def.params {
490 let type_name = canonical_routine_type_name(¶meter.type_name);
491 let is_input = matches!(
492 parameter.mode,
493 FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
494 );
495 if let Some(family) = polymorphic_family(&type_name) {
496 inputs.any |= is_input;
497 if is_input {
498 match family {
499 RoutinePolymorphicFamily::Simple => inputs.simple = true,
500 RoutinePolymorphicFamily::Compatible => inputs.compatible = true,
501 }
502 }
503 continue;
504 }
505 if ROUTINE_PARAMETER_PSEUDO_TYPES.contains(&type_name.as_str()) {
506 let supported = match type_name.as_str() {
507 "record" => !is_input || def.language == "plpgsql",
508 "refcursor" => true,
509 _ => false,
510 };
511 if !supported {
512 return Err(pseudo_type_error(
513 def,
514 format!("cannot have arguments of type {type_name}"),
515 format!("cannot accept type {type_name}"),
516 ));
517 }
518 }
519 }
520 Ok(inputs)
521}
522
523fn validate_routine_output_types(
524 def: &CreateFunction,
525 inputs: &PolymorphicInputs,
526) -> Result<(), SQLError> {
527 let mut output_types = def
528 .output_params()
529 .into_iter()
530 .map(|parameter| parameter.type_name.as_str())
531 .collect::<Vec<_>>();
532 if let FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } =
533 &def.returns
534 {
535 output_types.push(type_name);
536 }
537 for output_type in output_types {
538 let type_name = canonical_routine_type_name(output_type);
539 match polymorphic_family(&type_name) {
540 Some(RoutinePolymorphicFamily::Simple) if !inputs.simple => {
541 return Err(polymorphic_result_error(
542 &type_name,
543 "anyelement, anyarray, anynonarray, anyenum, anyrange, or anymultirange",
544 ));
545 }
546 Some(RoutinePolymorphicFamily::Compatible) if !inputs.compatible => {
547 return Err(polymorphic_result_error(&type_name, "anycompatible, anycompatiblearray, anycompatiblenonarray, anycompatiblerange, or anycompatiblemultirange"));
548 }
549 None if ROUTINE_RESULT_PSEUDO_TYPES.contains(&type_name.as_str())
550 && !matches!(type_name.as_str(), "record" | "refcursor" | "void")
551 && !(type_name == "trigger"
552 && def.language == "plpgsql"
553 && !def.is_procedure
554 && matches!(def.returns, FunctionReturns::Scalar { .. })) =>
555 {
556 let message = format!("cannot return type {type_name}");
557 return Err(pseudo_type_error(def, message.clone(), message));
558 }
559 Some(_) | None => {}
560 }
561 }
562 Ok(())
563}
564
565fn polymorphic_result_error(result: &str, inputs: &str) -> SQLError {
566 SQLError::Diagnostic {
567 sqlstate: "42P13".into(),
568 message: "cannot determine result data type".into(),
569 detail: Some(format!(
570 "A result of type {result} requires at least one input of type {inputs}."
571 )),
572 hint: None,
573 }
574}
575
576fn pseudo_type_error(def: &CreateFunction, sql: String, plpgsql: String) -> SQLError {
578 if def.language == "plpgsql" {
579 SQLError::Routine {
580 sqlstate: "0A000".into(),
581 message: format!("PL/pgSQL functions {plpgsql}"),
582 }
583 } else {
584 routine_definition_error(format!("SQL functions {sql}"))
585 }
586}
587
588#[derive(Debug, Clone, Copy, PartialEq, Eq)]
589enum RoutinePolymorphicFamily {
590 Simple,
591 Compatible,
592}
593
594fn polymorphic_family(type_name: &str) -> Option<RoutinePolymorphicFamily> {
595 if !POLYMORPHIC_PSEUDO_TYPES.contains(&type_name) {
596 return None;
597 }
598 Some(if type_name.starts_with("anycompatible") {
599 RoutinePolymorphicFamily::Compatible
600 } else {
601 RoutinePolymorphicFamily::Simple
602 })
603}
604
605fn routine_declaration_is_array(catalog: &dyn RoutineTypeCatalog, type_name: &str) -> bool {
606 let canonical = canonical_routine_type_name(type_name);
607 canonical.ends_with("[]")
608 || matches!(
609 canonical.as_str(),
610 "anyarray" | "anycompatiblearray" | "int2vector" | "oidvector"
611 )
612 || catalog
613 .resolve_catalog_column_type(&canonical)
614 .is_some_and(|ty| routine_column_type_is_array(&ty))
615}
616
617fn routine_column_type_is_array(ty: &ColumnType) -> bool {
618 match ty {
619 ColumnType::Array(_) | ColumnType::AnyArray => true,
620 ColumnType::Domain { base, .. } => routine_column_type_is_array(base),
621 _ => false,
622 }
623}
624
625pub(super) fn routine_definition_error(message: impl Into<String>) -> SQLError {
626 SQLError::Routine {
627 sqlstate: "42P13".into(),
628 message: message.into(),
629 }
630}