1use super::compilation::RoutineCompilationContext;
10use crate::{
11 ast::{CreateFunction, FunctionBody},
12 binding::{
13 bind_expression_plan_routines_for_storage,
14 stored_columns::StoredSourceCatalog,
15 stored_relations::bind_stored_statement_relations,
16 stored_routines::{bind_catalog_statement_routines, CatalogRoutineContext},
17 syntax_sites::expression_syntax_sites,
18 },
19 catalog::{resolution::RelationLookupMode, stored_ast},
20 plan::ExpressionPlan,
21 RowSchema, SQLError,
22};
23
24#[derive(Clone, Copy)]
25pub enum RoutineCompilationMode {
26 Definition,
27 Persisted,
28}
29
30pub fn bind_routine_definition_dependencies(
31 context: &RoutineCompilationContext<'_>,
32 sources: &dyn StoredSourceCatalog,
33 def: &mut CreateFunction,
34 mode: RoutineCompilationMode,
35) -> Result<bool, SQLError> {
36 let mut changed = bind_sql_standard_body_relations(context, sources, def, mode)?;
37 for parameter in &mut def.params {
38 let Some(default) = &mut parameter.default else {
39 continue;
40 };
41 let lowered = ExpressionPlan::lower_with(default.clone(), &|name: &str| {
42 context.catalog.has_registered_aggregate_function(name)
43 });
44 let mut plan = lowered.clone();
45 let binding = context.catalog.binding_snapshot()?;
46 let default_type = bind_expression_plan_routines_for_storage(
47 context.routines,
48 &mut plan,
49 &[],
50 &binding.context(),
51 &RowSchema::default(),
52 )?;
53 let default_type =
54 if crate::type_resolution::routine_polymorphic_type(¶meter.type_name).is_some() {
55 default_type
56 } else {
57 Some(
58 context
59 .types
60 .resolve_catalog_column_type_name(¶meter.type_name)?
61 .without_type_modifiers(),
62 )
63 };
64 let default_type =
65 super::defaults::default_expression_type(¶meter.type_name, default, default_type);
66 if parameter.default_type != default_type {
67 parameter.default_type = default_type;
68 changed = true;
69 }
70 let sites = expression_syntax_sites(&lowered, &plan)?;
71 changed |= stored_ast::bind_stored_expression_sites(default, &sites)?;
72 if let Some(ty) = user_defined_type(context, ¶meter.type_name)? {
74 changed |= stored_ast::fold_assigned_stored_literal(
75 default,
76 &ty,
77 context.routines.enum_labels(),
78 )?;
79 }
80 }
81 Ok(changed)
82}
83
84fn bind_sql_standard_body_relations(
85 context: &RoutineCompilationContext<'_>,
86 sources: &dyn StoredSourceCatalog,
87 def: &mut CreateFunction,
88 mode: RoutineCompilationMode,
89) -> Result<bool, SQLError> {
90 let FunctionBody::Statements(statements) = &mut def.body else {
91 return Ok(false);
92 };
93 let mut changed = false;
94 for statement in statements {
95 super::compilation::validate_sql_standard_statement(statement)?;
96 changed |= bind_stored_statement_relations(
97 context.relations,
98 statement,
99 match mode {
100 RoutineCompilationMode::Definition => RelationLookupMode::Dynamic,
101 RoutineCompilationMode::Persisted => RelationLookupMode::Bound,
102 },
103 matches!(mode, RoutineCompilationMode::Persisted),
104 "SQL routine body",
105 )?;
106 changed |=
107 super::merge_columns::bind_stored_merge_target_columns(context.merge, statement)?;
108 changed |= sources.stored_source_columns().bind_statement(statement)?;
109 }
110 Ok(changed)
111}
112
113pub fn bind_sql_standard_body_routines(
115 context: &RoutineCompilationContext<'_>,
116 def: &mut CreateFunction,
117 mode: RoutineCompilationMode,
118) -> Result<bool, SQLError> {
119 if !matches!(def.body, FunctionBody::Statements(_)) {
120 return Ok(false);
121 }
122 super::compilation::validate_routine_signature(context, def)?;
124 let positional = super::body_validation::routine_parameter_values(context.types, def);
125 let parameters = super::body_parameters::sql_body_parameter_scope(def, &positional)?;
126 let result_types = def_result_types(context, &def.params, &def.returns)?;
127 let FunctionBody::Statements(statements) = &mut def.body else {
128 return Ok(false);
129 };
130 let lowering = super::compilation::SQLRoutineLowering {
131 bind_catalog_dependencies: true,
132 persisted_definition: matches!(mode, RoutineCompilationMode::Persisted),
133 preserve_target_expressions: true,
134 };
135 let mut changed = false;
136 for statement in statements.iter_mut() {
137 let mut lowered =
138 super::compilation::lower_sql_routine_statement(context, statement.clone(), lowering)?;
139 let binding = context.catalog.binding_snapshot()?;
140 crate::binding::bind_routine_parameter_references(
141 context.routines,
142 &mut lowered,
143 &positional,
144 &binding.context(),
145 ¶meters,
146 )?;
147 let routines = bind_catalog_statement_routines(
148 &CatalogRoutineContext {
149 routines: context.routines,
150 binding: &binding.context(),
151 },
152 &lowered,
153 &positional,
154 )?;
155 changed |= stored_ast::bind_stored_statement_sites(statement, &routines.sites)?;
156 }
157 if let Some(statement) = statements.last_mut() {
158 changed |= fold_sql_function_result(context, result_types, statement)?;
159 }
160 Ok(changed)
161}
162
163fn user_defined_type(
165 context: &RoutineCompilationContext<'_>,
166 name: &str,
167) -> Result<Option<crate::ColumnType>, SQLError> {
168 if crate::ast::UserTypeIdentity::parse(name).is_none() {
169 return Ok(None);
170 }
171 context.routines.resolve_type_name(name)
172}
173
174fn def_result_types(
176 context: &RoutineCompilationContext<'_>,
177 params: &[crate::ast::FunctionParam],
178 returns: &crate::ast::FunctionReturns,
179) -> Result<Vec<Option<crate::ColumnType>>, SQLError> {
180 use crate::ast::{FunctionParamMode, FunctionReturns};
181 let outputs = params
182 .iter()
183 .filter(|parameter| {
184 matches!(
185 parameter.mode,
186 FunctionParamMode::Out | FunctionParamMode::InOut | FunctionParamMode::Table
187 )
188 })
189 .collect::<Vec<_>>();
190 let names = if outputs.is_empty() {
191 match returns {
192 FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
193 vec![type_name.as_str()]
194 }
195 FunctionReturns::None | FunctionReturns::Table => Vec::new(),
196 }
197 } else {
198 outputs
199 .iter()
200 .map(|parameter| parameter.type_name.as_str())
201 .collect()
202 };
203 names
204 .into_iter()
205 .map(|name| user_defined_type(context, name))
206 .collect()
207}
208
209fn fold_sql_function_result(
211 context: &RoutineCompilationContext<'_>,
212 types: Vec<Option<crate::ColumnType>>,
213 statement: &mut crate::ast::Statement,
214) -> Result<bool, SQLError> {
215 let crate::ast::Statement::Select(select) = statement else {
216 return Ok(false);
217 };
218 if types.is_empty()
219 || select.set_op.is_some()
220 || !select.values.is_empty()
221 || select.projections.len() != types.len()
222 {
223 return Ok(false);
224 }
225 let mut changed = false;
226 for (projection, ty) in select.projections.iter_mut().zip(&types) {
227 let Some(ty) = ty else {
228 continue;
229 };
230 changed |= stored_ast::fold_assigned_stored_literal(
231 &mut projection.expr,
232 ty,
233 context.routines.enum_labels(),
234 )?;
235 }
236 Ok(changed)
237}