Skip to main content

uqa_sql/routines/
declaration.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Routine type references, polymorphic signatures, and declaration validation.
8
9use 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            &parameter.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(&parameter.type_name)
282            .or_else(|| ColumnType::from_sql_name(&parameter.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, &parameter.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(&parameter.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}