use super::{
coerce_routine_value, result_row_values, Cell, CompiledFunctionBody, CreateFunction, Engine,
FunctionReturns, Interpreter, PLpgSQLDatum, RoutineOutcome, SQLError, SQLParam, SQLResult,
SQLUserFunction, UnifiedPlanExecutor, Value,
};
use crate::engine_user_functions::{canonical_routine_type_name, routine_returns_anonymous_record};
use uqa_sql::ast::RoutineInvocationBinding;
pub(in crate::sql) struct TriggerRoutineContext {
pub(in crate::sql) old: Value,
pub(in crate::sql) new: Value,
pub(in crate::sql) name: String,
pub(in crate::sql) when: String,
pub(in crate::sql) level: String,
pub(in crate::sql) operation: String,
pub(in crate::sql) relation_oid: i64,
pub(in crate::sql) table_name: String,
pub(in crate::sql) table_schema: String,
pub(in crate::sql) arguments: Vec<String>,
}
thread_local! {
static CALL_DEPTH: Cell<usize> = const { Cell::new(0) };
static STACK_BASE: Cell<usize> = const { Cell::new(0) };
}
const STACK_BYTE_BUDGET: usize = 1_000_000;
#[inline(never)]
fn stack_marker() -> usize {
let marker = 0u8;
std::ptr::from_ref(&marker) as usize
}
fn stack_depth_error() -> SQLError {
SQLError::Routine {
sqlstate: "54001".into(),
message: "stack depth limit exceeded".into(),
}
}
pub(super) struct DepthGuard;
impl DepthGuard {
pub(super) fn enter(engine: &Engine) -> Result<Self, SQLError> {
let depth = CALL_DEPTH.get();
if depth == 0 {
STACK_BASE.set(stack_marker());
} else if STACK_BASE.get().abs_diff(stack_marker()) > STACK_BYTE_BUDGET {
return Err(stack_depth_error());
}
if depth >= engine.sql_function_depth_limit() {
return Err(stack_depth_error());
}
CALL_DEPTH.set(depth + 1);
Ok(Self)
}
}
impl Drop for DepthGuard {
fn drop(&mut self) {
CALL_DEPTH.set(CALL_DEPTH.get().saturating_sub(1));
}
}
pub(super) fn execute_routine(
engine: &Engine,
function: &SQLUserFunction,
bound: Vec<Value>,
invocation: &RoutineInvocationBinding,
) -> Result<RoutineOutcome, SQLError> {
if matches!(
&function.def.returns,
FunctionReturns::Scalar { type_name }
if canonical_routine_type_name(type_name) == "trigger"
) {
return Err(SQLError::Routine {
sqlstate: "0A000".into(),
message: "trigger functions can only be called as triggers".into(),
});
}
let _guard = DepthGuard::enter(engine)?;
let specialized = specialized_definition(&function.def, invocation)?;
let definition = specialized.as_ref().unwrap_or(&function.def);
engine.ensure_routine_execute_privilege(definition)?;
engine.with_routine_context(definition, || match &function.compiled {
CompiledFunctionBody::PLpgSQL(parsed) => {
if specialized.is_some() {
let mut parsed = parsed.clone();
for (index, parameter) in definition.params.iter().enumerate() {
if let Some(PLpgSQLDatum::Var(variable)) = parsed.datums.get_mut(index) {
variable.type_name.clone_from(¶meter.type_name);
}
}
execute_plpgsql_language(engine, definition, &parsed, bound)
} else {
execute_plpgsql_language(engine, definition, parsed, bound)
}
}
CompiledFunctionBody::SQL(statements) => {
execute_sql_language(engine, definition, statements, &bound)
}
})
}
pub(in crate::sql) fn execute_trigger_routine(
engine: &Engine,
function: &SQLUserFunction,
context: &TriggerRoutineContext,
) -> Result<Value, SQLError> {
let _guard = DepthGuard::enter(engine)?;
engine.ensure_routine_execute_privilege(&function.def)?;
engine.with_routine_context(&function.def, || {
let CompiledFunctionBody::PLpgSQL(parsed) = &function.compiled else {
return Err(SQLError::Unsupported(
"only LANGUAGE plpgsql trigger functions are executable".into(),
));
};
let mut interpreter = Interpreter::new(engine, &function.def, parsed, Vec::new())?;
interpreter.initialize_trigger_context(parsed, context)?;
interpreter.run(&parsed.action)?;
Ok(interpreter.into_outcome().value)
})
}
fn execute_plpgsql_language(
engine: &Engine,
definition: &CreateFunction,
parsed: &uqa_sql::plpgsql::PLpgSQLFunction,
bound: Vec<Value>,
) -> Result<RoutineOutcome, SQLError> {
let mut interpreter = Interpreter::new(engine, definition, parsed, bound)?;
interpreter.run(&parsed.action)?;
Ok(interpreter.into_outcome())
}
fn specialized_definition(
definition: &CreateFunction,
invocation: &RoutineInvocationBinding,
) -> Result<Option<CreateFunction>, SQLError> {
if invocation.parameter_types.len() != definition.params.len() {
return Err(SQLError::Internal(format!(
"routine `{}` has {} concrete parameter types for {} parameters",
definition.name,
invocation.parameter_types.len(),
definition.params.len()
)));
}
let parameters_match = definition
.params
.iter()
.zip(&invocation.parameter_types)
.all(|(parameter, type_name)| parameter.type_name == *type_name);
let return_type_matches = match (&invocation.return_type, &definition.returns) {
(Some(concrete), FunctionReturns::Scalar { type_name })
| (Some(concrete), FunctionReturns::SetOf { type_name }) => concrete == type_name,
(None, _) | (Some(_), FunctionReturns::None | FunctionReturns::Table) => true,
};
if parameters_match && return_type_matches {
return Ok(None);
}
let mut specialized = definition.clone();
for (parameter, type_name) in specialized
.params
.iter_mut()
.zip(&invocation.parameter_types)
{
parameter.type_name.clone_from(type_name);
}
if let Some(return_type) = &invocation.return_type {
match &mut specialized.returns {
FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
type_name.clone_from(return_type);
}
FunctionReturns::None | FunctionReturns::Table => {}
}
}
Ok(Some(specialized))
}
fn execute_sql_language(
engine: &Engine,
def: &CreateFunction,
plans: &[uqa_planner::UnifiedPlan],
bound: &[Value],
) -> Result<RoutineOutcome, SQLError> {
let call_params = def.call_params();
if call_params.len() != bound.len() {
return Err(SQLError::Internal(format!(
"routine `{}` received {} values for {} concrete call parameters",
def.name,
bound.len(),
call_params.len()
)));
}
let params = bound
.iter()
.cloned()
.zip(call_params)
.map(|(value, parameter)| {
let ty = uqa_sql::ast::ColumnType::from_sql_name(¶meter.type_name)
.ok()
.or_else(|| crate::sql::resolve_catalog_column_type(engine, ¶meter.type_name))
.ok_or_else(|| {
SQLError::TypeMismatch(format!("unknown type `{}`", parameter.type_name))
})?;
Ok(SQLParam::typed_scalar(value, ty))
})
.collect::<Result<Vec<_>, SQLError>>()?;
let mut last = SQLResult::empty();
for plan in plans {
last = UnifiedPlanExecutor::new_nested(engine, ¶ms).execute(plan)?;
}
let out_params = def.output_params();
let returns_anonymous_record = routine_returns_anonymous_record(def);
let returns_void = matches!(
&def.returns,
FunctionReturns::Scalar { type_name } if type_name == "void"
);
let expected = if out_params.is_empty() {
1
} else {
out_params.len()
};
let shape_checked =
!returns_void && !returns_anonymous_record && (!def.is_procedure || !out_params.is_empty());
if shape_checked && last.columns.len() != expected {
return Err(sql_body_shape_error(def));
}
if def.returns_set() {
let mut set_rows = Vec::with_capacity(last.rows.len());
for row_index in 0..last.rows.len() {
let mut values = result_row_values(&last, row_index).unwrap_or_default();
if !returns_anonymous_record && values.len() != expected {
return Err(sql_body_shape_error(def));
}
if returns_anonymous_record {
values = vec![anonymous_record_value(&last.columns, values)];
} else if out_params.is_empty() {
if let FunctionReturns::SetOf { type_name } = &def.returns {
values[0] = coerce_routine_value(engine, &values[0], type_name)?;
}
} else {
for (value, parameter) in values.iter_mut().zip(&out_params) {
*value = coerce_routine_value(engine, value, ¶meter.type_name)?;
}
}
set_rows.push(values);
}
return Ok(RoutineOutcome {
value: Value::Null,
out_values: vec![Value::Null; out_params.len()],
set_rows,
anonymous_record_column_types: returns_anonymous_record
.then(|| last.column_types.clone()),
});
}
let first = result_row_values(&last, 0);
if !out_params.is_empty() {
let mut out_values = vec![Value::Null; out_params.len()];
if let Some(values) = first {
for (idx, value) in values.into_iter().take(out_values.len()).enumerate() {
out_values[idx] = coerce_routine_value(engine, &value, &out_params[idx].type_name)?;
}
}
return Ok(RoutineOutcome {
value: Value::Null,
out_values,
set_rows: Vec::new(),
anonymous_record_column_types: None,
});
}
let value = match first {
Some(_) if returns_void => Value::Null,
Some(values) if returns_anonymous_record => anonymous_record_value(&last.columns, values),
Some(mut values) => {
if values.is_empty() {
Value::Null
} else {
let value = values.remove(0);
match &def.returns {
FunctionReturns::Scalar { type_name } => {
coerce_routine_value(engine, &value, type_name)?
}
_ => value,
}
}
}
None => Value::Null,
};
Ok(RoutineOutcome {
value,
out_values: Vec::new(),
set_rows: Vec::new(),
anonymous_record_column_types: returns_anonymous_record.then(|| last.column_types.clone()),
})
}
fn anonymous_record_value(columns: &[String], values: Vec<Value>) -> Value {
Value::Record(columns.iter().cloned().zip(values).collect())
}
fn sql_body_shape_error(def: &CreateFunction) -> SQLError {
let declared = match &def.returns {
FunctionReturns::Scalar { type_name } | FunctionReturns::SetOf { type_name } => {
type_name.clone()
}
FunctionReturns::Table => "record".into(),
FunctionReturns::None => "record".into(),
};
SQLError::Routine {
sqlstate: "42P13".into(),
message: format!("return type mismatch in function declared to return {declared}"),
}
}