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 routine_local_name, CompiledFunctionBody, RoutineResolution,
16};
17use crate::{
18 ast::{ColumnType, CreateFunction, FunctionBody, Statement},
19 binding::{
20 snapshot::BindingSnapshot,
21 stored_relations::{
22 self, StoredQueryBindingContext, StoredQueryNamespace, StoredQuerySequences,
23 StoredRelationCatalog,
24 },
25 },
26 catalog::regrole_dependencies::StoredRegroleResolver,
27 plan::UnifiedPlan,
28 plpgsql::PlpgsqlCatalog,
29 SQLError, ScalarExpr,
30};
31
32pub trait RoutineParserCatalog {
33 fn plpgsql_catalog(&self) -> Result<PlpgsqlCatalog, SQLError>;
34}
35pub trait RoutineCompilationCatalog {
36 fn has_registered_aggregate_function(&self, name: &str) -> bool;
37 fn binding_snapshot(&self) -> Result<BindingSnapshot, SQLError>;
38 fn stored_query_namespace(&self) -> StoredQueryNamespace;
39}
40#[derive(Clone, Copy)]
41pub struct RoutineCompilationContext<'a> {
42 pub types: &'a dyn RoutineTypeCatalog,
43 pub parsers: &'a dyn RoutineParserCatalog,
44 pub catalog: &'a dyn RoutineCompilationCatalog,
45 pub routines: &'a dyn RoutineResolution,
46 pub relations: &'a dyn StoredRelationCatalog,
47 pub sequences: &'a dyn StoredQuerySequences,
48 pub merge: &'a dyn StoredMergeColumnCatalog,
49 pub regroles: &'a dyn StoredRegroleResolver,
50}
51
52pub fn compile_function_body(
53 context: &RoutineCompilationContext<'_>,
54 def: &CreateFunction,
55) -> Result<CompiledFunctionBody, SQLError> {
56 compile_function_body_inner(context, def, false, false)
57}
58
59pub fn compile_persisted_function_body(
60 context: &RoutineCompilationContext<'_>,
61 def: &CreateFunction,
62) -> Result<CompiledFunctionBody, SQLError> {
63 compile_function_body_inner(context, def, true, false)
64}
65
66pub fn compile_persisted_function_dependencies(
67 context: &RoutineCompilationContext<'_>,
68 def: &CreateFunction,
69) -> Result<CompiledFunctionBody, SQLError> {
70 compile_function_body_inner(context, def, true, true)
71}
72
73fn compile_function_body_inner(
74 context: &RoutineCompilationContext<'_>,
75 def: &CreateFunction,
76 persisted_definition: bool,
77 preserve_target_expressions: bool,
78) -> Result<CompiledFunctionBody, SQLError> {
79 if !matches!(def.language.as_str(), "plpgsql" | "sql") {
80 return Err(SQLError::Routine {
81 sqlstate: "42704".into(),
82 message: format!("language \"{}\" does not exist", def.language),
83 });
84 }
85 if def.language == "plpgsql" && matches!(def.body, FunctionBody::Statements(_)) {
86 return Err(routine_definition_error(
87 "inline SQL function body only valid for language SQL",
88 ));
89 }
90 let mut stored_regrole_constants = routine_parameter_regrole_constants(context.types, def);
91 stored_regrole_constants.validate_inputs_with(context.regroles)?;
92 validate_routine_declaration(context.types, def)?;
93 match def.language.as_str() {
94 "plpgsql" => {
95 stored_regrole_constants.reject_with(context.regroles)?;
96 let catalog = context.parsers.plpgsql_catalog()?;
97 let mut function = crate::plpgsql::parse_function_with_catalog(def, &catalog)?;
98 resolve_plpgsql_datum_types(context.types, &mut function)?;
99 Ok(CompiledFunctionBody::PLpgSQL(function))
100 }
101 "sql" => {
102 let (statements, bind_catalog_dependencies) = match &def.body {
103 FunctionBody::Source(source) => {
104 stored_regrole_constants.reject_with(context.regroles)?;
105 (crate::compile(source)?, false)
106 }
107 FunctionBody::Statements(statements) => (statements.clone(), true),
108 };
109 let mut plans = compile_sql_routine_plans(
110 context,
111 def,
112 statements,
113 bind_catalog_dependencies,
114 persisted_definition && matches!(def.body, FunctionBody::Statements(_)),
115 preserve_target_expressions,
116 )?;
117 if bind_catalog_dependencies {
118 for plan in &mut plans {
119 stored_regrole_constants.collect_plan(plan);
120 }
121 stored_regrole_constants.reject_with(context.regroles)?;
122 }
123 Ok(CompiledFunctionBody::SQL(plans))
124 }
125 _ => unreachable!("routine language was validated above"),
126 }
127}
128
129fn compile_sql_routine_plans(
130 context: &RoutineCompilationContext<'_>,
131 def: &CreateFunction,
132 statements: Vec<Statement>,
133 bind_catalog_dependencies: bool,
134 persisted_definition: bool,
135 preserve_target_expressions: bool,
136) -> Result<Vec<UnifiedPlan>, SQLError> {
137 let local_name = routine_local_name(&def.name)?;
138 let signature_params = def.signature_params();
139 let parameter_names: Vec<String> = signature_params
140 .iter()
141 .map(|parameter| parameter.name.clone())
142 .collect();
143 let parameter_types = signature_params
144 .iter()
145 .map(|parameter| {
146 context
147 .types
148 .resolve_catalog_column_type(¶meter.type_name)
149 .or_else(|| ColumnType::from_sql_name(¶meter.type_name).ok())
150 })
151 .collect::<Vec<_>>();
152 let positional_parameters = parameter_types
153 .iter()
154 .map(|parameter_type| match parameter_type {
155 Some(parameter_type) => {
156 crate::SQLParam::typed_scalar(uqa_core::Value::Null, parameter_type.clone())
157 }
158 None => crate::SQLParam::scalar(uqa_core::Value::Null),
159 })
160 .collect::<Vec<_>>();
161 let parameter_scope = crate::RowSchema::with_qualified_types(
162 &local_name,
163 parameter_names.clone(),
164 parameter_types,
165 );
166 statements
167 .into_iter()
168 .map(|mut statement| {
169 if bind_catalog_dependencies && !preserve_target_expressions {
170 super::merge_columns::normalize_stored_merge_target_columns(
171 context.merge,
172 &mut statement,
173 )?;
174 }
175 let mut plan = UnifiedPlan::lower_with(statement, &|name: &str| {
176 context.catalog.has_registered_aggregate_function(name)
177 });
178 if persisted_definition {
179 plan.rewrite_scalar_expressions(&mut |expression| {
180 let ScalarExpr::Func { name, binding, .. } = expression else {
181 return;
182 };
183 crate::ast::FunctionBinding::upgrade_legacy_serialized_dispatch(name, binding);
184 });
185 }
186 if bind_catalog_dependencies {
187 match &mut plan {
188 UnifiedPlan::Query(query) => {
189 let namespace = context.catalog.stored_query_namespace();
190 stored_relations::bind_stored_query_relations(
191 &StoredQueryBindingContext {
192 relations: context.relations,
193 sequences: context.sequences,
194 temporary_schema: &namespace.temporary_schema,
195 transition_relations: &namespace.transition_relations,
196 },
197 query,
198 "SQL routine body",
199 false,
200 persisted_definition,
201 )?;
202 let binding = context.catalog.binding_snapshot()?;
203 crate::binding::bind_query_plan_routines_for_storage(
204 context.routines,
205 query,
206 &positional_parameters,
207 &binding.context(),
208 Some(¶meter_scope),
209 )?;
210 }
211 UnifiedPlan::Command(_) => {
212 crate::binding::stored_routines::mark_catalog_statement_relations_bound(
213 &mut plan,
214 )?;
215 }
216 }
217 }
218 plan.rewrite_scalar_expressions(&mut |expression| {
219 let parameter = match expression {
220 ScalarExpr::Column(name) => parameter_names
221 .iter()
222 .position(|parameter| !parameter.is_empty() && parameter == name),
223 ScalarExpr::QualifiedColumn {
224 qualifier, column, ..
225 } if qualifier == &local_name => parameter_names
226 .iter()
227 .position(|parameter| !parameter.is_empty() && parameter == column),
228 _ => None,
229 };
230 if let Some(position) = parameter {
231 *expression = ScalarExpr::Param(position + 1);
232 }
233 });
234 Ok(plan)
237 })
238 .collect()
239}