use crate::{
ast::{
AlterRoutineStmt, ColumnDef, ColumnType, CreateFunction, FunctionBody, FunctionParamMode,
FunctionReturns, RoutineColumnTypeReference,
},
type_resolution::canonical_routine_type_name,
SQLError,
};
pub trait RoutineTypeCatalog {
fn try_describe_table(&self, reference: &str) -> Result<Option<Vec<ColumnDef>>, String>;
fn resolve_catalog_column_type(&self, name: &str) -> Option<ColumnType>;
fn resolve_catalog_column_type_name(&self, name: &str) -> Result<ColumnType, SQLError>;
fn resolve_catalog_domain_type_by_oid(&self, oid: u32) -> Option<ColumnType>;
}
pub fn resolve_routine_type_references(
catalog: &dyn RoutineTypeCatalog,
def: &mut CreateFunction,
) -> Result<(), SQLError> {
for parameter in &mut def.params {
parameter.type_name = resolve_routine_type_name_with_reference(
catalog,
¶meter.type_name,
ROUTINE_PARAMETER_PSEUDO_TYPES,
parameter.type_reference.as_ref(),
)?;
parameter.type_reference = None;
}
match &mut def.returns {
FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
*type_name = resolve_routine_type_name_with_reference(
catalog,
type_name,
ROUTINE_RESULT_PSEUDO_TYPES,
def.return_type_reference.as_ref(),
)?;
}
FunctionReturns::None | FunctionReturns::Table => {}
}
def.return_type_reference = None;
Ok(())
}
pub fn resolve_alter_routine_identity_types(
catalog: &dyn RoutineTypeCatalog,
stmt: &AlterRoutineStmt,
) -> Result<Option<Vec<String>>, SQLError> {
resolve_routine_identity_types(
catalog,
stmt.arg_types.as_deref(),
&stmt.arg_type_references,
"ALTER routine",
)
}
pub fn resolve_routine_identity_types(
catalog: &dyn RoutineTypeCatalog,
types: Option<&[String]>,
references: &[Option<RoutineColumnTypeReference>],
context: &str,
) -> Result<Option<Vec<String>>, SQLError> {
let Some(types) = types else {
if !references.is_empty() {
return Err(SQLError::Internal(format!(
"{context} omitted its identity types but retained type references"
)));
}
return Ok(None);
};
if !references.is_empty() && references.len() != types.len() {
return Err(SQLError::Internal(format!(
"{context} has {} identity types but {} type references",
types.len(),
references.len()
)));
}
types
.iter()
.enumerate()
.map(|(index, type_name)| {
resolve_routine_type_name_with_reference(
catalog,
type_name,
ROUTINE_PARAMETER_PSEUDO_TYPES,
references.get(index).and_then(Option::as_ref),
)
.map(|resolved| canonical_routine_type_name(&resolved))
})
.collect::<Result<Vec<_>, _>>()
.map(Some)
}
const POLYMORPHIC_PSEUDO_TYPES: &[&str] = &[
"anyelement",
"anyarray",
"anynonarray",
"anyenum",
"anyrange",
"anymultirange",
"anycompatible",
"anycompatiblearray",
"anycompatiblenonarray",
"anycompatiblerange",
"anycompatiblemultirange",
];
const ROUTINE_PARAMETER_PSEUDO_TYPES: &[&str] = &[
"record",
"refcursor",
"cstring",
"any",
"void",
"trigger",
"internal",
"event_trigger",
"anyelement",
"anyarray",
"anynonarray",
"anyenum",
"anyrange",
"anymultirange",
"anycompatible",
"anycompatiblearray",
"anycompatiblenonarray",
"anycompatiblerange",
"anycompatiblemultirange",
];
const ROUTINE_RESULT_PSEUDO_TYPES: &[&str] = &[
"record",
"refcursor",
"cstring",
"any",
"void",
"trigger",
"internal",
"event_trigger",
"anyelement",
"anyarray",
"anynonarray",
"anyenum",
"anyrange",
"anymultirange",
"anycompatible",
"anycompatiblearray",
"anycompatiblenonarray",
"anycompatiblerange",
"anycompatiblemultirange",
];
fn resolve_routine_type_name_with_reference(
catalog: &dyn RoutineTypeCatalog,
type_name: &str,
allowed_pseudo_types: &[&str],
structured_reference: Option<&RoutineColumnTypeReference>,
) -> Result<String, SQLError> {
let mut base = type_name.trim();
let mut array_dimensions = 0usize;
while let Some(element) = base.strip_suffix("[]") {
base = element.trim_end();
array_dimensions += 1;
}
let resolved = if base
.get(base.len().saturating_sub("%type".len())..)
.is_some_and(|suffix| suffix.eq_ignore_ascii_case("%type"))
{
let reference = structured_reference.ok_or_else(|| {
SQLError::Internal(format!(
"routine type reference `{type_name}` is missing structured relation-column identity"
))
})?;
let table = reference.relation_reference();
let columns = catalog
.try_describe_table(&table)
.map_err(|error| {
SQLError::Internal(format!(
"resolve routine type reference `{type_name}`: {error}"
))
})?
.ok_or_else(|| SQLError::UnknownTable(table.clone()))?;
columns
.into_iter()
.find(|definition| definition.name == reference.column)
.map(|definition| definition.ty)
.ok_or_else(|| SQLError::UnknownColumn(reference.type_reference()))?
} else {
let canonical = canonical_routine_type_name(base);
if allowed_pseudo_types.contains(&canonical.as_str()) {
if array_dimensions != 0 {
return Err(SQLError::Routine {
sqlstate: "42704".into(),
message: format!("type `{type_name}` does not exist"),
});
}
return Ok(canonical);
}
catalog.resolve_catalog_column_type_name(base)?
};
let mut resolved = resolved;
for _ in 0..array_dimensions {
resolved = ColumnType::Array(Box::new(resolved));
}
Ok(resolved.sql_name())
}
pub fn resolve_plpgsql_datum_types(
catalog: &dyn RoutineTypeCatalog,
function: &mut crate::plpgsql::PLpgSQLFunction,
) -> Result<(), SQLError> {
for datum in &mut function.datums {
let crate::plpgsql::PLpgSQLDatum::Var(variable) = datum else {
continue;
};
if variable.type_reference.is_none() {
if let Some(ty) = variable
.type_oid
.and_then(|oid| catalog.resolve_catalog_domain_type_by_oid(oid))
{
variable.type_name = ty.sql_name();
continue;
}
}
variable.type_name = resolve_routine_type_name_with_reference(
catalog,
&variable.type_name,
&[
"record",
"refcursor",
"anyelement",
"anyarray",
"anynonarray",
"anyenum",
"anyrange",
"anymultirange",
"anycompatible",
"anycompatiblearray",
"anycompatiblenonarray",
"anycompatiblerange",
"anycompatiblemultirange",
],
variable.type_reference.as_ref(),
)?;
variable.type_reference = None;
}
Ok(())
}
pub(super) fn validate_routine_declaration(
catalog: &dyn RoutineTypeCatalog,
def: &CreateFunction,
) -> Result<(), SQLError> {
validate_variadic_declaration(catalog, def)?;
let inputs = validate_routine_input_types(def)?;
if matches!(def.body, FunctionBody::Statements(_)) && inputs.any {
return Err(routine_definition_error(
"SQL function with unquoted function body cannot have polymorphic arguments",
));
}
validate_routine_output_types(def, &inputs)
}
pub(super) fn routine_parameter_regrole_constants(
catalog: &dyn RoutineTypeCatalog,
def: &CreateFunction,
) -> crate::catalog::regrole_dependencies::StoredRegroleConstants {
let mut constants = crate::catalog::regrole_dependencies::StoredRegroleConstants::default();
for parameter in &def.params {
let Some(default) = parameter.default.as_ref() else {
continue;
};
let target = catalog
.resolve_catalog_column_type(¶meter.type_name)
.or_else(|| ColumnType::from_sql_name(¶meter.type_name).ok());
constants.collect_expression(default, target.as_ref());
}
constants
}
fn validate_variadic_declaration(
catalog: &dyn RoutineTypeCatalog,
def: &CreateFunction,
) -> Result<(), SQLError> {
let variadic_positions = def
.params
.iter()
.enumerate()
.filter_map(|(index, parameter)| {
(parameter.mode == FunctionParamMode::Variadic).then_some(index)
})
.collect::<Vec<_>>();
if variadic_positions.len() > 1 {
return Err(routine_definition_error(
"VARIADIC parameter must be the last parameter",
));
}
if let Some(&variadic_index) = variadic_positions.first() {
let parameter = &def.params[variadic_index];
if !routine_declaration_is_array(catalog, ¶meter.type_name) {
return Err(routine_definition_error(
"VARIADIC parameter must be an array",
));
}
let has_later_input = def.params[variadic_index + 1..].iter().any(|parameter| {
matches!(
parameter.mode,
FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
)
});
if has_later_input || def.is_procedure && variadic_index + 1 != def.params.len() {
return Err(routine_definition_error(
"VARIADIC parameter must be the last parameter",
));
}
}
Ok(())
}
#[derive(Default)]
struct PolymorphicInputs {
simple: bool,
compatible: bool,
any: bool,
}
fn validate_routine_input_types(def: &CreateFunction) -> Result<PolymorphicInputs, SQLError> {
let mut inputs = PolymorphicInputs::default();
for parameter in &def.params {
let type_name = canonical_routine_type_name(¶meter.type_name);
let is_input = matches!(
parameter.mode,
FunctionParamMode::In | FunctionParamMode::InOut | FunctionParamMode::Variadic
);
if let Some(family) = polymorphic_family(&type_name) {
inputs.any |= is_input;
if is_input {
match family {
RoutinePolymorphicFamily::Simple => inputs.simple = true,
RoutinePolymorphicFamily::Compatible => inputs.compatible = true,
}
}
continue;
}
if ROUTINE_PARAMETER_PSEUDO_TYPES.contains(&type_name.as_str()) {
let supported = match type_name.as_str() {
"record" => !is_input || def.language == "plpgsql",
"refcursor" => true,
_ => false,
};
if !supported {
return Err(routine_definition_error(format!(
"{} routines cannot have arguments of type {type_name}",
def.language
)));
}
}
}
Ok(inputs)
}
fn validate_routine_output_types(
def: &CreateFunction,
inputs: &PolymorphicInputs,
) -> Result<(), SQLError> {
let mut output_types = def
.output_params()
.into_iter()
.map(|parameter| parameter.type_name.as_str())
.collect::<Vec<_>>();
if let FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } =
&def.returns
{
output_types.push(type_name);
}
for output_type in output_types {
let type_name = canonical_routine_type_name(output_type);
match polymorphic_family(&type_name) {
Some(RoutinePolymorphicFamily::Simple) if !inputs.simple => {
return Err(routine_definition_error(format!(
"cannot determine result data type: a result of type {type_name} requires at least one simple polymorphic input"
)));
}
Some(RoutinePolymorphicFamily::Compatible) if !inputs.compatible => {
return Err(routine_definition_error(format!(
"cannot determine result data type: a result of type {type_name} requires at least one compatible polymorphic input"
)));
}
None if ROUTINE_RESULT_PSEUDO_TYPES.contains(&type_name.as_str())
&& !matches!(type_name.as_str(), "record" | "refcursor" | "void")
&& !(type_name == "trigger"
&& def.language == "plpgsql"
&& !def.is_procedure
&& def.params.is_empty()
&& matches!(def.returns, FunctionReturns::Scalar { .. })) =>
{
return Err(routine_definition_error(format!(
"{} routines cannot return type {type_name}",
def.language
)));
}
Some(_) | None => {}
}
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RoutinePolymorphicFamily {
Simple,
Compatible,
}
fn polymorphic_family(type_name: &str) -> Option<RoutinePolymorphicFamily> {
if !POLYMORPHIC_PSEUDO_TYPES.contains(&type_name) {
return None;
}
Some(if type_name.starts_with("anycompatible") {
RoutinePolymorphicFamily::Compatible
} else {
RoutinePolymorphicFamily::Simple
})
}
fn routine_declaration_is_array(catalog: &dyn RoutineTypeCatalog, type_name: &str) -> bool {
let canonical = canonical_routine_type_name(type_name);
canonical.ends_with("[]")
|| matches!(
canonical.as_str(),
"anyarray" | "anycompatiblearray" | "int2vector" | "oidvector"
)
|| catalog
.resolve_catalog_column_type(&canonical)
.is_some_and(|ty| routine_column_type_is_array(&ty))
}
fn routine_column_type_is_array(ty: &ColumnType) -> bool {
match ty {
ColumnType::Array(_) | ColumnType::AnyArray => true,
ColumnType::Domain { base, .. } => routine_column_type_is_array(base),
_ => false,
}
}
pub(super) fn routine_definition_error(message: impl Into<String>) -> SQLError {
SQLError::Routine {
sqlstate: "42P13".into(),
message: message.into(),
}
}