From daecf901e6576df0ddd9d24dbc2aed6774b4599f Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Wed, 3 Apr 2024 16:09:51 +0200 Subject: [PATCH] feat: custom node (#138) * feat: custom node * add zen template, expose $nodes and $root * add custom handler to go and nodejs * add support for python, improve bindings, add error to zen template * fix: correct binding exports for nodejs and python * fix benchmark * improve rust api, trim template in zen templates * update cargo action * update expression version * compile action * fix action format --- Cargo.toml | 1 + actions/cargo-version-action/dist/index.js | 4 +- actions/cargo-version-action/src/cargo.ts | 4 +- .../cargo-version-action/src/index.spec.ts | 43 +++--- bindings/c/Cargo.toml | 2 + bindings/c/src/custom_node.rs | 65 +++++++++ bindings/c/src/decision.rs | 9 +- bindings/c/src/engine.rs | 26 ++-- bindings/c/src/error.rs | 5 + bindings/c/src/expression.rs | 62 ++++++--- bindings/c/src/helper.rs | 23 ++++ bindings/c/src/languages/go.rs | 59 ++++++++- bindings/c/src/languages/native.rs | 43 +++++- bindings/c/src/lib.rs | 2 + bindings/c/src/loader.rs | 13 +- bindings/c/zen_engine.h | 23 +++- bindings/nodejs/Cargo.toml | 9 +- bindings/nodejs/index.d.ts | 39 +++++- bindings/nodejs/index.js | 7 +- bindings/nodejs/src/custom_node.rs | 60 +++++++++ bindings/nodejs/src/decision.rs | 12 +- bindings/nodejs/src/engine.rs | 31 +++-- bindings/nodejs/src/expression.rs | 56 +++++--- bindings/nodejs/src/lib.rs | 2 + bindings/nodejs/src/loader.rs | 4 +- bindings/nodejs/src/types.rs | 125 ++++++++++++++++++ bindings/nodejs/test/decision.spec.ts | 50 ++++++- bindings/python/Cargo.toml | 8 +- bindings/python/index.py | 20 +++ bindings/python/src/custom_node.rs | 42 ++++++ bindings/python/src/decision.rs | 18 ++- bindings/python/src/engine.rs | 26 +++- bindings/python/src/expression.rs | 10 ++ bindings/python/src/lib.rs | 5 +- bindings/python/src/loader.rs | 6 + bindings/python/src/types.rs | 118 +++++++++++++++++ bindings/python/src/value.rs | 8 +- core/engine/Cargo.toml | 2 + core/engine/benches/engine.rs | 6 +- core/engine/src/decision.rs | 42 ++++-- core/engine/src/engine.rs | 74 +++++++---- .../engine/src/handler/custom_node_adapter.rs | 86 ++++++++++++ core/engine/src/handler/decision.rs | 18 ++- core/engine/src/handler/graph.rs | 58 +++++--- core/engine/src/handler/mod.rs | 5 +- core/engine/src/handler/node.rs | 2 +- core/engine/src/handler/traversal.rs | 30 ++++- core/engine/src/lib.rs | 2 +- core/engine/src/model/mod.rs | 85 +++++++++--- core/engine/tests/engine.rs | 14 +- core/expression/src/compiler/compiler.rs | 1 + core/expression/src/compiler/opcode.rs | 1 + core/expression/src/lexer/token.rs | 2 + core/expression/src/parser/ast.rs | 1 + core/expression/src/parser/parser.rs | 6 +- core/expression/src/parser/unary.rs | 1 + core/expression/src/vm/vm.rs | 3 + core/template/Cargo.toml | 15 +++ core/template/src/error.rs | 64 +++++++++ core/template/src/interpreter.rs | 97 ++++++++++++++ core/template/src/lexer.rs | 62 +++++++++ core/template/src/lib.rs | 19 +++ core/template/src/parser.rs | 77 +++++++++++ core/template/tests/template.rs | 100 ++++++++++++++ test-data/custom.json | 51 +++++++ 65 files changed, 1729 insertions(+), 235 deletions(-) create mode 100644 bindings/c/src/custom_node.rs create mode 100644 bindings/c/src/helper.rs create mode 100644 bindings/nodejs/src/custom_node.rs create mode 100644 bindings/nodejs/src/types.rs create mode 100644 bindings/python/src/custom_node.rs create mode 100644 bindings/python/src/types.rs create mode 100644 core/engine/src/handler/custom_node_adapter.rs create mode 100644 core/template/Cargo.toml create mode 100644 core/template/src/error.rs create mode 100644 core/template/src/interpreter.rs create mode 100644 core/template/src/lexer.rs create mode 100644 core/template/src/lib.rs create mode 100644 core/template/src/parser.rs create mode 100644 core/template/tests/template.rs create mode 100644 test-data/custom.json diff --git a/Cargo.toml b/Cargo.toml index 5499fdd4..6a796b8d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -25,6 +25,7 @@ serde_json = "1.0.111" hashbrown = "0.13.2" rust_decimal = "1.33.1" rust_decimal_macros = "1.33.1" +json_dotpath = "1.1.0" anyhow = "1.0.79" thiserror = "1.0.50" diff --git a/actions/cargo-version-action/dist/index.js b/actions/cargo-version-action/dist/index.js index cb5c99bb..f473737a 100644 --- a/actions/cargo-version-action/dist/index.js +++ b/actions/cargo-version-action/dist/index.js @@ -11270,10 +11270,12 @@ var toml = __nccwpck_require__(4920); const versionRegex = /version = "[0-9]+\.[0-9]+\.[0-9]+"$/im; const expressionDep = /zen-expression =.*$/im; +const templateDep = /zen-template =.*$/im; const updateCargoContents = (contents, { version }) => { return contents .replace(versionRegex, `version = "${version}"`) - .replace(expressionDep, `zen-expression = { path = "../expression", version = "${version}" }`); + .replace(expressionDep, `zen-expression = { path = "../expression", version = "${version}" }`) + .replace(templateDep, `zen-template = { path = "../template", version = "${version}" }`); }; const getCargoVersion = (contents) => { var _a; diff --git a/actions/cargo-version-action/src/cargo.ts b/actions/cargo-version-action/src/cargo.ts index a397bed0..5ab098c1 100644 --- a/actions/cargo-version-action/src/cargo.ts +++ b/actions/cargo-version-action/src/cargo.ts @@ -6,11 +6,13 @@ type UpdateCargoOptions = { const versionRegex = /version = "[0-9]+\.[0-9]+\.[0-9]+"$/im; const expressionDep = /zen-expression =.*$/im; +const templateDep = /zen-template =.*$/im; export const updateCargoContents = (contents: string, { version }: UpdateCargoOptions): string => { return contents .replace(versionRegex, `version = "${version}"`) - .replace(expressionDep, `zen-expression = { path = "../expression", version = "${version}" }`); + .replace(expressionDep, `zen-expression = { path = "../expression", version = "${version}" }`) + .replace(templateDep, `zen-template = { path = "../template", version = "${version}" }`); }; export const getCargoVersion = (contents: string): string => { diff --git a/actions/cargo-version-action/src/index.spec.ts b/actions/cargo-version-action/src/index.spec.ts index 359349e5..c95b6d49 100644 --- a/actions/cargo-version-action/src/index.spec.ts +++ b/actions/cargo-version-action/src/index.spec.ts @@ -7,28 +7,29 @@ import * as path from 'path'; // language=Toml const makeToml = ({ version }): string => ` -[package] -authors = ["GoRules Team "] -description = "Business rules engine" -name = "zen-engine" -license = "MIT" -version = "${version}" -edition = "2021" -repository = "https://github.com/gorules/zen.git" + [package] + authors = ["GoRules Team "] + description = "Business rules engine" + name = "zen-engine" + license = "MIT" + version = "${version}" + edition = "2021" + repository = "https://github.com/gorules/zen.git" -[dependencies] -async-recursion = "1.0.4" -anyhow = { workspace = true } -thiserror = { workspace = true } -async-trait = { workspace = true } -bincode = { workspace = true, optional = true } -serde_json = { workspace = true, features = ["arbitrary_precision"] } -serde = { version = "1", features = ["derive"] } -serde_v8 = { version = "0.88.0" } -once_cell = { version = "1.17.1" } -futures = "0.3.27" -v8 = { version = "0.66.0" } -zen-expression = { path = "../expression", version = "${version}" } + [dependencies] + async-recursion = "1.0.4" + anyhow = { workspace = true } + thiserror = { workspace = true } + async-trait = { workspace = true } + bincode = { workspace = true, optional = true } + serde_json = { workspace = true, features = ["arbitrary_precision"] } + serde = { version = "1", features = ["derive"] } + serde_v8 = { version = "0.88.0" } + once_cell = { version = "1.17.1" } + futures = "0.3.27" + v8 = { version = "0.66.0" } + zen-expression = { path = "../expression", version = "${version}" } + zen-template = { path = "../template", version = "${version}" } `; describe('GitHub Action', () => { diff --git a/bindings/c/Cargo.toml b/bindings/c/Cargo.toml index 0ac4c5e5..9b8f74f8 100644 --- a/bindings/c/Cargo.toml +++ b/bindings/c/Cargo.toml @@ -15,6 +15,7 @@ strum = { version = "0.26.1", features = ["derive"] } futures = "0.3.28" zen-engine = { path = "../../core/engine" } zen-expression = { path = "../../core/expression" } +zen-template = { path = "../../core/template" } [lib] crate-type = ["staticlib"] @@ -23,4 +24,5 @@ crate-type = ["staticlib"] cbindgen = "0.26.0" [features] +default = ["go"] go = [] \ No newline at end of file diff --git a/bindings/c/src/custom_node.rs b/bindings/c/src/custom_node.rs new file mode 100644 index 00000000..569fee82 --- /dev/null +++ b/bindings/c/src/custom_node.rs @@ -0,0 +1,65 @@ +use std::ffi::{c_char, CStr, CString}; + +use anyhow::anyhow; + +use zen_engine::handler::custom_node_adapter::{ + CustomNodeAdapter, CustomNodeRequest, NoopCustomNode, +}; +use zen_engine::handler::node::{NodeResponse, NodeResult}; + +use crate::languages::native::NativeCustomNode; + +#[derive(Debug)] +pub(crate) enum DynamicCustomNode { + Noop(NoopCustomNode), + Native(NativeCustomNode), + #[cfg(feature = "go")] + Go(crate::languages::go::GoCustomNode), +} + +impl Default for DynamicCustomNode { + fn default() -> Self { + Self::Noop(Default::default()) + } +} + +impl CustomNodeAdapter for DynamicCustomNode { + async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + match self { + DynamicCustomNode::Noop(cn) => cn.handle(request).await, + DynamicCustomNode::Native(cn) => cn.handle(request).await, + #[cfg(feature = "go")] + DynamicCustomNode::Go(cn) => cn.handle(request).await, + } + } +} + +#[repr(C)] +pub struct ZenCustomNodeResult { + content: *const c_char, + error: *mut c_char, +} + +impl ZenCustomNodeResult { + pub fn into_node_result(self) -> NodeResult { + let maybe_error = match self.error.is_null() { + false => Some(unsafe { CString::from_raw(self.error) }), + true => None, + }; + + if let Some(c_error) = maybe_error { + let maybe_str = c_error.to_str().unwrap_or("unknown error"); + return Err(anyhow!("{maybe_str}")); + } + + if self.content.is_null() { + return Err(anyhow!("response not provided")); + } + + let content_cstr = unsafe { CStr::from_ptr(self.content) }; + let node_response: NodeResponse = serde_json::from_slice(content_cstr.to_bytes()) + .map_err(|_| anyhow!("failed to deserialize"))?; + + Ok(node_response) + } +} diff --git a/bindings/c/src/decision.rs b/bindings/c/src/decision.rs index e69b60b0..7d4fc513 100644 --- a/bindings/c/src/decision.rs +++ b/bindings/c/src/decision.rs @@ -6,6 +6,7 @@ use futures::executor::block_on; use zen_engine::Decision; +use crate::custom_node::DynamicCustomNode; use crate::engine::ZenEngineEvaluationOptions; use crate::error::ZenError; use crate::loader::DynamicDecisionLoader; @@ -13,12 +14,12 @@ use crate::result::ZenResult; #[repr(C)] pub(crate) struct ZenDecision { - _data: Decision, + _data: Decision, _marker: PhantomData<(*mut c_void, PhantomPinned)>, } impl Deref for ZenDecision { - type Target = Decision; + type Target = Decision; fn deref(&self) -> &Self::Target { &self._data @@ -31,8 +32,8 @@ impl DerefMut for ZenDecision { } } -impl From> for ZenDecision { - fn from(value: Decision) -> Self { +impl From> for ZenDecision { + fn from(value: Decision) -> Self { Self { _data: value, _marker: PhantomData, diff --git a/bindings/c/src/engine.rs b/bindings/c/src/engine.rs index 05ecc031..c17c46b1 100644 --- a/bindings/c/src/engine.rs +++ b/bindings/c/src/engine.rs @@ -5,6 +5,7 @@ use std::sync::Arc; use futures::executor::block_on; +use crate::custom_node::DynamicCustomNode; use zen_engine::{DecisionEngine, EvaluationOptions}; use crate::decision::{ZenDecision, ZenDecisionStruct}; @@ -12,10 +13,10 @@ use crate::error::ZenError; use crate::loader::DynamicDecisionLoader; use crate::result::ZenResult; -pub(crate) struct ZenEngine(DecisionEngine); +pub(crate) struct ZenEngine(DecisionEngine); impl Deref for ZenEngine { - type Target = DecisionEngine; + type Target = DecisionEngine; fn deref(&self) -> &Self::Target { &self.0 @@ -30,13 +31,16 @@ impl DerefMut for ZenEngine { impl Default for ZenEngine { fn default() -> Self { - Self(DecisionEngine::new(DynamicDecisionLoader::default())) + Self(DecisionEngine::new( + Arc::new(DynamicDecisionLoader::default()), + Arc::new(DynamicCustomNode::default()), + )) } } impl ZenEngine { - pub fn with_loader(loader: DynamicDecisionLoader) -> Self { - Self(DecisionEngine::new(loader)) + pub fn new(loader: DynamicDecisionLoader, custom_node: DynamicCustomNode) -> Self { + Self(DecisionEngine::new(Arc::new(loader), Arc::new(custom_node))) } } @@ -88,11 +92,7 @@ pub extern "C" fn zen_engine_create_decision( } let cstr_content = unsafe { CStr::from_ptr(content) }; - let Ok(str_content) = cstr_content.to_str() else { - return ZenResult::error(ZenError::InvalidArgument); - }; - - let Ok(decision_content) = serde_json::from_str(str_content) else { + let Ok(decision_content) = serde_json::from_slice(cstr_content.to_bytes()) else { return ZenResult::error(ZenError::JsonDeserializationFailed); }; @@ -122,11 +122,7 @@ pub extern "C" fn zen_engine_evaluate( }; let cstr_context = unsafe { CStr::from_ptr(context) }; - let Ok(str_context) = cstr_context.to_str() else { - return ZenResult::error(ZenError::InvalidArgument); - }; - - let Ok(val_context) = serde_json::from_str(str_context) else { + let Ok(val_context) = serde_json::from_slice(cstr_context.to_bytes()) else { return ZenResult::error(ZenError::JsonDeserializationFailed); }; diff --git a/bindings/c/src/error.rs b/bindings/c/src/error.rs index a6390b1e..2345340a 100644 --- a/bindings/c/src/error.rs +++ b/bindings/c/src/error.rs @@ -19,6 +19,8 @@ pub enum ZenError { LoaderKeyNotFound { key: String }, LoaderInternalError { key: String, message: String }, + + TemplateEngineError { template: String, message: String }, } impl ZenError { @@ -30,6 +32,9 @@ impl ZenError { ZenError::LoaderInternalError { key, message } => { Some(json!({ "key": key, "message": message }).to_string()) } + ZenError::TemplateEngineError { template, message } => { + Some(json!({ "template": template, "message": message }).to_string()) + } _ => None, } } diff --git a/bindings/c/src/expression.rs b/bindings/c/src/expression.rs index 3011414b..352cbbd1 100644 --- a/bindings/c/src/expression.rs +++ b/bindings/c/src/expression.rs @@ -1,31 +1,27 @@ -use std::ffi::{CStr, CString}; +use std::ffi::CString; use libc::{c_char, c_int}; use zen_expression::{evaluate_expression, evaluate_unary_expression}; use crate::error::ZenError; +use crate::helper::{safe_cstr_from_ptr, safe_str_from_ptr}; use crate::result::ZenResult; -/// Evaluate expression, responsible for freeing expression and context #[no_mangle] pub extern "C" fn zen_evaluate_expression( expression: *const c_char, context: *const c_char, ) -> ZenResult { - if expression.is_null() || context.is_null() { - return ZenResult::error(ZenError::InvalidArgument); - } - - let Ok(expression_str) = unsafe { CStr::from_ptr(expression) }.to_str() else { + let Some(expression_str) = safe_str_from_ptr(expression) else { return ZenResult::error(ZenError::InvalidArgument); }; - let Ok(context_str) = unsafe { CStr::from_ptr(context) }.to_str() else { + let Some(context_cstr) = safe_cstr_from_ptr(context) else { return ZenResult::error(ZenError::InvalidArgument); }; - let Ok(context_val) = serde_json::from_str(context_str) else { + let Ok(context_val) = serde_json::from_slice(context_cstr.to_bytes()) else { return ZenResult::error(ZenError::JsonDeserializationFailed); }; @@ -51,19 +47,15 @@ pub extern "C" fn zen_evaluate_unary_expression( expression: *const c_char, context: *const c_char, ) -> ZenResult { - if expression.is_null() || context.is_null() { - return ZenResult::error(ZenError::InvalidArgument); - } - - let Ok(expression_str) = unsafe { CStr::from_ptr(expression) }.to_str() else { + let Some(expression_str) = safe_str_from_ptr(expression) else { return ZenResult::error(ZenError::InvalidArgument); }; - let Ok(context_str) = unsafe { CStr::from_ptr(context) }.to_str() else { + let Some(context_cstr) = safe_cstr_from_ptr(context) else { return ZenResult::error(ZenError::InvalidArgument); }; - let Ok(context_val) = serde_json::from_str(context_str) else { + let Ok(context_val) = serde_json::from_slice(context_cstr.to_bytes()) else { return ZenResult::error(ZenError::JsonDeserializationFailed); }; @@ -77,3 +69,41 @@ pub extern "C" fn zen_evaluate_unary_expression( ZenResult::ok(Box::into_raw(Box::new(c_result))) } + +/// Evaluate unary expression, responsible for freeing expression and context +/// True = 1 +/// False = 0 +#[no_mangle] +pub extern "C" fn zen_evaluate_template( + template: *const c_char, + context: *const c_char, +) -> ZenResult { + let Some(template_str) = safe_str_from_ptr(template) else { + return ZenResult::error(ZenError::InvalidArgument); + }; + + let Some(context_cstr) = safe_cstr_from_ptr(context) else { + return ZenResult::error(ZenError::InvalidArgument); + }; + + let Ok(context_val) = serde_json::from_slice(context_cstr.to_bytes()) else { + return ZenResult::error(ZenError::JsonDeserializationFailed); + }; + + let result = match zen_template::render(template_str, &context_val) { + Ok(r) => r, + Err(err) => { + return ZenResult::error(ZenError::TemplateEngineError { + message: err.to_string(), + template: template_str.to_string(), + }) + } + }; + + let Ok(s_result) = serde_json::to_string(&result) else { + return ZenResult::error(ZenError::JsonSerializationFailed); + }; + + let c_result = unsafe { CString::from_vec_unchecked(s_result.into_bytes()) }; + ZenResult::ok(c_result.into_raw()) +} diff --git a/bindings/c/src/helper.rs b/bindings/c/src/helper.rs new file mode 100644 index 00000000..d4dc9277 --- /dev/null +++ b/bindings/c/src/helper.rs @@ -0,0 +1,23 @@ +use std::ffi::CStr; + +use libc::c_char; + +pub(crate) fn safe_cstr_from_ptr<'a>(ptr: *const c_char) -> Option<&'a CStr> { + if ptr.is_null() { + None + } else { + // SAFETY: The caller must ensure the pointer is not null and points to a valid, null-terminated C string. + // This unsafe block is necessary because CStr::from_ptr inherently requires an unsafe operation. + Some(unsafe { CStr::from_ptr(ptr) }) + } +} + +pub(crate) fn safe_str_from_ptr<'a>(ptr: *const c_char) -> Option<&'a str> { + if ptr.is_null() { + None + } else { + // SAFETY: The caller must ensure the pointer is not null and points to a valid, null-terminated C string. + // This unsafe block is necessary because CStr::from_ptr inherently requires an unsafe operation. + Some(unsafe { CStr::from_ptr(ptr) }.to_str().ok()?) + } +} diff --git a/bindings/c/src/languages/go.rs b/bindings/c/src/languages/go.rs index 59f108b5..3b057fc0 100644 --- a/bindings/c/src/languages/go.rs +++ b/bindings/c/src/languages/go.rs @@ -1,15 +1,18 @@ -use std::ffi::c_char; +use std::ffi::{c_char, CString}; +use anyhow::anyhow; use async_trait::async_trait; +use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; +use zen_engine::handler::node::NodeResult; use zen_engine::loader::{DecisionLoader, LoaderError, LoaderResponse}; +use crate::custom_node::{DynamicCustomNode, ZenCustomNodeResult}; use crate::engine::{ZenEngine, ZenEngineStruct}; use crate::loader::{DynamicDecisionLoader, ZenDecisionLoaderResult}; #[derive(Debug, Default)] pub(crate) struct GoDecisionLoader { - #[allow(dead_code)] handler: Option, } @@ -26,7 +29,7 @@ impl DecisionLoader for GoDecisionLoader { return Err(LoaderError::NotFound(key.to_string()).into()); }; - let c_key = std::ffi::CString::new(key).unwrap(); + let c_key = CString::new(key).unwrap(); let c_content_ptr = unsafe { zen_engine_go_loader_callback(handler.clone(), c_key.as_ptr()) }; @@ -34,13 +37,53 @@ impl DecisionLoader for GoDecisionLoader { } } +#[derive(Debug, Default)] +pub(crate) struct GoCustomNode { + #[allow(dead_code)] + handler: Option, +} + +impl GoCustomNode { + pub fn new(handler: Option) -> Self { + Self { handler } + } +} + +impl CustomNodeAdapter for GoCustomNode { + async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + let Some(handler) = self.handler else { + return Err(anyhow!("go handler not found")); + }; + + let Ok(request_value) = serde_json::to_string(&request) else { + return Err(anyhow!("failed to serialize request json")); + }; + + let c_request = unsafe { CString::from_vec_unchecked(request_value.into_bytes()) }; + let c_node_result = + unsafe { zen_engine_go_custom_node_callback(handler, c_request.as_ptr()) }; + c_node_result.into_node_result() + } +} + +fn map_handler(i: Option) -> Option { + let Some(j) = i else { return None }; + + (j > 0).then_some(j) +} + /// Creates a DecisionEngine for using GoLang handler (optional). Caller is responsible for freeing DecisionEngine. #[no_mangle] -pub extern "C" fn zen_engine_new_with_go_loader( +pub extern "C" fn zen_engine_new_golang( maybe_loader: Option<&usize>, + maybe_custom_node: Option<&usize>, ) -> *mut ZenEngineStruct { - let loader = GoDecisionLoader::new(maybe_loader.cloned()); - let engine = ZenEngine::with_loader(DynamicDecisionLoader::Go(loader)); + let loader = GoDecisionLoader::new(map_handler(maybe_loader.cloned())); + let custom_node = GoCustomNode::new(map_handler(maybe_custom_node.cloned())); + let engine = ZenEngine::new( + DynamicDecisionLoader::Go(loader), + DynamicCustomNode::Go(custom_node), + ); Box::into_raw(Box::new(engine)) as *mut ZenEngineStruct } @@ -49,4 +92,8 @@ pub extern "C" fn zen_engine_new_with_go_loader( /// cbindgen:ignore extern "C" { fn zen_engine_go_loader_callback(cb_ptr: usize, key: *const c_char) -> ZenDecisionLoaderResult; + fn zen_engine_go_custom_node_callback( + cb_ptr: usize, + request: *const c_char, + ) -> ZenCustomNodeResult; } diff --git a/bindings/c/src/languages/native.rs b/bindings/c/src/languages/native.rs index 5c737ddd..e854fdc5 100644 --- a/bindings/c/src/languages/native.rs +++ b/bindings/c/src/languages/native.rs @@ -1,15 +1,21 @@ use std::ffi::{c_char, CString}; +use anyhow::anyhow; use async_trait::async_trait; +use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; +use zen_engine::handler::node::NodeResult; use zen_engine::loader::{DecisionLoader, LoaderResponse}; +use crate::custom_node::{DynamicCustomNode, ZenCustomNodeResult}; use crate::engine::{ZenEngine, ZenEngineStruct}; use crate::loader::{DynamicDecisionLoader, ZenDecisionLoaderResult}; pub type ZenDecisionLoaderNativeCallback = extern "C" fn(key: *const c_char) -> ZenDecisionLoaderResult; +pub type ZenCustomNodeNativeCallback = extern "C" fn(request: *const c_char) -> ZenCustomNodeResult; + #[derive(Debug)] pub(crate) struct NativeDecisionLoader { callback: ZenDecisionLoaderNativeCallback, @@ -31,14 +37,43 @@ impl DecisionLoader for NativeDecisionLoader { } } +#[derive(Debug)] +pub(crate) struct NativeCustomNode { + callback: ZenCustomNodeNativeCallback, +} + +impl NativeCustomNode { + pub fn new(callback: ZenCustomNodeNativeCallback) -> Self { + Self { callback } + } +} + +impl CustomNodeAdapter for NativeCustomNode { + async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + let Ok(request_value) = serde_json::to_string(&request) else { + return Err(anyhow!("failed to serialize request json")); + }; + + let c_request = unsafe { CString::from_vec_unchecked(request_value.into_bytes()) }; + let c_response_str = (&self.callback)(c_request.as_ptr()); + c_response_str.into_node_result() + } +} + /// Creates a new ZenEngine instance with loader, caller is responsible for freeing the returned reference /// by calling zen_engine_free. #[no_mangle] -pub extern "C" fn zen_engine_new_with_native_loader( - callback: ZenDecisionLoaderNativeCallback, +pub extern "C" fn zen_engine_new_native( + loader_callback: ZenDecisionLoaderNativeCallback, + custom_node_callback: ZenCustomNodeNativeCallback, ) -> *mut ZenEngineStruct { - let loader = NativeDecisionLoader::new(callback); - let engine = ZenEngine::with_loader(DynamicDecisionLoader::Native(loader)); + let loader = NativeDecisionLoader::new(loader_callback); + let custom_node = NativeCustomNode::new(custom_node_callback); + + let engine = ZenEngine::new( + DynamicDecisionLoader::Native(loader), + DynamicCustomNode::Native(custom_node), + ); Box::into_raw(Box::new(engine)) as *mut ZenEngineStruct } diff --git a/bindings/c/src/lib.rs b/bindings/c/src/lib.rs index 8bc8e4ac..013b08b5 100644 --- a/bindings/c/src/lib.rs +++ b/bindings/c/src/lib.rs @@ -1,9 +1,11 @@ extern crate zen_engine; +mod custom_node; mod decision; mod engine; mod error; mod expression; +mod helper; mod languages; mod loader; mod result; diff --git a/bindings/c/src/loader.rs b/bindings/c/src/loader.rs index d4bcae9e..6c045979 100644 --- a/bindings/c/src/loader.rs +++ b/bindings/c/src/loader.rs @@ -1,7 +1,7 @@ use crate::languages::native::NativeDecisionLoader; use anyhow::anyhow; use async_trait::async_trait; -use std::ffi::{c_char, CString}; +use std::ffi::{c_char, CStr, CString}; use std::sync::Arc; use zen_engine::loader::{DecisionLoader, LoaderError, LoaderResponse, NoopLoader}; use zen_engine::model::DecisionContent; @@ -54,7 +54,7 @@ impl ZenDecisionLoaderResult { } let maybe_content = match self.content.is_null() { - false => Some(unsafe { CString::from_raw(self.content) }), + false => Some(unsafe { CStr::from_ptr(self.content) }), true => None, }; @@ -63,13 +63,8 @@ impl ZenDecisionLoaderResult { return Err(LoaderError::NotFound(key.to_string()).into()); }; - let content = c_content.into_string().map_err(|e| LoaderError::Internal { - key: key.to_string(), - source: anyhow!(e), - })?; - - let decision_content: DecisionContent = - serde_json::from_str(&content).map_err(|e| LoaderError::Internal { + let decision_content: DecisionContent = serde_json::from_slice(c_content.to_bytes()) + .map_err(|e| LoaderError::Internal { key: key.to_string(), source: anyhow!(e), })?; diff --git a/bindings/c/zen_engine.h b/bindings/c/zen_engine.h index df7b6dd4..7217d758 100644 --- a/bindings/c/zen_engine.h +++ b/bindings/c/zen_engine.h @@ -53,6 +53,13 @@ typedef struct ZenDecisionLoaderResult { typedef struct ZenDecisionLoaderResult (*ZenDecisionLoaderNativeCallback)(const char *key); +typedef struct ZenCustomNodeResult { + const char *content; + char *error; +} ZenCustomNodeResult; + +typedef struct ZenCustomNodeResult (*ZenCustomNodeNativeCallback)(const char *request); + /** * Frees ZenDecision */ @@ -100,9 +107,6 @@ struct ZenResult_c_char zen_engine_evaluate(const struct ZenEngineStruct *engine struct ZenResult_ZenDecisionStruct zen_engine_get_decision(const struct ZenEngineStruct *engine, const char *key); -/** - * Evaluate expression, responsible for freeing expression and context - */ struct ZenResult_c_char zen_evaluate_expression(const char *expression, const char *context); /** @@ -112,13 +116,22 @@ struct ZenResult_c_char zen_evaluate_expression(const char *expression, const ch */ struct ZenResult_c_int zen_evaluate_unary_expression(const char *expression, const char *context); +/** + * Evaluate unary expression, responsible for freeing expression and context + * True = 1 + * False = 0 + */ +struct ZenResult_c_char zen_evaluate_template(const char *template_, const char *context); + /** * Creates a new ZenEngine instance with loader, caller is responsible for freeing the returned reference * by calling zen_engine_free. */ -struct ZenEngineStruct *zen_engine_new_with_native_loader(ZenDecisionLoaderNativeCallback callback); +struct ZenEngineStruct *zen_engine_new_native(ZenDecisionLoaderNativeCallback loader_callback, + ZenCustomNodeNativeCallback custom_node_callback); /** * Creates a DecisionEngine for using GoLang handler (optional). Caller is responsible for freeing DecisionEngine. */ -struct ZenEngineStruct *zen_engine_new_with_go_loader(const uintptr_t *maybe_loader); +struct ZenEngineStruct *zen_engine_new_golang(const uintptr_t *maybe_loader, + const uintptr_t *maybe_custom_node); diff --git a/bindings/nodejs/Cargo.toml b/bindings/nodejs/Cargo.toml index 4207dbe4..1ea90ff5 100644 --- a/bindings/nodejs/Cargo.toml +++ b/bindings/nodejs/Cargo.toml @@ -10,12 +10,15 @@ crate-type = ["cdylib"] [dependencies] async-trait = { workspace = true } -napi = { version = "2.14.4", features = ["serde-json", "error_anyhow", "tokio_rt"] } -napi-derive = "2.14.6" +napi = { version = "2.16", features = ["serde-json", "error_anyhow", "tokio_rt"] } +napi-derive = "2.16" serde_json = { workspace = true } futures = { workspace = true } zen-engine = { path = "../../core/engine" } zen-expression = { path = "../../core/expression" } +zen-template = { path = "../../core/template" } +serde = { workspace = true, features = ["derive"] } +json_dotpath = { workspace = true } [build-dependencies] -napi-build = "2.1.0" \ No newline at end of file +napi-build = "2.1.2" \ No newline at end of file diff --git a/bindings/nodejs/index.d.ts b/bindings/nodejs/index.d.ts index 0c744e8d..4207b556 100644 --- a/bindings/nodejs/index.d.ts +++ b/bindings/nodejs/index.d.ts @@ -9,17 +9,52 @@ export interface ZenEvaluateOptions { } export interface ZenEngineOptions { loader?: (key: string) => Promise + customHandler?: (request: ZenEngineHandlerRequest) => Promise } +export function evaluateExpressionSync(expression: string, context?: any | undefined | null): any +export function evaluateUnaryExpressionSync(expression: string, context: any): boolean +export function renderTemplateSync(template: string, context: any): any export function evaluateExpression(expression: string, context?: any | undefined | null): Promise export function evaluateUnaryExpression(expression: string, context: any): Promise +export function renderTemplate(template: string, context: any): Promise +export interface ZenEngineTrace { + id: string + name: string + input: any + output: any + performance?: string + traceData?: any +} +export interface ZenEngineResponse { + performance: string + result: any + trace?: Record +} +export interface ZenEngineHandlerResponse { + output: any + traceData?: any +} +export interface DecisionNode { + id: string + name: string + kind: string + config: any +} export class ZenDecision { constructor() - evaluate(context: any, opts?: ZenEvaluateOptions | undefined | null): Promise + evaluate(context: any, opts?: ZenEvaluateOptions | undefined | null): Promise validate(): void } export class ZenEngine { constructor(options?: ZenEngineOptions | undefined | null) - evaluate(key: string, context: any, opts?: ZenEvaluateOptions | undefined | null): Promise + evaluate(key: string, context: any, opts?: ZenEvaluateOptions | undefined | null): Promise createDecision(content: Buffer): ZenDecision getDecision(key: string): Promise } +export class ZenEngineHandlerRequest { + input: any + node: DecisionNode + constructor() + getField(path: string): unknown + getFieldRaw(path: string): unknown +} diff --git a/bindings/nodejs/index.js b/bindings/nodejs/index.js index 08c34739..15e057d2 100644 --- a/bindings/nodejs/index.js +++ b/bindings/nodejs/index.js @@ -281,9 +281,14 @@ if (!nativeBinding) { throw new Error(`Failed to load native binding`) } -const { ZenDecision, ZenEngine, evaluateExpression, evaluateUnaryExpression } = nativeBinding +const { ZenDecision, ZenEngine, evaluateExpressionSync, evaluateUnaryExpressionSync, renderTemplateSync, evaluateExpression, evaluateUnaryExpression, renderTemplate, ZenEngineHandlerRequest } = nativeBinding module.exports.ZenDecision = ZenDecision module.exports.ZenEngine = ZenEngine +module.exports.evaluateExpressionSync = evaluateExpressionSync +module.exports.evaluateUnaryExpressionSync = evaluateUnaryExpressionSync +module.exports.renderTemplateSync = renderTemplateSync module.exports.evaluateExpression = evaluateExpression module.exports.evaluateUnaryExpression = evaluateUnaryExpression +module.exports.renderTemplate = renderTemplate +module.exports.ZenEngineHandlerRequest = ZenEngineHandlerRequest diff --git a/bindings/nodejs/src/custom_node.rs b/bindings/nodejs/src/custom_node.rs new file mode 100644 index 00000000..445f7641 --- /dev/null +++ b/bindings/nodejs/src/custom_node.rs @@ -0,0 +1,60 @@ +use napi::anyhow::anyhow; +use napi::bindgen_prelude::Promise; +use napi::threadsafe_function::{ErrorStrategy, ThreadSafeCallContext, ThreadsafeFunction}; +use napi::{Env, JsFunction}; + +use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; +use zen_engine::handler::node::{NodeResponse, NodeResult}; + +use crate::types::{ZenEngineHandlerRequest, ZenEngineHandlerResponse}; + +pub(crate) struct CustomNode { + function: Option>, +} + +impl Default for CustomNode { + fn default() -> Self { + Self { function: None } + } +} + +impl CustomNode { + pub fn try_new(env: &mut Env, function: JsFunction) -> napi::Result { + let mut tsf = function.create_threadsafe_function( + 0, + |cx: ThreadSafeCallContext| Ok(vec![cx.value]), + )?; + + tsf.unref(env)?; + + Ok(Self { + function: Some(tsf), + }) + } +} + +impl CustomNodeAdapter for CustomNode { + async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + let Some(function) = &self.function else { + return Err(anyhow!("Custom function is undefined")); + }; + + let node_data = crate::types::DecisionNode::from(request.node); + + let promise: Promise = function + .clone() + .call_async(ZenEngineHandlerRequest { + input: request.input.clone(), + node: node_data, + }) + .await + .map_err(|err| anyhow!(err.reason))?; + + let result = promise.await.map_err(|err| anyhow!(err.reason))?; + + Ok(NodeResponse { + output: result.output, + trace_data: result.trace_data, + }) + } +} diff --git a/bindings/nodejs/src/decision.rs b/bindings/nodejs/src/decision.rs index e6d7f638..6d02add2 100644 --- a/bindings/nodejs/src/decision.rs +++ b/bindings/nodejs/src/decision.rs @@ -1,5 +1,7 @@ +use crate::custom_node::CustomNode; use crate::engine::ZenEvaluateOptions; use crate::loader::DecisionLoader; +use crate::types::ZenEngineResponse; use napi::anyhow::anyhow; use napi::tokio; use napi_derive::napi; @@ -8,10 +10,10 @@ use std::sync::Arc; use zen_engine::{Decision, EvaluationOptions}; #[napi] -pub struct ZenDecision(pub(crate) Arc>); +pub struct ZenDecision(pub(crate) Arc>); -impl From> for ZenDecision { - fn from(value: Decision) -> Self { +impl From> for ZenDecision { + fn from(value: Decision) -> Self { Self(value.into()) } } @@ -28,7 +30,7 @@ impl ZenDecision { &self, context: Value, opts: Option, - ) -> napi::Result { + ) -> napi::Result { let decision = self.0.clone(); let result = tokio::spawn(async move { let options = opts.unwrap_or_default(); @@ -46,7 +48,7 @@ impl ZenDecision { anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) })?; - Ok(serde_json::to_value(&result)?) + Ok(ZenEngineResponse::from(result)) } #[napi] diff --git a/bindings/nodejs/src/engine.rs b/bindings/nodejs/src/engine.rs index 8eecd98f..3a41173c 100644 --- a/bindings/nodejs/src/engine.rs +++ b/bindings/nodejs/src/engine.rs @@ -1,5 +1,7 @@ +use crate::custom_node::CustomNode; use crate::decision::ZenDecision; use crate::loader::DecisionLoader; +use crate::types::ZenEngineResponse; use napi::anyhow::{anyhow, Context}; use napi::bindgen_prelude::Buffer; use napi::{tokio, Env, JsFunction}; @@ -11,7 +13,7 @@ use zen_engine::{DecisionEngine, EvaluationOptions}; #[napi] pub struct ZenEngine { - graph: Arc>, + graph: Arc>, } #[napi(object)] @@ -33,6 +35,9 @@ impl Default for ZenEvaluateOptions { pub struct ZenEngineOptions { #[napi(ts_type = "(key: string) => Promise")] pub loader: Option, + + #[napi(ts_type = "(request: ZenEngineHandlerRequest) => Promise")] + pub custom_handler: Option, } #[napi] @@ -41,18 +46,26 @@ impl ZenEngine { pub fn new(mut env: Env, options: Option) -> napi::Result { let Some(opts) = options else { return Ok(Self { - graph: DecisionEngine::new(DecisionLoader::default()).into(), + graph: DecisionEngine::new( + DecisionLoader::default().into(), + CustomNode::default().into(), + ) + .into(), }); }; - let Some(loader_fn) = opts.loader else { - return Ok(Self { - graph: DecisionEngine::new(DecisionLoader::default()).into(), - }); + let loader = match opts.loader { + None => DecisionLoader::default(), + Some(loader_fn) => DecisionLoader::try_new(&mut env, loader_fn)?, + }; + + let custom_handler = match opts.custom_handler { + None => CustomNode::default(), + Some(custom_fn) => CustomNode::try_new(&mut env, custom_fn)?, }; Ok(Self { - graph: DecisionEngine::new(DecisionLoader::try_new(&mut env, loader_fn)?).into(), + graph: DecisionEngine::new(loader.into(), custom_handler.into()).into(), }) } @@ -62,7 +75,7 @@ impl ZenEngine { key: String, context: Value, opts: Option, - ) -> napi::Result { + ) -> napi::Result { let graph = self.graph.clone(); let result = tokio::spawn(async move { let options = opts.unwrap_or_default(); @@ -82,7 +95,7 @@ impl ZenEngine { anyhow!(serde_json::to_string(e.as_ref()).unwrap_or_else(|_| e.to_string())) })?; - Ok(serde_json::to_value(&result)?) + Ok(ZenEngineResponse::from(result)) } #[napi] diff --git a/bindings/nodejs/src/expression.rs b/bindings/nodejs/src/expression.rs index 08e8bd87..0f258155 100644 --- a/bindings/nodejs/src/expression.rs +++ b/bindings/nodejs/src/expression.rs @@ -2,33 +2,55 @@ use napi::anyhow::anyhow; use napi_derive::napi; use serde_json::Value; +#[napi] +pub fn evaluate_expression_sync(expression: String, context: Option) -> napi::Result { + let ctx = context.unwrap_or(Value::Null); + + Ok( + zen_expression::evaluate_expression(expression.as_str(), &ctx) + .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?, + ) +} + +#[allow(dead_code)] +#[napi] +pub fn evaluate_unary_expression_sync(expression: String, context: Value) -> napi::Result { + Ok( + zen_expression::evaluate_unary_expression(expression.as_str(), &context) + .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?, + ) +} + +#[allow(dead_code)] +#[napi] +pub fn render_template_sync(template: String, context: Value) -> napi::Result { + Ok(zen_template::render(template.as_str(), &context) + .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?) +} + #[allow(dead_code)] #[napi] pub async fn evaluate_expression( expression: String, context: Option, ) -> napi::Result { - let ctx = context.unwrap_or(Value::Null); - - let result: Value = napi::tokio::spawn(async move { - zen_expression::evaluate_expression(expression.as_str(), &ctx) - }) - .await - .map_err(|_| anyhow!("Hook timed out"))? - .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?; - - Ok(result) + napi::tokio::spawn(async move { evaluate_expression_sync(expression, context) }) + .await + .map_err(|_| anyhow!("Hook timed out"))? } #[allow(dead_code)] #[napi] pub async fn evaluate_unary_expression(expression: String, context: Value) -> napi::Result { - let result: bool = napi::tokio::spawn(async move { - zen_expression::evaluate_unary_expression(expression.as_str(), &context) - }) - .await - .map_err(|_| anyhow!("Hook timed out"))? - .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?; + napi::tokio::spawn(async move { evaluate_unary_expression_sync(expression, context) }) + .await + .map_err(|_| anyhow!("Hook timed out"))? +} - Ok(result) +#[allow(dead_code)] +#[napi] +pub async fn render_template(template: String, context: Value) -> napi::Result { + napi::tokio::spawn(async move { render_template_sync(template, context) }) + .await + .map_err(|_| anyhow!("Hook timed out"))? } diff --git a/bindings/nodejs/src/lib.rs b/bindings/nodejs/src/lib.rs index 9068031c..4bcead12 100644 --- a/bindings/nodejs/src/lib.rs +++ b/bindings/nodejs/src/lib.rs @@ -1,4 +1,6 @@ +mod custom_node; mod decision; mod engine; mod expression; mod loader; +mod types; diff --git a/bindings/nodejs/src/loader.rs b/bindings/nodejs/src/loader.rs index 799d99d1..d493b887 100644 --- a/bindings/nodejs/src/loader.rs +++ b/bindings/nodejs/src/loader.rs @@ -48,12 +48,12 @@ impl DecisionLoader { .await .map_err(|e| LoaderError::Internal { key: key.to_string(), - source: e.into(), + source: anyhow!(e.reason), })?; let result = promise.await.map_err(|e| LoaderError::Internal { key: key.to_string(), - source: e.into(), + source: anyhow!(e.reason), })?; let Some(buffer) = result else { diff --git a/bindings/nodejs/src/types.rs b/bindings/nodejs/src/types.rs new file mode 100644 index 00000000..518dfb5f --- /dev/null +++ b/bindings/nodejs/src/types.rs @@ -0,0 +1,125 @@ +use std::collections::HashMap; + +use json_dotpath::DotPaths; +use napi::anyhow::{anyhow, Context}; +use napi_derive::napi; +use serde_json::Value; + +use zen_engine::handler::custom_node_adapter::CustomDecisionNode; +use zen_engine::{DecisionGraphResponse, DecisionGraphTrace}; + +#[napi(object)] +pub struct ZenEngineTrace { + pub id: String, + pub name: String, + pub input: Value, + pub output: Value, + pub performance: Option, + pub trace_data: Option, +} + +impl From for ZenEngineTrace { + fn from(value: DecisionGraphTrace) -> Self { + Self { + id: value.id, + name: value.name, + input: value.input, + output: value.output, + performance: value.performance, + trace_data: value.trace_data, + } + } +} + +#[napi(object)] +pub struct ZenEngineResponse { + pub performance: String, + pub result: Value, + pub trace: Option>, +} + +impl From for ZenEngineResponse { + fn from(value: DecisionGraphResponse) -> Self { + Self { + performance: value.performance, + result: value.result, + trace: value.trace.map(|opt| { + opt.into_iter() + .map(|(key, value)| (key, ZenEngineTrace::from(value))) + .collect() + }), + } + } +} + +#[napi(object)] +pub struct ZenEngineHandlerResponse { + pub output: Value, + pub trace_data: Option, +} + +#[derive(Clone)] +#[napi(object)] +pub struct DecisionNode { + pub id: String, + pub name: String, + pub kind: String, + pub config: Value, +} + +impl From> for DecisionNode { + fn from(value: CustomDecisionNode<'_>) -> Self { + Self { + id: value.id.to_string(), + name: value.name.to_string(), + kind: value.kind.to_string(), + config: value.config.clone(), + } + } +} + +#[napi] +pub struct ZenEngineHandlerRequest { + pub input: Value, + pub node: DecisionNode, +} + +#[napi] +impl ZenEngineHandlerRequest { + #[napi(constructor)] + pub fn new() -> napi::Result { + Err(anyhow!("Private constructor").into()) + } + + #[napi(ts_return_type = "unknown")] + pub fn get_field(&self, path: String) -> napi::Result { + let node_config = &self.node.config; + + let selected_value: Value = node_config + .dot_get(path.as_str()) + .ok() + .flatten() + .context("Failed to find JSON path")?; + let Value::String(template) = selected_value else { + return Ok(selected_value); + }; + + let template_value = zen_template::render(template.as_str(), &self.input) + .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?; + + Ok(template_value) + } + + #[napi(ts_return_type = "unknown")] + pub fn get_field_raw(&self, path: String) -> napi::Result { + let node_config = &self.node.config; + + let selected_value: Value = node_config + .dot_get(path.as_str()) + .ok() + .flatten() + .context("Failed to find JSON path")?; + + Ok(selected_value.clone()) + } +} diff --git a/bindings/nodejs/test/decision.spec.ts b/bindings/nodejs/test/decision.spec.ts index 0e1a2752..2bf9868f 100644 --- a/bindings/nodejs/test/decision.spec.ts +++ b/bindings/nodejs/test/decision.spec.ts @@ -1,4 +1,11 @@ -import {ZenEngine, evaluateExpression, evaluateUnaryExpression} from "../index"; +import { + ZenEngine, + evaluateExpression, + evaluateUnaryExpression, + renderTemplate, + evaluateExpressionSync, + evaluateUnaryExpressionSync, renderTemplateSync +} from "../index"; import fs from 'fs/promises'; import path from 'path'; import {describe, expect, it, jest} from "@jest/globals"; @@ -49,6 +56,31 @@ describe('ZenEngine', () => { const r = await functionDecision.evaluate({input: 15}); expect(r.result.output).toEqual(30); }, 10000) + + it('Evaluate custom nodes with a handler', async () => { + const engine = new ZenEngine({ + loader, + customHandler: async (request) => { + const prop1 = request.getField('prop1') as number; + const prop1Raw = request.getFieldRaw('prop1'); + + expect(prop1).toEqual(15); + expect(prop1Raw).toEqual('{{ a + 10 }}') + expect(request.node).toMatchObject({ + id: '138b3b11-ff46-450f-9704-3f3c712067b2', + name: 'customNode1', + kind: 'sum', + config: { + prop1: '{{ a + 10 }}' + } + }); + return {output: {data: prop1 + 10}} + } + }); + + const r = await engine.evaluate('custom.json', {a: 5}); + expect(r.result.data).toEqual(25); + }); }) describe('Expressions', () => { @@ -63,6 +95,7 @@ describe('Expressions', () => { for (const {expression, result, context} of expressions) { expect(await evaluateExpression(expression, context)).toEqual(result); + expect(evaluateExpressionSync(expression, context)).toEqual(result); } }); @@ -76,6 +109,21 @@ describe('Expressions', () => { for (const {expression, result, context} of expressions) { expect(await evaluateUnaryExpression(expression, context)).toEqual(result); + expect(evaluateUnaryExpressionSync(expression, context)).toEqual(result); + } + }); + + it('Renders templates', async () => { + const templateCases = [ + {template: '{{ a + 10 }}', context: {a: 10}, result: 20}, + {template: '{{ a + 10 }}', context: {a: 15}, result: 25}, + {template: '{{ a + 10 }}', context: {a: 20}, result: 30}, + {template: '{{ a + 10 }}', context: {a: 25}, result: 35}, + ]; + + for (const {template, context, result} of templateCases) { + expect(await renderTemplate(template, context)).toEqual(result); + expect(renderTemplateSync(template, context)).toEqual(result); } }); }); \ No newline at end of file diff --git a/bindings/python/Cargo.toml b/bindings/python/Cargo.toml index 353c682c..392d8eb4 100644 --- a/bindings/python/Cargo.toml +++ b/bindings/python/Cargo.toml @@ -12,10 +12,12 @@ crate-type = ["cdylib"] [dependencies] async-trait = { workspace = true } anyhow = { workspace = true } -pyo3 = { version = "0.20.2", features = ["anyhow", "serde"] } -pythonize = "0.20.0" +pyo3 = { version = "0.20", features = ["anyhow", "serde"] } +pythonize = "0.20" +json_dotpath = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } futures = { workspace = true } zen-engine = { path = "../../core/engine" } -zen-expression = { path = "../../core/expression" } \ No newline at end of file +zen-expression = { path = "../../core/expression" } +zen-template = { path = "../../core/template" } \ No newline at end of file diff --git a/bindings/python/index.py b/bindings/python/index.py index 8cc45c51..bf8bad8a 100644 --- a/bindings/python/index.py +++ b/bindings/python/index.py @@ -6,6 +6,12 @@ def loader(key): with open("../../test-data/" + key, "r") as f: return f.read() +def custom_handler(request): + p1 = request.get_field("prop1") + return { + "output": { "sum": p1 } + } + # The test based on unittest module class ZenEngine(unittest.TestCase): def test_decision_using_loader(self): @@ -41,6 +47,16 @@ class ZenEngine(unittest.TestCase): r = functionDecision.evaluate({"input": 15}) self.assertEqual(r["result"]["output"], 30) + def test_engine_custom_handler(self): + engine = zen.ZenEngine({ "loader": loader, "customHandler": custom_handler }) + r1 = engine.evaluate("custom.json", {"a": 10}) + r2 = engine.evaluate("custom.json", {"a": 20}) + r3 = engine.evaluate("custom.json", {"a": 30}) + + self.assertEqual(r1["result"]["sum"], 20) + self.assertEqual(r2["result"]["sum"], 30) + self.assertEqual(r3["result"]["sum"], 40) + def test_evaluate_expression(self): result = zen.evaluate_expression("sum(a)", { "a": [1, 2, 3, 4] }) self.assertEqual(result, 10) @@ -49,5 +65,9 @@ class ZenEngine(unittest.TestCase): result = zen.evaluate_unary_expression("'FR', 'ES', 'GB'", { "$": "GB" }) self.assertEqual(result, True) + def test_render_template(self): + result = zen.render_template("{{ a + b }}", { "a": 10, "b": 20 }) + self.assertEqual(result, 30) + # run the test unittest.main() diff --git a/bindings/python/src/custom_node.rs b/bindings/python/src/custom_node.rs new file mode 100644 index 00000000..64bc41ee --- /dev/null +++ b/bindings/python/src/custom_node.rs @@ -0,0 +1,42 @@ +use anyhow::anyhow; +use pyo3::types::PyDict; +use pyo3::{PyObject, Python}; +use pythonize::depythonize; + +use zen_engine::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; +use zen_engine::handler::node::{NodeResponse, NodeResult}; + +use crate::types::PyNodeRequest; + +#[derive(Default)] +pub(crate) struct PyCustomNode(Option); + +impl From for PyCustomNode { + fn from(value: PyObject) -> Self { + Self(Some(value)) + } +} + +impl From> for PyCustomNode { + fn from(value: Option) -> Self { + Self(value) + } +} + +impl CustomNodeAdapter for PyCustomNode { + async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + let Some(callable) = &self.0 else { + return Err(anyhow!("Custom node handler not provided")); + }; + + let content: NodeResponse = Python::with_gil(|py| { + let req = PyNodeRequest::from_request(py, request)?; + let result = callable.call1(py, (req,))?; + + let dict = result.extract::<&PyDict>(py)?; + depythonize(dict) + })?; + + Ok(content) + } +} diff --git a/bindings/python/src/decision.rs b/bindings/python/src/decision.rs index 0687d5d6..900c736f 100644 --- a/bindings/python/src/decision.rs +++ b/bindings/python/src/decision.rs @@ -1,19 +1,23 @@ -use crate::engine::PyZenEvaluateOptions; -use crate::loader::PyDecisionLoader; -use crate::value::PyValue; +use std::sync::Arc; + use anyhow::{anyhow, Context}; use pyo3::types::PyDict; use pyo3::{pyclass, pymethods, PyObject, PyResult, Python, ToPyObject}; use pythonize::depythonize; -use std::sync::Arc; + use zen_engine::{Decision, EvaluationOptions}; +use crate::custom_node::PyCustomNode; +use crate::engine::PyZenEvaluateOptions; +use crate::loader::PyDecisionLoader; +use crate::value::PyValue; + #[pyclass] #[pyo3(name = "ZenDecision")] -pub struct PyZenDecision(pub(crate) Arc>); +pub struct PyZenDecision(pub(crate) Arc>); -impl From> for PyZenDecision { - fn from(value: Decision) -> Self { +impl From> for PyZenDecision { + fn from(value: Decision) -> Self { Self(value.into()) } } diff --git a/bindings/python/src/engine.rs b/bindings/python/src/engine.rs index d0e2b5cd..338b7a2e 100644 --- a/bindings/python/src/engine.rs +++ b/bindings/python/src/engine.rs @@ -1,3 +1,4 @@ +use crate::custom_node::PyCustomNode; use crate::decision::PyZenDecision; use crate::loader::PyDecisionLoader; use crate::value::PyValue; @@ -13,7 +14,7 @@ use zen_engine::{DecisionEngine, EvaluationOptions}; #[pyclass] #[pyo3(name = "ZenEngine")] pub struct PyZenEngine { - graph: Arc>, + graph: Arc>, } #[derive(Serialize, Deserialize)] @@ -34,7 +35,11 @@ impl Default for PyZenEvaluateOptions { impl Default for PyZenEngine { fn default() -> Self { Self { - graph: DecisionEngine::new(PyDecisionLoader::default()).into(), + graph: DecisionEngine::new( + Arc::new(PyDecisionLoader::default()), + Arc::new(PyCustomNode::default()), + ) + .into(), } } } @@ -47,13 +52,22 @@ impl PyZenEngine { return Ok(Default::default()); }; - let Some(loader_any) = options.get_item("loader")? else { - return Ok(Default::default()); + let loader = match options.get_item("loader")? { + Some(loader) => Some(Python::with_gil(|py| loader.to_object(py))), + None => None, + }; + + let custom_node = match options.get_item("customHandler")? { + Some(custom_node) => Some(Python::with_gil(|py| custom_node.to_object(py))), + None => None, }; - let loader = Python::with_gil(|py| loader_any.to_object(py)); Ok(Self { - graph: DecisionEngine::new(PyDecisionLoader::from(loader)).into(), + graph: DecisionEngine::new( + Arc::new(PyDecisionLoader::from(loader)), + Arc::new(PyCustomNode::from(custom_node)), + ) + .into(), }) } diff --git a/bindings/python/src/expression.rs b/bindings/python/src/expression.rs index c93d0b77..c8c53ecb 100644 --- a/bindings/python/src/expression.rs +++ b/bindings/python/src/expression.rs @@ -33,3 +33,13 @@ pub fn evaluate_unary_expression(expression: String, ctx: &PyDict) -> PyResult PyResult { + let context: Value = depythonize(ctx).context("Failed to convert context")?; + + let result = zen_template::render(template.as_str(), &context) + .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?; + + Ok(PyValue(result).to_object(py)) +} diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 20798e47..29b377a4 100644 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -1,13 +1,15 @@ use crate::decision::PyZenDecision; use crate::engine::PyZenEngine; -use crate::expression::{evaluate_expression, evaluate_unary_expression}; +use crate::expression::{evaluate_expression, evaluate_unary_expression, render_template}; use pyo3::types::PyModule; use pyo3::{pymodule, wrap_pyfunction, PyResult, Python}; +mod custom_node; mod decision; mod engine; mod expression; mod loader; +mod types; mod value; #[pymodule] @@ -16,6 +18,7 @@ fn zen(_py: Python, m: &PyModule) -> PyResult<()> { m.add_class::()?; m.add_function(wrap_pyfunction!(evaluate_expression, m)?)?; m.add_function(wrap_pyfunction!(evaluate_unary_expression, m)?)?; + m.add_function(wrap_pyfunction!(render_template, m)?)?; Ok(()) } diff --git a/bindings/python/src/loader.rs b/bindings/python/src/loader.rs index 31f2b838..3707502e 100644 --- a/bindings/python/src/loader.rs +++ b/bindings/python/src/loader.rs @@ -14,6 +14,12 @@ impl From for PyDecisionLoader { } } +impl From> for PyDecisionLoader { + fn from(value: Option) -> Self { + Self(value) + } +} + impl PyDecisionLoader { fn load_element(&self, key: &str) -> Result, anyhow::Error> { let Some(object) = &self.0 else { diff --git a/bindings/python/src/types.rs b/bindings/python/src/types.rs new file mode 100644 index 00000000..42a6852c --- /dev/null +++ b/bindings/python/src/types.rs @@ -0,0 +1,118 @@ +use anyhow::{anyhow, Context}; +use json_dotpath::DotPaths; +use pyo3::{pyclass, pymethods, PyObject, PyResult, Python, ToPyObject}; +use serde::Serialize; +use serde_json::Value; + +use zen_engine::handler::custom_node_adapter::{ + CustomDecisionNode as BaseCustomDecisionNode, CustomNodeRequest, +}; +use zen_engine::handler::node::NodeResponse; + +use crate::value::{value_to_object, PyValue}; + +#[derive(Serialize)] +struct CustomDecisionNode { + pub id: String, + pub name: String, + pub kind: String, + pub config: Value, +} + +impl From> for CustomDecisionNode { + fn from(value: BaseCustomDecisionNode) -> Self { + Self { + id: value.id.to_string(), + name: value.name.to_string(), + kind: value.kind.to_string(), + config: value.config.clone(), + } + } +} + +#[pyclass] +pub struct PyNodeRequest { + inner_node: CustomDecisionNode, + inner_input: Value, + + #[pyo3(get)] + pub input: PyObject, + #[pyo3(get)] + pub node: PyObject, +} + +impl PyNodeRequest { + pub fn from_request( + py: Python, + value: CustomNodeRequest<'_>, + ) -> pythonize::Result { + let inner_node = value.node.into(); + let node_val = serde_json::to_value(&inner_node).unwrap(); + + Ok(Self { + input: value_to_object(py, &value.input), + node: value_to_object(py, &node_val), + + inner_input: value.input.clone(), + inner_node, + }) + } +} + +#[pymethods] +impl PyNodeRequest { + fn get_field(&self, py: Python, path: String) -> PyResult { + let node_config = &self.inner_node.config; + + let selected_value: Value = node_config + .dot_get(path.as_str()) + .ok() + .flatten() + .context("Failed to find JSON path")?; + let Value::String(template) = selected_value else { + return Ok(PyValue(selected_value).to_object(py)); + }; + + let template_value = zen_template::render(template.as_str(), &self.inner_input) + .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?; + + Ok(PyValue(template_value).to_object(py)) + } + + fn get_field_raw(&self, py: Python, path: String) -> PyResult { + let node_config = &self.inner_node.config; + + let selected_value: Value = node_config + .dot_get(path.as_str()) + .ok() + .flatten() + .context("Failed to find JSON path")?; + + Ok(PyValue(selected_value).to_object(py)) + } +} + +#[pyclass] +#[derive(Clone)] +pub struct PyNodeResponse { + pub output: Value, + pub trace_data: Option, +} + +impl From for PyNodeResponse { + fn from(value: NodeResponse) -> Self { + Self { + output: value.output, + trace_data: value.trace_data, + } + } +} + +impl From for NodeResponse { + fn from(value: PyNodeResponse) -> Self { + Self { + output: value.output, + trace_data: value.trace_data, + } + } +} diff --git a/bindings/python/src/value.rs b/bindings/python/src/value.rs index 0eeeec13..fca04171 100644 --- a/bindings/python/src/value.rs +++ b/bindings/python/src/value.rs @@ -6,7 +6,7 @@ use std::collections::HashMap; #[derive(Clone, Debug)] pub struct PyValue(pub Value); -fn value_to_object(val: &Value, py: Python<'_>) -> PyObject { +pub fn value_to_object(py: Python<'_>, val: &Value) -> PyObject { match val { Value::Null => py.None(), Value::Bool(b) => b.to_object(py), @@ -18,11 +18,11 @@ fn value_to_object(val: &Value, py: Python<'_>) -> PyObject { } Value::String(s) => s.to_object(py), Value::Array(v) => { - let inner: Vec<_> = v.iter().map(|x| value_to_object(x, py)).collect(); + let inner: Vec<_> = v.iter().map(|x| value_to_object(py, x)).collect(); inner.to_object(py) } Value::Object(m) => { - let inner: HashMap<_, _> = m.iter().map(|(k, v)| (k, value_to_object(v, py))).collect(); + let inner: HashMap<_, _> = m.iter().map(|(k, v)| (k, value_to_object(py, v))).collect(); inner.to_object(py) } } @@ -30,6 +30,6 @@ fn value_to_object(val: &Value, py: Python<'_>) -> PyObject { impl ToPyObject for PyValue { fn to_object(&self, py: Python<'_>) -> PyObject { - value_to_object(&self.0, py) + value_to_object(py, &self.0) } } diff --git a/core/engine/Cargo.toml b/core/engine/Cargo.toml index d77beef2..fcb3b725 100644 --- a/core/engine/Cargo.toml +++ b/core/engine/Cargo.toml @@ -20,11 +20,13 @@ petgraph = { workspace = true } serde_json = { workspace = true, features = ["arbitrary_precision"] } serde = { workspace = true, features = ["derive"] } once_cell = { workspace = true } +json_dotpath = { workspace = true } fixedbitset = "0.4.2" futures = { workspace = true } rquickjs = { version = "0.4.3", features = ["macro", "loader", "rust-alloc"] } itertools = "0.12.1" zen-expression = { path = "../expression", version = "0.19.0" } +zen-template = { path = "../template", version = "0.19.0" } [dev-dependencies] tokio = { version = "1.35.1", features = ["rt", "macros"] } diff --git a/core/engine/benches/engine.rs b/core/engine/benches/engine.rs index 9f2144da..24cde5a6 100644 --- a/core/engine/benches/engine.rs +++ b/core/engine/benches/engine.rs @@ -3,10 +3,12 @@ use criterion::{criterion_group, criterion_main, Bencher, Criterion}; use futures::executor::block_on; use serde_json::{json, Value}; use std::path::Path; +use std::sync::Arc; +use zen_engine::handler::custom_node_adapter::NoopCustomNode; use zen_engine::loader::{FilesystemLoader, FilesystemLoaderOptions}; use zen_engine::DecisionEngine; -fn create_graph() -> DecisionEngine { +fn create_graph() -> DecisionEngine { let cargo_root = Path::new(env!("CARGO_MANIFEST_DIR")); let loader = FilesystemLoader::new(FilesystemLoaderOptions { keep_in_memory: true, @@ -17,7 +19,7 @@ fn create_graph() -> DecisionEngine { .unwrap(), }); - DecisionEngine::new(loader) + DecisionEngine::new(Arc::new(loader), Arc::new(NoopCustomNode::default())) } fn bench_decision(b: &mut Bencher, key: &str, context: Value) { diff --git a/core/engine/src/decision.rs b/core/engine/src/decision.rs index 58899456..c8ea9d92 100644 --- a/core/engine/src/decision.rs +++ b/core/engine/src/decision.rs @@ -1,49 +1,69 @@ +use std::sync::Arc; + +use serde_json::Value; + use crate::engine::EvaluationOptions; +use crate::handler::custom_node_adapter::{CustomNodeAdapter, NoopCustomNode}; use crate::handler::graph::{DecisionGraph, DecisionGraphConfig, DecisionGraphResponse}; use crate::loader::{DecisionLoader, NoopLoader}; use crate::model::DecisionContent; use crate::{DecisionGraphValidationError, EvaluationError}; -use serde_json::Value; -use std::sync::Arc; /// Represents a JDM decision which can be evaluated #[derive(Debug, Clone)] -pub struct Decision +pub struct Decision where - L: DecisionLoader, + Loader: DecisionLoader, + CustomNode: CustomNodeAdapter, { content: Arc, - loader: Arc, + loader: Arc, + adapter: Arc, } -impl From for Decision { +impl From for Decision { fn from(value: DecisionContent) -> Self { Self { content: value.into(), loader: NoopLoader::default().into(), + adapter: NoopCustomNode::default().into(), } } } -impl From> for Decision { +impl From> for Decision { fn from(value: Arc) -> Self { Self { content: value, loader: NoopLoader::default().into(), + adapter: NoopCustomNode::default().into(), } } } -impl Decision +impl Decision where L: DecisionLoader, + A: CustomNodeAdapter, { - pub fn with_loader(self, loader: Arc) -> Decision + pub fn with_loader(self, loader: Arc) -> Decision where - NL: DecisionLoader, + Loader: DecisionLoader, { Decision { loader, + adapter: self.adapter, + content: self.content, + } + } + + pub fn with_adapter(self, adapter: Arc) -> Decision + where + Adapter: CustomNodeAdapter, + { + Decision { + loader: self.loader, + adapter, content: self.content, } } @@ -66,6 +86,7 @@ where max_depth: options.max_depth.unwrap_or(5), trace: options.trace.unwrap_or_default(), loader: self.loader.clone(), + adapter: self.adapter.clone(), iteration: 0, content: &self.content, })?; @@ -78,6 +99,7 @@ where max_depth: 1, trace: false, loader: self.loader.clone(), + adapter: self.adapter.clone(), iteration: 0, content: &self.content, })?; diff --git a/core/engine/src/engine.rs b/core/engine/src/engine.rs index 90e30b78..216d55f5 100644 --- a/core/engine/src/engine.rs +++ b/core/engine/src/engine.rs @@ -1,21 +1,24 @@ -use crate::decision::Decision; -use crate::loader::{ClosureLoader, DecisionLoader, LoaderResponse, LoaderResult, NoopLoader}; -use crate::model::DecisionContent; +use std::future::Future; +use std::sync::Arc; use serde_json::Value; -use std::future::Future; +use crate::decision::Decision; +use crate::handler::custom_node_adapter::{CustomNodeAdapter, NoopCustomNode}; use crate::handler::graph::DecisionGraphResponse; +use crate::loader::{ClosureLoader, DecisionLoader, LoaderResponse, LoaderResult, NoopLoader}; +use crate::model::DecisionContent; use crate::EvaluationError; -use std::sync::Arc; /// Structure used for generating and evaluating JDM decisions #[derive(Debug, Clone)] -pub struct DecisionEngine +pub struct DecisionEngine where - L: DecisionLoader, + Loader: DecisionLoader, + CustomNode: CustomNodeAdapter, { - loader: Arc, + loader: Arc, + adapter: Arc, } #[derive(Debug, Default)] @@ -24,38 +27,49 @@ pub struct EvaluationOptions { pub max_depth: Option, } -impl Default for DecisionEngine { +impl Default for DecisionEngine { fn default() -> Self { Self { loader: Arc::new(NoopLoader::default()), + adapter: Arc::new(NoopCustomNode::default()), } } } -impl DecisionEngine> -where - F: Fn(String) -> O + Sync + Send, - O: Future + Send, -{ - pub fn async_loader(loader: F) -> Self { - Self { - loader: Arc::new(ClosureLoader::new(loader)), - } +impl DecisionEngine { + pub fn new(loader: Arc, adapter: Arc) -> Self { + Self { loader, adapter } } -} -impl DecisionEngine { - pub fn new(loader: Loader) -> Self + pub fn with_adapter(self, adapter: Arc) -> DecisionEngine where - Loader: Into>, + CustomNode: CustomNodeAdapter, { - Self { - loader: loader.into(), + DecisionEngine { + loader: self.loader, + adapter, } } - pub fn new_arc(loader: Arc) -> Self { - Self { loader } + pub fn with_loader(self, loader: Arc) -> DecisionEngine + where + Loader: DecisionLoader, + { + DecisionEngine { + loader, + adapter: self.adapter, + } + } + + pub fn with_closure_loader(self, loader: F) -> DecisionEngine, A> + where + F: Fn(String) -> O + Sync + Send, + O: Future + Send, + { + DecisionEngine { + loader: Arc::new(ClosureLoader::new(loader)), + adapter: self.adapter, + } } /// Evaluates a decision through loader using a key @@ -87,12 +101,14 @@ impl DecisionEngine { } /// Creates a decision from DecisionContent, exists for easier binding creation - pub fn create_decision(&self, content: Arc) -> Decision { - Decision::from(content).with_loader(self.loader.clone()) + pub fn create_decision(&self, content: Arc) -> Decision { + Decision::from(content) + .with_loader(self.loader.clone()) + .with_adapter(self.adapter.clone()) } /// Retrieves a decision based on the loader - pub async fn get_decision(&self, key: &str) -> LoaderResult> { + pub async fn get_decision(&self, key: &str) -> LoaderResult> { let content = self.loader.load(key).await?; Ok(self.create_decision(content)) } diff --git a/core/engine/src/handler/custom_node_adapter.rs b/core/engine/src/handler/custom_node_adapter.rs new file mode 100644 index 00000000..3279c63a --- /dev/null +++ b/core/engine/src/handler/custom_node_adapter.rs @@ -0,0 +1,86 @@ +use crate::handler::node::{NodeRequest, NodeResult}; +use crate::model::{DecisionNode, DecisionNodeKind}; +use anyhow::anyhow; +use json_dotpath::DotPaths; +use serde::Serialize; +use serde_json::Value; +use zen_template::TemplateRenderError; + +pub trait CustomNodeAdapter { + fn handle( + &self, + request: CustomNodeRequest<'_>, + ) -> impl std::future::Future + Send; +} + +#[derive(Default, Debug)] +pub struct NoopCustomNode; + +impl CustomNodeAdapter for NoopCustomNode { + async fn handle(&self, _: CustomNodeRequest<'_>) -> NodeResult { + Err(anyhow!("Custom node handler not provided")) + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +pub struct CustomNodeRequest<'a> { + pub input: &'a Value, + pub node: CustomDecisionNode<'a>, +} + +impl<'a> TryFrom<&'a NodeRequest<'a>> for CustomNodeRequest<'a> { + type Error = (); + + fn try_from(value: &'a NodeRequest<'a>) -> Result { + Ok(Self { + input: &value.input, + node: value.node.try_into()?, + }) + } +} + +impl<'a> CustomNodeRequest<'a> { + pub fn get_field(&self, path: &str) -> Result, TemplateRenderError> { + let Some(selected_value) = self.get_field_raw(path) else { + return Ok(None); + }; + + let Value::String(template) = selected_value else { + return Ok(Some(selected_value)); + }; + + let template_value = zen_template::render(template.as_str(), &self.input)?; + Ok(Some(template_value)) + } + + fn get_field_raw(&self, path: &str) -> Option { + self.node.config.dot_get(path).ok().flatten() + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +pub struct CustomDecisionNode<'a> { + pub id: &'a str, + pub name: &'a str, + pub kind: &'a str, + pub config: &'a Value, +} + +impl<'a> TryFrom<&'a DecisionNode> for CustomDecisionNode<'a> { + type Error = (); + + fn try_from(value: &'a DecisionNode) -> Result { + let DecisionNodeKind::CustomNode { content } = &value.kind else { + return Err(()); + }; + + Ok(Self { + id: &value.id, + name: &value.name, + kind: &content.kind, + config: &content.config, + }) + } +} diff --git a/core/engine/src/handler/decision.rs b/core/engine/src/handler/decision.rs index 6037c263..1aa760b0 100644 --- a/core/engine/src/handler/decision.rs +++ b/core/engine/src/handler/decision.rs @@ -1,3 +1,4 @@ +use crate::handler::custom_node_adapter::CustomNodeAdapter; use crate::handler::graph::{DecisionGraph, DecisionGraphConfig}; use crate::handler::node::{NodeRequest, NodeResponse, NodeResult}; use crate::loader::DecisionLoader; @@ -8,18 +9,26 @@ use rquickjs::Runtime; use std::ops::Deref; use std::sync::Arc; -pub struct DecisionHandler { +pub struct DecisionHandler { trace: bool, - loader: Arc, + loader: Arc, + adapter: Arc, max_depth: u8, js_runtime: Option, } -impl DecisionHandler { - pub fn new(trace: bool, max_depth: u8, loader: Arc, js_runtime: Option) -> Self { +impl DecisionHandler { + pub fn new( + trace: bool, + max_depth: u8, + loader: Arc, + adapter: Arc, + js_runtime: Option, + ) -> Self { Self { trace, loader, + adapter, max_depth, js_runtime, } @@ -37,6 +46,7 @@ impl DecisionHandler { content: sub_decision.deref(), max_depth: self.max_depth, loader: self.loader.clone(), + adapter: self.adapter.clone(), iteration: request.iteration + 1, trace: self.trace, })? diff --git a/core/engine/src/handler/graph.rs b/core/engine/src/handler/graph.rs index 90b13a91..06affd49 100644 --- a/core/engine/src/handler/graph.rs +++ b/core/engine/src/handler/graph.rs @@ -10,6 +10,7 @@ use serde::{Deserialize, Serialize, Serializer}; use serde_json::Value; use thiserror::Error; +use crate::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest}; use crate::handler::decision::DecisionHandler; use crate::handler::expression::ExpressionHandler; use crate::handler::function::runtime::create_runtime; @@ -21,8 +22,9 @@ use crate::loader::DecisionLoader; use crate::model::{DecisionContent, DecisionNodeKind}; use crate::{EvaluationError, NodeError}; -pub struct DecisionGraph<'a, L: DecisionLoader> { +pub struct DecisionGraph<'a, L: DecisionLoader, A: CustomNodeAdapter> { graph: StableDiDecisionGraph<'a>, + adapter: Arc, loader: Arc, trace: bool, max_depth: u8, @@ -30,17 +32,18 @@ pub struct DecisionGraph<'a, L: DecisionLoader> { runtime: Option, } -pub struct DecisionGraphConfig<'a, T: DecisionLoader> { - pub loader: Arc, +pub struct DecisionGraphConfig<'a, L: DecisionLoader, A: CustomNodeAdapter> { + pub loader: Arc, + pub adapter: Arc, pub content: &'a DecisionContent, pub trace: bool, pub iteration: u8, pub max_depth: u8, } -impl<'a, L: DecisionLoader> DecisionGraph<'a, L> { +impl<'a, L: DecisionLoader, A: CustomNodeAdapter> DecisionGraph<'a, L, A> { pub fn try_new( - config: DecisionGraphConfig<'a, L>, + config: DecisionGraphConfig<'a, L, A>, ) -> Result { let content = config.content; let mut graph = StableDiDecisionGraph::new(); @@ -69,7 +72,8 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> { graph, iteration: config.iteration, trace: config.trace, - loader: config.loader.clone(), + loader: config.loader, + adapter: config.adapter, max_depth: config.max_depth, runtime: None, }) @@ -154,13 +158,13 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> { }; } - let node_request = NodeRequest { + let mut node_request = NodeRequest { node, iteration: self.iteration, - input: walker.incoming_node_data(&self.graph, nid), + input: walker.incoming_node_data(&self.graph, nid, true), }; - match node.kind { + match &node.kind { DecisionNodeKind::InputNode => { walker.set_node_data(nid, context.clone()); trace!({ @@ -181,6 +185,9 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> { performance: None, trace_data: None, }); + if let Some(obj) = node_request.input.as_object_mut() { + obj.remove("$nodes"); + } return Ok(DecisionGraphResponse { result: node_request.input, @@ -228,6 +235,7 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> { self.trace, self.max_depth, self.loader.clone(), + self.adapter.clone(), self.runtime.clone(), ) .handle(&node_request) @@ -285,6 +293,26 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> { trace_data: res.trace_data, }); } + DecisionNodeKind::CustomNode { .. } => { + let res = self + .adapter + .handle(CustomNodeRequest::try_from(&node_request).unwrap()) + .await + .map_err(|e| NodeError { + node_id: node.id.clone(), + source: e.into(), + })?; + + walker.set_node_data(nid, res.output.clone()); + trace!({ + input: node_request.input, + output: res.output, + name: node.name.clone(), + id: node.id.clone(), + performance: Some(format!("{:?}", start.elapsed())), + trace_data: res.trace_data, + }); + } } } @@ -351,10 +379,10 @@ pub struct DecisionGraphResponse { #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct DecisionGraphTrace { - input: Value, - output: Value, - name: String, - id: String, - performance: Option, - trace_data: Option, + pub input: Value, + pub output: Value, + pub name: String, + pub id: String, + pub performance: Option, + pub trace_data: Option, } diff --git a/core/engine/src/handler/mod.rs b/core/engine/src/handler/mod.rs index c41450c5..22596d91 100644 --- a/core/engine/src/handler/mod.rs +++ b/core/engine/src/handler/mod.rs @@ -3,6 +3,7 @@ pub mod expression; pub mod function; pub mod table; -pub(crate) mod graph; -pub(crate) mod node; +pub mod custom_node_adapter; +pub mod graph; +pub mod node; pub(crate) mod traversal; diff --git a/core/engine/src/handler/node.rs b/core/engine/src/handler/node.rs index 75205768..3ef864d5 100644 --- a/core/engine/src/handler/node.rs +++ b/core/engine/src/handler/node.rs @@ -11,7 +11,7 @@ pub struct NodeResponse { pub trace_data: Option, } -#[derive(Debug)] +#[derive(Debug, Serialize)] pub struct NodeRequest<'a> { pub input: Value, pub iteration: u8, diff --git a/core/engine/src/handler/traversal.rs b/core/engine/src/handler/traversal.rs index 791a0175..a8c027de 100644 --- a/core/engine/src/handler/traversal.rs +++ b/core/engine/src/handler/traversal.rs @@ -62,12 +62,36 @@ impl GraphWalker { self.node_data.get(&node_id) } + pub fn get_all_node_data(&self, g: &StableDiDecisionGraph) -> Value { + let node_values: Map = self + .node_data + .iter() + .map(|(idx, value)| { + let weight = g.node_weight(*idx).unwrap(); + (weight.name.clone(), value.clone()) + }) + .collect(); + + Value::Object(node_values) + } + pub fn set_node_data(&mut self, node_id: NodeIndex, value: Value) { self.node_data.insert(node_id, value); } - pub fn incoming_node_data(&self, g: &StableDiDecisionGraph, node_id: NodeIndex) -> Value { - self.merge_node_data(g.neighbors_directed(node_id, Incoming)) + pub fn incoming_node_data( + &self, + g: &StableDiDecisionGraph, + node_id: NodeIndex, + with_nodes: bool, + ) -> Value { + let mut value = self.merge_node_data(g.neighbors_directed(node_id, Incoming)); + + if let Some(object) = with_nodes.then_some(value.as_object_mut()).flatten() { + object.insert("$nodes".to_string(), self.get_all_node_data(g)); + } + + value } pub fn merge_node_data(&self, iter: I) -> Value @@ -98,7 +122,7 @@ impl GraphWalker { self.ordered.visit(nid); if let DecisionNodeKind::SwitchNode { content } = &decision_node.kind { - let mut input_data = self.incoming_node_data(g, nid); + let mut input_data = self.incoming_node_data(g, nid, true); let input_context = json!({ "$": &input_data }); merge_json(&mut input_data, &input_context, true); diff --git a/core/engine/src/lib.rs b/core/engine/src/lib.rs index c68be552..797335d9 100644 --- a/core/engine/src/lib.rs +++ b/core/engine/src/lib.rs @@ -125,7 +125,7 @@ mod decision; mod engine; mod error; -mod handler; +pub mod handler; mod util; pub mod loader; diff --git a/core/engine/src/model/mod.rs b/core/engine/src/model/mod.rs index a610f48d..1bd968ae 100644 --- a/core/engine/src/model/mod.rs +++ b/core/engine/src/model/mod.rs @@ -1,12 +1,11 @@ -use serde::{Deserialize, Serialize}; use std::collections::HashMap; -#[cfg(feature = "bincode")] -use bincode::{Decode, Encode}; +use serde::{Deserialize, Serialize}; +use serde_json::Value; /// JDM Decision model #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionContent { pub nodes: Vec, @@ -14,7 +13,7 @@ pub struct DecisionContent { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionEdge { pub id: String, @@ -24,7 +23,7 @@ pub struct DecisionEdge { } #[derive(Clone, Debug, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionNode { pub id: String, @@ -41,7 +40,7 @@ impl PartialEq for DecisionNode { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(tag = "type")] #[serde(rename_all = "camelCase")] pub enum DecisionNodeKind { @@ -52,17 +51,18 @@ pub enum DecisionNodeKind { DecisionTableNode { content: DecisionTableContent }, ExpressionNode { content: ExpressionNodeContent }, SwitchNode { content: SwitchNodeContent }, + CustomNode { content: CustomNodeContent }, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionNodeContent { pub key: String, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionTableContent { pub rules: Vec>, @@ -72,7 +72,7 @@ pub struct DecisionTableContent { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub enum DecisionTableHitPolicy { First, @@ -80,7 +80,7 @@ pub enum DecisionTableHitPolicy { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionTableInputField { pub id: String, @@ -89,7 +89,7 @@ pub struct DecisionTableInputField { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionTableOutputField { pub id: String, @@ -98,14 +98,14 @@ pub struct DecisionTableOutputField { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct ExpressionNodeContent { pub expressions: Vec, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct Expression { pub id: String, @@ -114,7 +114,7 @@ pub struct Expression { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct SwitchNodeContent { #[serde(default)] @@ -123,7 +123,7 @@ pub struct SwitchNodeContent { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct SwitchStatement { pub id: String, @@ -131,10 +131,61 @@ pub struct SwitchStatement { } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Default)] -#[cfg_attr(feature = "bincode", derive(Encode, Decode))] +#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub enum SwitchStatementHitPolicy { #[default] First, Collect, } + +#[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Default)] +#[serde(rename_all = "camelCase")] +pub struct CustomNodeContent { + pub kind: String, + pub config: Value, +} + +#[cfg(feature = "bincode")] +impl ::bincode::Encode for CustomNodeContent { + fn encode<__E: ::bincode::enc::Encoder>( + &self, + encoder: &mut __E, + ) -> Result<(), ::bincode::error::EncodeError> { + let config_string = self.config.to_string(); + + ::bincode::Encode::encode(&self.kind, encoder)?; + ::bincode::Encode::encode(config_string.as_bytes(), encoder)?; + Ok(()) + } +} + +#[cfg(feature = "bincode")] +impl ::bincode::Decode for CustomNodeContent { + fn decode<__D: ::bincode::de::Decoder>( + decoder: &mut __D, + ) -> Result { + let kind: String = ::bincode::Decode::decode(decoder)?; + let config_string: String = ::bincode::Decode::decode(decoder)?; + + let config = serde_json::from_str(config_string.as_str()) + .map_err(|_| ::bincode::error::DecodeError::Other("failed to deserialize value"))?; + + Ok(Self { kind, config }) + } +} + +#[cfg(feature = "bincode")] +impl<'__de> ::bincode::BorrowDecode<'__de> for CustomNodeContent { + fn borrow_decode<__D: ::bincode::de::BorrowDecoder<'__de>>( + decoder: &mut __D, + ) -> Result { + let kind: String = ::bincode::BorrowDecode::borrow_decode(decoder)?; + let config_string: String = ::bincode::BorrowDecode::borrow_decode(decoder)?; + + let config = serde_json::from_str(config_string.as_str()) + .map_err(|_| ::bincode::error::DecodeError::Other("failed to deserialize value"))?; + + Ok(Self { kind, config }) + } +} diff --git a/core/engine/tests/engine.rs b/core/engine/tests/engine.rs index ac84a1f8..aaa5260b 100644 --- a/core/engine/tests/engine.rs +++ b/core/engine/tests/engine.rs @@ -22,7 +22,7 @@ async fn engine_memory_loader() { memory_loader.add("table", load_test_data("table.json")); memory_loader.add("function", load_test_data("function.json")); - let engine = DecisionEngine::new_arc(memory_loader.clone()); + let engine = DecisionEngine::default().with_loader(memory_loader.clone()); let table = engine.evaluate("table", &json!({ "input": 12 })).await; let function = engine.evaluate("function", &json!({ "input": 12 })).await; @@ -37,7 +37,7 @@ async fn engine_memory_loader() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_filesystem_loader() { - let engine = DecisionEngine::new(create_fs_loader()); + let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); let table = engine.evaluate("table.json", &json!({ "input": 12 })).await; let function = engine .evaluate("function.json", &json!({ "input": 12 })) @@ -52,7 +52,7 @@ async fn engine_filesystem_loader() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_closure_loader() { - let engine = DecisionEngine::async_loader(|key| async { + let engine = DecisionEngine::default().with_closure_loader(|key| async { match key.as_str() { "function" => Ok(Arc::new(load_test_data("function.json"))), "table" => Ok(Arc::new(load_test_data("table.json"))), @@ -80,7 +80,7 @@ async fn engine_noop_loader() { #[tokio::test] async fn engine_get_decision() { - let engine = DecisionEngine::new(create_fs_loader()); + let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); assert!(engine.get_decision("table.json").await.is_ok()); assert!(engine.get_decision("any.json").await.is_err()); @@ -95,7 +95,7 @@ async fn engine_create_decision() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_errors() { - let engine = DecisionEngine::new(create_fs_loader()); + let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); let infinite_fn = engine.evaluate("infinite-function.json", &json!({})).await; match infinite_fn.unwrap_err().deref() { @@ -117,7 +117,7 @@ async fn engine_errors() { #[tokio::test] async fn engine_with_trace() { - let engine = DecisionEngine::new(create_fs_loader()); + let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); let table_r = engine.evaluate("table.json", &json!({ "input": 12 })).await; let table_opt_r = engine @@ -179,7 +179,7 @@ async fn engine_function_imports() { #[tokio::test] async fn engine_switch_node() { - let engine = DecisionEngine::new(create_fs_loader()); + let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); let switch_node_r = engine .evaluate("switch-node.json", &json!({ "color": "yellow" })) diff --git a/core/expression/src/compiler/compiler.rs b/core/expression/src/compiler/compiler.rs index cf732daf..adf339a7 100644 --- a/core/expression/src/compiler/compiler.rs +++ b/core/expression/src/compiler/compiler.rs @@ -112,6 +112,7 @@ impl<'arena, 'bytecode_ref> CompilerInner<'arena, 'bytecode_ref> { Node::Number(v) => Ok(self.emit(Opcode::Push(Variable::Number(*v)))), Node::String(v) => Ok(self.emit(Opcode::Push(Variable::String(v)))), Node::Pointer => Ok(self.emit(Opcode::Pointer)), + Node::Root => Ok(self.emit(Opcode::FetchRootEnv)), Node::Array(v) => { v.iter() .try_for_each(|&n| self.compile_node(n).map(|_| ()))?; diff --git a/core/expression/src/compiler/opcode.rs b/core/expression/src/compiler/opcode.rs index 5654da20..35a6dc7e 100644 --- a/core/expression/src/compiler/opcode.rs +++ b/core/expression/src/compiler/opcode.rs @@ -8,6 +8,7 @@ pub enum Opcode<'a> { Pop, Rot, Fetch, + FetchRootEnv, FetchEnv(&'a str), Negate, Not, diff --git a/core/expression/src/lexer/token.rs b/core/expression/src/lexer/token.rs index f8cf5636..93e7e669 100644 --- a/core/expression/src/lexer/token.rs +++ b/core/expression/src/lexer/token.rs @@ -25,6 +25,7 @@ pub enum TokenKind { #[derive(Debug, PartialEq, Eq, Clone, Copy, Display)] pub enum Identifier { ContextReference, // $ + RootReference, // $root CallbackReference, // # Null, // null Variable, @@ -34,6 +35,7 @@ impl From<&str> for Identifier { fn from(value: &str) -> Self { match value { "$" => Identifier::ContextReference, + "$root" => Identifier::RootReference, "#" => Identifier::CallbackReference, "null" => Identifier::Null, _ => Identifier::Variable, diff --git a/core/expression/src/parser/ast.rs b/core/expression/src/parser/ast.rs index 0720dc21..36e157a2 100644 --- a/core/expression/src/parser/ast.rs +++ b/core/expression/src/parser/ast.rs @@ -13,6 +13,7 @@ pub enum Node<'a> { Array(&'a [&'a Node<'a>]), Identifier(&'a str), Closure(&'a Node<'a>), + Root, Member { node: &'a Node<'a>, property: &'a Node<'a>, diff --git a/core/expression/src/parser/parser.rs b/core/expression/src/parser/parser.rs index 3306284f..bfc55c7c 100644 --- a/core/expression/src/parser/parser.rs +++ b/core/expression/src/parser/parser.rs @@ -280,7 +280,11 @@ impl<'arena, 'token_ref, Flavor> Parser<'arena, 'token_ref, Flavor> { self.next()?; let current_token = self.current(); if current_token.kind != TokenKind::Bracket(Bracket::LeftParenthesis) { - let identifier_node = self.node(Node::Identifier(identifier_token.value)); + let identifier_node = match identifier_token.kind { + TokenKind::Identifier(Identifier::RootReference) => self.node(Node::Root), + _ => self.node(Node::Identifier(identifier_token.value)), + }; + return self .with_postfix(identifier_node, expression_parser) .map(Some); diff --git a/core/expression/src/parser/unary.rs b/core/expression/src/parser/unary.rs index 8cc62fce..b531b7ae 100644 --- a/core/expression/src/parser/unary.rs +++ b/core/expression/src/parser/unary.rs @@ -214,6 +214,7 @@ impl From<&Node<'_>> for UnaryNodeBehaviour { match value { Node::Null => CompareWithReference(Equal), + Node::Root => CompareWithReference(Equal), Node::Bool(_) => CompareWithReference(Equal), Node::Number(_) => CompareWithReference(Equal), Node::String(_) => CompareWithReference(Equal), diff --git a/core/expression/src/vm/vm.rs b/core/expression/src/vm/vm.rs index ca8da7a5..7f84eec0 100644 --- a/core/expression/src/vm/vm.rs +++ b/core/expression/src/vm/vm.rs @@ -165,6 +165,9 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'arena, 'parent_ref, 'bytecode_ }); } }, + Opcode::FetchRootEnv => { + self.push(env.clone()); + } Opcode::Negate => { let a = self.pop()?; match a { diff --git a/core/template/Cargo.toml b/core/template/Cargo.toml new file mode 100644 index 00000000..04bdda12 --- /dev/null +++ b/core/template/Cargo.toml @@ -0,0 +1,15 @@ +[package] +authors = ["GoRules Team "] +description = "Zen Template Language" +name = "zen-template" +license = "MIT" +version = "0.19.0" +edition = "2021" +repository = "https://github.com/gorules/zen.git" + +[dependencies] +zen-expression = { path = "../expression", version = "0.19.0" } +itertools = "0.12.1" +thiserror = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } \ No newline at end of file diff --git a/core/template/src/error.rs b/core/template/src/error.rs new file mode 100644 index 00000000..0afceb29 --- /dev/null +++ b/core/template/src/error.rs @@ -0,0 +1,64 @@ +use serde::ser::SerializeMap; +use serde::{Serialize, Serializer}; +use thiserror::Error; +use zen_expression::IsolateError; + +#[derive(Debug, Error)] +pub enum TemplateRenderError { + #[error("isolate error: {0}")] + IsolateError(IsolateError), + + #[error("parser error: {0}")] + ParserError(ParserError), +} + +impl From for TemplateRenderError { + fn from(value: IsolateError) -> Self { + Self::IsolateError(value) + } +} + +impl From for TemplateRenderError { + fn from(value: ParserError) -> Self { + Self::ParserError(value) + } +} + +impl Serialize for TemplateRenderError { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + match self { + TemplateRenderError::IsolateError(isolate) => isolate.serialize(serializer), + TemplateRenderError::ParserError(parser) => parser.serialize(serializer), + } + } +} + +#[derive(Debug, Error)] +pub enum ParserError { + #[error("Open bracket")] + OpenBracket, + + #[error("Close bracket")] + CloseBracket, +} + +impl Serialize for ParserError { + fn serialize(&self, serializer: S) -> Result + where + S: Serializer, + { + let mut map = serializer.serialize_map(None)?; + + map.serialize_entry("type", "templateParserError")?; + + match self { + ParserError::OpenBracket => map.serialize_entry("value", "openBracket")?, + ParserError::CloseBracket => map.serialize_entry("value", "closeBracket")?, + } + + map.end() + } +} diff --git a/core/template/src/interpreter.rs b/core/template/src/interpreter.rs new file mode 100644 index 00000000..5cfd27de --- /dev/null +++ b/core/template/src/interpreter.rs @@ -0,0 +1,97 @@ +use std::iter::Peekable; +use std::slice::Iter; + +use crate::error::TemplateRenderError; +use serde_json::Value; +use zen_expression::Isolate; + +use crate::parser::Node; + +#[derive(Debug, PartialEq)] +pub(crate) enum InterpreterResult<'a> { + String(&'a str), + Value(Value), +} + +#[derive(Debug)] +pub(crate) struct Interpreter<'source, 'nodes> { + cursor: Peekable>>, + isolate: Isolate<'source>, + results: Vec>, +} + +impl<'source, 'nodes, T> From for Interpreter<'source, 'nodes> +where + T: Into<&'nodes [Node<'source>]>, +{ + fn from(value: T) -> Self { + let nodes = value.into(); + let cursor = nodes.iter().peekable(); + + Self { + cursor, + isolate: Isolate::new(), + results: Default::default(), + } + } +} + +impl<'source, 'nodes> Interpreter<'source, 'nodes> { + pub(crate) fn collect_for(mut self, context: &Value) -> Result { + self.isolate.set_environment(context); + + while let Some(node) = self.cursor.next() { + match node { + Node::Text(data) => self.text(data), + Node::Expression(data) => self.expression(data)?, + } + } + + match self.results.len() { + 0 => Ok(Value::Null), + 1 => { + let item = self.results.remove(0); + match item { + InterpreterResult::Value(val) => Ok(val), + InterpreterResult::String(str) => Ok(Value::String(str.to_string())), + } + } + _ => { + let string_data = self + .results + .into_iter() + .map(|item| match item { + InterpreterResult::String(str) => str.to_string(), + InterpreterResult::Value(value) => value_to_string(value), + }) + .collect::(); + + Ok(Value::String(string_data)) + } + } + } + + fn text(&mut self, data: &'source str) { + self.results.push(InterpreterResult::String(data)); + } + + fn expression(&mut self, data: &'source str) -> Result<(), TemplateRenderError> { + let result = self.isolate.run_standard(data)?; + self.results.push(InterpreterResult::Value(result)); + Ok(()) + } +} + +fn value_to_string(value: Value) -> String { + match value { + Value::Null => String::from("null"), + Value::Bool(b) => match b { + true => String::from("true"), + false => String::from("false"), + }, + Value::Number(n) => n.to_string(), + Value::String(s) => s, + Value::Array(arr) => Value::Array(arr).to_string(), + Value::Object(obj) => Value::Object(obj).to_string(), + } +} diff --git a/core/template/src/lexer.rs b/core/template/src/lexer.rs new file mode 100644 index 00000000..fe1fbde5 --- /dev/null +++ b/core/template/src/lexer.rs @@ -0,0 +1,62 @@ +use std::iter::{Enumerate, Peekable}; +use std::str::Chars; + +#[derive(Debug, PartialOrd, PartialEq)] +pub(crate) enum Token<'source> { + Text(&'source str), + OpenBracket, + CloseBracket, +} + +pub(crate) struct Lexer<'source> { + cursor: Peekable>>, + source: &'source str, + tokens: Vec>, + text_start: Option, +} + +impl<'source, T> From for Lexer<'source> +where + T: Into<&'source str>, +{ + fn from(value: T) -> Self { + let source: &'source str = value.into(); + + Self { + source, + cursor: source.chars().enumerate().peekable(), + tokens: Default::default(), + text_start: None, + } + } +} + +impl<'source> Lexer<'source> { + pub fn collect(mut self) -> Vec> { + while let Some((index, char)) = self.cursor.next() { + if char == '{' && matches!(self.cursor.peek(), Some((_, '{'))) { + self.flush(index); + + self.cursor.next(); + self.tokens.push(Token::OpenBracket); + } else if char == '}' && matches!(self.cursor.peek(), Some((_, '}'))) { + self.flush(index); + + self.cursor.next(); + self.tokens.push(Token::CloseBracket); + } else { + self.text_start.get_or_insert(index); + } + } + + self.flush(self.source.len()); + self.tokens + } + + fn flush(&mut self, index: usize) { + if let Some(start) = self.text_start { + self.tokens.push(Token::Text(&self.source[start..index])); + self.text_start = None; + } + } +} diff --git a/core/template/src/lib.rs b/core/template/src/lib.rs new file mode 100644 index 00000000..f08dc266 --- /dev/null +++ b/core/template/src/lib.rs @@ -0,0 +1,19 @@ +mod error; +mod interpreter; +mod lexer; +mod parser; + +use serde_json::Value; + +use crate::interpreter::Interpreter; +use crate::lexer::Lexer; +use crate::parser::Parser; + +pub use crate::error::{ParserError, TemplateRenderError}; + +pub fn render(template: &str, context: &Value) -> Result { + let tokens = Lexer::from(template.trim()).collect(); + let nodes = Parser::from(tokens.as_slice()).collect()?; + + Interpreter::from(nodes.as_slice()).collect_for(context) +} diff --git a/core/template/src/parser.rs b/core/template/src/parser.rs new file mode 100644 index 00000000..124a0744 --- /dev/null +++ b/core/template/src/parser.rs @@ -0,0 +1,77 @@ +use crate::error::{ParserError, TemplateRenderError}; +use crate::lexer::Token; +use std::iter::Peekable; +use std::slice::Iter; + +#[derive(Debug, PartialOrd, PartialEq)] +pub(crate) enum Node<'a> { + Text(&'a str), + Expression(&'a str), +} + +#[derive(Debug, PartialOrd, PartialEq)] +enum ParserState { + Text, + Expression, +} + +pub(crate) struct Parser<'source, 'tokens> { + cursor: Peekable>>, + state: ParserState, + nodes: Vec>, +} + +impl<'source, 'tokens, T> From for Parser<'source, 'tokens> +where + T: Into<&'tokens [Token<'source>]>, +{ + fn from(value: T) -> Self { + let tokens = value.into(); + let cursor = tokens.iter().peekable(); + + Self { + cursor, + nodes: Default::default(), + state: ParserState::Text, + } + } +} + +impl<'source, 'tokens> Parser<'source, 'tokens> { + pub(crate) fn collect(mut self) -> Result>, TemplateRenderError> { + while let Some(token) = self.cursor.next() { + match token { + Token::Text(text) => self.text(text), + Token::OpenBracket => self.open_bracket()?, + Token::CloseBracket => self.close_bracket()?, + } + } + + Ok(self.nodes) + } + + fn text(&mut self, data: &'source str) { + match self.state { + ParserState::Text => self.nodes.push(Node::Text(data)), + ParserState::Expression => self.nodes.push(Node::Expression(data)), + } + } + + fn open_bracket(&mut self) -> Result<(), ParserError> { + if self.state == ParserState::Expression { + return Err(ParserError::OpenBracket); + } + + self.state = ParserState::Expression; + Ok(()) + } + + fn close_bracket(&mut self) -> Result<(), ParserError> { + if self.state != ParserState::Expression { + return Err(ParserError::CloseBracket); + } + + self.state = ParserState::Text; + Ok(()) + } +} diff --git a/core/template/tests/template.rs b/core/template/tests/template.rs new file mode 100644 index 00000000..43646c5e --- /dev/null +++ b/core/template/tests/template.rs @@ -0,0 +1,100 @@ +use serde_json::{json, Value}; +use zen_template::render; + +#[test] +fn test_values_types() { + struct TestCase { + template: &'static str, + context: Value, + expected: Value, + } + + let test_cases = vec![ + TestCase { + template: "{{ null }}", + context: json!(null), + expected: json!(null), + }, + TestCase { + template: "{{ 1 + 1 }}", + context: json!(null), + expected: json!(2), + }, + TestCase { + template: "{{ 'hello' }}", + context: json!(null), + expected: json!("hello"), + }, + TestCase { + template: "{{ true or false }}", + context: json!(null), + expected: json!(true), + }, + TestCase { + template: "{{ [1, 2, 3] }}", + context: json!(null), + expected: json!([1, 2, 3]), + }, + TestCase { + template: "{{ customer }}", + context: json!({ "customer": { "firstName": "John", "lastName": "Doe" } }), + expected: json!({ "firstName": "John", "lastName": "Doe" }), + }, + ]; + + for test_case in test_cases { + assert_eq!( + render(test_case.template, &test_case.context).unwrap(), + test_case.expected + ); + } +} + +#[test] +fn test_interpolation() { + struct TestCase { + template: &'static str, + context: Value, + expected: Value, + } + + let test_cases = vec![ + TestCase { + template: "{{ null }}s", + context: json!(null), + expected: json!("nulls"), + }, + TestCase { + template: "{{ 1 + 1 }} hello", + context: json!(null), + expected: json!("2 hello"), + }, + TestCase { + template: "{{ 'hello' }} world", + context: json!(null), + expected: json!("hello world"), + }, + TestCase { + template: "{{ true or false }} test", + context: json!(null), + expected: json!("true test"), + }, + TestCase { + template: "{{ [1, 2, 3] }} array", + context: json!(null), + expected: json!("[1,2,3] array"), + }, + TestCase { + template: "Customer: {{ customer }}", + context: json!({ "customer": { "firstName": "John", "lastName": "Doe" } }), + expected: json!(r#"Customer: {"firstName":"John","lastName":"Doe"}"#), + }, + ]; + + for test_case in test_cases { + assert_eq!( + render(test_case.template, &test_case.context).unwrap(), + test_case.expected + ); + } +} diff --git a/test-data/custom.json b/test-data/custom.json new file mode 100644 index 00000000..2ee9e78b --- /dev/null +++ b/test-data/custom.json @@ -0,0 +1,51 @@ +{ + "nodes": [ + { + "id": "115975ef-2f43-4e22-b553-0da6f4cc7f68", + "type": "inputNode", + "position": { + "x": 180, + "y": 240 + }, + "name": "Request" + }, + { + "id": "138b3b11-ff46-450f-9704-3f3c712067b2", + "type": "customNode", + "position": { + "x": 470, + "y": 240 + }, + "name": "customNode1", + "content": { + "kind": "sum", + "config": { + "prop1": "{{ a + 10 }}" + } + } + }, + { + "id": "db8797b1-bcc1-4fbf-a5d8-e7d43a181d5e", + "type": "outputNode", + "position": { + "x": 780, + "y": 240 + }, + "name": "Response" + } + ], + "edges": [ + { + "id": "05740fa7-3755-4756-b85e-bc1af2f6773b", + "sourceId": "115975ef-2f43-4e22-b553-0da6f4cc7f68", + "type": "edge", + "targetId": "138b3b11-ff46-450f-9704-3f3c712067b2" + }, + { + "id": "5d89c1d6-e894-4e8a-bd13-22368c2a6bc7", + "sourceId": "138b3b11-ff46-450f-9704-3f3c712067b2", + "type": "edge", + "targetId": "db8797b1-bcc1-4fbf-a5d8-e7d43a181d5e" + } + ] +} \ No newline at end of file