use std::collections::HashMap;
use std::sync::Arc;
use surrealdb_strand::Strand;
use surrealdb_types::{SqlFormat, ToSql};
use super::helpers::{
args_access_mode, args_required_context, check_permission, evaluate_args, validate_arg_count,
validate_return,
};
use crate::catalog::providers::DatabaseProvider;
use crate::err::Error;
use crate::exec::physical_expr::{BlockPhysicalExpr, EvalContext, PhysicalExpr};
use crate::exec::{AccessMode, BoxFut};
use crate::expr::{ControlFlow, FlowResult};
use crate::val::Value;
#[derive(Debug, Clone)]
pub struct UserDefinedFunctionExec {
pub(crate) name: String,
pub(crate) arguments: Vec<Arc<dyn PhysicalExpr>>,
pub(crate) plan_depth: u32,
}
impl PhysicalExpr for UserDefinedFunctionExec {
fn name(&self) -> &'static str {
"UserDefinedFunction"
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn required_context(&self) -> crate::exec::ContextLevel {
args_required_context(&self.arguments).max(crate::exec::ContextLevel::Database)
}
fn evaluate<'a>(&'a self, ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
let func_name = format!("fn::{}", self.name);
let db_ctx =
ctx.exec_ctx.database().map_err(|e| ControlFlow::Err(anyhow::Error::new(e)))?;
if !ctx.capabilities().allows_function_name(&func_name) {
return Err(Error::FunctionNotAllowed(func_name).into());
}
let ns_id = db_ctx.ns_ctx.ns.namespace_id;
let db_id = db_ctx.db.database_id;
let func_def = ctx
.txn()
.get_db_function(ns_id, db_id, &self.name, ctx.exec_ctx.version_stamp())
.await
.map_err(|e| anyhow::anyhow!("Function '{}' not found: {}", func_name, e))?;
let auth_limit =
crate::iam::AuthLimit::try_from(&func_def.auth_limit).map_err(|e| {
anyhow::anyhow!("Invalid auth limit on function '{}': {}", func_name, e)
})?;
let limited_ctx = ctx.exec_ctx.with_limited_auth(&auth_limit);
let ctx = EvalContext {
exec_ctx: &limited_ctx,
current_value: ctx.current_value,
local_params: ctx.local_params,
recursion_ctx: ctx.recursion_ctx,
document_root: ctx.document_root,
skip_fetch_perms: ctx.skip_fetch_perms,
computing_record: ctx.computing_record,
plan_depth: ctx.plan_depth,
};
if ctx.exec_ctx.should_check_perms(crate::iam::Action::View)? {
check_permission(&func_def.permissions, &func_name, &ctx).await?;
}
let evaluated_args = evaluate_args(&self.arguments, ctx.clone()).await?;
validate_arg_count(&func_name, evaluated_args.len(), &func_def.args)?;
let mut local_params: HashMap<Strand, Value> = HashMap::new();
for ((param_name, kind), arg_value) in func_def.args.iter().zip(evaluated_args) {
let coerced = arg_value.coerce_to_kind(kind).map_err(|e| {
Error::InvalidFunctionArguments {
name: func_name.clone(),
message: format!("Failed to coerce argument `${param_name}`: {e}"),
}
})?;
local_params.insert(param_name.as_str().into(), coerced);
}
let mut isolated_ctx = limited_ctx.clone();
for (name, value) in &local_params {
isolated_ctx = isolated_ctx.with_param(name.clone(), value.clone());
}
let block_expr = BlockPhysicalExpr {
block: func_def.block.clone(),
};
let eval_ctx = EvalContext {
exec_ctx: &isolated_ctx,
current_value: ctx.current_value,
local_params: Some(&local_params),
recursion_ctx: None,
document_root: None,
skip_fetch_perms: ctx.skip_fetch_perms,
computing_record: ctx.computing_record.clone(),
plan_depth: self.plan_depth + 1,
};
let result = match block_expr.evaluate(eval_ctx).await {
Ok(v) => v,
Err(ControlFlow::Return(v)) => v,
Err(ControlFlow::Break) | Err(ControlFlow::Continue) => {
return Err(Error::InvalidControlFlow.into());
}
Err(e) => return Err(e),
};
Ok(validate_return(&func_name, func_def.returns.as_ref(), result)?)
})
}
fn access_mode(&self) -> AccessMode {
AccessMode::ReadWrite.combine(args_access_mode(&self.arguments))
}
}
impl ToSql for UserDefinedFunctionExec {
fn fmt_sql(&self, f: &mut String, _fmt: SqlFormat) {
f.push_str("fn::");
f.push_str(&self.name);
f.push_str("(...)");
}
}