use super::compilation::{self, StoredRoutineCompilationContext};
use std::sync::Arc;
use uqa_sql::{
ast::{CreateFunction, FunctionBody},
binding::stored_columns::StoredSourceCatalog,
routines::{
body_parameters::record_sql_standard_body_parameters,
compilation::{compile_function_body, defer_function_body},
dependencies::{self, RoutineCompilationMode},
regclass::{self, RoutineRegclassCatalog},
CompiledFunctionBody, RoutineBody,
},
SQLError,
};
pub struct RoutineDefinitionContext<'a> {
pub compilation: StoredRoutineCompilationContext<'a>,
pub sources: &'a dyn StoredSourceCatalog,
pub regclasses: &'a dyn RoutineRegclassCatalog,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RoutineBodyCompilation {
Checked,
Unchecked,
Stored,
}
impl RoutineBodyCompilation {
const fn mode(self) -> RoutineCompilationMode {
match self {
Self::Checked | Self::Unchecked => RoutineCompilationMode::Definition,
Self::Stored => RoutineCompilationMode::Persisted,
}
}
}
pub struct BoundRoutine {
pub body: RoutineBody,
pub validated: Option<CompiledFunctionBody>,
pub changed: bool,
}
pub fn compile_catalog_bound_routine(
context: &RoutineDefinitionContext<'_>,
def: &mut CreateFunction,
bodies: RoutineBodyCompilation,
) -> Result<BoundRoutine, SQLError> {
let mode = bodies.mode();
if matches!(mode, RoutineCompilationMode::Definition) {
if matches!(def.body, FunctionBody::Statements(_))
|| def
.params
.iter()
.any(|parameter| parameter.default.is_some())
{
def.creation_search_path = context.compilation.session.routine_search_path();
} else {
def.creation_search_path.clear();
}
}
let mut changed = bind_routine_definition_dependencies(context, def, mode)?;
if matches!(mode, RoutineCompilationMode::Definition) {
changed |= record_sql_standard_body_parameters(&context.compilation.analysis, def)?;
}
let body_changed = with_creation_search_path(context, def, |def| {
dependencies::bind_sql_standard_body_routines(&context.compilation.analysis, def, mode)
})?;
changed |= body_changed;
let mut compiled = compile_routine_body(&context.compilation, def, bodies)?;
let regclass_changed = bind_routine_regclass_constants(context, def)?;
changed |= regclass_changed;
if regclass_changed {
compiled = compile_routine_body(&context.compilation, def, bodies)?;
}
if !matches!(def.body, FunctionBody::Statements(_)) {
return Ok(BoundRoutine {
body: RoutineBody::Source,
validated: compiled,
changed,
});
}
let compiled = compiled.ok_or_else(|| {
SQLError::Internal(format!(
"SQL-standard body of routine `{}` was not compiled",
def.name
))
})?;
Ok(BoundRoutine {
body: RoutineBody::Bound(Arc::new(compiled)),
validated: None,
changed,
})
}
fn compile_routine_body(
context: &StoredRoutineCompilationContext<'_>,
def: &CreateFunction,
bodies: RoutineBodyCompilation,
) -> Result<Option<CompiledFunctionBody>, SQLError> {
match (bodies, &def.body) {
(
RoutineBodyCompilation::Checked | RoutineBodyCompilation::Unchecked,
FunctionBody::Statements(_),
) => compile_function_body(&context.analysis, def).map(Some),
(RoutineBodyCompilation::Stored, FunctionBody::Statements(_)) => {
compilation::compile_persisted_sql_function(context, def).map(Some)
}
(RoutineBodyCompilation::Checked, FunctionBody::Source(_)) => {
compilation::with_routine_settings(context, def, || {
compile_function_body(&context.analysis, def)
})
.map(Some)
}
(RoutineBodyCompilation::Unchecked, FunctionBody::Source(_)) => {
defer_function_body(&context.analysis, def)
}
(RoutineBodyCompilation::Stored, FunctionBody::Source(_)) => Ok(None),
}
}
fn bind_routine_definition_dependencies(
context: &RoutineDefinitionContext<'_>,
def: &mut CreateFunction,
mode: RoutineCompilationMode,
) -> Result<bool, SQLError> {
with_creation_search_path(context, def, |def| {
dependencies::bind_routine_definition_dependencies(
&context.compilation.analysis,
context.sources,
def,
mode,
)
})
}
fn with_creation_search_path<T>(
context: &RoutineDefinitionContext<'_>,
def: &mut CreateFunction,
bind: impl FnOnce(&mut CreateFunction) -> Result<T, SQLError>,
) -> Result<T, SQLError> {
if def.creation_search_path.is_empty() {
return bind(def);
}
let previous = context
.compilation
.session
.replace_routine_search_path(def.creation_search_path.clone());
let result = bind(def);
context
.compilation
.session
.restore_routine_search_path(previous);
result
}
fn bind_routine_regclass_constants(
context: &RoutineDefinitionContext<'_>,
definition: &mut CreateFunction,
) -> Result<bool, SQLError> {
let previous = context
.compilation
.session
.replace_routine_search_path(definition.creation_search_path.clone());
let result = regclass::bind_routine_regclass_constants(
context.compilation.analysis.types,
context.regclasses,
definition,
);
context
.compilation
.session
.restore_routine_search_path(previous);
result
}