1mod 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
31pub 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 ¶meters,
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, ¶meter_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(¶meter.type_name),
477 column_type: declared_parameter_type(catalog, ¶meter.type_name),
478 has_default: parameter.default.is_some(),
479 variadic: parameter.mode == FunctionParamMode::Variadic,
480 })
481 .collect::<Vec<_>>();
482 match_routine_signature(
483 ¶meters,
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(¶meter.type_name))
539 .unwrap_or_else(|| canonical_routine_type_name(¶meter.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}