1use super::{
10 declaration::{
11 resolve_plpgsql_datum_types, routine_definition_error, routine_parameter_regrole_constants,
12 validate_routine_declaration, RoutineTypeCatalog,
13 },
14 merge_columns::StoredMergeColumnCatalog,
15 CompiledFunctionBody, RoutineResolution,
16};
17use crate::{
18 ast::{ColumnType, CreateFunction, FunctionBody, FunctionReturns, Statement},
19 binding::{
20 snapshot::BindingSnapshot,
21 stored_relations::{
22 self, StoredQueryBindingContext, StoredQueryNamespace, StoredQuerySequences,
23 StoredRelationCatalog,
24 },
25 },
26 catalog::{
27 regrole_dependencies::{StoredRegroleConstants, StoredRegroleResolver},
28 resolution::RelationLookupMode,
29 },
30 plan::UnifiedPlan,
31 plpgsql::PlpgsqlCatalog,
32 type_resolution::canonical_routine_type_name,
33 SQLError, ScalarExpr,
34};
35
36pub trait RoutineParserCatalog {
37 fn plpgsql_catalog(&self) -> Result<PlpgsqlCatalog, SQLError>;
38 fn parser_settings(&self) -> crate::parser::ParserSettings {
39 crate::parser::ParserSettings::default()
40 }
41 fn parser_notice(&self, _notice: crate::SQLNotice) {}
42}
43pub trait RoutineCompilationCatalog: crate::schema::dependencies::oid_alias::OidAliasInput {
44 fn has_registered_aggregate_function(&self, name: &str) -> bool;
45 fn binding_snapshot(&self) -> Result<BindingSnapshot, SQLError>;
46 fn stored_query_namespace(&self) -> StoredQueryNamespace;
47}
48#[derive(Clone, Copy)]
49pub struct RoutineCompilationContext<'a> {
50 pub types: &'a dyn RoutineTypeCatalog,
51 pub parsers: &'a dyn RoutineParserCatalog,
52 pub catalog: &'a dyn RoutineCompilationCatalog,
53 pub routines: &'a dyn RoutineResolution,
54 pub relations: &'a dyn StoredRelationCatalog,
55 pub sequences: &'a dyn StoredQuerySequences,
56 pub merge: &'a dyn StoredMergeColumnCatalog,
57 pub regroles: &'a dyn StoredRegroleResolver,
58}
59
60pub fn compile_function_body(
61 context: &RoutineCompilationContext<'_>,
62 def: &CreateFunction,
63) -> Result<CompiledFunctionBody, SQLError> {
64 compile_function_body_inner(
65 context,
66 def,
67 false,
68 crate::plpgsql::PLpgSQLCompileMode::Validate,
69 )
70}
71
72pub fn compile_function_body_for_execution(
74 context: &RoutineCompilationContext<'_>,
75 def: &CreateFunction,
76) -> Result<CompiledFunctionBody, SQLError> {
77 compile_function_body_inner(
78 context,
79 def,
80 false,
81 crate::plpgsql::PLpgSQLCompileMode::Runtime,
82 )
83}
84
85pub fn compile_persisted_function_body(
86 context: &RoutineCompilationContext<'_>,
87 def: &CreateFunction,
88) -> Result<CompiledFunctionBody, SQLError> {
89 compile_function_body_inner(
90 context,
91 def,
92 true,
93 crate::plpgsql::PLpgSQLCompileMode::Validate,
94 )
95}
96
97pub fn defer_function_body(
99 context: &RoutineCompilationContext<'_>,
100 def: &CreateFunction,
101) -> Result<Option<CompiledFunctionBody>, SQLError> {
102 if matches!(def.body, FunctionBody::Statements(_)) {
103 return compile_function_body(context, def).map(Some);
104 }
105 validate_routine_signature(context, def)?.reject_with(context.regroles)?;
106 Ok(None)
107}
108
109pub(super) fn validate_routine_signature(
111 context: &RoutineCompilationContext<'_>,
112 def: &CreateFunction,
113) -> Result<StoredRegroleConstants, SQLError> {
114 if !matches!(def.language.as_str(), "plpgsql" | "sql") {
115 return Err(SQLError::Routine {
116 sqlstate: "42704".into(),
117 message: format!("language \"{}\" does not exist", def.language),
118 });
119 }
120 if def.language == "plpgsql" && matches!(def.body, FunctionBody::Statements(_)) {
121 return Err(routine_definition_error(
122 "inline SQL function body only valid for language SQL",
123 ));
124 }
125 let stored_regrole_constants = routine_parameter_regrole_constants(context.types, def);
126 stored_regrole_constants.validate_inputs_with(context.regroles)?;
127 validate_routine_declaration(def)?;
128 Ok(stored_regrole_constants)
129}
130
131fn reject_trigger_function_arguments(def: &CreateFunction) -> Result<(), SQLError> {
133 let returns_trigger = matches!(
134 &def.returns,
135 FunctionReturns::Scalar { type_name } if canonical_routine_type_name(type_name) == "trigger"
136 );
137 if returns_trigger && def.identity_arity() != 0 {
138 return Err(SQLError::Diagnostic {
139 sqlstate: "42P13".into(),
140 message: "trigger functions cannot have declared arguments".into(),
141 detail: None,
142 hint: Some(
143 "The arguments of the trigger can be accessed through TG_NARGS and TG_ARGV instead."
144 .into(),
145 ),
146 });
147 }
148 Ok(())
149}
150
151fn compile_function_body_inner(
152 context: &RoutineCompilationContext<'_>,
153 def: &CreateFunction,
154 persisted_definition: bool,
155 mode: crate::plpgsql::PLpgSQLCompileMode,
156) -> Result<CompiledFunctionBody, SQLError> {
157 with_parser_context(context.parsers, || {
158 compile_body(context, def, persisted_definition, mode)
159 })
160}
161
162pub fn with_parser_context<T>(
164 parsers: &dyn RoutineParserCatalog,
165 compile: impl FnOnce() -> Result<T, SQLError>,
166) -> Result<T, SQLError> {
167 let (result, metadata) = crate::parser::with_settings(parsers.parser_settings(), compile);
168 for notice in metadata.notices.iter() {
169 parsers.parser_notice(notice.clone());
170 }
171 result
172}
173
174fn compile_body(
175 context: &RoutineCompilationContext<'_>,
176 def: &CreateFunction,
177 persisted_definition: bool,
178 mode: crate::plpgsql::PLpgSQLCompileMode,
179) -> Result<CompiledFunctionBody, SQLError> {
180 let mut stored_regrole_constants = validate_routine_signature(context, def)?;
181 match def.language.as_str() {
182 "plpgsql" => {
183 stored_regrole_constants.reject_with(context.regroles)?;
184 reject_trigger_function_arguments(def)?;
185 let catalog = context.parsers.plpgsql_catalog()?;
186 let mut function =
187 crate::plpgsql::parse_function_with_catalog_mode(def, &catalog, mode)?;
188 resolve_plpgsql_datum_types(context.types, &mut function)?;
189 Ok(CompiledFunctionBody::PLpgSQL(function))
190 }
191 "sql" => {
192 let (statements, bind_catalog_dependencies) = match &def.body {
193 FunctionBody::Source(source) => {
194 stored_regrole_constants.reject_with(context.regroles)?;
195 (crate::compile(source)?, false)
196 }
197 FunctionBody::Statements(statements) => (statements.clone(), true),
198 };
199 let mut plans = compile_sql_routine_plans(
200 context,
201 def,
202 statements,
203 bind_catalog_dependencies,
204 persisted_definition && matches!(def.body, FunctionBody::Statements(_)),
205 )?;
206 if bind_catalog_dependencies {
207 for plan in &mut plans {
208 stored_regrole_constants.collect_plan(plan);
209 }
210 stored_regrole_constants.reject_with(context.regroles)?;
211 }
212 Ok(CompiledFunctionBody::SQL(plans))
213 }
214 _ => unreachable!("routine language was validated above"),
215 }
216}
217
218fn compile_sql_routine_plans(
219 context: &RoutineCompilationContext<'_>,
220 def: &CreateFunction,
221 statements: Vec<Statement>,
222 bind_catalog_dependencies: bool,
223 persisted_definition: bool,
224) -> Result<Vec<UnifiedPlan>, SQLError> {
225 let positional_parameters =
226 super::body_validation::routine_parameter_values(context.types, def);
227 let parameters = super::body_parameters::sql_body_parameter_scope(def, &positional_parameters)?;
228 statements
229 .into_iter()
230 .map(|statement| {
231 let mut plan = lower_sql_routine_statement(
232 context,
233 statement,
234 SQLRoutineLowering {
235 bind_catalog_dependencies,
236 persisted_definition,
237 preserve_target_expressions: false,
238 },
239 )?;
240 if bind_catalog_dependencies {
244 let binding = context.catalog.binding_snapshot()?;
245 crate::binding::bind_routine_parameter_references(
246 context.routines,
247 &mut plan,
248 &positional_parameters,
249 &binding.context(),
250 ¶meters,
251 )?;
252 if let UnifiedPlan::Query(query) = &mut plan {
253 crate::binding::bind_query_plan_routines_for_storage(
254 context.routines,
255 query,
256 &positional_parameters,
257 &binding.context(),
258 None,
259 )?;
260 }
261 }
262 Ok(plan)
265 })
266 .collect()
267}
268
269#[derive(Clone, Copy)]
271pub struct SQLRoutineLowering {
272 pub bind_catalog_dependencies: bool,
274 pub persisted_definition: bool,
276 pub preserve_target_expressions: bool,
278}
279
280pub fn lower_sql_routine_statement(
282 context: &RoutineCompilationContext<'_>,
283 mut statement: Statement,
284 lowering: SQLRoutineLowering,
285) -> Result<UnifiedPlan, SQLError> {
286 if lowering.bind_catalog_dependencies {
287 validate_sql_standard_statement(&statement)?;
288 }
289 if lowering.bind_catalog_dependencies && !lowering.preserve_target_expressions {
290 super::merge_columns::normalize_stored_merge_target_columns(context.merge, &mut statement)?;
291 }
292 let mut plan = UnifiedPlan::lower_with(statement, &|name: &str| {
293 context.catalog.has_registered_aggregate_function(name)
294 });
295 if lowering.persisted_definition {
296 plan.rewrite_scalar_expressions(&mut |expression| {
297 let ScalarExpr::Func { name, binding, .. } = expression else {
298 return;
299 };
300 crate::ast::FunctionBinding::upgrade_legacy_serialized_dispatch(name, binding);
301 });
302 }
303 if lowering.bind_catalog_dependencies {
304 match &mut plan {
305 UnifiedPlan::Query(query) => {
306 let namespace = context.catalog.stored_query_namespace();
307 stored_relations::bind_stored_query_relations(
308 &StoredQueryBindingContext {
309 relations: context.relations,
310 lookup_mode: if lowering.persisted_definition {
311 RelationLookupMode::Bound
312 } else {
313 RelationLookupMode::Dynamic
314 },
315 sequences: context.sequences,
316 temporary_schema: &namespace.temporary_schema,
317 transition_relations: &namespace.transition_relations,
318 },
319 query,
320 "SQL routine body",
321 false,
322 lowering.persisted_definition,
323 )?;
324 }
325 UnifiedPlan::Command(_) => {
326 crate::binding::stored_routines::mark_catalog_statement_relations_bound(&mut plan)?;
327 }
328 }
329 }
330 Ok(plan)
331}
332
333pub(super) fn validate_sql_standard_statement(statement: &Statement) -> Result<(), SQLError> {
335 if matches!(statement, Statement::Drop(drop) if drop.kind == crate::ast::DropKind::ForeignServer)
336 {
337 return Err(SQLError::Unsupported(
338 "DROP SERVER is not yet supported in unquoted SQL function body".into(),
339 ));
340 }
341 if matches!(statement, Statement::Drop(drop) if drop.kind == crate::ast::DropKind::ForeignWrapper)
342 {
343 return Err(SQLError::Unsupported(
344 "DROP FOREIGN DATA WRAPPER is not yet supported in unquoted SQL function body".into(),
345 ));
346 }
347 if matches!(statement, Statement::CreateForeignWrapper(_)) {
348 return Err(SQLError::Unsupported(
349 "CREATE FOREIGN DATA WRAPPER is not yet supported in unquoted SQL function body".into(),
350 ));
351 }
352 if matches!(statement, Statement::CreateForeignServer(_)) {
353 return Err(SQLError::Unsupported(
354 "CREATE SERVER is not yet supported in unquoted SQL function body".into(),
355 ));
356 }
357 if matches!(
358 statement,
359 Statement::CreateForeignTable(_) | Statement::CreateForeignTableDefinition(_)
360 ) {
361 return Err(SQLError::Unsupported(
362 "CREATE FOREIGN TABLE is not yet supported in unquoted SQL function body".into(),
363 ));
364 }
365 Ok(())
366}
367
368pub fn bind_analyzed_sql_body_types(
371 types: &dyn RoutineTypeCatalog,
372 plan: &mut UnifiedPlan,
373) -> Result<(), SQLError> {
374 crate::binding::stored_types::bind_unified_plan_type_identities(plan, &mut |name| {
375 let resolved = types.resolve_catalog_column_type(name);
376 let mut element = resolved.as_ref();
377 while let Some(ColumnType::Array(inner)) = element {
378 element = Some(inner.as_ref());
379 }
380 let domain = matches!(element, Some(ColumnType::Domain { .. }));
381 Ok(resolved.filter(|_| !domain))
382 })
383}
384
385#[cfg(test)]
386mod tests;