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(name, argument_types, kind, &matched.function.def)?;
305            return Ok(Some(matched));
306        }
307        let Some(overloads) = self.catalog.lookup_sql_routine_candidates(name)? else {
308            return Ok(None);
309        };
310        resolve_static_routine_overload(
311            &self.catalog.routine_type_snapshot(),
312            name,
313            overloads,
314            argument_names,
315            argument_types,
316            explicit_variadic,
317            kind,
318        )
319        .map(Some)
320    }
321}
322
323fn resolve_static_routine_overload(
324    catalog: &RoutineTypeSnapshot,
325    name: &str,
326    overloads: Vec<Arc<SQLUserFunction>>,
327    argument_names: &[Option<String>],
328    argument_types: &[Option<ColumnType>],
329    explicit_variadic: bool,
330    kind: RoutineCallKind,
331) -> Result<StaticFunctionMatch, SQLError> {
332    let mut candidates = Vec::new();
333    let mut match_error = None;
334    for function in overloads {
335        match static_routine_match(
336            catalog,
337            function,
338            argument_names,
339            argument_types,
340            explicit_variadic,
341            kind,
342        ) {
343            Ok(Some(candidate)) => candidates.push(candidate),
344            Ok(None) => {}
345            Err(error) => {
346                match_error.get_or_insert(error);
347            }
348        }
349    }
350    retain_earliest_effective_signatures(&mut candidates);
351    if candidates.is_empty() {
352        if let Some(error) = match_error {
353            return Err(static_signature_error(kind, name, error));
354        }
355        return Err(static_routine_resolution_error(
356            kind,
357            "42883",
358            name,
359            argument_types,
360            "does not exist",
361        ));
362    }
363
364    if !rank_function_matches(&mut candidates, argument_types) || candidates.len() != 1 {
365        return Err(static_routine_resolution_error(
366            kind,
367            "42725",
368            name,
369            argument_types,
370            "is not unique",
371        ));
372    }
373    let matched = candidates
374        .pop()
375        .ok_or_else(|| SQLError::Internal("resolved routine candidate disappeared".into()))?;
376    ensure_routine_kind(name, argument_types, kind, &matched.function.def)?;
377    Ok(matched)
378}
379
380pub(super) fn retain_earliest_effective_signatures(candidates: &mut Vec<StaticFunctionMatch>) {
381    let mut visible = Vec::<(Vec<String>, String)>::new();
382    candidates.retain(|candidate| {
383        let schema = RelationIdentity::parse_reference(&candidate.function.def.name)
384            .ok()
385            .and_then(|(schema, _)| schema)
386            .unwrap_or_default();
387        if let Some((_, first_schema)) = visible
388            .iter()
389            .find(|(signature, _)| signature == &candidate.argument_types)
390        {
391            return first_schema == &schema;
392        }
393        visible.push((candidate.argument_types.clone(), schema));
394        true
395    });
396}
397
398fn static_routine_match(
399    catalog: &RoutineTypeSnapshot,
400    function: Arc<SQLUserFunction>,
401    argument_names: &[Option<String>],
402    argument_types: &[Option<ColumnType>],
403    explicit_variadic: bool,
404    kind: RoutineCallKind,
405) -> Result<Option<StaticFunctionMatch>, RoutineSignatureMatchError> {
406    let parameter_indices = routine_call_parameter_indices(&function.def, kind);
407    let parameters = parameter_indices
408        .iter()
409        .map(|index| &function.def.params[*index])
410        .collect::<Vec<_>>();
411    let Some(matched) = match_static_function_signature(
412        catalog,
413        &parameters,
414        argument_names,
415        argument_types,
416        explicit_variadic,
417    )?
418    else {
419        return Ok(None);
420    };
421    let invocation = routine_invocation_binding(&function.def, &parameter_indices, &matched);
422    Ok(Some(StaticFunctionMatch {
423        function,
424        argument_types: matched.argument_targets,
425        raw_exact_matches: matched.raw_exact_matches,
426        exact_matches: matched.exact_matches,
427        preferred_matches: matched.preferred_matches,
428        variadic_expansion: matches!(
429            matched.variadic_mode,
430            crate::type_resolution::RoutineVariadicMode::Pack
431        ),
432        invocation: Box::new(invocation),
433    }))
434}
435
436pub(super) fn static_function_match(
437    catalog: &RoutineTypeSnapshot,
438    function: Arc<SQLUserFunction>,
439    argument_names: &[Option<String>],
440    argument_types: &[Option<ColumnType>],
441    explicit_variadic: bool,
442) -> Result<Option<StaticFunctionMatch>, RoutineSignatureMatchError> {
443    static_routine_match(
444        catalog,
445        function,
446        argument_names,
447        argument_types,
448        explicit_variadic,
449        RoutineCallKind::Function,
450    )
451}
452
453fn match_static_function_signature(
454    catalog: &RoutineTypeSnapshot,
455    signature: &[&crate::ast::FunctionParam],
456    argument_names: &[Option<String>],
457    argument_types: &[Option<ColumnType>],
458    explicit_variadic: bool,
459) -> Result<Option<MatchedRoutineSignature>, RoutineSignatureMatchError> {
460    let parameters = signature
461        .iter()
462        .map(|parameter| RoutineParameterDescriptor {
463            name: Some(parameter.name.clone()),
464            type_name: canonical_routine_type_name(&parameter.type_name),
465            column_type: declared_parameter_type(catalog, &parameter.type_name),
466            has_default: parameter.default.is_some(),
467            variadic: parameter.mode == FunctionParamMode::Variadic,
468        })
469        .collect::<Vec<_>>();
470    match_routine_signature(
471        &parameters,
472        RoutineCallDescriptor {
473            argument_names,
474            argument_types,
475            explicit_variadic,
476        },
477    )
478}
479
480fn declared_parameter_type(catalog: &RoutineTypeSnapshot, type_name: &str) -> Option<ColumnType> {
481    if let Some(element) = type_name.strip_suffix("[]") {
482        return declared_parameter_type(catalog, element).map(|ty| ColumnType::Array(Box::new(ty)));
483    }
484    ColumnType::from_sql_name(type_name).ok().or_else(|| {
485        catalog.values().map(StoredDomain::column_type).find(|ty| {
486            canonical_routine_type_name(&ty.sql_name()) == canonical_routine_type_name(type_name)
487        })
488    })
489}
490
491fn routine_call_parameter_indices(def: &CreateFunction, kind: RoutineCallKind) -> Vec<usize> {
492    def.params
493        .iter()
494        .enumerate()
495        .filter_map(|(index, parameter)| {
496            let participates = if kind == RoutineCallKind::Procedure && !def.is_procedure {
497                true
498            } else {
499                match parameter.mode {
500                    FunctionParamMode::In
501                    | FunctionParamMode::InOut
502                    | FunctionParamMode::Variadic => true,
503                    FunctionParamMode::Out => def.is_procedure,
504                    FunctionParamMode::Table => false,
505                }
506            };
507            participates.then_some(index)
508        })
509        .collect()
510}
511
512fn routine_invocation_binding(
513    def: &CreateFunction,
514    parameter_indices: &[usize],
515    matched: &MatchedRoutineSignature,
516) -> RoutineInvocationBinding {
517    let parameter_types = def
518        .params
519        .iter()
520        .enumerate()
521        .map(|(definition_index, parameter)| {
522            parameter_indices
523                .iter()
524                .position(|index| *index == definition_index)
525                .and_then(|call_index| matched.parameter_types.get(call_index).cloned())
526                .or_else(|| matched.substitute_type_name(&parameter.type_name))
527                .unwrap_or_else(|| canonical_routine_type_name(&parameter.type_name))
528        })
529        .collect::<Vec<_>>();
530    let argument_positions = matched
531        .argument_positions
532        .iter()
533        .map(|call_index| parameter_indices[*call_index])
534        .collect();
535    let variadic_mode = match &matched.variadic_plan {
536        crate::type_resolution::RoutineVariadicPlan::Pack {
537            parameter_index, ..
538        } => RoutineVariadicMode::Expanded {
539            parameter_index: parameter_indices[*parameter_index],
540        },
541        crate::type_resolution::RoutineVariadicPlan::PassThrough {
542            parameter_index, ..
543        } => RoutineVariadicMode::Explicit {
544            parameter_index: parameter_indices[*parameter_index],
545        },
546        crate::type_resolution::RoutineVariadicPlan::None
547        | crate::type_resolution::RoutineVariadicPlan::Default { .. } => RoutineVariadicMode::None,
548    };
549    let output_indices = def
550        .params
551        .iter()
552        .enumerate()
553        .filter_map(|(index, parameter)| {
554            matches!(
555                parameter.mode,
556                FunctionParamMode::Out | FunctionParamMode::InOut | FunctionParamMode::Table
557            )
558            .then_some(index)
559        })
560        .collect::<Vec<_>>();
561    let return_type = match &def.returns {
562        FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => matched
563            .substitute_type_name(type_name)
564            .or_else(|| Some(canonical_routine_type_name(type_name))),
565        FunctionReturns::Table => Some("record".into()),
566        FunctionReturns::None => match output_indices.as_slice() {
567            [] => None,
568            [index] => parameter_types.get(*index).cloned(),
569            _ => Some("record".into()),
570        },
571    };
572    RoutineInvocationBinding {
573        argument_positions,
574        argument_targets: matched.argument_targets.clone(),
575        argument_sources: matched.argument_sources.clone(),
576        parameter_types,
577        return_type,
578        variadic_mode,
579    }
580}
581
582pub(super) fn static_signature_error(
583    kind: RoutineCallKind,
584    name: &str,
585    error: RoutineSignatureMatchError,
586) -> SQLError {
587    let sqlstate = error.sqlstate().to_string();
588    let message = match error {
589        RoutineSignatureMatchError::InvalidVariadicSignature { reason } => reason,
590        RoutineSignatureMatchError::IndeterminatePolymorphicType { .. } => {
591            format!(
592                "could not determine polymorphic type for {} `{name}` because an input has type unknown",
593                kind.name()
594            )
595        }
596    };
597    SQLError::Routine { sqlstate, message }
598}
599
600pub(super) fn static_function_return_type(
601    resolver: &RoutineOverloadContext<'_>,
602    name: &str,
603    def: &CreateFunction,
604    invocation: Option<&RoutineInvocationBinding>,
605) -> Result<ColumnType, SQLError> {
606    let invocation_return = invocation.and_then(|invocation| invocation.return_type.as_deref());
607    let declared_return = match &def.returns {
608        FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
609            Some(type_name.as_str())
610        }
611        FunctionReturns::Table | FunctionReturns::None => None,
612    };
613    if invocation_return
614        .or(declared_return)
615        .is_some_and(|type_name| canonical_routine_type_name(type_name) == "trigger")
616    {
617        return Err(SQLError::Routine {
618            sqlstate: "0A000".into(),
619            message: "trigger functions can only be called as triggers".into(),
620        });
621    }
622    if matches!(def.returns, FunctionReturns::Table) || def.output_params().len() > 1 {
623        return Ok(ColumnType::Record);
624    }
625    if let Some(type_name) = invocation_return {
626        return resolver
627            .catalog
628            .resolve_catalog_column_type(type_name)
629            .or_else(|| ColumnType::from_sql_name(type_name).ok())
630            .ok_or_else(|| {
631                SQLError::TypeMismatch(format!(
632                    "function `{name}` has unresolved return type `{type_name}`"
633                ))
634            });
635    }
636    let type_name = match &def.returns {
637        FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => type_name,
638        FunctionReturns::None => {
639            let outputs = def.output_params();
640            if outputs.len() > 1 {
641                return Ok(ColumnType::Record);
642            }
643            &outputs
644                .first()
645                .ok_or_else(|| {
646                    SQLError::TypeMismatch(format!("function `{name}` does not return a value"))
647                })?
648                .type_name
649        }
650        FunctionReturns::Table => unreachable!("table result handled above"),
651    };
652    resolver
653        .catalog
654        .resolve_catalog_column_type(type_name)
655        .or_else(|| ColumnType::from_sql_name(type_name).ok())
656        .ok_or_else(|| SQLError::TypeMismatch(format!("unknown type `{type_name}`")))
657}
658
659fn ensure_routine_kind(
660    name: &str,
661    argument_types: &[Option<ColumnType>],
662    expected: RoutineCallKind,
663    definition: &CreateFunction,
664) -> Result<(), SQLError> {
665    if definition.is_procedure == expected.is_procedure() {
666        return Ok(());
667    }
668    let arguments = static_routine_argument_types(argument_types);
669    let suffix = if definition.is_procedure {
670        "is a procedure"
671    } else {
672        "is not a procedure"
673    };
674    Err(SQLError::Routine {
675        sqlstate: "42809".into(),
676        message: format!("{name}({arguments}) {suffix}"),
677    })
678}
679
680fn static_bound_routine_error(kind: RoutineCallKind, binding: &FunctionBinding) -> SQLError {
681    SQLError::Routine {
682        sqlstate: "42883".into(),
683        message: format!(
684            "bound {} {}({}) does not exist",
685            kind.name(),
686            binding.name,
687            binding.argument_types.join(", ")
688        ),
689    }
690}
691
692fn static_routine_resolution_error(
693    kind: RoutineCallKind,
694    sqlstate: &str,
695    name: &str,
696    argument_types: &[Option<ColumnType>],
697    suffix: &str,
698) -> SQLError {
699    let arguments = static_routine_argument_types(argument_types);
700    SQLError::Routine {
701        sqlstate: sqlstate.into(),
702        message: format!("{} {name}({arguments}) {suffix}", kind.name()),
703    }
704}
705
706fn static_routine_argument_types(argument_types: &[Option<ColumnType>]) -> String {
707    argument_types
708        .iter()
709        .map(|ty| {
710            ty.as_ref()
711                .map_or_else(|| "unknown".into(), ColumnType::sql_name)
712        })
713        .collect::<Vec<_>>()
714        .join(", ")
715}