use super::context::RoutineInvocationSession;
use parking_lot::Mutex;
use std::collections::BTreeMap;
use std::sync::Arc;
use uqa_sql::ast::CreateFunction;
use uqa_sql::routines::{
compilation::{
compile_function_body, compile_function_body_for_execution, RoutineCompilationContext,
},
CompiledFunctionBody, RoutineBody, SQLUserFunction,
};
use uqa_sql::SQLError;
pub trait RoutineBodySession {
fn retain_routine_body(
&self,
function: &SQLUserFunction,
body: CompiledFunctionBody,
) -> Result<(), SQLError>;
}
pub(super) struct CompiledRoutine {
pub(super) version: u64,
pub(super) body: Arc<CompiledFunctionBody>,
pub(super) plpgsql: Vec<super::plpgsql_cache::PreparedBody>,
}
#[derive(Default)]
pub struct SessionRoutineBodies {
pub(super) compiled: Mutex<BTreeMap<[u8; 16], CompiledRoutine>>,
sql_inputs: crate::routines::sql_body::inputs::SQLRoutineInputs,
pub(in crate::routines) procedural: crate::routines::preparation::PLpgSQLPreparationRegistry,
}
impl SessionRoutineBodies {
pub(crate) fn invalidate(
&self,
affects: impl Fn(&uqa_sql::prepared::dependencies::PreparedAnalysisDependencies) -> bool,
) {
self.sql_inputs.invalidate(&affects);
self.procedural.invalidate(&affects);
}
pub fn sql_inputs(&self) -> &crate::routines::sql_body::inputs::SQLRoutineInputs {
&self.sql_inputs
}
pub fn body(
&self,
function: &SQLUserFunction,
compile: impl FnOnce(&CreateFunction) -> Result<CompiledFunctionBody, SQLError>,
) -> Result<Arc<CompiledFunctionBody>, SQLError> {
if let Some(body) = self.retained(function)? {
return Ok(body);
}
let identity = routine_identity(function)?;
let version = function.definition_version()?;
let body = Arc::new(compile(&function.def)?);
self.compiled
.lock()
.insert(identity, CompiledRoutine::new(version, Arc::clone(&body)));
Ok(body)
}
pub fn inspect(
&self,
function: &SQLUserFunction,
compile: impl FnOnce(&CreateFunction) -> Result<CompiledFunctionBody, SQLError>,
) -> Result<Arc<CompiledFunctionBody>, SQLError> {
self.retained(function)?
.map_or_else(|| compile(&function.def).map(Arc::new), Ok)
}
fn retained(
&self,
function: &SQLUserFunction,
) -> Result<Option<Arc<CompiledFunctionBody>>, SQLError> {
if let RoutineBody::Bound(body) = &function.body {
return Ok(Some(Arc::clone(body)));
}
let identity = routine_identity(function)?;
let version = function.definition_version()?;
Ok(self
.compiled
.lock()
.get(&identity)
.filter(|compiled| compiled.version == version)
.map(|compiled| Arc::clone(&compiled.body)))
}
pub fn retain(
&self,
function: &SQLUserFunction,
body: CompiledFunctionBody,
) -> Result<(), SQLError> {
let identity = routine_identity(function)?;
let version = function.definition_version()?;
self.compiled
.lock()
.insert(identity, CompiledRoutine::new(version, Arc::new(body)));
Ok(())
}
}
pub(super) fn routine_identity(function: &SQLUserFunction) -> Result<[u8; 16], SQLError> {
function.def.object_id.ok_or_else(|| {
SQLError::Internal(format!(
"routine `{}` has no object identity",
function.def.name
))
})
}
pub fn compile_session_body(
session: &dyn RoutineInvocationSession,
compilation: &RoutineCompilationContext<'_>,
def: &CreateFunction,
) -> Result<CompiledFunctionBody, SQLError> {
let compile = || {
let mut body = compile_function_body_for_execution(compilation, def)?;
if let CompiledFunctionBody::PLpgSQL(parsed) = &mut body {
crate::routines::compilation::apply_session_compile_options(session, parsed);
}
Ok(body)
};
if def.config.is_empty() && !def.security.security_definer {
return compile();
}
super::scopes::with_routine_context(session, def, compile)
}
pub fn compile_analysis_body(
session: &dyn RoutineInvocationSession,
compilation: &RoutineCompilationContext<'_>,
def: &CreateFunction,
) -> Result<CompiledFunctionBody, SQLError> {
struct AnalysisParser<'a>(&'a dyn uqa_sql::routines::compilation::RoutineParserCatalog);
impl uqa_sql::routines::compilation::RoutineParserCatalog for AnalysisParser<'_> {
fn plpgsql_catalog(&self) -> Result<uqa_sql::plpgsql::PlpgsqlCatalog, SQLError> {
self.0.plpgsql_catalog()
}
fn parser_settings(&self) -> uqa_sql::parser::ParserSettings {
self.0.parser_settings()
}
}
let parsers = AnalysisParser(compilation.parsers);
let analysis = RoutineCompilationContext {
parsers: &parsers,
..*compilation
};
let compile = || {
let mut body = compile_function_body(&analysis, def)?;
if let CompiledFunctionBody::PLpgSQL(parsed) = &mut body {
crate::routines::compilation::apply_session_compile_options(session, parsed);
}
Ok(body)
};
if def.config.is_empty() && !def.security.security_definer {
compile()
} else {
super::scopes::with_routine_context(session, def, compile)
}
}
#[cfg(test)]
mod tests;