uqa-engine 0.1.12

Engine: schema-aware table store, catalog restore, transactions
//
// Unified Query Algebra
//
// Copyright (c) 2023-2026 Cognica, Inc.
//

//! Routine execution, recursion limits, and `LANGUAGE sql` result shaping.

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) };
}

/// Native stack budget for nested routine calls, measured from the
/// outermost routine entry. The `PostgreSQL` `max_stack_depth`
/// setting plays the same role (default 2MB there); this budget is
/// sized so the guard
/// always fires before a 2MB thread stack (the Rust test-runner
/// default) is exhausted, even in debug builds.
const STACK_BYTE_BUDGET: usize = 1_000_000;

/// Approximate current stack position.
#[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(),
    }
}

/// RAII guard for the user-routine nesting caps: a configurable
/// frame-count limit plus a native stack-byte budget.
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 _transition_scope = crate::sql::triggers::enter_empty_transition_relation_scope();
    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(&parameter.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))
}

/// `LANGUAGE sql` body: run every statement; the last statement's
/// result shapes the routine output.
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(&parameter.type_name)
                .ok()
                .or_else(|| crate::sql::resolve_catalog_column_type(engine, &parameter.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, &params).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()
    };
    // PostgreSQL enforces the final statement's column shape at
    // CREATE time; the engine has no schema binding there, so the
    // same 42P13 error surfaces on the first call instead.
    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, &parameter.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}"),
    }
}