Skip to main content

uqa_sql/routines/
resolution.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Static routine signature matching, invocation binding, and declared return-type resolution.
8
9mod combined_overloads;
10
11use super::{
12    declaration::RoutineTypeCatalog, routine_signature_types, SQLUserFunction, StaticFunctionMatch,
13};
14use crate::type_resolution::{
15    canonical_routine_type_name, match_routine_signature, rank_function_matches,
16    BuiltinFunctionOverload, FunctionTypeResolver, MatchedRoutineSignature,
17    ResolvedFunctionOverload, RoutineCallDescriptor, RoutineParameterDescriptor,
18    RoutineSignatureMatchError,
19};
20use crate::{
21    ast::{
22        ColumnType, CreateFunction, FunctionBinding, FunctionParamMode, FunctionReturns,
23        RoutineInvocationBinding, RoutineVariadicMode,
24    },
25    catalog::domain::StoredDomain,
26    SQLError,
27};
28use std::{collections::BTreeMap, sync::Arc};
29use uqa_core::RelationIdentity;
30
31/// The original immutable domain allocation captured for one signature-matching pass.
32pub type RoutineTypeSnapshot = Arc<BTreeMap<String, StoredDomain>>;
33pub trait RoutineOverloadCatalog: RoutineTypeCatalog + Send + Sync {
34    fn routine_type_snapshot(&self) -> RoutineTypeSnapshot;
35    fn routine_search_path(&self) -> Vec<String>;
36    fn has_registered_scalar_function(&self, name: &str) -> bool;
37    fn lookup_sql_routine_candidates(
38        &self,
39        name: &str,
40    ) -> Result<Option<Vec<Arc<SQLUserFunction>>>, SQLError>;
41    fn lookup_bound_sql_routine_candidates_by_binding(
42        &self,
43        binding: &FunctionBinding,
44    ) -> Option<Vec<Arc<SQLUserFunction>>>;
45    fn lookup_bound_sql_functions_by_binding(
46        &self,
47        binding: &FunctionBinding,
48    ) -> Option<Vec<Arc<SQLUserFunction>>>;
49}
50pub struct RoutineOverloadContext<'a> {
51    pub catalog: &'a dyn RoutineOverloadCatalog,
52}
53
54impl FunctionTypeResolver for RoutineOverloadContext<'_> {
55    fn has_untyped_function(&self, name: &str) -> bool {
56        self.catalog.has_registered_scalar_function(name)
57    }
58
59    fn resolve_type_name(&self, name: &str) -> Result<Option<ColumnType>, SQLError> {
60        self.catalog
61            .resolve_catalog_column_type_name(name)
62            .map(Some)
63    }
64
65    fn resolve_function_type(
66        &self,
67        name: &str,
68        binding: Option<&FunctionBinding>,
69        argument_names: &[Option<String>],
70        argument_types: &[Option<ColumnType>],
71        explicit_variadic: bool,
72    ) -> Result<Option<ColumnType>, SQLError> {
73        self.resolve_function_overload(
74            name,
75            binding,
76            argument_names,
77            argument_types,
78            explicit_variadic,
79        )
80        .map(|resolved| resolved.map(|resolved| resolved.return_type))
81    }
82
83    fn resolve_function_overload(
84        &self,
85        name: &str,
86        binding: Option<&FunctionBinding>,
87        argument_names: &[Option<String>],
88        argument_types: &[Option<ColumnType>],
89        explicit_variadic: bool,
90    ) -> Result<Option<ResolvedFunctionOverload>, SQLError> {
91        let Some(matched) = self.resolve_static_sql_routine_match(
92            name,
93            binding,
94            argument_names,
95            argument_types,
96            explicit_variadic,
97            RoutineCallKind::Function,
98        )?
99        else {
100            return Ok(None);
101        };
102        let function = &matched.function;
103        Ok(Some(ResolvedFunctionOverload {
104            binding: matched.binding(),
105            return_type: static_function_return_type(
106                self,
107                name,
108                &function.def,
109                Some(&matched.invocation),
110            )?,
111            exact_matches: matched.exact_matches,
112            known_arguments: argument_types.iter().flatten().count(),
113            preferred_matches: matched.preferred_matches,
114            precedes_pg_catalog: self.user_function_precedes_pg_catalog(&function.def.name),
115        }))
116    }
117
118    fn is_scalar_function_binding(&self, binding: &FunctionBinding) -> Result<bool, SQLError> {
119        if binding.builtin {
120            return Ok(false);
121        }
122        let function = self
123            .catalog
124            .lookup_bound_sql_functions_by_binding(binding)
125            .and_then(|overloads| {
126                overloads.into_iter().find(|function| {
127                    !function.def.is_procedure
128                        && routine_signature_types(&function.def) == binding.argument_types
129                })
130            });
131        Ok(function.is_some_and(|function| !function.def.returns_set()))
132    }
133
134    fn resolve_function_overload_with_builtins(
135        &self,
136        name: &str,
137        binding: Option<&FunctionBinding>,
138        argument_names: &[Option<String>],
139        argument_types: &[Option<ColumnType>],
140        explicit_variadic: bool,
141        builtins: &[BuiltinFunctionOverload],
142    ) -> Result<Option<ResolvedFunctionOverload>, SQLError> {
143        combined_overloads::resolve(
144            self,
145            name,
146            binding,
147            argument_names,
148            argument_types,
149            explicit_variadic,
150            builtins,
151        )
152        .map(Some)
153    }
154}
155
156#[derive(Debug, Clone, Copy, PartialEq, Eq)]
157pub enum RoutineCallKind {
158    Function,
159    Procedure,
160}
161
162impl RoutineCallKind {
163    fn is_procedure(self) -> bool {
164        self == Self::Procedure
165    }
166
167    fn name(self) -> &'static str {
168        match self {
169            Self::Function => "function",
170            Self::Procedure => "procedure",
171        }
172    }
173}
174
175impl RoutineOverloadContext<'_> {
176    pub(super) fn user_function_precedes_pg_catalog(&self, name: &str) -> bool {
177        let Ok((Some(schema), _)) = RelationIdentity::parse_reference(name) else {
178            return false;
179        };
180        let search_path = self.catalog.routine_search_path();
181        let Some(user_position) = search_path.iter().position(|entry| entry == &schema) else {
182            return false;
183        };
184        search_path
185            .iter()
186            .position(|entry| entry == "pg_catalog")
187            .is_some_and(|catalog_position| user_position < catalog_position)
188    }
189
190    pub fn resolve_static_sql_function(
191        &self,
192        name: &str,
193        binding: Option<&FunctionBinding>,
194        argument_names: &[Option<String>],
195        argument_types: &[Option<ColumnType>],
196        explicit_variadic: bool,
197    ) -> Result<Option<Arc<SQLUserFunction>>, SQLError> {
198        self.resolve_static_sql_function_match(
199            name,
200            binding,
201            argument_names,
202            argument_types,
203            explicit_variadic,
204        )
205        .map(|matched| matched.map(|matched| matched.function))
206    }
207
208    pub fn resolve_static_sql_function_match(
209        &self,
210        name: &str,
211        binding: Option<&FunctionBinding>,
212        argument_names: &[Option<String>],
213        argument_types: &[Option<ColumnType>],
214        explicit_variadic: bool,
215    ) -> Result<Option<StaticFunctionMatch>, SQLError> {
216        self.resolve_static_sql_routine_match(
217            name,
218            binding,
219            argument_names,
220            argument_types,
221            explicit_variadic,
222            RoutineCallKind::Function,
223        )
224    }
225
226    pub fn resolve_table_function_overload_with_builtins(
227        &self,
228        name: &str,
229        binding: Option<&FunctionBinding>,
230        argument_names: &[Option<String>],
231        argument_types: &[Option<ColumnType>],
232        explicit_variadic: bool,
233        builtins: &[BuiltinFunctionOverload],
234    ) -> Result<Option<ResolvedFunctionOverload>, SQLError> {
235        combined_overloads::resolve_table(
236            self,
237            name,
238            binding,
239            argument_names,
240            argument_types,
241            explicit_variadic,
242            builtins,
243        )
244        .map(Some)
245    }
246
247    pub fn resolve_static_sql_routine_match(
248        &self,
249        name: &str,
250        binding: Option<&FunctionBinding>,
251        argument_names: &[Option<String>],
252        argument_types: &[Option<ColumnType>],
253        explicit_variadic: bool,
254        kind: RoutineCallKind,
255    ) -> Result<Option<StaticFunctionMatch>, SQLError> {
256        if let Some(binding) = binding {
257            if binding.builtin {
258                return Ok(None);
259            }
260            let function = self
261                .catalog
262                .lookup_bound_sql_routine_candidates_by_binding(binding)
263                .and_then(|overloads| {
264                    overloads.into_iter().find(|function| {
265                        function.def.is_procedure == kind.is_procedure()
266                            && routine_signature_types(&function.def) == binding.argument_types
267                    })
268                })
269                .ok_or_else(|| static_bound_routine_error(kind, binding))?;
270            let matched = if let Some(invocation) = &binding.invocation {
271                let invocation_is_explicit = matches!(
272                    invocation.variadic_mode,
273                    RoutineVariadicMode::Explicit { .. }
274                );
275                if invocation.argument_positions.len() != argument_types.len()
276                    || invocation_is_explicit != explicit_variadic
277                {
278                    return Err(static_bound_routine_error(kind, binding));
279                }
280                StaticFunctionMatch {
281                    function,
282                    invocation: invocation.clone(),
283                    argument_types: invocation.argument_targets.clone(),
284                    raw_exact_matches: 0,
285                    exact_matches: 0,
286                    preferred_matches: 0,
287                    variadic_expansion: matches!(
288                        invocation.variadic_mode,
289                        RoutineVariadicMode::Expanded { .. }
290                    ),
291                }
292            } else {
293                static_routine_match(
294                    &self.catalog.routine_type_snapshot(),
295                    function,
296                    argument_names,
297                    argument_types,
298                    explicit_variadic,
299                    kind,
300                )
301                .map_err(|error| static_signature_error(kind, name, error))?
302                .ok_or_else(|| static_bound_routine_error(kind, binding))?
303            };
304            ensure_routine_kind(
305                name,
306                argument_names,
307                argument_types,
308                kind,
309                &matched.function.def,
310            )?;
311            return Ok(Some(matched));
312        }
313        let Some(overloads) = self.catalog.lookup_sql_routine_candidates(name)? else {
314            return Ok(None);
315        };
316        resolve_static_routine_overload(
317            &self.catalog.routine_type_snapshot(),
318            name,
319            overloads,
320            argument_names,
321            argument_types,
322            explicit_variadic,
323            kind,
324        )
325        .map(Some)
326    }
327}
328
329fn resolve_static_routine_overload(
330    catalog: &RoutineTypeSnapshot,
331    name: &str,
332    overloads: Vec<Arc<SQLUserFunction>>,
333    argument_names: &[Option<String>],
334    argument_types: &[Option<ColumnType>],
335    explicit_variadic: bool,
336    kind: RoutineCallKind,
337) -> Result<StaticFunctionMatch, SQLError> {
338    let mut candidates = Vec::new();
339    let mut match_error = None;
340    for function in overloads {
341        match static_routine_match(
342            catalog,
343            function,
344            argument_names,
345            argument_types,
346            explicit_variadic,
347            kind,
348        ) {
349            Ok(Some(candidate)) => candidates.push(candidate),
350            Ok(None) => {}
351            Err(error) => {
352                match_error.get_or_insert(error);
353            }
354        }
355    }
356    retain_earliest_effective_signatures(&mut candidates);
357    if candidates.is_empty() {
358        if let Some(error) = match_error {
359            return Err(static_signature_error(kind, name, error));
360        }
361        return Err(static_routine_resolution_error(
362            kind,
363            "42883",
364            name,
365            argument_types,
366            "does not exist",
367        ));
368    }
369
370    if !rank_function_matches(&mut candidates, argument_types) || candidates.len() != 1 {
371        return Err(static_routine_resolution_error(
372            kind,
373            "42725",
374            name,
375            argument_types,
376            "is not unique",
377        ));
378    }
379    let matched = candidates
380        .pop()
381        .ok_or_else(|| SQLError::Internal("resolved routine candidate disappeared".into()))?;
382    ensure_routine_kind(
383        name,
384        argument_names,
385        argument_types,
386        kind,
387        &matched.function.def,
388    )?;
389    Ok(matched)
390}
391
392pub(super) fn retain_earliest_effective_signatures(candidates: &mut Vec<StaticFunctionMatch>) {
393    let mut visible = Vec::<(Vec<String>, String)>::new();
394    candidates.retain(|candidate| {
395        let schema = RelationIdentity::parse_reference(&candidate.function.def.name)
396            .ok()
397            .and_then(|(schema, _)| schema)
398            .unwrap_or_default();
399        if let Some((_, first_schema)) = visible
400            .iter()
401            .find(|(signature, _)| signature == &candidate.argument_types)
402        {
403            return first_schema == &schema;
404        }
405        visible.push((candidate.argument_types.clone(), schema));
406        true
407    });
408}
409
410fn static_routine_match(
411    catalog: &RoutineTypeSnapshot,
412    function: Arc<SQLUserFunction>,
413    argument_names: &[Option<String>],
414    argument_types: &[Option<ColumnType>],
415    explicit_variadic: bool,
416    kind: RoutineCallKind,
417) -> Result<Option<StaticFunctionMatch>, RoutineSignatureMatchError> {
418    let parameter_indices = routine_call_parameter_indices(&function.def, kind);
419    let parameters = parameter_indices
420        .iter()
421        .map(|index| &function.def.params[*index])
422        .collect::<Vec<_>>();
423    let Some(matched) = match_static_function_signature(
424        catalog,
425        &parameters,
426        argument_names,
427        argument_types,
428        explicit_variadic,
429    )?
430    else {
431        return Ok(None);
432    };
433    let invocation = routine_invocation_binding(&function.def, &parameter_indices, &matched);
434    Ok(Some(StaticFunctionMatch {
435        function,
436        argument_types: matched.argument_targets,
437        raw_exact_matches: matched.raw_exact_matches,
438        exact_matches: matched.exact_matches,
439        preferred_matches: matched.preferred_matches,
440        variadic_expansion: matches!(
441            matched.variadic_mode,
442            crate::type_resolution::RoutineVariadicMode::Pack
443        ),
444        invocation: Box::new(invocation),
445    }))
446}
447
448pub(super) fn static_function_match(
449    catalog: &RoutineTypeSnapshot,
450    function: Arc<SQLUserFunction>,
451    argument_names: &[Option<String>],
452    argument_types: &[Option<ColumnType>],
453    explicit_variadic: bool,
454) -> Result<Option<StaticFunctionMatch>, RoutineSignatureMatchError> {
455    static_routine_match(
456        catalog,
457        function,
458        argument_names,
459        argument_types,
460        explicit_variadic,
461        RoutineCallKind::Function,
462    )
463}
464
465fn match_static_function_signature(
466    catalog: &RoutineTypeSnapshot,
467    signature: &[&crate::ast::FunctionParam],
468    argument_names: &[Option<String>],
469    argument_types: &[Option<ColumnType>],
470    explicit_variadic: bool,
471) -> Result<Option<MatchedRoutineSignature>, RoutineSignatureMatchError> {
472    let parameters = signature
473        .iter()
474        .map(|parameter| RoutineParameterDescriptor {
475            name: Some(parameter.name.clone()),
476            type_name: canonical_routine_type_name(&parameter.type_name),
477            column_type: declared_parameter_type(catalog, &parameter.type_name),
478            has_default: parameter.default.is_some(),
479            variadic: parameter.mode == FunctionParamMode::Variadic,
480        })
481        .collect::<Vec<_>>();
482    match_routine_signature(
483        &parameters,
484        RoutineCallDescriptor {
485            argument_names,
486            argument_types,
487            explicit_variadic,
488        },
489    )
490}
491
492fn declared_parameter_type(catalog: &RoutineTypeSnapshot, type_name: &str) -> Option<ColumnType> {
493    if let Some(element) = type_name.strip_suffix("[]") {
494        return declared_parameter_type(catalog, element).map(|ty| ColumnType::Array(Box::new(ty)));
495    }
496    ColumnType::from_sql_name(type_name).ok().or_else(|| {
497        catalog.values().map(StoredDomain::column_type).find(|ty| {
498            canonical_routine_type_name(&ty.sql_name()) == canonical_routine_type_name(type_name)
499        })
500    })
501}
502
503fn routine_call_parameter_indices(def: &CreateFunction, kind: RoutineCallKind) -> Vec<usize> {
504    def.params
505        .iter()
506        .enumerate()
507        .filter_map(|(index, parameter)| {
508            let participates = if kind == RoutineCallKind::Procedure && !def.is_procedure {
509                true
510            } else {
511                match parameter.mode {
512                    FunctionParamMode::In
513                    | FunctionParamMode::InOut
514                    | FunctionParamMode::Variadic => true,
515                    FunctionParamMode::Out => def.is_procedure,
516                    FunctionParamMode::Table => false,
517                }
518            };
519            participates.then_some(index)
520        })
521        .collect()
522}
523
524fn routine_invocation_binding(
525    def: &CreateFunction,
526    parameter_indices: &[usize],
527    matched: &MatchedRoutineSignature,
528) -> RoutineInvocationBinding {
529    let parameter_types = def
530        .params
531        .iter()
532        .enumerate()
533        .map(|(definition_index, parameter)| {
534            parameter_indices
535                .iter()
536                .position(|index| *index == definition_index)
537                .and_then(|call_index| matched.parameter_types.get(call_index).cloned())
538                .or_else(|| matched.substitute_type_name(&parameter.type_name))
539                .unwrap_or_else(|| canonical_routine_type_name(&parameter.type_name))
540        })
541        .collect::<Vec<_>>();
542    let argument_positions = matched
543        .argument_positions
544        .iter()
545        .map(|call_index| parameter_indices[*call_index])
546        .collect();
547    let variadic_mode = match &matched.variadic_plan {
548        crate::type_resolution::RoutineVariadicPlan::Pack {
549            parameter_index, ..
550        } => RoutineVariadicMode::Expanded {
551            parameter_index: parameter_indices[*parameter_index],
552        },
553        crate::type_resolution::RoutineVariadicPlan::PassThrough {
554            parameter_index, ..
555        } => RoutineVariadicMode::Explicit {
556            parameter_index: parameter_indices[*parameter_index],
557        },
558        crate::type_resolution::RoutineVariadicPlan::None
559        | crate::type_resolution::RoutineVariadicPlan::Default { .. } => RoutineVariadicMode::None,
560    };
561    let output_indices = def
562        .params
563        .iter()
564        .enumerate()
565        .filter_map(|(index, parameter)| {
566            matches!(
567                parameter.mode,
568                FunctionParamMode::Out | FunctionParamMode::InOut | FunctionParamMode::Table
569            )
570            .then_some(index)
571        })
572        .collect::<Vec<_>>();
573    let return_type = match &def.returns {
574        FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => matched
575            .substitute_type_name(type_name)
576            .or_else(|| Some(canonical_routine_type_name(type_name))),
577        FunctionReturns::Table => Some("record".into()),
578        FunctionReturns::None => match output_indices.as_slice() {
579            [] => None,
580            [index] => parameter_types.get(*index).cloned(),
581            _ => Some("record".into()),
582        },
583    };
584    RoutineInvocationBinding {
585        argument_positions,
586        argument_targets: matched.argument_targets.clone(),
587        argument_sources: matched.argument_sources.clone(),
588        parameter_types,
589        return_type,
590        variadic_mode,
591    }
592}
593
594pub(super) fn static_signature_error(
595    kind: RoutineCallKind,
596    name: &str,
597    error: RoutineSignatureMatchError,
598) -> SQLError {
599    let sqlstate = error.sqlstate().to_string();
600    let message = match error {
601        RoutineSignatureMatchError::InvalidVariadicSignature { reason } => reason,
602        RoutineSignatureMatchError::IndeterminatePolymorphicType { .. } => {
603            format!(
604                "could not determine polymorphic type for {} `{name}` because an input has type unknown",
605                kind.name()
606            )
607        }
608    };
609    SQLError::Routine { sqlstate, message }
610}
611
612pub(super) fn static_function_return_type(
613    resolver: &RoutineOverloadContext<'_>,
614    name: &str,
615    def: &CreateFunction,
616    invocation: Option<&RoutineInvocationBinding>,
617) -> Result<ColumnType, SQLError> {
618    let invocation_return = invocation.and_then(|invocation| invocation.return_type.as_deref());
619    let declared_return = match &def.returns {
620        FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
621            Some(type_name.as_str())
622        }
623        FunctionReturns::Table | FunctionReturns::None => None,
624    };
625    if invocation_return
626        .or(declared_return)
627        .is_some_and(|type_name| canonical_routine_type_name(type_name) == "trigger")
628    {
629        return Err(SQLError::Routine {
630            sqlstate: "0A000".into(),
631            message: "trigger functions can only be called as triggers".into(),
632        });
633    }
634    if matches!(def.returns, FunctionReturns::Table) || def.output_params().len() > 1 {
635        return Ok(ColumnType::Record);
636    }
637    if let Some(type_name) = invocation_return {
638        return resolver
639            .catalog
640            .resolve_catalog_column_type(type_name)
641            .or_else(|| ColumnType::from_sql_name(type_name).ok())
642            .ok_or_else(|| {
643                SQLError::TypeMismatch(format!(
644                    "function `{name}` has unresolved return type `{type_name}`"
645                ))
646            });
647    }
648    let type_name = match &def.returns {
649        FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => type_name,
650        FunctionReturns::None => {
651            let outputs = def.output_params();
652            if outputs.len() > 1 {
653                return Ok(ColumnType::Record);
654            }
655            &outputs
656                .first()
657                .ok_or_else(|| {
658                    SQLError::TypeMismatch(format!("function `{name}` does not return a value"))
659                })?
660                .type_name
661        }
662        FunctionReturns::Table => unreachable!("table result handled above"),
663    };
664    resolver
665        .catalog
666        .resolve_catalog_column_type(type_name)
667        .or_else(|| ColumnType::from_sql_name(type_name).ok())
668        .ok_or_else(|| SQLError::TypeMismatch(format!("unknown type `{type_name}`")))
669}
670
671fn ensure_routine_kind(
672    name: &str,
673    argument_names: &[Option<String>],
674    argument_types: &[Option<ColumnType>],
675    expected: RoutineCallKind,
676    definition: &CreateFunction,
677) -> Result<(), SQLError> {
678    if definition.is_procedure == expected.is_procedure() {
679        return Ok(());
680    }
681    let arguments = static_routine_argument_types(argument_names, argument_types);
682    let suffix = if definition.is_procedure {
683        "is a procedure"
684    } else {
685        "is not a procedure"
686    };
687    Err(SQLError::Diagnostic {
688        sqlstate: "42809".into(),
689        message: format!("{name}({arguments}) {suffix}"),
690        detail: None,
691        hint: Some(if definition.is_procedure {
692            "To call a procedure, use CALL.".into()
693        } else {
694            "To call a function, use SELECT.".into()
695        }),
696    })
697}
698
699fn static_bound_routine_error(kind: RoutineCallKind, binding: &FunctionBinding) -> SQLError {
700    SQLError::Routine {
701        sqlstate: "42883".into(),
702        message: format!(
703            "bound {} {}({}) does not exist",
704            kind.name(),
705            binding.name,
706            binding.argument_types.join(", ")
707        ),
708    }
709}
710
711fn static_routine_resolution_error(
712    kind: RoutineCallKind,
713    sqlstate: &str,
714    name: &str,
715    argument_types: &[Option<ColumnType>],
716    suffix: &str,
717) -> SQLError {
718    let arguments = static_routine_argument_types(&[], argument_types);
719    SQLError::Routine {
720        sqlstate: sqlstate.into(),
721        message: format!("{} {name}({arguments}) {suffix}", kind.name()),
722    }
723}
724
725fn static_routine_argument_types(
726    argument_names: &[Option<String>],
727    argument_types: &[Option<ColumnType>],
728) -> String {
729    argument_types
730        .iter()
731        .enumerate()
732        .map(|(index, ty)| {
733            let ty = ty
734                .as_ref()
735                .map_or_else(|| "unknown".into(), ColumnType::regtype_name);
736            argument_names
737                .get(index)
738                .and_then(Option::as_ref)
739                .map_or_else(|| ty.clone(), |name| format!("{name} => {ty}"))
740        })
741        .collect::<Vec<_>>()
742        .join(", ")
743}