use super::{
builtin_routine_support_oid, lifecycle::require_routine_ownership, routine_kind,
routine_local_name,
};
use crate::catalog::roles::identity::RoleSubject;
use crate::{
ast::{AlterRoutineStmt, CreateFunction},
catalog::roles::{role_inherits, RoleDefinition, RoleMembership, RoleMembershipKey},
type_resolution::canonical_routine_type_name,
SQLError,
};
use std::collections::BTreeMap;
pub trait RoutineSupportAuthority {
fn current_user_is_superuser(&self) -> bool;
}
pub fn validate_routine_support(
authority: &dyn RoutineSupportAuthority,
support: &str,
) -> Result<(), SQLError> {
if builtin_routine_support_oid(support).is_none() {
return Err(SQLError::Routine {
sqlstate: "42883".into(),
message: format!("function {support}(internal) does not exist"),
});
}
if !authority.current_user_is_superuser() {
return Err(SQLError::Routine {
sqlstate: "42501".into(),
message: "must be superuser to specify a support function".into(),
});
}
Ok(())
}
pub fn validate_routine_security_attributes(
def: &CreateFunction,
current_user_is_superuser: bool,
) -> Result<(), SQLError> {
if current_user_is_superuser {
return Ok(());
}
if def.support.is_some() {
return Err(SQLError::Routine {
sqlstate: "42501".into(),
message: "must be superuser to specify a support function".into(),
});
}
if def.security.leakproof {
return Err(SQLError::Routine {
sqlstate: "42501".into(),
message: "only superuser can define a leakproof function".into(),
});
}
Ok(())
}
pub fn prepare_routine_replacement(
existing: &CreateFunction,
def: &mut CreateFunction,
current_user: &(impl RoleSubject + ?Sized),
roles: &BTreeMap<String, RoleDefinition>,
memberships: &BTreeMap<RoleMembershipKey, RoleMembership>,
signature: &str,
) -> Result<(), SQLError> {
if !def.or_replace {
let kind = routine_kind(def);
return Err(SQLError::Routine {
sqlstate: "42723".into(),
message: format!(
"{kind} \"{}\" already exists with same argument types",
routine_local_name(&existing.name)?
),
});
}
require_routine_ownership(
"function",
&routine_local_name(&existing.name)?,
role_inherits(
roles,
memberships,
current_user,
&crate::routines::security::bound_routine_owner(existing)?,
),
)?;
if existing.is_procedure != def.is_procedure {
return Err(SQLError::Routine {
sqlstate: "42809".into(),
message: "cannot change routine kind".into(),
});
}
validate_replacement_result(existing, def, signature)?;
validate_replacement_defaults(existing, def, signature)?;
def.object_id = Some(existing.object_id.ok_or_else(|| {
SQLError::Internal(format!(
"existing routine `{}` has no catalog object identity",
existing.name,
))
})?);
def.catalog_oid = existing.catalog_oid;
def.owner = existing.owner;
def.execute_acl.clone_from(&existing.execute_acl);
Ok(())
}
fn validate_replacement_defaults(
existing: &CreateFunction,
replacement: &CreateFunction,
signature: &str,
) -> Result<(), SQLError> {
let defaults = |definition: &CreateFunction| {
definition
.params
.iter()
.filter(|parameter| parameter.default.is_some())
.map(|parameter| {
parameter.default_type.as_ref().map_or(705, |ty| match ty {
crate::ast::RoutineDefaultType::Concrete(ty) => {
crate::catalog::type_metadata::pg_type_oid(ty)
}
crate::ast::RoutineDefaultType::Polymorphic(name) => {
crate::catalog::type_metadata::routine_type_oid(name)
}
})
})
.collect::<Vec<_>>()
};
let existing_defaults = defaults(existing);
let replacement_defaults = defaults(replacement);
let message = if replacement_defaults.len() < existing_defaults.len() {
"cannot remove parameter defaults from existing function"
} else if !existing_defaults
.iter()
.rev()
.zip(replacement_defaults.iter().rev())
.all(|(existing, replacement)| existing == replacement)
{
"cannot change data type of existing parameter default value"
} else {
return Ok(());
};
Err(SQLError::Diagnostic {
sqlstate: "42P13".into(),
message: message.into(),
detail: None,
hint: Some(format!(
"Use DROP {} {signature} first.",
if existing.is_procedure {
"PROCEDURE"
} else {
"FUNCTION"
}
)),
})
}
pub fn alter_routine_attributes(
existing: &CreateFunction,
stmt: &AlterRoutineStmt,
current_user_is_superuser: bool,
authority: &dyn RoutineSupportAuthority,
) -> Result<CreateFunction, SQLError> {
super::attributes::check_attribute_clauses(&stmt.attribute_clauses, existing.is_procedure)?;
let mut def = existing.clone();
if let Some(volatility) = stmt.volatility {
def.volatility = volatility;
}
if let Some(strict) = stmt.strict {
def.strict = strict;
}
if let Some(security_definer) = stmt.security_definer {
def.security.security_definer = security_definer;
}
if let Some(leakproof) = stmt.leakproof {
if leakproof && !current_user_is_superuser {
return Err(SQLError::Routine {
sqlstate: "42501".into(),
message: "only superuser can define a leakproof function".into(),
});
}
def.security.leakproof = leakproof;
}
if let Some(cost) = stmt.cost {
super::attributes::validate_cost(Some(cost))?;
def.cost = Some(cost);
}
if let Some(rows) = stmt.rows {
super::attributes::validate_rows(Some(rows))?;
super::attributes::validate_rows_applicability(Some(rows), existing.returns_set())?;
def.rows = Some(rows);
}
if let Some(support) = &stmt.support {
validate_routine_support(authority, support)?;
def.support = Some(support.clone());
}
super::attributes::validate_parallel(&stmt.attribute_clauses)?;
if let Some(parallel) = stmt.parallel {
def.parallel = parallel;
}
def.config_actions.clone_from(&stmt.config_actions);
Ok(def)
}
fn validate_replacement_result(
existing: &CreateFunction,
replacement: &CreateFunction,
signature: &str,
) -> Result<(), SQLError> {
let existing_type = canonical_routine_type_name(super::declaration::result_type_name(existing));
let replacement_type =
canonical_routine_type_name(super::declaration::result_type_name(replacement));
let same_type =
existing_type == replacement_type && existing.returns_set() == replacement.returns_set();
let same_record = existing_type != "record"
|| record_output_shape(existing) == record_output_shape(replacement);
if same_type && same_record {
return Ok(());
}
Err(SQLError::Diagnostic {
sqlstate: "42P13".into(),
message: if existing.is_procedure && !same_type {
"cannot change whether a procedure has output parameters"
} else {
"cannot change return type of existing function"
}
.into(),
detail: same_type.then(|| "Row type defined by OUT parameters is different.".into()),
hint: Some(format!(
"Use DROP {} {signature} first.",
if existing.is_procedure {
"PROCEDURE"
} else {
"FUNCTION"
}
)),
})
}
fn record_output_shape(definition: &CreateFunction) -> Vec<(String, String)> {
let outputs = definition.output_params();
if !definition.is_procedure && outputs.len() < 2 {
return Vec::new();
}
outputs
.iter()
.enumerate()
.map(|(index, parameter)| {
(
if parameter.name.is_empty() {
format!("column{}", index + 1)
} else {
parameter.name.clone()
},
canonical_routine_type_name(¶meter.type_name),
)
})
.collect()
}
#[cfg(test)]
mod tests;