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::dbs::capabilities::Error as CapabilitiesError;
use crate::exec::physical_expr::{BlockPhysicalExpr, EvalContext, PhysicalExpr};
use crate::exec::{AccessMode, BoxFut, Error as ExecError};
use crate::expr::{ControlFlow, Error as ExprError, 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,
pub(crate) body_access_mode: AccessMode,
}
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(
anyhow::Error::new(CapabilitiesError::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| {
ExprError::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(),
matches_scope: None,
};
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(ExecError::InvalidControlFlow.into());
}
Err(e) => return Err(e),
};
Ok(validate_return(&func_name, func_def.returns.as_ref(), result)?)
})
}
fn access_mode(&self) -> AccessMode {
self.body_access_mode.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("(...)");
}
}
#[cfg(all(test, feature = "kv-mem"))]
#[allow(clippy::unwrap_used)]
mod tests {
use crate::exec::AccessMode;
use crate::exec::operators::test_util::TestDb;
use crate::exec::planner::Planner;
async fn access_mode_of(db: &TestDb, expression: &str) -> AccessMode {
let ctx = db.exec_ctx().await;
let crate::exec::ExecutionContext::Database(db_ctx) = &ctx else {
panic!("exec_ctx builds a Database context");
};
let txn = ctx.txn();
let planner = Planner::with_txn(
ctx.ctx(),
&db_ctx.ns_ctx.root.function_registry,
txn,
Some("test".to_owned()),
Some("test".to_owned()),
);
let expr: crate::expr::Expr = crate::syn::expr(expression).unwrap().into();
planner.physical_expr(expr).await.unwrap().access_mode()
}
#[tokio::test]
async fn a_read_only_udf_resolves_read_only() {
let db = TestDb::new(
"DEFINE FUNCTION fn::pure() { RETURN 1; };
DEFINE FUNCTION fn::relay() { RETURN fn::pure(); };",
)
.await;
assert_eq!(access_mode_of(&db, "fn::pure()").await, AccessMode::ReadOnly);
assert_eq!(access_mode_of(&db, "fn::relay()").await, AccessMode::ReadOnly);
}
#[tokio::test]
async fn a_writing_udf_resolves_read_write() {
let db = TestDb::new(
"DEFINE TABLE log SCHEMALESS;
DEFINE FUNCTION fn::sink() { CREATE log; RETURN 1; };",
)
.await;
assert_eq!(access_mode_of(&db, "fn::sink()").await, AccessMode::ReadWrite);
}
#[tokio::test]
async fn an_undefined_callee_resolves_read_write() {
let db = TestDb::new("").await;
assert_eq!(access_mode_of(&db, "fn::ghost()").await, AccessMode::ReadWrite);
}
#[tokio::test]
async fn a_body_invoking_a_closure_field_resolves_read_write() {
let db =
TestDb::new("DEFINE FUNCTION fn::call_field($o: object) { RETURN $o.w(); };").await;
assert_eq!(access_mode_of(&db, "fn::call_field({ w: || 1 })").await, AccessMode::ReadWrite);
}
#[tokio::test]
async fn a_read_only_body_with_a_writing_argument_is_read_write() {
let db = TestDb::new(
"DEFINE TABLE log SCHEMALESS;
DEFINE FUNCTION fn::id($x: any) { RETURN $x; };
DEFINE FUNCTION fn::sink() { CREATE log; RETURN 1; };",
)
.await;
assert_eq!(access_mode_of(&db, "fn::id(fn::sink())").await, AccessMode::ReadWrite);
}
}