diff --git a/Cargo.lock b/Cargo.lock index 55a96a2c33..3a3993f485 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1849,7 +1849,7 @@ dependencies = [ [[package]] name = "daggy" version = "0.8.0" -source = "git+https://github.com/getdozer/daggy?branch=feat/try_map#4ff8ebbdba979ecd7f638b6636c33ec9c0d27ccc" +source = "git+https://github.com/getdozer/daggy?branch=feat%2Ftry_map#4ff8ebbdba979ecd7f638b6636c33ec9c0d27ccc" dependencies = [ "petgraph 0.6.3", "serde", @@ -2768,6 +2768,12 @@ version = "0.3.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fea41bba32d969b513997752735605054bc0dfa92b4c56bf1189f2e174be7a10" +[[package]] +name = "downcast-rs" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75b325c5dbd37f80359721ad39aca5a29fb04c89279657cffdda8736d0c0b9d2" + [[package]] name = "dozer-cli" version = "0.4.0" @@ -3046,7 +3052,9 @@ dependencies = [ "multimap 0.9.1", "proptest", "regex", + "tempfile", "tokio", + "wat", ] [[package]] @@ -3068,7 +3076,10 @@ dependencies = [ "ort", "proptest", "sqlparser 0.35.0", + "tempfile", "tokio", + "wasmi", + "wat", ] [[package]] @@ -4285,7 +4296,7 @@ dependencies = [ "httpdate", "itoa", "pin-project-lite", - "socket2 0.4.10", + "socket2 0.5.6", "tokio", "tower-service", "tracing", @@ -4535,6 +4546,12 @@ dependencies = [ "serde", ] +[[package]] +name = "indexmap-nostd" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e04e2fd2b8188ea827b32ef11de88377086d690286ab35747ef7f9bf3ccb590" + [[package]] name = "indicatif" version = "0.17.8" @@ -4808,6 +4825,12 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" +[[package]] +name = "leb128" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c83bff1d572d6b9aeef67ddfc8448e4a3737909cb28e81f97c791b9018703e52" + [[package]] name = "lexical-core" version = "0.8.5" @@ -6202,7 +6225,7 @@ dependencies = [ [[package]] name = "petgraph" version = "0.6.3" -source = "git+https://github.com/getdozer/petgraph?branch=feat/try_map#2422a3a21d92e7a6446f0d0f172c8aa24258846e" +source = "git+https://github.com/getdozer/petgraph?branch=feat%2Ftry_map#2422a3a21d92e7a6446f0d0f172c8aa24258846e" dependencies = [ "fixedbitset", "indexmap 1.9.3", @@ -9870,6 +9893,15 @@ version = "0.2.92" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "af190c94f2773fdb3729c55b007a722abb5384da03bc0986df4c289bf5567e96" +[[package]] +name = "wasm-encoder" +version = "0.41.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "972f97a5d8318f908dded23594188a90bcd09365986b1163e66d70170e5287ae" +dependencies = [ + "leb128", +] + [[package]] name = "wasm-streams" version = "0.3.0" @@ -9883,6 +9915,68 @@ dependencies = [ "web-sys", ] +[[package]] +name = "wasmi" +version = "0.31.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77a8281d1d660cdf54c76a3efa9ddd0c270cada1383a995db3ccb43d166456c7" +dependencies = [ + "smallvec", + "spin 0.9.8", + "wasmi_arena", + "wasmi_core", + "wasmparser-nostd", +] + +[[package]] +name = "wasmi_arena" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "104a7f73be44570cac297b3035d76b169d6599637631cf37a1703326a0727073" + +[[package]] +name = "wasmi_core" +version = "0.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dcf1a7db34bff95b85c261002720c00c3a6168256dcb93041d3fa2054d19856a" +dependencies = [ + "downcast-rs", + "libm", + "num-traits", + "paste", +] + +[[package]] +name = "wasmparser-nostd" +version = "0.100.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d5a015fe95f3504a94bb1462c717aae75253e39b9dd6c3fb1062c934535c64aa" +dependencies = [ + "indexmap-nostd", +] + +[[package]] +name = "wast" +version = "70.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3d5061300042ff5065123dae1e27d00c03f567d34a2937c8472255148a216dc" +dependencies = [ + "bumpalo", + "leb128", + "memchr", + "unicode-width", + "wasm-encoder", +] + +[[package]] +name = "wat" +version = "1.0.85" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "afd7357b6cc46d46a2509c43dcb1dd4131dafbf4e75562d87017b5a05ffad2d6" +dependencies = [ + "wast", +] + [[package]] name = "web-sys" version = "0.3.69" diff --git a/dozer-cli/Cargo.toml b/dozer-cli/Cargo.toml index cb0519af7e..c0f79bfcab 100644 --- a/dozer-cli/Cargo.toml +++ b/dozer-cli/Cargo.toml @@ -64,4 +64,5 @@ mongodb = ["dozer-ingestion/mongodb"] onnx = ["dozer-sql/onnx"] tokio-console = ["dozer-tracing/tokio-console"] javascript = ["dozer-ingestion/javascript", "dozer-sql/javascript"] +wasm = ["dozer-sql/wasm"] datafusion = ["dozer-ingestion/datafusion"] diff --git a/dozer-sql/Cargo.toml b/dozer-sql/Cargo.toml index 28bb354f41..f8d9c391a0 100644 --- a/dozer-sql/Cargo.toml +++ b/dozer-sql/Cargo.toml @@ -23,8 +23,11 @@ tokio = { version = "1", features = ["rt", "macros"] } [dev-dependencies] proptest = "1.3.1" +tempfile = "3.10.1" +wat = "=1.0.85" [features] python = ["dozer-sql-expression/python"] onnx = ["dozer-sql-expression/onnx"] javascript = ["dozer-sql-expression/javascript"] +wasm = ["dozer-sql-expression/wasm"] diff --git a/dozer-sql/expression/Cargo.toml b/dozer-sql/expression/Cargo.toml index 3423d7904a..5dcb02757a 100644 --- a/dozer-sql/expression/Cargo.toml +++ b/dozer-sql/expression/Cargo.toml @@ -19,15 +19,23 @@ jsonpath = { path = "../jsonpath" } bincode = { workspace = true } tokio = "1.34.0" async-recursion = "1.0.5" +wasmi = { version = "0.31.2", optional = true } dozer-deno = { path = "../../dozer-deno", optional = true } deno_core = { workspace = true, optional = true } [dev-dependencies] proptest = "1.2.0" +wat = "=1.0.85" +tempfile = "3.10.1" [features] bigdecimal = ["dep:bigdecimal", "sqlparser/bigdecimal"] python = ["dozer-types/python-auto-initialize"] onnx = ["dep:ort", "dep:ndarray", "dep:half"] javascript = ["dep:dozer-deno", "dep:deno_core"] +wasm = ["dep:wasmi"] + +[[example]] +name = "wasm_udf" +required-features = ["wasm"] diff --git a/dozer-sql/expression/examples/wasm_udf.rs b/dozer-sql/expression/examples/wasm_udf.rs new file mode 100644 index 0000000000..b930f5915a --- /dev/null +++ b/dozer-sql/expression/examples/wasm_udf.rs @@ -0,0 +1,91 @@ +//! Evaluate the AssemblyScript example through Dozer's SQL parser and expression planner. +use std::{error::Error, sync::Arc}; + +use dozer_sql_expression::{ + builder::ExpressionBuilder, + sqlparser::{ + ast::{SelectItem, SetExpr, Statement}, + dialect::DozerDialect, + parser::Parser, + }, +}; +use dozer_types::{ + models::config::Config, + serde_yaml, + types::{Field, FieldDefinition, FieldType, Record, Schema, SourceDefinition}, +}; + +fn main() -> Result<(), Box> { + let config_path = std::env::args() + .nth(1) + .ok_or("usage: wasm_udf ")?; + let config: Config = serde_yaml::from_str(&std::fs::read_to_string(config_path)?)?; + let sql = config.sql.as_deref().ok_or("config must include sql")?; + let statements = Parser::parse_sql(&DozerDialect {}, sql)?; + let Statement::Query(query) = &statements[0] else { + return Err("expected SELECT".into()); + }; + let SetExpr::Select(select) = &*query.body else { + return Err("expected SELECT body".into()); + }; + let schema = Schema::default() + .field( + FieldDefinition::new( + "value".into(), + FieldType::Int, + true, + SourceDefinition::Dynamic, + ), + false, + ) + .field( + FieldDefinition::new( + "amount".into(), + FieldType::Float, + true, + SourceDefinition::Dynamic, + ), + false, + ) + .to_owned(); + let runtime = Arc::new(tokio::runtime::Builder::new_current_thread().build()?); + let mut builder = ExpressionBuilder::new(schema.fields.len(), runtime.clone()); + println!("SQL: {}", sql.trim()); + let mut expressions = Vec::new(); + for item in &select.projection { + let SelectItem::UnnamedExpr(expr) = item else { + return Err("example expects expressions without aliases".into()); + }; + let expression = runtime.block_on(builder.build(false, expr, &schema, &config.udfs))?; + expression.get_type(&schema)?; + expressions.push(expression); + } + for input in [Field::Int(41), Field::Null, Field::Int(-1)] { + let amount = match &input { + Field::Int(value) => Field::Float((*value as f64).into()), + _ => Field::Null, + }; + println!( + "\nInput value: {}, amount: {}", + display(&input), + display(&amount) + ); + let record = Record::new(vec![input, amount]); + for expression in &mut expressions { + let label = expression.to_string(&schema); + match expression.evaluate(&record, &schema) { + Ok(value) => println!(" {label} = {}", display(&value)), + Err(error) => println!(" {label}: {error}"), + } + } + } + Ok(()) +} + +fn display(value: &Field) -> String { + if matches!(value, Field::Null) { + "NULL".into() + } else { + value.to_string() + } +} diff --git a/dozer-sql/expression/src/builder.rs b/dozer-sql/expression/src/builder.rs index cefcda0039..e75c1440ea 100644 --- a/dozer-sql/expression/src/builder.rs +++ b/dozer-sql/expression/src/builder.rs @@ -573,6 +573,35 @@ impl ExpressionBuilder { Err(Error::JavaScriptNotEnabled) } } + UdfType::Wasm(config) => { + #[cfg(feature = "wasm")] + { + let mut args = Vec::with_capacity(sql_function.args.len()); + for argument in &sql_function.args { + args.push( + self.parse_sql_function_arg( + parse_aggregations, + argument, + schema, + udfs, + ) + .await?, + ); + } + let udf = crate::wasm::Udf::new(function_name, config, args)?; + // Selection factories do not request expression types separately. + // Aggregate projections must wait for the post-aggregation schema. + if !parse_aggregations { + udf.get_type(schema)?; + } + Ok(Expression::WasmUdf(Box::new(udf))) + } + #[cfg(not(feature = "wasm"))] + { + let _ = config; + Err(Error::WasmNotEnabled) + } + } }; } diff --git a/dozer-sql/expression/src/error.rs b/dozer-sql/expression/src/error.rs index 7078229713..166940ab19 100644 --- a/dozer-sql/expression/src/error.rs +++ b/dozer-sql/expression/src/error.rs @@ -105,6 +105,13 @@ pub enum Error { #[error("Javascript is not enabled")] JavaScriptNotEnabled, + #[error("WASM UDF support is not enabled; build with --features wasm")] + WasmNotEnabled, + + #[cfg(feature = "wasm")] + #[error("WASM UDF error: {0}")] + Wasm(String), + #[cfg(feature = "javascript")] #[error("JavaScript UDF error: {0}")] JavaScript(#[from] crate::javascript::Error), diff --git a/dozer-sql/expression/src/execution.rs b/dozer-sql/expression/src/execution.rs index 1c2891b234..62e7ca7e44 100644 --- a/dozer-sql/expression/src/execution.rs +++ b/dozer-sql/expression/src/execution.rs @@ -105,6 +105,8 @@ pub enum Expression { }, #[cfg(feature = "javascript")] JavaScriptUdf(crate::javascript::Udf), + #[cfg(feature = "wasm")] + WasmUdf(Box), } impl Expression { @@ -285,6 +287,8 @@ impl Expression { } #[cfg(feature = "javascript")] Expression::JavaScriptUdf(udf) => udf.to_string(schema), + #[cfg(feature = "wasm")] + Expression::WasmUdf(udf) => udf.to_string(schema), Expression::IsNull { arg } => arg.to_string(schema) + " IS NULL ", Expression::IsNotNull { arg } => arg.to_string(schema) + " IS NOT NULL ", } @@ -378,6 +382,8 @@ impl Expression { Expression::IsNotNull { arg } => evaluate_is_not_null(schema, arg, record), #[cfg(feature = "javascript")] Expression::JavaScriptUdf(udf) => udf.evaluate(record, schema), + #[cfg(feature = "wasm")] + Expression::WasmUdf(udf) => udf.evaluate(record, schema), } } @@ -487,6 +493,8 @@ impl Expression { )), #[cfg(feature = "javascript")] Expression::JavaScriptUdf(udf) => Ok(udf.get_type()), + #[cfg(feature = "wasm")] + Expression::WasmUdf(udf) => udf.get_type(schema), Expression::IsNull { arg: _ } => Ok(ExpressionType::new( FieldType::Boolean, false, diff --git a/dozer-sql/expression/src/lib.rs b/dozer-sql/expression/src/lib.rs index c94463f9d9..da24bae06e 100644 --- a/dozer-sql/expression/src/lib.rs +++ b/dozer-sql/expression/src/lib.rs @@ -23,6 +23,8 @@ mod javascript; mod onnx; #[cfg(feature = "python")] mod python_udf; +#[cfg(feature = "wasm")] +mod wasm; pub use num_traits; pub use sqlparser; diff --git a/dozer-sql/expression/src/wasm.rs b/dozer-sql/expression/src/wasm.rs new file mode 100644 index 0000000000..2b7dfd26b6 --- /dev/null +++ b/dozer-sql/expression/src/wasm.rs @@ -0,0 +1,425 @@ +//! Scalar WebAssembly UDFs. Modules are compiled once; each row gets fresh instance state. + +use std::sync::Arc; + +use dozer_types::{ + models::udf_config::WasmConfig, + ordered_float::OrderedFloat, + types::{Field, FieldType, Record, Schema, SourceDefinition}, +}; +use wasmi::{ + core::{Trap, ValueType}, + Config, Engine, ExternType, Instance, Linker, Module, Store, StoreLimits, StoreLimitsBuilder, + Value, +}; + +use crate::{ + error::Error, + execution::{Expression, ExpressionType}, +}; + +const FUEL_PER_CALL: u64 = 10_000_000; +const MEMORY_LIMIT_BYTES: usize = 64 * 1024 * 1024; + +#[derive(Debug)] +struct CompiledModule { + engine: Engine, + module: Module, +} + +#[derive(Clone, Debug)] +pub struct Udf { + name: String, + path: String, + export: String, + args: Vec, + compiled: Arc, + params: Vec, + result: ValueType, +} + +impl PartialEq for Udf { + fn eq(&self, other: &Self) -> bool { + self.name == other.name + && self.path == other.path + && self.export == other.export + && self.args == other.args + } +} + +fn wasm_error(error: impl std::fmt::Display) -> Error { + Error::Wasm(error.to_string()) +} + +fn argument_types(ty: ValueType) -> Vec { + match ty { + ValueType::I32 | ValueType::I64 => vec![ + FieldType::Int8, + FieldType::Int, + FieldType::UInt, + FieldType::I128, + FieldType::U128, + FieldType::Boolean, + ], + ValueType::F32 | ValueType::F64 => vec![FieldType::Float], + _ => vec![], + } +} + +impl Udf { + pub fn new(name: String, config: &WasmConfig, args: Vec) -> Result { + let bytes = std::fs::read(&config.path) + .map_err(|error| wasm_error(format!("cannot read {}: {error}", config.path)))?; + let mut engine_config = Config::default(); + engine_config.consume_fuel(true); + let engine = Engine::new(&engine_config); + let module = Module::new(&engine, &mut &bytes[..]).map_err(wasm_error)?; + let export = config.function.clone().unwrap_or_else(|| name.clone()); + let function_type = match module.get_export(&export) { + Some(ExternType::Func(ty)) => ty, + _ => return Err(wasm_error(format!("{export} is not an exported function"))), + }; + let params = function_type.params().to_vec(); + if params.len() != args.len() { + return Err(Error::InvalidNumberOfArguments { + function_name: name, + expected: params.len()..params.len() + 1, + actual: args.len(), + }); + } + let result = match function_type.results() { + [ty @ (ValueType::I32 | ValueType::I64 | ValueType::F32 | ValueType::F64)] => *ty, + _ => { + return Err(wasm_error( + "a scalar UDF must return exactly one numeric value", + )) + } + }; + for ty in ¶ms { + if argument_types(*ty).is_empty() { + return Err(wasm_error(format!( + "unsupported WASM parameter type {ty:?}" + ))); + } + } + let udf = Self { + name, + path: config.path.clone(), + export, + args, + compiled: Arc::new(CompiledModule { engine, module }), + params, + result, + }; + // Reject unsupported imports, oversized memory and failing start functions at planning time. + udf.instantiate()?; + Ok(udf) + } + + fn instantiate(&self) -> Result<(Store, Instance), Error> { + let limits = StoreLimitsBuilder::new() + .memory_size(MEMORY_LIMIT_BYTES) + .memories(1) + .table_elements(10_000) + .tables(1) + .instances(1) + .build(); + let mut store = Store::new(&self.compiled.engine, limits); + store.limiter(|limits| limits); + store.add_fuel(FUEL_PER_CALL).map_err(wasm_error)?; + let mut linker = Linker::new(&self.compiled.engine); + // AssemblyScript uses this import for assertions/runtime checks. It is not a host I/O API. + linker + .func_wrap( + "env", + "abort", + |_message: i32, _file: i32, line: i32, column: i32| -> Result<(), Trap> { + Err(Trap::new(format!( + "AssemblyScript abort at {line}:{column}" + ))) + }, + ) + .map_err(wasm_error)?; + let instance = linker + .instantiate(&mut store, &self.compiled.module) + .map_err(wasm_error)? + .start(&mut store) + .map_err(wasm_error)?; + Ok((store, instance)) + } + + pub fn get_type(&self, schema: &Schema) -> Result { + // Aggregate columns are appended by the planner after expression construction. + let mut nullable = false; + for (index, (arg, ty)) in self.args.iter().zip(&self.params).enumerate() { + if matches!(arg, Expression::Literal(Field::Null)) { + nullable = true; + continue; + } + let argument_type = arg.get_type(schema)?; + nullable |= argument_type.nullable; + let expected = argument_types(*ty); + if !expected.contains(&argument_type.return_type) { + return Err(Error::InvalidFunctionArgumentType { + function_name: self.name.clone(), + argument_index: index, + expected, + actual: argument_type.return_type, + }); + } + } + Ok(ExpressionType { + return_type: match self.result { + ValueType::I32 | ValueType::I64 => FieldType::Int, + _ => FieldType::Float, + }, + nullable, + source: SourceDefinition::Dynamic, + is_primary_key: false, + }) + } + + pub fn to_string(&self, schema: &Schema) -> String { + format!( + "{}({})", + self.name, + self.args + .iter() + .map(|arg| arg.to_string(schema)) + .collect::>() + .join(",") + ) + } + + pub fn evaluate(&mut self, record: &Record, schema: &Schema) -> Result { + let fields = self + .args + .iter_mut() + .map(|arg| arg.evaluate(record, schema)) + .collect::, _>>()?; + if fields.iter().any(|field| matches!(field, Field::Null)) { + return Ok(Field::Null); + } + let values = fields + .into_iter() + .zip(&self.params) + .enumerate() + .map(|(index, (field, ty))| to_wasm_value(field, *ty, &self.name, index)) + .collect::, _>>()?; + let (mut store, instance) = self.instantiate()?; + let function = instance + .get_func(&store, &self.export) + .ok_or_else(|| wasm_error("configured function export disappeared"))?; + let mut output = [Value::default(self.result)]; + function + .call(&mut store, &values, &mut output) + .map_err(wasm_error)?; + match output[0] { + Value::I32(value) => Ok(Field::Int(value.into())), + Value::I64(value) => Ok(Field::Int(value)), + Value::F32(value) => Ok(Field::Float(OrderedFloat(f32::from(value).into()))), + Value::F64(value) => Ok(Field::Float(OrderedFloat(f64::from(value)))), + _ => Err(wasm_error("unsupported WASM result")), + } + } +} + +fn to_wasm_value( + field: Field, + ty: ValueType, + function: &str, + index: usize, +) -> Result { + let integer = match &field { + Field::Int8(value) => Some(i128::from(*value)), + Field::Int(value) => Some(i128::from(*value)), + Field::UInt(value) => Some(i128::from(*value)), + Field::I128(value) => Some(*value), + Field::U128(value) => i128::try_from(*value).ok(), + Field::Boolean(value) => Some(i128::from(*value)), + _ => None, + }; + let value = match (ty, &field) { + (ValueType::I32, _) => integer.and_then(|v| i32::try_from(v).ok()).map(Value::I32), + (ValueType::I64, _) => integer.and_then(|v| i64::try_from(v).ok()).map(Value::I64), + (ValueType::F32, Field::Float(value)) => { + let converted = value.0 as f32; + if value.0.is_finite() && !converted.is_finite() { + None + } else { + Some(Value::F32(converted.into())) + } + } + (ValueType::F64, Field::Float(value)) => Some(Value::F64(value.0.into())), + _ => None, + }; + value.ok_or_else(|| Error::InvalidFunctionArgument { + function_name: function.to_owned(), + argument_index: index, + argument: field, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn udf(wat: &str, args: Vec) -> Result { + let file = tempfile::NamedTempFile::new().unwrap(); + std::fs::write(file.path(), wat::parse_str(wat).unwrap()).unwrap(); + let function = Udf::new( + "test".into(), + &WasmConfig { + path: file.path().to_str().unwrap().into(), + function: Some("run".into()), + }, + args, + )?; + function.get_type(&Schema::default())?; + Ok(function) + } + + fn evaluate(udf: &mut Udf) -> Result { + udf.evaluate(&Record::new(vec![]), &Schema::default()) + } + + #[test] + fn numeric_exports_and_checked_integer_conversions() { + for (ty, field, expected) in [ + ("i32", Field::Int8(i8::MIN), Field::Int(i8::MIN.into())), + ("i64", Field::Int8(i8::MAX), Field::Int(i8::MAX.into())), + ("i32", Field::Int(42), Field::Int(42)), + ("i64", Field::UInt(42), Field::Int(42)), + ("i32", Field::Boolean(true), Field::Int(1)), + ( + "f32", + Field::Float(OrderedFloat(1.25)), + Field::Float(OrderedFloat(1.25)), + ), + ( + "f64", + Field::Float(OrderedFloat(2.5)), + Field::Float(OrderedFloat(2.5)), + ), + ] { + let code = + format!("(module (func (export \"run\") (param {ty}) (result {ty}) local.get 0))"); + let mut function = udf(&code, vec![Expression::Literal(field)]).unwrap(); + assert_eq!(evaluate(&mut function).unwrap(), expected); + } + for field in [ + Field::Int(i64::MAX), + Field::UInt(u64::MAX), + Field::U128(u128::MAX), + ] { + let mut function = udf( + "(module (func (export \"run\") (param i32) (result i32) local.get 0))", + vec![Expression::Literal(field)], + ) + .unwrap(); + assert!(matches!( + evaluate(&mut function), + Err(Error::InvalidFunctionArgument { .. }) + )); + } + } + + #[test] + fn null_skips_execution_and_preserves_nullability() { + let mut function = udf( + "(module (func (export \"run\") (param i64) (result i64) unreachable))", + vec![Expression::Literal(Field::Null)], + ) + .unwrap(); + assert!(function.get_type(&Schema::default()).unwrap().nullable); + assert_eq!(evaluate(&mut function).unwrap(), Field::Null); + } + + #[test] + fn validates_arity_types_exports_and_results() { + assert!(udf( + "(module (func (export \"run\") (param i64) (result i64) local.get 0))", + vec![] + ) + .is_err()); + assert!(udf( + "(module (func (export \"run\") (param i64) (result i64) local.get 0))", + vec![Expression::Literal(Field::String("no".into()))] + ) + .is_err()); + assert!(udf("(module (memory (export \"run\") 1))", vec![]).is_err()); + assert!(udf("(module (func (export \"run\")))", vec![]).is_err()); + assert!(udf( + "(module (func (export \"run\") (result i32 i32) i32.const 1 i32.const 2))", + vec![] + ) + .is_err()); + assert!(udf( + "(module (func (export \"run\") (param externref) (result i32) i32.const 1))", + vec![Expression::Literal(Field::Null)] + ) + .is_err()); + } + + #[test] + fn bounds_execution_and_memory_and_reports_traps() { + let mut loop_forever = udf( + "(module (func (export \"run\") (result i32) (loop br 0) i32.const 0))", + vec![], + ) + .unwrap(); + assert!(evaluate(&mut loop_forever).is_err()); + let mut trap = udf( + "(module (func (export \"run\") (result i32) unreachable))", + vec![], + ) + .unwrap(); + assert!(evaluate(&mut trap).is_err()); + assert!(udf( + "(module (memory 1025) (func (export \"run\") (result i32) i32.const 0))", + vec![] + ) + .is_err()); + assert!(udf( + "(module (table 10001 funcref) (func (export \"run\") (result i32) i32.const 0))", + vec![] + ) + .is_err()); + } + + #[test] + fn unsupported_imports_and_failing_start_functions_are_rejected() { + assert!(udf("(module (import \"env\" \"read_file\" (func)) (func (export \"run\") (result i32) i32.const 0))", vec![]).is_err()); + assert!(udf("(module (func $start unreachable) (start $start) (func (export \"run\") (result i32) i32.const 0))", vec![]).is_err()); + assert!(udf("(module (func $start (loop br 0)) (start $start) (func (export \"run\") (result i32) i32.const 0))", vec![]).is_err()); + } + + #[test] + fn overflowing_float_argument_is_rejected() { + let mut function = udf( + "(module (func (export \"run\") (param f32) (result f32) local.get 0))", + vec![Expression::Literal(Field::Float(OrderedFloat(f64::MAX)))], + ) + .unwrap(); + assert!(matches!( + evaluate(&mut function), + Err(Error::InvalidFunctionArgument { .. }) + )); + } + + #[test] + fn each_record_has_fresh_module_state() { + let mut function = udf("(module (global $counter (mut i32) (i32.const 0)) (func (export \"run\") (result i32) global.get $counter i32.const 1 i32.add global.set $counter global.get $counter))", vec![]).unwrap(); + assert_eq!(evaluate(&mut function).unwrap(), Field::Int(1)); + assert_eq!(evaluate(&mut function).unwrap(), Field::Int(1)); + } + + #[test] + fn assemblyscript_abort_is_reported_without_host_io() { + let mut function = udf("(module (import \"env\" \"abort\" (func $abort (param i32 i32 i32 i32))) (func (export \"run\") (result i32) i32.const 0 i32.const 0 i32.const 9 i32.const 3 call $abort i32.const 0))", vec![]).unwrap(); + assert!(evaluate(&mut function) + .unwrap_err() + .to_string() + .contains("AssemblyScript abort at 9:3")); + } +} diff --git a/dozer-sql/expression/tests/wasm_udf.rs b/dozer-sql/expression/tests/wasm_udf.rs new file mode 100644 index 0000000000..f865599243 --- /dev/null +++ b/dozer-sql/expression/tests/wasm_udf.rs @@ -0,0 +1,215 @@ +use std::sync::Arc; + +use dozer_sql_expression::{ + builder::ExpressionBuilder, + error::Error, + execution::Expression, + sqlparser::{ + ast::{SelectItem, SetExpr, Statement}, + dialect::DozerDialect, + parser::Parser, + }, +}; +use dozer_types::{models::config::Config, serde_yaml, types::Schema}; + +fn build(sql: &str, schema: &Schema, yaml: &str) -> Result { + let (expression, _) = plan(sql, schema, yaml, false)?; + expression.get_type(schema)?; + Ok(expression) +} + +fn plan( + sql: &str, + schema: &Schema, + yaml: &str, + aggregates: bool, +) -> Result<(Expression, Vec), Error> { + let config: Config = serde_yaml::from_str(yaml).unwrap(); + let statements = Parser::parse_sql(&DozerDialect {}, sql).unwrap(); + let Statement::Query(query) = &statements[0] else { + panic!("expected a SELECT query"); + }; + let SetExpr::Select(select) = &*query.body else { + panic!("expected a SELECT body"); + }; + let SelectItem::UnnamedExpr(expr) = &select.projection[0] else { + panic!("expected an expression"); + }; + let runtime = Arc::new( + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap(), + ); + let mut builder = ExpressionBuilder::new(schema.fields.len(), runtime.clone()); + let expression = runtime.block_on(builder.build(aggregates, expr, schema, &config.udfs))?; + Ok((expression, builder.aggregations)) +} + +#[cfg(not(feature = "wasm"))] +#[test] +fn disabled_feature_is_reported_before_opening_the_module() { + let error = build( + "SELECT increment(1)", + &Schema::default(), + "version: 1\napp_name: test\nudfs:\n - name: increment\n config: !Wasm\n path: missing.wasm\n", + ) + .unwrap_err(); + assert!(matches!(error, Error::WasmNotEnabled)); +} + +#[cfg(feature = "wasm")] +mod enabled { + use super::*; + use dozer_types::types::{Field, FieldDefinition, FieldType, Record, SourceDefinition}; + + fn fixture() -> (tempfile::NamedTempFile, String) { + let file = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + file.path(), + wat::parse_str( + r#"(module + (func (export "addOne") (param i64) (result i64) + local.get 0 i64.const 1 i64.add) + (func (export "double") (param i64) (result i64) + local.get 0 i64.const 2 i64.mul))"#, + ) + .unwrap(), + ) + .unwrap(); + let yaml = format!( + "version: 1\napp_name: test\nudfs:\n - name: increment\n config: !Wasm\n path: {}\n function: addOne\n - name: double\n config: !Wasm\n path: {}\n", + file.path().display(), file.path().display(), + ); + (file, yaml) + } + + #[test] + fn yaml_configured_exports_work_in_nested_sql_and_on_nullable_rows() { + let (_module, yaml) = fixture(); + let schema = Schema::default() + .field( + FieldDefinition::new( + "value".into(), + FieldType::Int, + true, + SourceDefinition::Dynamic, + ), + false, + ) + .to_owned(); + let mut expression = + build("SELECT double(increment(value)) FROM input", &schema, &yaml).unwrap(); + let expression_type = expression.get_type(&schema).unwrap(); + assert_eq!(expression_type.return_type, FieldType::Int); + assert!(expression_type.nullable); + for (input, expected) in [ + (Field::Int(10), Field::Int(22)), + (Field::Int(-2), Field::Int(-2)), + (Field::Null, Field::Null), + ] { + assert_eq!( + expression + .evaluate(&Record::new(vec![input]), &schema) + .unwrap(), + expected + ); + } + let mut arithmetic = build("SELECT increment(40 + 1)", &schema, &yaml).unwrap(); + assert_eq!( + arithmetic.evaluate(&Record::new(vec![]), &schema).unwrap(), + Field::Int(42) + ); + } + + #[test] + fn sql_planning_validates_types_arity_and_exports() { + let (_module, yaml) = fixture(); + for sql in [ + "SELECT increment('wrong')", + "SELECT increment()", + "SELECT increment(1, 2)", + ] { + assert!(build(sql, &Schema::default(), &yaml).is_err(), "{sql}"); + } + let invalid = yaml.replace("function: addOne", "function: missing"); + assert!(build("SELECT increment(1)", &Schema::default(), &invalid).is_err()); + } + + #[test] + fn aggregate_arguments_are_typed_against_the_post_aggregation_schema() { + let (_module, yaml) = fixture(); + let schema = Schema::default() + .field( + FieldDefinition::new( + "value".into(), + FieldType::Int, + true, + SourceDefinition::Dynamic, + ), + false, + ) + .to_owned(); + let (mut expression, aggregates) = plan( + "SELECT double(increment(SUM(value))) FROM input", + &schema, + &yaml, + true, + ) + .unwrap(); + assert_eq!(aggregates.len(), 1); + let mut output_schema = schema.clone(); + for aggregate in &aggregates { + let ty = aggregate.get_type(&schema).unwrap(); + output_schema.field( + FieldDefinition::new( + aggregate.to_string(&schema), + ty.return_type, + ty.nullable, + ty.source, + ), + false, + ); + } + assert_eq!( + expression.get_type(&output_schema).unwrap().return_type, + FieldType::Int + ); + assert_eq!( + expression + .evaluate( + &Record::new(vec![Field::Int(5), Field::Int(12)]), + &output_schema + ) + .unwrap(), + Field::Int(26) + ); + } + + #[test] + fn planned_expression_keeps_the_compiled_module_and_clone_semantics() { + let (module, yaml) = fixture(); + let schema = Schema::default(); + let expression = build("SELECT increment(41)", &schema, &yaml).unwrap(); + assert_eq!( + expression, + build("SELECT increment(41)", &schema, &yaml).unwrap() + ); + let mut clone = expression.clone(); + assert_eq!(expression, clone); + drop(module); + assert_eq!( + clone.evaluate(&Record::new(vec![]), &schema).unwrap(), + Field::Int(42) + ); + assert_eq!(clone.to_string(&schema), "increment(41)"); + } + + #[test] + fn missing_and_invalid_module_files_fail_at_planning_time() { + let (module, yaml) = fixture(); + std::fs::write(module.path(), b"not a wasm module").unwrap(); + assert!(build("SELECT increment(1)", &Schema::default(), &yaml).is_err()); + drop(module); + assert!(build("SELECT increment(1)", &Schema::default(), &yaml).is_err()); + } +} diff --git a/dozer-sql/src/selection/mod.rs b/dozer-sql/src/selection/mod.rs index 9a12ba12cc..225a89667b 100644 --- a/dozer-sql/src/selection/mod.rs +++ b/dozer-sql/src/selection/mod.rs @@ -1,2 +1,5 @@ pub mod factory; pub mod processor; + +#[cfg(all(test, feature = "wasm"))] +mod wasm_tests; diff --git a/dozer-sql/src/selection/wasm_tests.rs b/dozer-sql/src/selection/wasm_tests.rs new file mode 100644 index 0000000000..311407d1d8 --- /dev/null +++ b/dozer-sql/src/selection/wasm_tests.rs @@ -0,0 +1,142 @@ +use std::collections::HashMap; + +use dozer_core::{ + channels::ProcessorChannelForwarder, event::EventHub, node::ProcessorFactory, + DEFAULT_PORT_HANDLE, +}; +use dozer_types::{ + models::udf_config::{UdfConfig, UdfType, WasmConfig}, + types::{ + Field, FieldDefinition, FieldType, Operation, Record, Schema, SourceDefinition, + TableOperation, + }, +}; + +use crate::tests::utils::{create_test_runtime, get_select}; + +use super::factory::SelectionProcessorFactory; + +fn fixture() -> (tempfile::NamedTempFile, UdfConfig) { + let module = tempfile::NamedTempFile::new().unwrap(); + std::fs::write( + module.path(), + wat::parse_str( + r#"(module + (func (export "increment") (param i32) (result i32) + local.get 0 i32.const 1 i32.add))"#, + ) + .unwrap(), + ) + .unwrap(); + let config = UdfConfig { + name: "increment".into(), + config: UdfType::Wasm(WasmConfig { + path: module.path().to_str().unwrap().into(), + function: None, + }), + }; + (module, config) +} + +fn schema(typ: FieldType) -> Schema { + Schema::default() + .field( + FieldDefinition::new("value".into(), typ, true, SourceDefinition::Dynamic), + false, + ) + .to_owned() +} + +#[test] +fn where_factory_rejects_invalid_wasm_arguments_before_processing_rows() { + let (_module, config) = fixture(); + let runtime = create_test_runtime(); + let select = get_select("SELECT value FROM input WHERE increment(value) > 0").unwrap(); + let factory = SelectionProcessorFactory::new( + "selection".into(), + select.selection.unwrap(), + vec![config], + runtime.clone(), + ); + let result = runtime.block_on(factory.build( + HashMap::from([(DEFAULT_PORT_HANDLE, schema(FieldType::String))]), + HashMap::new(), + EventHub::new(1), + )); + let Err(error) = result else { + panic!("a string argument must fail during factory build") + }; + assert!(error.to_string().contains("increment")); +} + +#[derive(Default)] +struct Forwarder(Vec); + +impl ProcessorChannelForwarder for Forwarder { + fn send(&mut self, operation: TableOperation) { + self.0.push(operation); + } +} + +#[test] +fn where_wasm_filters_int8_and_null_rows_and_handles_updates() { + let (_module, config) = fixture(); + let runtime = create_test_runtime(); + let select = get_select("SELECT value FROM input WHERE increment(value) > 0").unwrap(); + let factory = SelectionProcessorFactory::new( + "selection".into(), + select.selection.unwrap(), + vec![config], + runtime.clone(), + ); + let mut processor = runtime + .block_on(factory.build( + HashMap::from([(DEFAULT_PORT_HANDLE, schema(FieldType::Int8))]), + HashMap::new(), + EventHub::new(1), + )) + .unwrap(); + let record = |field| Record::new(vec![field]); + let mut forwarder = Forwarder::default(); + for op in [ + Operation::Insert { + new: record(Field::Int8(-1)), + }, + Operation::Insert { + new: record(Field::Null), + }, + Operation::Insert { + new: record(Field::Int8(0)), + }, + Operation::Update { + old: record(Field::Int8(-1)), + new: record(Field::Int8(2)), + }, + Operation::Update { + old: record(Field::Int8(2)), + new: record(Field::Int8(-1)), + }, + ] { + processor + .process( + TableOperation::without_id(op, DEFAULT_PORT_HANDLE), + &mut forwarder, + ) + .unwrap(); + } + let operations: Vec<_> = forwarder.0.into_iter().map(|op| op.op).collect(); + assert_eq!( + operations, + vec![ + Operation::Insert { + new: record(Field::Int8(0)) + }, + Operation::Insert { + new: record(Field::Int8(2)) + }, + Operation::Delete { + old: record(Field::Int8(2)) + }, + ] + ); +} diff --git a/dozer-tests/wasm_udf/.gitignore b/dozer-tests/wasm_udf/.gitignore new file mode 100644 index 0000000000..d49756dc7d --- /dev/null +++ b/dozer-tests/wasm_udf/.gitignore @@ -0,0 +1,2 @@ +/node_modules/ +/build/ diff --git a/dozer-tests/wasm_udf/README.md b/dozer-tests/wasm_udf/README.md new file mode 100644 index 0000000000..04fb0db3f0 --- /dev/null +++ b/dozer-tests/wasm_udf/README.md @@ -0,0 +1,45 @@ +# WebAssembly scalar UDFs + +Build Dozer with `--features wasm` to enable scalar WebAssembly functions. Configure each function in the existing `udfs` list: + +```yaml +udfs: + - name: increment + config: !Wasm + path: ./math.wasm + function: addOne +``` + +`name` is the SQL function name; use lowercase names, as Dozer normalizes SQL function names to lowercase. `function` is the case-sensitive module export and defaults to `name`. `path` is resolved from the process working directory. Call it like any other scalar function: `SELECT increment(value) FROM input`. + +## AssemblyScript example + +From the repository root, with Node.js 16+ and the normal Rust/protobuf build prerequisites installed: + +```sh +npm ci --prefix dozer-tests/wasm_udf +npm run build --prefix dozer-tests/wasm_udf +cargo run -p dozer-sql-expression --features wasm --example wasm_udf -- dozer-tests/wasm_udf/dozer-config.yaml +``` + +This compiles `assembly/index.ts` into a real `.wasm` module, reads the example YAML as a Dozer configuration, parses its SQL with the Dozer dialect, and evaluates the expressions on three in-memory rows with nullable integer `value` and floating-point `amount` columns. No connector or database is required for this example. The row `(41, 41.0)` produces `42`, `41.5`, and `41`; a null row produces three nulls; `(-1, -1.0)` demonstrates an AssemblyScript assertion reported as a SQL evaluation error. + +To use the same UDFs in an ingestion pipeline, add the `udfs` entries and SQL to the pipeline's normal configuration, then run a CLI built with `cargo build -p dozer-cli --features wasm`. + +## Supported ABI and execution + +Each exported function must return exactly one `i32`, `i64`, `f32`, or `f64` value. Integer parameters accept Dozer integer fields when the value fits the signed WebAssembly parameter width, and booleans as `0`/`1`. Floating-point parameters accept Dozer `Float`; use an explicit SQL `CAST(... AS FLOAT)` for integer inputs. Integer results become Dozer `Int`, floating-point results become `Float`. An `f32` parameter uses normal IEEE rounding; finite values that overflow are rejected. + +Any null argument returns SQL null without invoking the exported function. Wrong argument counts, unsupported types, missing exports, invalid modules, unsupported imports, and failing module start functions are rejected during SQL planning. Module paths are read and compiled during planning; evaluation reuses that compiled module even if the file changes later. Each row receives a fresh instance so mutable globals and memory do not leak between records. + +Execution allows up to 10 million fuel units per instantiation/call, one memory of up to 64 MiB, and one table of up to 10,000 elements. Traps and fuel exhaustion become SQL errors. The only host import is AssemblyScript's `env.abort`, which reports the source line and column. WASI, host I/O, strings/objects passed through linear memory, multiple results, and reference types are outside this scalar ABI. + +Without the `wasm` feature, configuration still deserializes and calling a configured WASM UDF returns a clear feature-disabled error. + +## Tests + +```sh +cargo test -p dozer-sql-expression --features wasm +cargo test -p dozer-sql-expression --no-default-features --test wasm_udf +cargo test -p dozer-types udf_yaml_deserialize +``` diff --git a/dozer-tests/wasm_udf/assembly/index.ts b/dozer-tests/wasm_udf/assembly/index.ts new file mode 100644 index 0000000000..29338687da --- /dev/null +++ b/dozer-tests/wasm_udf/assembly/index.ts @@ -0,0 +1,12 @@ +export function addOne(value: i64): i64 { + return value + 1; +} + +export function add(left: f64, right: f64): f64 { + return left + right; +} + +export function nonnegative(value: i64): i64 { + assert(value >= 0, "value must be non-negative"); + return value; +} diff --git a/dozer-tests/wasm_udf/dozer-config.yaml b/dozer-tests/wasm_udf/dozer-config.yaml new file mode 100644 index 0000000000..f8474225e8 --- /dev/null +++ b/dozer-tests/wasm_udf/dozer-config.yaml @@ -0,0 +1,18 @@ +version: 1 +app_name: wasm_udf_example + +udfs: + - name: increment + config: !Wasm + path: dozer-tests/wasm_udf/build/math.wasm + function: addOne + - name: add + config: !Wasm + path: dozer-tests/wasm_udf/build/math.wasm + - name: nonnegative + config: !Wasm + path: dozer-tests/wasm_udf/build/math.wasm + +sql: | + SELECT increment(value), add(amount, 0.5), nonnegative(value) + FROM input; diff --git a/dozer-tests/wasm_udf/package-lock.json b/dozer-tests/wasm_udf/package-lock.json new file mode 100644 index 0000000000..8f206c7ba5 --- /dev/null +++ b/dozer-tests/wasm_udf/package-lock.json @@ -0,0 +1,56 @@ +{ + "name": "dozer-wasm-udf-example", + "version": "1.0.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "dozer-wasm-udf-example", + "version": "1.0.0", + "devDependencies": { + "assemblyscript": "0.27.31" + } + }, + "node_modules/assemblyscript": { + "version": "0.27.31", + "resolved": "https://registry.npmjs.org/assemblyscript/-/assemblyscript-0.27.31.tgz", + "integrity": "sha512-Ra8kiGhgJQGZcBxjtMcyVRxOEJZX64kd+XGpjWzjcjgxWJVv+CAQO0aDBk4GQVhjYbOkATarC83mHjAVGtwPBQ==", + "dev": true, + "license": "Apache-2.0", + "dependencies": { + "binaryen": "116.0.0-nightly.20240114", + "long": "^5.2.1" + }, + "bin": { + "asc": "bin/asc.js", + "asinit": "bin/asinit.js" + }, + "engines": { + "node": ">=16", + "npm": ">=7" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/assemblyscript" + } + }, + "node_modules/binaryen": { + "version": "116.0.0-nightly.20240114", + "resolved": "https://registry.npmjs.org/binaryen/-/binaryen-116.0.0-nightly.20240114.tgz", + "integrity": "sha512-0GZrojJnuhoe+hiwji7QFaL3tBlJoA+KFUN7ouYSDGZLSo9CKM8swQX8n/UcbR0d1VuZKU+nhogNzv423JEu5A==", + "dev": true, + "license": "Apache-2.0", + "bin": { + "wasm-opt": "bin/wasm-opt", + "wasm2js": "bin/wasm2js" + } + }, + "node_modules/long": { + "version": "5.3.2", + "resolved": "https://registry.npmjs.org/long/-/long-5.3.2.tgz", + "integrity": "sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==", + "dev": true, + "license": "Apache-2.0" + } + } +} diff --git a/dozer-tests/wasm_udf/package.json b/dozer-tests/wasm_udf/package.json new file mode 100644 index 0000000000..2b4b989af9 --- /dev/null +++ b/dozer-tests/wasm_udf/package.json @@ -0,0 +1,11 @@ +{ + "name": "dozer-wasm-udf-example", + "version": "1.0.0", + "private": true, + "scripts": { + "build": "asc assembly/index.ts --outFile build/math.wasm --textFile build/math.wat --runtime stub --optimize" + }, + "devDependencies": { + "assemblyscript": "0.27.31" + } +} diff --git a/dozer-types/src/models/udf_config.rs b/dozer-types/src/models/udf_config.rs index 512fba3a57..9741808c20 100644 --- a/dozer-types/src/models/udf_config.rs +++ b/dozer-types/src/models/udf_config.rs @@ -16,6 +16,7 @@ pub struct UdfConfig { pub enum UdfType { Onnx(OnnxConfig), JavaScript(JavaScriptConfig), + Wasm(WasmConfig), } #[derive(Debug, Serialize, Deserialize, JsonSchema, Eq, PartialEq, Clone)] @@ -31,3 +32,13 @@ pub struct JavaScriptConfig { /// path to the module file pub module: String, } + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)] +#[serde(deny_unknown_fields)] +pub struct WasmConfig { + /// Path to a WebAssembly module, relative to the process working directory. + pub path: String, + /// Exported function name. Defaults to the configured SQL UDF name. + #[serde(default)] + pub function: Option, +} diff --git a/dozer-types/src/tests/udf_yaml_deserialize.rs b/dozer-types/src/tests/udf_yaml_deserialize.rs index 6cec087a5b..af329e5766 100644 --- a/dozer-types/src/tests/udf_yaml_deserialize.rs +++ b/dozer-types/src/tests/udf_yaml_deserialize.rs @@ -1,4 +1,4 @@ -use crate::models::udf_config::{OnnxConfig, UdfConfig, UdfType}; +use crate::models::udf_config::{OnnxConfig, UdfConfig, UdfType, WasmConfig}; #[test] fn standard() { @@ -17,3 +17,46 @@ fn standard() { let expected = udf_conf; assert_eq!(expected, deserializer_result); } + +#[test] +fn wasm_defaults_to_the_sql_function_name() { + let config: UdfConfig = + serde_yaml::from_str("name: increment\nconfig: !Wasm\n path: ./math.wasm\n").unwrap(); + assert_eq!( + config.config, + UdfType::Wasm(WasmConfig { + path: "./math.wasm".into(), + function: None, + }) + ); +} + +#[test] +fn wasm_can_use_a_different_export_name() { + let config: UdfConfig = serde_yaml::from_str( + "name: increment\nconfig: !Wasm\n path: ./math.wasm\n function: addOne\n", + ) + .unwrap(); + assert_eq!( + config.config, + UdfType::Wasm(WasmConfig { + path: "./math.wasm".into(), + function: Some("addOne".into()), + }) + ); + let serialized = serde_yaml::to_string(&config).unwrap(); + assert_eq!( + serde_yaml::from_str::(&serialized).unwrap(), + config + ); +} + +#[test] +fn wasm_requires_a_path_and_rejects_unknown_options() { + for yaml in [ + "name: f\nconfig: !Wasm {}\n", + "name: f\nconfig: !Wasm\n path: ./math.wasm\n typo: ignored\n", + ] { + assert!(serde_yaml::from_str::(yaml).is_err()); + } +} diff --git a/json_schemas/dozer.json b/json_schemas/dozer.json index 3bac6f286e..0c3c1570d7 100644 --- a/json_schemas/dozer.json +++ b/json_schemas/dozer.json @@ -2084,9 +2084,42 @@ } }, "additionalProperties": false + }, + { + "type": "object", + "required": [ + "Wasm" + ], + "properties": { + "Wasm": { + "$ref": "#/definitions/WasmConfig" + } + }, + "additionalProperties": false } ] }, + "WasmConfig": { + "type": "object", + "required": [ + "path" + ], + "properties": { + "function": { + "description": "Exported function name. Defaults to the configured SQL UDF name.", + "default": null, + "type": [ + "string", + "null" + ] + }, + "path": { + "description": "Path to a WebAssembly module, relative to the process working directory.", + "type": "string" + } + }, + "additionalProperties": false + }, "WebhookConfig": { "examples": [ {