1use super::{
10 builtin_routine_support_oid, lifecycle::require_routine_ownership, routine_kind,
11 routine_local_name,
12};
13use crate::catalog::roles::identity::RoleSubject;
14use crate::{
15 ast::{AlterRoutineStmt, CreateFunction},
16 catalog::roles::{role_inherits, RoleDefinition, RoleMembership, RoleMembershipKey},
17 type_resolution::canonical_routine_type_name,
18 SQLError,
19};
20use std::collections::BTreeMap;
21
22pub trait RoutineSupportAuthority {
23 fn current_user_is_superuser(&self) -> bool;
24}
25pub fn validate_routine_support(
26 authority: &dyn RoutineSupportAuthority,
27 support: &str,
28) -> Result<(), SQLError> {
29 if builtin_routine_support_oid(support).is_none() {
30 return Err(SQLError::Routine {
31 sqlstate: "42883".into(),
32 message: format!("function {support}(internal) does not exist"),
33 });
34 }
35 if !authority.current_user_is_superuser() {
36 return Err(SQLError::Routine {
37 sqlstate: "42501".into(),
38 message: "must be superuser to specify a support function".into(),
39 });
40 }
41 Ok(())
42}
43
44pub fn validate_routine_security_attributes(
45 def: &CreateFunction,
46 current_user_is_superuser: bool,
47) -> Result<(), SQLError> {
48 if current_user_is_superuser {
49 return Ok(());
50 }
51 if def.support.is_some() {
53 return Err(SQLError::Routine {
54 sqlstate: "42501".into(),
55 message: "must be superuser to specify a support function".into(),
56 });
57 }
58 if def.security.leakproof {
59 return Err(SQLError::Routine {
60 sqlstate: "42501".into(),
61 message: "only superuser can define a leakproof function".into(),
62 });
63 }
64 Ok(())
65}
66
67pub fn prepare_routine_replacement(
68 existing: &CreateFunction,
69 def: &mut CreateFunction,
70 current_user: &(impl RoleSubject + ?Sized),
71 roles: &BTreeMap<String, RoleDefinition>,
72 memberships: &BTreeMap<RoleMembershipKey, RoleMembership>,
73 signature: &str,
74) -> Result<(), SQLError> {
75 if !def.or_replace {
76 let kind = routine_kind(def);
78 return Err(SQLError::Routine {
79 sqlstate: "42723".into(),
80 message: format!(
81 "{kind} \"{}\" already exists with same argument types",
82 routine_local_name(&existing.name)?
83 ),
84 });
85 }
86 require_routine_ownership(
88 "function",
89 &routine_local_name(&existing.name)?,
90 role_inherits(
91 roles,
92 memberships,
93 current_user,
94 &crate::routines::security::bound_routine_owner(existing)?,
95 ),
96 )?;
97 if existing.is_procedure != def.is_procedure {
98 return Err(SQLError::Routine {
99 sqlstate: "42809".into(),
100 message: "cannot change routine kind".into(),
101 });
102 }
103 validate_replacement_result(existing, def, signature)?;
104 validate_replacement_defaults(existing, def, signature)?;
105 def.object_id = Some(existing.object_id.ok_or_else(|| {
107 SQLError::Internal(format!(
108 "existing routine `{}` has no catalog object identity",
109 existing.name,
110 ))
111 })?);
112 def.catalog_oid = existing.catalog_oid;
113 def.owner = existing.owner;
114 def.execute_acl.clone_from(&existing.execute_acl);
115 Ok(())
116}
117
118fn validate_replacement_defaults(
120 existing: &CreateFunction,
121 replacement: &CreateFunction,
122 signature: &str,
123) -> Result<(), SQLError> {
124 let defaults = |definition: &CreateFunction| {
125 definition
126 .params
127 .iter()
128 .filter(|parameter| parameter.default.is_some())
129 .map(|parameter| {
130 parameter.default_type.as_ref().map_or(705, |ty| match ty {
131 crate::ast::RoutineDefaultType::Concrete(ty) => {
132 crate::catalog::type_metadata::pg_type_oid(ty)
133 }
134 crate::ast::RoutineDefaultType::Polymorphic(name) => {
135 crate::catalog::type_metadata::routine_type_oid(name)
136 }
137 })
138 })
139 .collect::<Vec<_>>()
140 };
141 let existing_defaults = defaults(existing);
142 let replacement_defaults = defaults(replacement);
143 let message = if replacement_defaults.len() < existing_defaults.len() {
144 "cannot remove parameter defaults from existing function"
145 } else if !existing_defaults
146 .iter()
147 .rev()
148 .zip(replacement_defaults.iter().rev())
149 .all(|(existing, replacement)| existing == replacement)
150 {
151 "cannot change data type of existing parameter default value"
152 } else {
153 return Ok(());
154 };
155 Err(SQLError::Diagnostic {
156 sqlstate: "42P13".into(),
157 message: message.into(),
158 detail: None,
159 hint: Some(format!(
160 "Use DROP {} {signature} first.",
161 if existing.is_procedure {
162 "PROCEDURE"
163 } else {
164 "FUNCTION"
165 }
166 )),
167 })
168}
169
170pub fn alter_routine_attributes(
172 existing: &CreateFunction,
173 stmt: &AlterRoutineStmt,
174 current_user_is_superuser: bool,
175 authority: &dyn RoutineSupportAuthority,
176) -> Result<CreateFunction, SQLError> {
177 super::attributes::check_attribute_clauses(&stmt.attribute_clauses, existing.is_procedure)?;
178 let mut def = existing.clone();
179 if let Some(volatility) = stmt.volatility {
180 def.volatility = volatility;
181 }
182 if let Some(strict) = stmt.strict {
183 def.strict = strict;
184 }
185 if let Some(security_definer) = stmt.security_definer {
186 def.security.security_definer = security_definer;
187 }
188 if let Some(leakproof) = stmt.leakproof {
189 if leakproof && !current_user_is_superuser {
190 return Err(SQLError::Routine {
191 sqlstate: "42501".into(),
192 message: "only superuser can define a leakproof function".into(),
193 });
194 }
195 def.security.leakproof = leakproof;
196 }
197 if let Some(cost) = stmt.cost {
198 super::attributes::validate_cost(Some(cost))?;
199 def.cost = Some(cost);
200 }
201 if let Some(rows) = stmt.rows {
202 super::attributes::validate_rows(Some(rows))?;
203 super::attributes::validate_rows_applicability(Some(rows), existing.returns_set())?;
204 def.rows = Some(rows);
205 }
206 if let Some(support) = &stmt.support {
207 validate_routine_support(authority, support)?;
208 def.support = Some(support.clone());
209 }
210 super::attributes::validate_parallel(&stmt.attribute_clauses)?;
211 if let Some(parallel) = stmt.parallel {
212 def.parallel = parallel;
213 }
214 def.config_actions.clone_from(&stmt.config_actions);
215 Ok(def)
216}
217
218fn validate_replacement_result(
219 existing: &CreateFunction,
220 replacement: &CreateFunction,
221 signature: &str,
222) -> Result<(), SQLError> {
223 let existing_type = canonical_routine_type_name(super::declaration::result_type_name(existing));
224 let replacement_type =
225 canonical_routine_type_name(super::declaration::result_type_name(replacement));
226 let same_type =
227 existing_type == replacement_type && existing.returns_set() == replacement.returns_set();
228 let same_record = existing_type != "record"
229 || record_output_shape(existing) == record_output_shape(replacement);
230 if same_type && same_record {
231 return Ok(());
232 }
233 Err(SQLError::Diagnostic {
234 sqlstate: "42P13".into(),
235 message: if existing.is_procedure && !same_type {
236 "cannot change whether a procedure has output parameters"
237 } else {
238 "cannot change return type of existing function"
239 }
240 .into(),
241 detail: same_type.then(|| "Row type defined by OUT parameters is different.".into()),
242 hint: Some(format!(
243 "Use DROP {} {signature} first.",
244 if existing.is_procedure {
245 "PROCEDURE"
246 } else {
247 "FUNCTION"
248 }
249 )),
250 })
251}
252
253fn record_output_shape(definition: &CreateFunction) -> Vec<(String, String)> {
257 let outputs = definition.output_params();
258 if !definition.is_procedure && outputs.len() < 2 {
259 return Vec::new();
260 }
261 outputs
262 .iter()
263 .enumerate()
264 .map(|(index, parameter)| {
265 (
266 if parameter.name.is_empty() {
267 format!("column{}", index + 1)
268 } else {
269 parameter.name.clone()
270 },
271 canonical_routine_type_name(¶meter.type_name),
272 )
273 })
274 .collect()
275}
276
277#[cfg(test)]
278mod tests;