use std::sync::Arc;
use surrealdb_types::{SqlFormat, ToSql};
use super::helpers::{args_access_mode, args_required_context};
#[cfg(feature = "ml")]
use super::helpers::{check_permission, evaluate_args};
use crate::err::Error;
use crate::exec::physical_expr::{EvalContext, PhysicalExpr};
use crate::exec::{AccessMode, BoxFut};
use crate::expr::{FlowResult, Model};
use crate::val::Value;
#[derive(Debug, Clone)]
pub struct ModelFunctionExec {
pub(crate) model: Model,
pub(crate) arguments: Vec<Arc<dyn PhysicalExpr>>,
}
impl PhysicalExpr for ModelFunctionExec {
fn name(&self) -> &'static str {
"ModelFunction"
}
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)
}
#[cfg(feature = "ml")]
fn evaluate<'a>(&'a self, ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
use surrealml_core::errors::error::SurrealError;
use surrealml_core::execution::compute::ModelComputation;
use surrealml_core::ndarray as mlNdarray;
use surrealml_core::storage::surml_file::SurMlFile;
use crate::catalog::providers::DatabaseProvider;
use crate::expr::model::get_model_path;
use crate::iam::Action;
use crate::val::Number;
const ARGUMENTS: &str = "The model expects 1 argument. The argument can be either a number, an object, or an array of numbers.";
let name = format!("ml::{}", self.model.name);
ctx.check_allowed_function(&name)?;
let db_ctx = ctx.exec_ctx.database().map_err(|_| {
anyhow::anyhow!("Model function '{}' requires database context", name)
})?;
let ns_id = db_ctx.ns_ctx.ns.namespace_id;
let db_id = db_ctx.db.database_id;
let val = ctx
.txn()
.get_db_model(
ns_id,
db_id,
&self.model.name,
&self.model.version,
ctx.exec_ctx.version_stamp(),
)
.await?
.ok_or_else(|| Error::MlNotFound {
name: format!("{}<{}>", self.model.name, self.model.version),
})?;
let ns_name = db_ctx.ns_name();
let db_name = db_ctx.db_name();
let path =
get_model_path(ns_name, db_name, &self.model.name, &self.model.version, &val.hash);
if ctx.exec_ctx.should_check_perms(Action::View)? {
check_permission(&val.permissions, &self.model.name, &ctx).await?;
}
let mut args = evaluate_args(&self.arguments, ctx.clone()).await?;
if args.len() != 1 {
return Err(Error::InvalidFunctionArguments {
name: format!("ml::{}<{}>", self.model.name, self.model.version),
message: ARGUMENTS.into(),
}
.into());
}
let argument = args.pop().expect("single argument validated above");
match argument {
Value::Object(v) => {
let mut args = v
.into_iter()
.map(|(k, v)| {
v.coerce_to::<f64>().map(|f| (k.into_string(), f as f32)).map_err(
|_| Error::InvalidFunctionArguments {
name: format!(
"ml::{}<{}>",
self.model.name, self.model.version
),
message: ARGUMENTS.into(),
},
)
})
.collect::<Result<std::collections::HashMap<String, f32>, _>>()?;
let bytes = crate::obs::get(&path).await?;
let outcome: Vec<f32> = tokio::task::spawn_blocking(move || {
let mut file =
SurMlFile::from_bytes(bytes).map_err(|err: SurrealError| {
anyhow::anyhow!("Failed to load model: {}", err.message)
})?;
let compute_unit = ModelComputation {
surml_file: &mut file,
};
compute_unit.buffered_compute(&mut args).map_err(|err: SurrealError| {
anyhow::anyhow!("Model computation failed: {}", err.message)
})
})
.await
.map_err(|e| anyhow::anyhow!("ML task failed: {e}"))??;
Ok(outcome
.into_iter()
.map(|x| Value::Number(Number::Float(x as f64)))
.collect())
}
Value::Number(v) => {
let args: f32 = Value::Number(v).coerce_to::<f64>().map_err(|_| {
Error::InvalidFunctionArguments {
name: format!("ml::{}<{}>", self.model.name, self.model.version),
message: ARGUMENTS.into(),
}
})? as f32;
let bytes = crate::obs::get(&path).await?;
let tensor = mlNdarray::arr1::<f32>(&[args]).into_dyn();
let outcome: Vec<f32> = tokio::task::spawn_blocking(move || {
let mut file =
SurMlFile::from_bytes(bytes).map_err(|err: SurrealError| {
anyhow::anyhow!("Failed to load model: {}", err.message)
})?;
let compute_unit = ModelComputation {
surml_file: &mut file,
};
compute_unit.raw_compute(tensor, None).map_err(|err: SurrealError| {
anyhow::anyhow!("Model computation failed: {}", err.message)
})
})
.await
.map_err(|e| anyhow::anyhow!("ML task failed: {e}"))??;
Ok(outcome
.into_iter()
.map(|x| Value::Number(Number::Float(x as f64)))
.collect())
}
Value::Array(v) => {
let args = v
.into_iter()
.map(|x| x.coerce_to::<f64>().map(|x| x as f32))
.collect::<Result<Vec<f32>, _>>()
.map_err(|_| Error::InvalidFunctionArguments {
name: format!("ml::{}<{}>", self.model.name, self.model.version),
message: ARGUMENTS.into(),
})?;
let bytes = crate::obs::get(&path).await?;
let tensor = mlNdarray::arr1::<f32>(&args).into_dyn();
let outcome: Vec<f32> = tokio::task::spawn_blocking(move || {
let mut file =
SurMlFile::from_bytes(bytes).map_err(|err: SurrealError| {
anyhow::anyhow!("Failed to load model: {}", err.message)
})?;
let compute_unit = ModelComputation {
surml_file: &mut file,
};
compute_unit.raw_compute(tensor, None).map_err(|err: SurrealError| {
anyhow::anyhow!("Model computation failed: {}", err.message)
})
})
.await
.map_err(|e| anyhow::anyhow!("ML task failed: {e}"))??;
Ok(outcome
.into_iter()
.map(|x| Value::Number(Number::Float(x as f64)))
.collect())
}
_ => Err(Error::InvalidFunctionArguments {
name: format!("ml::{}<{}>", self.model.name, self.model.version),
message: ARGUMENTS.into(),
}
.into()),
}
})
}
#[cfg(not(feature = "ml"))]
fn evaluate<'a>(&'a self, _ctx: EvalContext<'a>) -> BoxFut<'a, FlowResult<Value>> {
Box::pin(async move {
Err(Error::InvalidModel {
message: String::from("Machine learning computation is not enabled."),
}
.into())
})
}
fn access_mode(&self) -> AccessMode {
AccessMode::ReadOnly.combine(args_access_mode(&self.arguments))
}
}
impl ToSql for ModelFunctionExec {
fn fmt_sql(&self, f: &mut String, fmt: SqlFormat) {
self.model.fmt_sql(f, fmt);
f.push_str("(...)");
}
}