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(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 ¶meters,
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, ¶meter_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(¶meter.type_name),
465 column_type: declared_parameter_type(catalog, ¶meter.type_name),
466 has_default: parameter.default.is_some(),
467 variadic: parameter.mode == FunctionParamMode::Variadic,
468 })
469 .collect::<Vec<_>>();
470 match_routine_signature(
471 ¶meters,
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(¶meter.type_name))
527 .unwrap_or_else(|| canonical_routine_type_name(¶meter.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}