diff --git a/.github/workflows/rust.yaml b/.github/workflows/rust.yaml index c7088035..97ba831e 100644 --- a/.github/workflows/rust.yaml +++ b/.github/workflows/rust.yaml @@ -52,8 +52,8 @@ jobs: - uses: actions/checkout@v3 - name: Install Rust run: rustup install 1.80 - - run: cargo test --workspace --all-features --exclude zen-ffi --exclude zen-nodejs - - run: cargo test --workspace --all-features --exclude zen-ffi --exclude zen-nodejs --release + - run: cargo test --workspace --all-features --exclude zen-ffi --exclude zen-nodejs --exclude zen-python + - run: cargo test --workspace --all-features --exclude zen-ffi --exclude zen-nodejs --exclude zen-python --release build: name: cargo +${{ matrix.rust }} build diff --git a/bindings/c/src/custom_node.rs b/bindings/c/src/custom_node.rs index f642a3c6..26140484 100644 --- a/bindings/c/src/custom_node.rs +++ b/bindings/c/src/custom_node.rs @@ -24,7 +24,7 @@ impl Default for DynamicCustomNode { } impl CustomNodeAdapter for DynamicCustomNode { - async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + async fn handle(&self, request: CustomNodeRequest) -> NodeResult { match self { DynamicCustomNode::Noop(cn) => cn.handle(request).await, DynamicCustomNode::Native(cn) => cn.handle(request).await, diff --git a/bindings/c/src/languages/go.rs b/bindings/c/src/languages/go.rs index 8e5ecd8e..8d48e4bd 100644 --- a/bindings/c/src/languages/go.rs +++ b/bindings/c/src/languages/go.rs @@ -50,7 +50,7 @@ impl GoCustomNode { } impl CustomNodeAdapter for GoCustomNode { - async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + async fn handle(&self, request: CustomNodeRequest) -> NodeResult { let Some(handler) = self.handler else { return Err(anyhow!("go handler not found")); }; diff --git a/bindings/c/src/languages/native.rs b/bindings/c/src/languages/native.rs index 9c9dafad..e2a7368e 100644 --- a/bindings/c/src/languages/native.rs +++ b/bindings/c/src/languages/native.rs @@ -49,7 +49,7 @@ impl NativeCustomNode { } impl CustomNodeAdapter for NativeCustomNode { - async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + 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")); }; diff --git a/bindings/nodejs/.gitignore b/bindings/nodejs/.gitignore index e551c717..e8cf6bd9 100644 --- a/bindings/nodejs/.gitignore +++ b/bindings/nodejs/.gitignore @@ -1,2 +1,3 @@ *.node -*.internal.js \ No newline at end of file +*.internal.js +temp.d.ts \ No newline at end of file diff --git a/bindings/nodejs/package.json b/bindings/nodejs/package.json index 83e24ec1..e2059973 100644 --- a/bindings/nodejs/package.json +++ b/bindings/nodejs/package.json @@ -71,7 +71,7 @@ "url": "git+https://github.com/gorules/zen.git" }, "scripts": { - "build": "napi build --dts false --platform --release", + "build": "napi build --dts temp.d.ts --platform --release", "build:debug": "napi build --platform --js index.js --dts index.d.ts", "watch": "cargo watch --ignore '{index.js,index.d.ts}' -- npm run build:debug", "test": "jest", diff --git a/bindings/nodejs/src/custom_node.rs b/bindings/nodejs/src/custom_node.rs index e8426b9e..dedec446 100644 --- a/bindings/nodejs/src/custom_node.rs +++ b/bindings/nodejs/src/custom_node.rs @@ -21,7 +21,7 @@ impl CustomNode { } impl CustomNodeAdapter for CustomNode { - async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + async fn handle(&self, request: CustomNodeRequest) -> NodeResult { let Some(function) = &self.function else { return Err(anyhow!("Custom function is undefined")); }; diff --git a/bindings/nodejs/src/types.rs b/bindings/nodejs/src/types.rs index 4f5e6d50..1b4d18d5 100644 --- a/bindings/nodejs/src/types.rs +++ b/bindings/nodejs/src/types.rs @@ -1,9 +1,9 @@ -use std::collections::HashMap; - use json_dotpath::DotPaths; use napi::anyhow::{anyhow, Context}; use napi_derive::napi; use serde_json::Value; +use std::collections::HashMap; +use std::sync::Arc; use zen_engine::handler::custom_node_adapter::CustomDecisionNode; use zen_engine::{DecisionGraphResponse, DecisionGraphTrace}; @@ -65,16 +65,16 @@ pub struct DecisionNode { pub id: String, pub name: String, pub kind: String, - pub config: Value, + pub config: Arc, } -impl From> for DecisionNode { - fn from(value: CustomDecisionNode<'_>) -> Self { +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(), + id: value.id, + name: value.name, + kind: value.kind, + config: value.config, } } } diff --git a/bindings/python/src/custom_node.rs b/bindings/python/src/custom_node.rs index fd23aff9..293e4399 100644 --- a/bindings/python/src/custom_node.rs +++ b/bindings/python/src/custom_node.rs @@ -32,7 +32,7 @@ fn extract_custom_node_response(py: Python<'_>, result: PyObject) -> NodeResult } impl CustomNodeAdapter for PyCustomNode { - async fn handle(&self, request: CustomNodeRequest<'_>) -> NodeResult { + async fn handle(&self, request: CustomNodeRequest) -> NodeResult { let Some(callable) = &self.0 else { return Err(anyhow!("Custom node handler not provided")); }; diff --git a/bindings/python/src/types.rs b/bindings/python/src/types.rs index eb2dd106..95691f4a 100644 --- a/bindings/python/src/types.rs +++ b/bindings/python/src/types.rs @@ -3,6 +3,7 @@ use json_dotpath::DotPaths; use pyo3::{pyclass, pymethods, PyObject, PyResult, Python, ToPyObject}; use serde::Serialize; use serde_json::Value; +use std::sync::Arc; use crate::value::{value_to_object, PyValue}; use zen_engine::handler::custom_node_adapter::{ @@ -16,15 +17,15 @@ struct CustomDecisionNode { pub id: String, pub name: String, pub kind: String, - pub config: Value, + pub config: Arc, } -impl From> for CustomDecisionNode { +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(), + id: value.id, + name: value.name, + kind: value.kind, config: value.config.clone(), } } @@ -42,10 +43,7 @@ pub struct PyNodeRequest { } impl PyNodeRequest { - pub fn from_request( - py: Python, - value: CustomNodeRequest<'_>, - ) -> pythonize::Result { + 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(); diff --git a/core/engine/Cargo.toml b/core/engine/Cargo.toml index 248af4dc..34bdfbe8 100644 --- a/core/engine/Cargo.toml +++ b/core/engine/Cargo.toml @@ -17,7 +17,7 @@ thiserror = { workspace = true } bincode = { workspace = true, optional = true } petgraph = { workspace = true } serde_json = { workspace = true, features = ["arbitrary_precision"] } -serde = { workspace = true, features = ["derive"] } +serde = { workspace = true, features = ["derive", "rc"] } once_cell = { workspace = true } json_dotpath = { workspace = true } rust_decimal = { workspace = true, features = ["maths-nopanic"] } diff --git a/core/engine/src/decision.rs b/core/engine/src/decision.rs index aaf8a2bb..4b8610bf 100644 --- a/core/engine/src/decision.rs +++ b/core/engine/src/decision.rs @@ -82,12 +82,12 @@ where options: EvaluationOptions, ) -> Result> { let mut decision_graph = DecisionGraph::try_new(DecisionGraphConfig { + content: self.content.clone(), max_depth: options.max_depth.unwrap_or(5), trace: options.trace.unwrap_or_default(), loader: Arc::new(CachedLoader::from(self.loader.clone())), adapter: self.adapter.clone(), iteration: 0, - content: &self.content, })?; Ok(decision_graph.evaluate(context).await?) @@ -95,12 +95,12 @@ where pub fn validate(&self) -> Result<(), DecisionGraphValidationError> { let decision_graph = DecisionGraph::try_new(DecisionGraphConfig { + content: self.content.clone(), max_depth: 1, trace: false, loader: Arc::new(CachedLoader::from(self.loader.clone())), adapter: self.adapter.clone(), iteration: 0, - content: &self.content, })?; decision_graph.validate() diff --git a/core/engine/src/handler/custom_node_adapter.rs b/core/engine/src/handler/custom_node_adapter.rs index c9a019ab..17b29c78 100644 --- a/core/engine/src/handler/custom_node_adapter.rs +++ b/core/engine/src/handler/custom_node_adapter.rs @@ -4,44 +4,43 @@ use anyhow::anyhow; use json_dotpath::DotPaths; use serde::Serialize; use serde_json::Value; +use std::ops::Deref; +use std::sync::Arc; use zen_expression::variable::Variable; use zen_tmpl::TemplateRenderError; pub trait CustomNodeAdapter { - fn handle( - &self, - request: CustomNodeRequest<'_>, - ) -> impl std::future::Future; + fn handle(&self, request: CustomNodeRequest) -> impl std::future::Future; } #[derive(Default, Debug)] pub struct NoopCustomNode; impl CustomNodeAdapter for NoopCustomNode { - async fn handle(&self, _: CustomNodeRequest<'_>) -> NodeResult { + async fn handle(&self, _: CustomNodeRequest) -> NodeResult { Err(anyhow!("Custom node handler not provided")) } } #[derive(Serialize)] #[serde(rename_all = "camelCase")] -pub struct CustomNodeRequest<'a> { +pub struct CustomNodeRequest { pub input: Variable, - pub node: CustomDecisionNode<'a>, + pub node: CustomDecisionNode, } -impl<'a> TryFrom<&'a NodeRequest<'a>> for CustomNodeRequest<'a> { +impl TryFrom for CustomNodeRequest { type Error = (); - fn try_from(value: &'a NodeRequest<'a>) -> Result { + fn try_from(value: NodeRequest) -> Result { Ok(Self { input: value.input.clone(), - node: value.node.try_into()?, + node: value.node.deref().try_into()?, }) } } -impl<'a> CustomNodeRequest<'a> { +impl CustomNodeRequest { pub fn get_field(&self, path: &str) -> Result, TemplateRenderError> { let Some(selected_value) = self.get_field_raw(path) else { return Ok(None); @@ -62,26 +61,26 @@ impl<'a> CustomNodeRequest<'a> { #[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, +pub struct CustomDecisionNode { + pub id: String, + pub name: String, + pub kind: String, + pub config: Arc, } -impl<'a> TryFrom<&'a DecisionNode> for CustomDecisionNode<'a> { +impl TryFrom<&DecisionNode> for CustomDecisionNode { type Error = (); - fn try_from(value: &'a DecisionNode) -> Result { + fn try_from(value: &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, + id: value.id.clone(), + name: value.name.clone(), + kind: content.kind.clone(), + config: content.config.clone(), }) } } diff --git a/core/engine/src/handler/decision.rs b/core/engine/src/handler/decision.rs index 1ddff8a7..05d8f710 100644 --- a/core/engine/src/handler/decision.rs +++ b/core/engine/src/handler/decision.rs @@ -3,15 +3,13 @@ use crate::handler::function::function::Function; use crate::handler::graph::{DecisionGraph, DecisionGraphConfig}; use crate::handler::node::{NodeRequest, NodeResponse, NodeResult}; use crate::loader::DecisionLoader; -use crate::model::{DecisionNodeKind, TransformExecutionMode}; -use anyhow::{anyhow, Context}; -use serde_json::Value; +use crate::model::DecisionNodeKind; +use anyhow::anyhow; use std::future::Future; -use std::ops::Deref; use std::pin::Pin; use std::rc::Rc; use std::sync::Arc; -use zen_expression::{Isolate, Variable}; +use tokio::sync::Mutex; pub struct DecisionHandler { trace: bool, @@ -40,7 +38,7 @@ impl DecisionHandle pub fn handle<'s, 'arg, 'recursion>( &'s self, - request: &'arg NodeRequest<'_>, + request: NodeRequest, ) -> Pin + 'recursion>> where 's: 'recursion, @@ -52,11 +50,9 @@ impl DecisionHandle _ => Err(anyhow!("Unexpected node type")), }?; - let mut isolate = Isolate::new(); - let sub_decision = self.loader.load(&content.key).await?; - let mut sub_tree = DecisionGraph::try_new(DecisionGraphConfig { - content: sub_decision.deref(), + let sub_tree = DecisionGraph::try_new(DecisionGraphConfig { + content: sub_decision, max_depth: self.max_depth, loader: self.loader.clone(), adapter: self.adapter.clone(), @@ -65,69 +61,27 @@ impl DecisionHandle })? .with_function(self.js_function.clone()); - let input_data = match &content.transform_attributes.input_field { - None => request.input.clone(), - Some(input_field) => { - isolate.set_environment(request.input.clone()); - isolate.run_standard(input_field.as_str())? - } - }; + let sub_tree_mutex = Arc::new(Mutex::new(sub_tree)); - let mut trace_data: Option = None; - let mut output_data = match &content.transform_attributes.execution_mode { - TransformExecutionMode::Single => { - let response = sub_tree - .evaluate(request.input.clone()) - .await - .map_err(|e| e.source)?; + content + .transform_attributes + .run_with(request.input, |input| { + let sub_tree_mutex = sub_tree_mutex.clone(); - if self.trace { - trace_data.replace( - serde_json::to_value(response.trace) - .context("Failed to serialize trace")?, - ); - } + async move { + let mut sub_tree_ref = sub_tree_mutex.lock().await; - response.result - } - TransformExecutionMode::Loop => { - let input_array_ref = input_data.as_array().context("Expected an array")?; - let input_array = input_array_ref.borrow(); - - let mut output_array = Vec::with_capacity(input_array.len()); - let mut trace_datum = Vec::with_capacity(input_array.len()); - for input in input_array.iter() { - let response = sub_tree - .evaluate(input.clone()) + sub_tree_ref + .evaluate(input) .await - .map_err(|e| e.source)?; - - output_array.push(response.result); - trace_datum.push(response.trace); + .map(|r| NodeResponse { + output: r.result, + trace_data: serde_json::to_value(r.trace).ok(), + }) + .map_err(|e| e.source) } - - if self.trace { - trace_data.replace( - serde_json::to_value(trace_datum) - .context("Failed to parse trace data")?, - ); - } - - Variable::from_array(output_array) - } - }; - - if let Some(output_path) = &content.transform_attributes.output_path { - let new_output_data = Variable::empty_object(); - new_output_data.dot_insert(output_path.as_str(), output_data); - - output_data = new_output_data; - } - - Ok(NodeResponse { - output: output_data, - trace_data, - }) + }) + .await }) } } diff --git a/core/engine/src/handler/expression/mod.rs b/core/engine/src/handler/expression/mod.rs index a78295f0..4621dc31 100644 --- a/core/engine/src/handler/expression/mod.rs +++ b/core/engine/src/handler/expression/mod.rs @@ -1,15 +1,16 @@ use crate::handler::node::{NodeRequest, NodeResponse, NodeResult}; -use crate::model::DecisionNodeKind; +use crate::model::{DecisionNodeKind, ExpressionNodeContent}; use ahash::{HashMap, HashMapExt}; +use std::sync::Arc; use anyhow::{anyhow, Context}; use serde::Serialize; +use tokio::sync::Mutex; use zen_expression::variable::Variable; use zen_expression::Isolate; -pub struct ExpressionHandler<'a> { +pub struct ExpressionHandler { trace: bool, - isolate: Isolate<'a>, } #[derive(Debug, Serialize)] @@ -17,26 +18,62 @@ struct ExpressionTrace { result: String, } -impl<'a> ExpressionHandler<'a> { +impl ExpressionHandler { pub fn new(trace: bool) -> Self { - Self { - trace, - isolate: Isolate::new(), - } + Self { trace } } - pub async fn handle(&mut self, request: &'a NodeRequest<'_>) -> NodeResult { + pub async fn handle(&mut self, request: NodeRequest) -> NodeResult { let content = match &request.node.kind { DecisionNodeKind::ExpressionNode { content } => Ok(content), _ => Err(anyhow!("Unexpected node type")), }?; + let inner_handler_mutex = Arc::new(Mutex::new(ExpressionHandlerInner::new(self.trace))); + + content + .transform_attributes + .run_with(request.input, |input| { + let inner_handler_mutex = inner_handler_mutex.clone(); + + async move { + let mut inner_handler_ref = inner_handler_mutex.lock().await; + inner_handler_ref.handle(input, content).await + } + }) + .await + } +} + +struct ExpressionHandlerInner<'a> { + isolate: Isolate<'a>, + trace: bool, +} + +impl<'a> ExpressionHandlerInner<'a> { + pub fn new(trace: bool) -> Self { + Self { + isolate: Isolate::new(), + trace, + } + } + + async fn handle(&mut self, input: Variable, content: &'a ExpressionNodeContent) -> NodeResult { let result = Variable::empty_object(); let mut trace_map = self.trace.then(|| HashMap::<&str, ExpressionTrace>::new()); - self.isolate.set_environment(request.input.depth_clone(1)); + self.isolate.set_environment(input.depth_clone(1)); for expression in &content.expressions { - let value = self.evaluate_expression(&expression.value)?; + if expression.key.is_empty() || expression.value.is_empty() { + continue; + } + + let value = self + .isolate + .run_standard(&expression.value) + .with_context(|| { + format!(r#"Failed to evaluate expression: "{}""#, &expression.value) + })?; if let Some(tmap) = &mut trace_map { tmap.insert( &expression.key, @@ -66,10 +103,4 @@ impl<'a> ExpressionHandler<'a> { .context("Failed to serialize trace data")?, }) } - - fn evaluate_expression(&mut self, expression: &'a str) -> anyhow::Result { - self.isolate - .run_standard(expression) - .with_context(|| format!(r#"Failed to evaluate expression: "{expression}""#)) - } } diff --git a/core/engine/src/handler/function/mod.rs b/core/engine/src/handler/function/mod.rs index 77cfba99..e98519ea 100644 --- a/core/engine/src/handler/function/mod.rs +++ b/core/engine/src/handler/function/mod.rs @@ -43,7 +43,7 @@ impl FunctionHandler { } } - pub async fn handle(&self, request: &NodeRequest<'_>) -> NodeResult { + pub async fn handle(&self, request: NodeRequest) -> NodeResult { let content = match &request.node.kind { DecisionNodeKind::FunctionNode { content } => match content { FunctionNodeContent::Version2(content) => Ok(content), diff --git a/core/engine/src/handler/function/module/zen.rs b/core/engine/src/handler/function/module/zen.rs index 2b32f69f..ad3541f5 100644 --- a/core/engine/src/handler/function/module/zen.rs +++ b/core/engine/src/handler/function/module/zen.rs @@ -58,7 +58,7 @@ impl Run let load_result = loader.load(key.as_str()).await; let decision_content = load_result.or_throw(&ctx)?; let mut sub_tree = DecisionGraph::try_new(DecisionGraphConfig { - content: &decision_content, + content: decision_content, max_depth, loader, adapter, diff --git a/core/engine/src/handler/function_v1/mod.rs b/core/engine/src/handler/function_v1/mod.rs index fcc31345..1dde7059 100644 --- a/core/engine/src/handler/function_v1/mod.rs +++ b/core/engine/src/handler/function_v1/mod.rs @@ -22,7 +22,7 @@ impl FunctionHandler { Self { trace, runtime } } - pub async fn handle(&self, request: &NodeRequest<'_>) -> NodeResult { + pub async fn handle(&self, request: NodeRequest) -> NodeResult { let content = match &request.node.kind { DecisionNodeKind::FunctionNode { content } => match content { FunctionNodeContent::Version1(content) => Ok(content), diff --git a/core/engine/src/handler/graph.rs b/core/engine/src/handler/graph.rs index de1912a6..4d20d857 100644 --- a/core/engine/src/handler/graph.rs +++ b/core/engine/src/handler/graph.rs @@ -25,8 +25,8 @@ use std::time::Instant; use thiserror::Error; use zen_expression::variable::Variable; -pub struct DecisionGraph<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> { - graph: StableDiDecisionGraph<'a>, +pub struct DecisionGraph { + graph: StableDiDecisionGraph, adapter: Arc, loader: Arc, trace: bool, @@ -35,18 +35,18 @@ pub struct DecisionGraph<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + runtime: Option>, } -pub struct DecisionGraphConfig<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> { +pub struct DecisionGraphConfig { pub loader: Arc, pub adapter: Arc, - pub content: &'a DecisionContent, + pub content: Arc, pub trace: bool, pub iteration: u8, pub max_depth: u8, } -impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGraph<'a, L, A> { +impl DecisionGraph { pub fn try_new( - config: DecisionGraphConfig<'a, L, A>, + config: DecisionGraphConfig, ) -> Result { let content = config.content; let mut graph = StableDiDecisionGraph::new(); @@ -54,7 +54,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr for node in &content.nodes { let node_id = node.id.clone(); - let node_index = graph.add_node(node); + let node_index = graph.add_node(node.clone()); index_map.insert(node_id, node_index); } @@ -68,7 +68,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr DecisionGraphValidationError::MissingNode(edge.target_id.to_string()) })?; - graph.add_edge(source_index.clone(), target_index.clone(), edge); + graph.add_edge(source_index.clone(), target_index.clone(), edge.clone()); } Ok(Self { @@ -117,13 +117,6 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr )); } - let output_count = self.node_kind_count(DecisionNodeKind::OutputNode); - if output_count < 1 { - return Err(DecisionGraphValidationError::InvalidOutputCount( - output_count as u32, - )); - } - if is_cyclic_directed(&self.graph) { return Err(DecisionGraphValidationError::CyclicGraph); } @@ -171,7 +164,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr continue; } - let node = self.graph[nid]; + let node = (&self.graph[nid]).clone(); let start = Instant::now(); macro_rules! trace { @@ -222,7 +215,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr })?; let node_request = NodeRequest { - node, + node: node.clone(), iteration: self.iteration, input: walker.incoming_node_data(&self.graph, nid, true), }; @@ -233,7 +226,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr self.iteration, self.max_depth, ) - .handle(&node_request) + .handle(node_request.clone()) .await .map_err(|e| NodeError { source: e.into(), @@ -246,7 +239,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr })?; function_v1::FunctionHandler::new(self.trace, runtime) - .handle(&node_request) + .handle(node_request.clone()) .await .map_err(|e| NodeError { source: e.into(), @@ -270,7 +263,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr } DecisionNodeKind::DecisionNode { .. } => { let node_request = NodeRequest { - node, + node: node.clone(), iteration: self.iteration, input: walker.incoming_node_data(&self.graph, nid, true), }; @@ -282,7 +275,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr self.adapter.clone(), self.runtime.clone(), ) - .handle(&node_request) + .handle(node_request.clone()) .await .map_err(|e| NodeError { source: e.into(), @@ -304,13 +297,13 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr } DecisionNodeKind::DecisionTableNode { .. } => { let node_request = NodeRequest { - node, + node: node.clone(), iteration: self.iteration, input: walker.incoming_node_data(&self.graph, nid, true), }; let res = DecisionTableHandler::new(self.trace) - .handle(&node_request) + .handle(node_request.clone()) .await .map_err(|e| NodeError { node_id: node.id.clone(), @@ -319,6 +312,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr node_request.input.dot_remove("$nodes"); res.output.dot_remove("$nodes"); + res.output.dot_remove("$"); trace!({ input: node_request.input, @@ -332,13 +326,13 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr } DecisionNodeKind::ExpressionNode { .. } => { let node_request = NodeRequest { - node, + node: node.clone(), iteration: self.iteration, input: walker.incoming_node_data(&self.graph, nid, true), }; let res = ExpressionHandler::new(self.trace) - .handle(&node_request) + .handle(node_request.clone()) .await .map_err(|e| NodeError { node_id: node.id.clone(), @@ -360,14 +354,14 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr } DecisionNodeKind::CustomNode { .. } => { let node_request = NodeRequest { - node, + node: node.clone(), iteration: self.iteration, input: walker.incoming_node_data(&self.graph, nid, true), }; let res = self .adapter - .handle(CustomNodeRequest::try_from(&node_request).unwrap()) + .handle(CustomNodeRequest::try_from(node_request.clone()).unwrap()) .await .map_err(|e| NodeError { node_id: node.id.clone(), @@ -390,9 +384,10 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr } } - Err(NodeError { - node_id: "".to_string(), - source: anyhow!("Graph did not halt. Missing output node."), + Ok(DecisionGraphResponse { + result: walker.ending_variables(&self.graph), + performance: format!("{:?}", root_start.elapsed()), + trace: node_traces, }) } } diff --git a/core/engine/src/handler/node.rs b/core/engine/src/handler/node.rs index 2d0513f3..c87e8bf1 100644 --- a/core/engine/src/handler/node.rs +++ b/core/engine/src/handler/node.rs @@ -2,6 +2,7 @@ use crate::model::DecisionNode; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::fmt::{Display, Formatter}; +use std::sync::Arc; use thiserror::Error; use zen_expression::variable::Variable; @@ -12,11 +13,11 @@ pub struct NodeResponse { pub trace_data: Option, } -#[derive(Debug, Serialize)] -pub struct NodeRequest<'a> { +#[derive(Debug, Serialize, Clone)] +pub struct NodeRequest { pub input: Variable, pub iteration: u8, - pub node: &'a DecisionNode, + pub node: Arc, } #[derive(Error, Debug)] diff --git a/core/engine/src/handler/table/zen.rs b/core/engine/src/handler/table/zen.rs index da8ddefb..ef551221 100644 --- a/core/engine/src/handler/table/zen.rs +++ b/core/engine/src/handler/table/zen.rs @@ -1,10 +1,12 @@ use ahash::HashMap; use anyhow::{anyhow, Context}; +use std::sync::Arc; use crate::handler::node::{NodeRequest, NodeResponse, NodeResult}; use crate::handler::table::{RowOutput, RowOutputKind}; use crate::model::{DecisionNodeKind, DecisionTableContent, DecisionTableHitPolicy}; use serde::Serialize; +use tokio::sync::Mutex; use zen_expression::variable::Variable; use zen_expression::Isolate; @@ -18,12 +20,34 @@ struct RowResult { } #[derive(Debug)] -pub struct DecisionTableHandler<'a> { +pub struct DecisionTableHandler { + trace: bool, +} + +impl DecisionTableHandler { + pub fn new(trace: bool) -> Self { + Self { trace } + } + + pub async fn handle(&mut self, request: NodeRequest) -> NodeResult { + let content = match &request.node.kind { + DecisionNodeKind::DecisionTableNode { content } => Ok(content), + _ => Err(anyhow!("Unexpected node type")), + }?; + + let inner_handler = DecisionTableHandlerInner::new(self.trace); + inner_handler + .handle(request.input.depth_clone(1), content) + .await + } +} + +struct DecisionTableHandlerInner<'a> { isolate: Isolate<'a>, trace: bool, } -impl<'a> DecisionTableHandler<'a> { +impl<'a> DecisionTableHandlerInner<'a> { pub fn new(trace: bool) -> Self { Self { isolate: Isolate::new(), @@ -31,18 +55,24 @@ impl<'a> DecisionTableHandler<'a> { } } - pub async fn handle(&mut self, request: &'a NodeRequest<'_>) -> NodeResult { - let content = match &request.node.kind { - DecisionNodeKind::DecisionTableNode { content } => Ok(content), - _ => Err(anyhow!("Unexpected node type")), - }?; + pub async fn handle(self, input: Variable, content: &'a DecisionTableContent) -> NodeResult { + let self_mutex = Arc::new(Mutex::new(self)); - self.isolate.set_environment(request.input.depth_clone(1)); + content + .transform_attributes + .run_with(input, |input| { + let self_mutex = self_mutex.clone(); + async move { + let mut self_ref = self_mutex.lock().await; - match &content.hit_policy { - DecisionTableHitPolicy::First => self.handle_first_hit(&content).await, - DecisionTableHitPolicy::Collect => self.handle_collect(&content).await, - } + self_ref.isolate.set_environment(input); + match &content.hit_policy { + DecisionTableHitPolicy::First => self_ref.handle_first_hit(&content).await, + DecisionTableHitPolicy::Collect => self_ref.handle_collect(&content).await, + } + } + }) + .await } async fn handle_first_hit(&mut self, content: &'a DecisionTableContent) -> NodeResult { diff --git a/core/engine/src/handler/traversal.rs b/core/engine/src/handler/traversal.rs index 7bb089a7..fe0e94f6 100644 --- a/core/engine/src/handler/traversal.rs +++ b/core/engine/src/handler/traversal.rs @@ -3,11 +3,12 @@ use fixedbitset::FixedBitSet; use petgraph::data::DataMap; use petgraph::matrix_graph::Zero; use petgraph::prelude::{EdgeIndex, NodeIndex, StableDiGraph}; -use petgraph::visit::{EdgeRef, IntoNeighbors, IntoNodeIdentifiers, Reversed, VisitMap, Visitable}; +use petgraph::visit::{EdgeRef, IntoNodeIdentifiers, VisitMap, Visitable}; use petgraph::{Incoming, Outgoing}; use serde_json::json; use std::rc::Rc; use std::sync::atomic::Ordering; +use std::sync::Arc; use std::time::Instant; use crate::config::ZEN_CONFIG; @@ -18,7 +19,7 @@ use crate::DecisionGraphTrace; use zen_expression::variable::Variable; use zen_expression::Isolate; -pub(crate) type StableDiDecisionGraph<'a> = StableDiGraph<&'a DecisionNode, &'a DecisionEdge>; +pub(crate) type StableDiDecisionGraph = StableDiGraph, Arc>; pub(crate) struct GraphWalker { ordered: FixedBitSet, @@ -43,8 +44,8 @@ impl GraphWalker { // find all initial nodes (nodes without incoming edges) self.to_visit .extend(g.node_identifiers().filter(move |&nid| { - g.neighbors_directed(nid, Incoming).count().is_zero() - && !g.neighbors_directed(nid, Outgoing).count().is_zero() + g.node_weight(nid) + .is_some_and(|n| n.kind == DecisionNodeKind::InputNode) })); } @@ -72,6 +73,20 @@ impl GraphWalker { self.node_data.get(&node_id).cloned() } + pub fn ending_variables(&self, g: &StableDiDecisionGraph) -> Variable { + g.node_indices() + .filter(|nid| { + self.ordered.is_visited(nid) + && g.neighbors_directed(*nid, Outgoing).count().is_zero() + }) + .fold(Variable::empty_object(), |mut acc, curr| { + match self.node_data.get(&curr) { + None => acc, + Some(data) => acc.merge(data), + } + }) + } + pub fn get_all_node_data(&self, g: &StableDiDecisionGraph) -> Variable { let node_values = self .node_data @@ -130,11 +145,18 @@ impl GraphWalker { } // Take an unvisited element and find which of its neighbors are next while let Some(nid) = self.to_visit.pop() { - let decision_node = *g.node_weight(nid)?; + let decision_node = g.node_weight(nid)?.clone(); if self.ordered.is_visited(&nid) { continue; } + if !self.all_dependencies_resolved(g, nid) { + self.to_visit.push(nid); + self.to_visit + .extend(self.get_unresolved_dependencies(g, nid)); + continue; + } + self.ordered.visit(nid); if let DecisionNodeKind::SwitchNode { content } = &decision_node.kind { @@ -211,22 +233,29 @@ impl GraphWalker { } } - for neigh in g.neighbors(nid) { - // Look at each neighbor, and those that only have incoming edges - // from the already ordered list, they are the next to visit. - if Reversed(&*g) - .neighbors(neigh) - .all(|b| self.ordered.is_visited(&b)) - { - self.to_visit.push(neigh); - } - } + let successors = g.neighbors_directed(nid, Outgoing); + self.to_visit.extend(successors); return Some(nid); } None } + + fn all_dependencies_resolved(&self, g: &StableDiDecisionGraph, nid: NodeIndex) -> bool { + g.neighbors_directed(nid, Incoming) + .all(|dep| self.ordered.is_visited(&dep)) + } + + fn get_unresolved_dependencies( + &self, + g: &StableDiDecisionGraph, + nid: NodeIndex, + ) -> Vec { + g.neighbors_directed(nid, Incoming) + .filter(|dep| !self.ordered.is_visited(dep)) + .collect() + } } fn switch_statement_evaluate<'a>( diff --git a/core/engine/src/lib.rs b/core/engine/src/lib.rs index 38ddbf39..80804090 100644 --- a/core/engine/src/lib.rs +++ b/core/engine/src/lib.rs @@ -130,6 +130,7 @@ pub mod handler; pub mod loader; #[path = "model/mod.rs"] pub mod model; +mod util; pub use config::ZEN_CONFIG; pub use decision::Decision; diff --git a/core/engine/src/model/mod.rs b/core/engine/src/model/mod.rs index 34524efd..c19d283e 100644 --- a/core/engine/src/model/mod.rs +++ b/core/engine/src/model/mod.rs @@ -1,14 +1,15 @@ use ahash::HashMap; -use serde::{Deserialize, Serialize}; +use serde::{Deserialize, Deserializer, Serialize}; use serde_json::Value; +use std::sync::Arc; /// JDM Decision model #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] #[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct DecisionContent { - pub nodes: Vec, - pub edges: Vec, + pub nodes: Vec>, + pub edges: Vec>, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -86,6 +87,8 @@ pub struct DecisionTableContent { pub inputs: Vec, pub outputs: Vec, pub hit_policy: DecisionTableHitPolicy, + #[serde(flatten)] + pub transform_attributes: TransformAttributes, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -102,6 +105,7 @@ pub enum DecisionTableHitPolicy { pub struct DecisionTableInputField { pub id: String, pub name: String, + #[serde(default, deserialize_with = "empty_string_is_none")] pub field: Option, } @@ -119,6 +123,8 @@ pub struct DecisionTableOutputField { #[serde(rename_all = "camelCase")] pub struct ExpressionNodeContent { pub expressions: Vec, + #[serde(flatten)] + pub transform_attributes: TransformAttributes, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -160,10 +166,14 @@ pub enum SwitchStatementHitPolicy { #[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))] #[serde(rename_all = "camelCase")] pub struct TransformAttributes { + #[serde(default, deserialize_with = "empty_string_is_none")] pub input_field: Option, + #[serde(default, deserialize_with = "empty_string_is_none")] pub output_path: Option, #[serde(default)] pub execution_mode: TransformExecutionMode, + #[serde(default)] + pub pass_through: bool, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Default)] @@ -179,7 +189,7 @@ pub enum TransformExecutionMode { #[serde(rename_all = "camelCase")] pub struct CustomNodeContent { pub kind: String, - pub config: Value, + pub config: Arc, } #[cfg(feature = "bincode")] @@ -225,3 +235,21 @@ impl<'__de> ::bincode::BorrowDecode<'__de> for CustomNodeContent { Ok(Self { kind, config }) } } + +fn empty_string_is_none<'de, D>(deserializer: D) -> Result, D::Error> +where + D: Deserializer<'de>, +{ + #[derive(Deserialize)] + #[serde(untagged)] + enum StringOrNull { + String(String), + Null, + } + + match StringOrNull::deserialize(deserializer)? { + StringOrNull::String(s) if s.trim().is_empty() => Ok(None), + StringOrNull::String(s) => Ok(Some(s)), + StringOrNull::Null => Ok(None), + } +} diff --git a/core/engine/src/util/mod.rs b/core/engine/src/util/mod.rs new file mode 100644 index 00000000..59a47e0e --- /dev/null +++ b/core/engine/src/util/mod.rs @@ -0,0 +1 @@ +mod transform_attribute; diff --git a/core/engine/src/util/transform_attribute.rs b/core/engine/src/util/transform_attribute.rs new file mode 100644 index 00000000..b68bed1c --- /dev/null +++ b/core/engine/src/util/transform_attribute.rs @@ -0,0 +1,95 @@ +use crate::handler::node::{NodeResponse, NodeResult}; +use crate::model::{TransformAttributes, TransformExecutionMode}; +use anyhow::Context; +use serde_json::Value; +use std::future::Future; +use zen_expression::{Isolate, Variable}; + +impl TransformAttributes { + pub(crate) async fn run_with(&self, node_input: Variable, evaluate: F) -> NodeResult + where + F: Fn(Variable) -> Fut, + Fut: Future, + { + let input = match &self.input_field { + None => node_input.clone(), + Some(input_field) => { + let mut isolate = Isolate::new(); + isolate.set_environment(node_input.clone()); + let calculated_input = isolate.run_standard(input_field.as_str())?; + + let nodes = node_input.dot("$nodes").unwrap_or(Variable::Null); + match &calculated_input { + Variable::Array(arr) => { + let arr = arr.borrow(); + let s: Vec<_> = arr + .iter() + .map(|v| { + let new_v = v.depth_clone(1); + new_v.dot_insert("$nodes", nodes.clone()); + new_v + }) + .collect(); + + Variable::from_array(s) + } + _ => { + let new_input = calculated_input.depth_clone(1); + new_input.dot_insert("$nodes", nodes); + new_input + } + } + } + }; + + let mut trace_data: Option = None; + let mut output = match self.execution_mode { + TransformExecutionMode::Single => { + let response = evaluate(input).await?; + if let Some(td) = response.trace_data { + trace_data.replace(td); + } + + response.output.dot_remove("$nodes"); + response.output + } + TransformExecutionMode::Loop => { + let input_array_ref = input.as_array().context("Expected an array")?; + let input_array = input_array_ref.borrow(); + + let mut output_array = Vec::with_capacity(input_array.len()); + let mut trace_datum = Vec::with_capacity(input_array.len()); + for input in input_array.iter() { + let mut response = evaluate(input.clone()).await?; + if let Some(td) = response.trace_data { + trace_datum.push(td); + } + + if self.pass_through { + response.output = input.clone().merge_clone(&response.output); + } + + response.output.dot_remove("$nodes"); + output_array.push(response.output); + } + + trace_data.replace(Value::Array(trace_datum)); + Variable::from_array(output_array) + } + }; + + if let Some(output_path) = &self.output_path { + let new_output = Variable::empty_object(); + new_output.dot_insert(output_path.as_str(), output); + + output = new_output; + } + + if self.pass_through { + let mut node_input = node_input; + output = node_input.merge_clone(&output) + } + + Ok(NodeResponse { output, trace_data }) + } +} diff --git a/core/engine/tests/decision.rs b/core/engine/tests/decision.rs index 66ca463d..cc5609b0 100644 --- a/core/engine/tests/decision.rs +++ b/core/engine/tests/decision.rs @@ -87,11 +87,4 @@ fn decision_validation() { missing_input_error, DecisionGraphValidationError::InvalidInputCount(_) )); - - let missing_output_decision = Decision::from(load_test_data("error-missing-output.json")); - let missing_output_error = missing_output_decision.validate().unwrap_err(); - assert!(matches!( - missing_output_error, - DecisionGraphValidationError::InvalidOutputCount(_) - )); } diff --git a/core/engine/tests/engine.rs b/core/engine/tests/engine.rs index f42a3671..961453a4 100644 --- a/core/engine/tests/engine.rs +++ b/core/engine/tests/engine.rs @@ -8,7 +8,7 @@ use std::path::Path; use std::sync::Arc; use tokio::runtime::Builder; use zen_engine::loader::{LoaderError, MemoryLoader}; -use zen_engine::model::{DecisionContent, DecisionNodeKind, FunctionNodeContent}; +use zen_engine::model::{DecisionContent, DecisionNode, DecisionNodeKind, FunctionNodeContent}; use zen_engine::Variable; use zen_engine::{DecisionEngine, EvaluationError, EvaluationOptions}; @@ -158,24 +158,36 @@ fn engine_with_trace() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_function_imports() { - let mut function_content = load_test_data("function.json"); + let function_content = load_test_data("function.json"); let imports_js_path = Path::new("js").join("imports.js"); let mut replace_buffer = load_raw_test_data(imports_js_path.to_str().unwrap()); let mut replace_data = String::new(); replace_buffer.read_to_string(&mut replace_data).unwrap(); - function_content.nodes.iter_mut().for_each(|node| { - if let DecisionNodeKind::FunctionNode { content, .. } = &mut node.kind { - match content { - FunctionNodeContent::Version1(content) => { - let _ = std::mem::replace(content, replace_data.clone()); - } - _ => {} - } - } - }); + let new_nodes = function_content + .nodes + .into_iter() + .map(|node| match &node.kind { + DecisionNodeKind::FunctionNode { .. } => { + let new_kind = DecisionNodeKind::FunctionNode { + content: FunctionNodeContent::Version1(replace_data.clone()), + }; + Arc::new(DecisionNode { + id: node.id.clone(), + name: node.name.clone(), + kind: new_kind, + }) + } + _ => node, + }) + .collect::>(); + + let function_content = DecisionContent { + edges: function_content.edges, + nodes: new_nodes, + }; let decision = DecisionEngine::default().create_decision(function_content.into()); let response = decision.evaluate(json!({}).into()).await.unwrap(); diff --git a/core/expression/src/intellisense/types/provider.rs b/core/expression/src/intellisense/types/provider.rs index fca038c9..b6f30797 100644 --- a/core/expression/src/intellisense/types/provider.rs +++ b/core/expression/src/intellisense/types/provider.rs @@ -264,7 +264,7 @@ impl TypesProvider { }, Operator::Comparison(comp) => match comp { ComparisonOperator::Equal => { - if !left_type.omit_const().satisfies(&right_type.omit_const()) { + if !left_type.omit_const().satisfies(&right_type.omit_const()) && !left_type.is_null() && !right_type.is_null() { on_fly_error.replace(format!( "Hint: Expression will always evaluate to `false` because `{left_type}` != `{right_type}`." )); @@ -273,7 +273,7 @@ impl TypesProvider { V(VariableType::Bool) }, ComparisonOperator::NotEqual => { - if !left_type.omit_const().satisfies(&right_type.omit_const()) { + if !left_type.omit_const().satisfies(&right_type.omit_const()) && !left_type.is_null() && !right_type.is_null() { on_fly_error.replace(format!( "Hint: Expression will always evaluate to `true` because `{left_type}` != `{right_type}`." )); diff --git a/core/expression/src/parser/parser.rs b/core/expression/src/parser/parser.rs index 696656d4..1a880493 100644 --- a/core/expression/src/parser/parser.rs +++ b/core/expression/src/parser/parser.rs @@ -128,7 +128,11 @@ impl<'arena, 'token_ref, Flavor> Parser<'arena, 'token_ref, Flavor> { } pub(crate) fn prev_token_end(&self) -> u32 { - match self.tokens.get(self.position() - 1) { + let Some(pos) = self.position().checked_sub(1) else { + return self.token_start(); + }; + + match self.tokens.get(pos) { None => self.token_start(), Some(t) => t.span.1, } diff --git a/core/expression/src/variable/mod.rs b/core/expression/src/variable/mod.rs index 73f7932e..2c60b024 100644 --- a/core/expression/src/variable/mod.rs +++ b/core/expression/src/variable/mod.rs @@ -166,11 +166,18 @@ impl Variable { } pub fn merge(&mut self, patch: &Variable) -> Variable { - merge_variables(self, patch, true); + let _ = merge_variables(self, patch, true, MergeStrategy::InPlace); self.shallow_clone() } + pub fn merge_clone(&mut self, patch: &Variable) -> Variable { + let mut new_self = self.shallow_clone(); + + let _ = merge_variables(&mut new_self, patch, true, MergeStrategy::CloneOnWrite); + new_self + } + pub fn shallow_clone(&self) -> Self { match self { Variable::Null => Variable::Null, @@ -228,40 +235,96 @@ impl Clone for Variable { } } -fn merge_variables(doc: &mut Variable, patch: &Variable, top_level: bool) { - if !patch.is_object() && !patch.is_array() && top_level { - return; +#[derive(Copy, Clone)] +enum MergeStrategy { + InPlace, + CloneOnWrite, +} + +fn merge_variables( + doc: &mut Variable, + patch: &Variable, + top_level: bool, + strategy: MergeStrategy, +) -> bool { + if patch.is_array() && top_level { + *doc = patch.shallow_clone(); + return true; + } + + if !patch.is_object() && top_level { + return false; } if doc.is_object() && patch.is_object() { - let map_ref = doc.as_object().unwrap(); + let doc_ref = doc.as_object().unwrap(); let patch_ref = patch.as_object().unwrap(); - if Rc::ptr_eq(&map_ref, &patch_ref) { - return; + if Rc::ptr_eq(&doc_ref, &patch_ref) { + return false; } - let mut map = map_ref.borrow_mut(); let patch = patch_ref.borrow(); - for (key, value) in patch.deref() { - if value == &Variable::Null { - map.remove(key.as_str()); - } else { - let entry = map.entry(key.to_string()).or_insert(Variable::Null); - merge_variables(entry, value, false) + match strategy { + MergeStrategy::InPlace => { + let mut map = doc_ref.borrow_mut(); + for (key, value) in patch.deref() { + if value == &Variable::Null { + map.remove(key.as_str()); + } else { + let entry = map.entry(key.to_string()).or_insert(Variable::Null); + merge_variables(entry, value, false, strategy); + } + } + + return true; + } + MergeStrategy::CloneOnWrite => { + let mut changed = false; + let mut new_map = None; + + for (key, value) in patch.deref() { + // Get or create the new map if we haven't yet + let map = if let Some(ref mut m) = new_map { + m + } else { + let m = doc_ref.borrow().clone(); + new_map = Some(m); + new_map.as_mut().unwrap() + }; + + if value == &Variable::Null { + // Remove null values + if map.remove(key.as_str()).is_some() { + changed = true; + } + } else { + // Handle nested merging + let entry = map.entry(key.to_string()).or_insert(Variable::Null); + if merge_variables(entry, value, false, strategy) { + changed = true; + } + } + } + + // Only update doc if changes were made + if changed { + if let Some(new_map) = new_map { + *doc = Variable::Object(Rc::new(RefCell::new(new_map))); + } + return true; + } + + return false; } } - } else if doc.is_array() && patch.is_array() { - let arr_ref = doc.as_array().unwrap(); - let patch_ref = patch.as_array().unwrap(); - if Rc::ptr_eq(&arr_ref, &patch_ref) { - return; + } else { + let new_value = patch.shallow_clone(); + if *doc != new_value { + *doc = new_value; + return true; } - let mut arr = arr_ref.borrow_mut(); - let patch = patch_ref.borrow(); - arr.extend(patch.iter().map(|s| s.shallow_clone())); - } else { - *doc = patch.shallow_clone(); + return false; } } diff --git a/core/expression/src/variable/types/util.rs b/core/expression/src/variable/types/util.rs index ce1d341e..2816b1d3 100644 --- a/core/expression/src/variable/types/util.rs +++ b/core/expression/src/variable/types/util.rs @@ -62,7 +62,7 @@ impl VariableType { (VariableType::Bool, VariableType::Bool) => true, (VariableType::String, VariableType::String) => true, (VariableType::Number, VariableType::Number) => true, - (VariableType::Array(a1), VariableType::Array(a2)) => a1 == a2, + (VariableType::Array(a1), VariableType::Array(a2)) => a1.satisfies(a2), (VariableType::Object(o1), VariableType::Object(o2)) => o1 .iter() .all(|(k, v)| o2.get(k).is_some_and(|tv| v.satisfies(tv))), @@ -149,6 +149,13 @@ impl VariableType { (_, _) => VariableType::Any, } } + + pub fn is_null(&self) -> bool { + match self { + VariableType::Null => true, + _ => false, + } + } } #[cfg(test)] diff --git a/test-data/graphs/expression-default.json b/test-data/graphs/expression-default.json new file mode 100644 index 00000000..c5fb8940 --- /dev/null +++ b/test-data/graphs/expression-default.json @@ -0,0 +1,93 @@ +{ + "tests": [ + { + "input": { + "a": 1, + "b": 2, + "extra": "This should not pass through" + }, + "output": { + "sum": 3 + } + }, + { + "input": { + "a": 5, + "b": 3, + "sum": 100 + }, + "output": { + "sum": 8 + } + }, + { + "input": { + "a": 10, + "b": 20, + "c": 30, + "d": 40 + }, + "output": { + "sum": 30 + } + }, + { + "input": { + "a": 7, + "b": 3, + "passThrough": "This should not appear in output" + }, + "output": { + "sum": 10 + } + }, + { + "input": { + "a": -5, + "b": 5, + "negative": true + }, + "output": { + "sum": 0 + } + } + ], + "nodes": [ + { + "type": "inputNode", + "id": "deced339-bace-452a-8db0-777f038bffe8", + "name": "request", + "position": { + "x": 145, + "y": 235 + } + }, + { + "type": "expressionNode", + "content": { + "expressions": [ + { + "id": "d54641c2-5f24-4140-a9c4-8453542a622a", + "key": "sum", + "value": "a + b" + } + ], + "passThrough": false + }, + "id": "6b9cfc7e-4776-4b3f-8a19-c2a4d7874770", + "name": "expression1", + "position": { + "x": 475, + "y": 235 + } + } + ], + "edges": [ + { + "id": "639f3a3a-7545-4e98-aa79-7a9bd9644af1", + "sourceId": "deced339-bace-452a-8db0-777f038bffe8", + "type": "edge", + "targetId": "6b9cfc7e-4776-4b3f-8a19-c2a4d7874770" + } + ] +} \ No newline at end of file diff --git a/test-data/graphs/expression-fields.json b/test-data/graphs/expression-fields.json new file mode 100644 index 00000000..18df3783 --- /dev/null +++ b/test-data/graphs/expression-fields.json @@ -0,0 +1,94 @@ +{ + "tests": [ + { + "input": { + "customer": { + "firstName": "John", + "lastName": "Doe", + "age": 30 + }, + "order": { + "id": "ORD-001", + "total": 100 + } + }, + "output": { + "customer": { + "firstName": "John", + "lastName": "Doe", + "age": 30, + "fullName": "John Doe" + }, + "order": { + "id": "ORD-001", + "total": 100 + } + } + }, + { + "input": { + "customer": { + "firstName": "Jane", + "lastName": "Smith", + "age": 25 + }, + "order": { + "id": "ORD-002", + "total": 150 + } + }, + "output": { + "customer": { + "firstName": "Jane", + "lastName": "Smith", + "age": 25, + "fullName": "Jane Smith" + }, + "order": { + "id": "ORD-002", + "total": 150 + } + } + } + ], + "nodes": [ + { + "type": "inputNode", + "id": "input-node", + "name": "request", + "position": { + "x": 100, + "y": 100 + } + }, + { + "type": "expressionNode", + "id": "expression-node-1", + "name": "customerFullName", + "position": { + "x": 300, + "y": 100 + }, + "content": { + "inputField": "customer", + "outputPath": "customer", + "expressions": [ + { + "id": "07795ded-cb9b-4165-9b5e-783b066dda61", + "key": "fullName", + "value": "`${firstName} ${lastName}`" + } + ], + "passThrough": true + } + } + ], + "edges": [ + { + "id": "edge-1", + "sourceId": "input-node", + "targetId": "expression-node-1", + "type": "edge" + } + ] +} \ No newline at end of file diff --git a/test-data/graphs/expression-loop.json b/test-data/graphs/expression-loop.json new file mode 100644 index 00000000..29d48e2b --- /dev/null +++ b/test-data/graphs/expression-loop.json @@ -0,0 +1,94 @@ +{ + "tests": [ + { + "input": { + "cart": { + "items": [ + { + "name": "Apple", + "price": 0.5, + "quantity": 3 + }, + { + "name": "Banana", + "price": 0.3, + "quantity": 5 + }, + { + "name": "Orange", + "price": 0.7, + "quantity": 2 + } + ], + "customerId": "CUST-001" + } + }, + "output": { + "cart": { + "items": [ + { + "name": "Apple", + "price": 0.5, + "quantity": 3, + "totalPrice": 1.5 + }, + { + "name": "Banana", + "price": 0.3, + "quantity": 5, + "totalPrice": 1.5 + }, + { + "name": "Orange", + "price": 0.7, + "quantity": 2, + "totalPrice": 1.4 + } + ], + "customerId": "CUST-001" + } + } + } + ], + "nodes": [ + { + "type": "inputNode", + "id": "input-node", + "name": "request", + "position": { + "x": 100, + "y": 100 + } + }, + { + "type": "expressionNode", + "id": "expression-node-1", + "name": "calculateItemTotal", + "position": { + "x": 300, + "y": 100 + }, + "content": { + "inputField": "cart.items", + "outputPath": "cart.items", + "executionMode": "loop", + "expressions": [ + { + "id": "total-price-exp", + "key": "totalPrice", + "value": "price * quantity" + } + ], + "passThrough": true + } + } + ], + "edges": [ + { + "id": "edge-1", + "sourceId": "input-node", + "targetId": "expression-node-1", + "type": "edge" + } + ] +} \ No newline at end of file diff --git a/test-data/graphs/expression-passthrough.json b/test-data/graphs/expression-passthrough.json new file mode 100644 index 00000000..1b8bf907 --- /dev/null +++ b/test-data/graphs/expression-passthrough.json @@ -0,0 +1,108 @@ +{ + "tests": [ + { + "input": { + "a": 1, + "b": 2, + "sum": [] + }, + "output": { + "a": 1, + "b": 2, + "sum": 3 + } + }, + { + "input": { + "a": 5, + "b": 3, + "sum": 100 + }, + "output": { + "a": 5, + "b": 3, + "sum": 8 + } + }, + { + "input": { + "a": 10, + "b": 20, + "extra": "This should pass through" + }, + "output": { + "a": 10, + "b": 20, + "sum": 30, + "extra": "This should pass through" + } + }, + { + "input": { + "a": 7, + "b": 3, + "sum": "Original", + "passThrough": "Test" + }, + "output": { + "a": 7, + "b": 3, + "sum": 10, + "passThrough": "Test" + } + }, + { + "input": { + "a": 1, + "b": 1, + "c": 3, + "d": 4 + }, + "output": { + "a": 1, + "b": 1, + "c": 3, + "d": 4, + "sum": 2 + } + } + ], + "nodes": [ + { + "type": "inputNode", + "id": "deced339-bace-452a-8db0-777f038bffe8", + "name": "request", + "position": { + "x": 145, + "y": 235 + } + }, + { + "type": "expressionNode", + "content": { + "expressions": [ + { + "id": "d54641c2-5f24-4140-a9c4-8453542a622a", + "key": "sum", + "value": "a + b" + } + ], + "passThrough": true + }, + "id": "6b9cfc7e-4776-4b3f-8a19-c2a4d7874770", + "name": "expression1", + "position": { + "x": 475, + "y": 235 + } + } + ], + "edges": [ + { + "id": "639f3a3a-7545-4e98-aa79-7a9bd9644af1", + "sourceId": "deced339-bace-452a-8db0-777f038bffe8", + "type": "edge", + "targetId": "6b9cfc7e-4776-4b3f-8a19-c2a4d7874770" + } + ] +} \ No newline at end of file diff --git a/test-data/graphs/expression-table-map.json b/test-data/graphs/expression-table-map.json new file mode 100644 index 00000000..1b216b93 --- /dev/null +++ b/test-data/graphs/expression-table-map.json @@ -0,0 +1,313 @@ +{ + "tests": [ + { + "input": { + "customer": { + "name": "John Doe", + "country": "US" + }, + "items": [ + { + "name": "Laptop", + "amount": 1000, + "group": "electronics" + }, + { + "name": "Mouse", + "amount": 50, + "group": "accessories" + } + ] + }, + "output": { + "customer": { + "country": "US", + "name": "John Doe" + }, + "exprItems": [ + { + "amount": 50, + "discount": 5, + "group": "accessories", + "name": "Mouse" + } + ], + "exprItemsSum": 50, + "items": [ + { + "amount": 1000, + "group": "electronics", + "name": "Laptop" + }, + { + "amount": 50, + "group": "accessories", + "name": "Mouse" + } + ], + "tblItems": [ + { + "amount": 1000, + "group": "electronics", + "name": "Laptop" + }, + { + "amount": 50, + "discount": 5, + "group": "accessories", + "name": "Mouse" + } + ], + "tblItemsSum": 1050 + } + }, + { + "input": { + "customer": { + "name": "John Doe", + "country": "US" + }, + "items": [ + { + "name": "Laptop", + "amount": 1000, + "group": "electronics" + }, + { + "name": "Mouse", + "amount": 50, + "group": "accessories_new" + } + ] + }, + "output": { + "customer": { + "country": "US", + "name": "John Doe" + }, + "exprItems": [], + "exprItemsSum": 0, + "items": [ + { + "amount": 1000, + "group": "electronics", + "name": "Laptop" + }, + { + "amount": 50, + "group": "accessories_new", + "name": "Mouse" + } + ], + "tblItems": [ + { + "amount": 1000, + "group": "electronics", + "name": "Laptop" + }, + { + "amount": 50, + "group": "accessories_new", + "name": "Mouse" + } + ], + "tblItemsSum": 1050 + } + }, + { + "input": { + "customer": { + "name": "John Doe", + "country": "US" + }, + "items": [ + { + "name": "Laptop", + "amount": 1000, + "group": "electronics" + }, + { + "name": "Mouse", + "amount": 51, + "group": "accessories" + } + ] + }, + "output": { + "customer": { + "country": "US", + "name": "John Doe" + }, + "exprItems": [ + { + "amount": 51, + "discount": 5, + "group": "accessories", + "name": "Mouse" + } + ], + "exprItemsSum": 51, + "items": [ + { + "amount": 1000, + "group": "electronics", + "name": "Laptop" + }, + { + "amount": 51, + "group": "accessories", + "name": "Mouse" + } + ], + "tblItems": [ + { + "amount": 1000, + "group": "electronics", + "name": "Laptop" + }, + { + "amount": 51, + "discount": 5, + "group": "accessories", + "name": "Mouse" + } + ], + "tblItemsSum": 1051 + } + } + ], + "contentType": "application/vnd.gorules.decision", + "nodes": [ + { + "id": "be2a55f4-cdd8-482e-b53c-a947cbb7d7a0", + "name": "request", + "type": "inputNode", + "position": { + "x": 125, + "y": 265 + } + }, + { + "id": "b8ebb212-e290-4131-9703-0bb7b1fd2328", + "name": "table", + "type": "decisionTableNode", + "content": { + "rules": [ + { + "_id": "fd81538a-7451-4eb7-a25c-01c3afbebbdf", + "e06432d2-4609-425b-ad99-9ca27db08130": "5", + "e9090e94-051c-4c27-be58-baac03041c2b": "amount > 40 and group == 'accessories' and $nodes.request.customer.country == 'US'" + } + ], + "inputs": [ + { + "id": "e9090e94-051c-4c27-be58-baac03041c2b", + "name": "Input" + } + ], + "outputs": [ + { + "id": "e06432d2-4609-425b-ad99-9ca27db08130", + "name": "Discount", + "field": "discount" + } + ], + "hitPolicy": "first", + "inputField": "items", + "outputPath": "tblItems", + "passThrough": true, + "executionMode": "loop" + }, + "position": { + "x": 490, + "y": 170 + } + }, + { + "id": "5f56e61b-5687-4758-8524-8121df13b2ba", + "name": "expr", + "type": "expressionNode", + "content": { + "inputField": "filter(items,#.group == 'accessories')", + "outputPath": "exprItems", + "expressions": [ + { + "id": "fcec4ba0-0279-46f5-99c4-fe4bfb6d842c", + "key": "discount", + "value": "amount > 40 and group == 'accessories' ? 5 : 0" + } + ], + "passThrough": true, + "executionMode": "loop" + }, + "position": { + "x": 490, + "y": 265 + } + }, + { + "type": "expressionNode", + "content": { + "expressions": [ + { + "id": "504955cc-8ea1-4cf1-a53b-1bf874dd51b1", + "key": "tblItemsSum", + "value": "sum(map(tblItems, #.amount))" + } + ], + "passThrough": true + }, + "id": "8ceca45e-e6ff-4613-984c-5846e53c5a39", + "name": "sum", + "position": { + "x": 785, + "y": 170 + } + }, + { + "type": "expressionNode", + "content": { + "expressions": [ + { + "id": "504955cc-8ea1-4cf1-a53b-1bf874dd51b1", + "key": "exprItemsSum", + "value": "sum(map(exprItems, #.amount))" + } + ], + "passThrough": true + }, + "id": "3814877d-0169-45dc-b4f5-5f9eade2e4f1", + "name": "sum", + "position": { + "x": 785, + "y": 265 + } + } + ], + "edges": [ + { + "id": "565a3028-2cd3-4d35-b46a-c5a776776d9e", + "type": "edge", + "sourceId": "be2a55f4-cdd8-482e-b53c-a947cbb7d7a0", + "targetId": "5f56e61b-5687-4758-8524-8121df13b2ba" + }, + { + "id": "b2f6aff4-2e3e-4e49-a8ff-c4b63f65e71e", + "type": "edge", + "sourceId": "be2a55f4-cdd8-482e-b53c-a947cbb7d7a0", + "targetId": "b8ebb212-e290-4131-9703-0bb7b1fd2328" + }, + { + "id": "69267b59-5f70-4cd2-825b-c1920901ac56", + "sourceId": "b8ebb212-e290-4131-9703-0bb7b1fd2328", + "type": "edge", + "targetId": "8ceca45e-e6ff-4613-984c-5846e53c5a39" + }, + { + "id": "01303153-a825-43d0-8127-73e4bac1514e", + "sourceId": "5f56e61b-5687-4758-8524-8121df13b2ba", + "type": "edge", + "targetId": "3814877d-0169-45dc-b4f5-5f9eade2e4f1" + } + ] +} \ No newline at end of file diff --git a/test-data/graphs/set-fee.json b/test-data/graphs/set-fee.json new file mode 100644 index 00000000..04225f34 --- /dev/null +++ b/test-data/graphs/set-fee.json @@ -0,0 +1,243 @@ +{ + "tests": [ + { + "input": { + "country": "US", + "currency": "EUR", + "amount": 100 + }, + "output": { + "cat_2": 4, + "currency": "EUR", + "percent": 4, + "cat_1": 3, + "country": "US", + "fee": 0.6, + "amount": 100 + } + }, + { + "input": { + "country": "US", + "currency": "USD", + "amount": 110 + }, + "output": { + "percent": 1, + "amount": 110, + "currency": "USD", + "fee": 0.5, + "country": "US" + } + }, + { + "input": { + "country": "CA", + "currency": "AUD", + "amount": 110 + }, + "output": { + "currency": "AUD", + "percent": 3, + "country": "CA", + "cat_1": 3, + "fee": 0.7, + "amount": 110, + "cat_2": 4 + } + } + ], + "contentType": "application/vnd.gorules.decision", + "nodes": [ + { + "id": "e0f3612a-6019-4647-a682-4f3019b67fce", + "name": "request", + "type": "inputNode", + "position": { + "x": 165, + "y": 120 + } + }, + { + "id": "3a801d19-b611-4163-bfdd-f247a689a275", + "name": "pricing", + "type": "decisionTableNode", + "content": { + "rules": [ + { + "_id": "6d08b48b-fe57-4efe-8d97-1bce5fcc75e8", + "1a981ce1-aa27-460c-a875-250d93d5cdbb": "0.5", + "6c6ac20d-cb40-4af9-b8f2-b5fc409467b6": "country == 'US' and currency == 'USD'", + "d4cb5b78-e7a2-4214-8e0c-0f5499d51395": "1" + }, + { + "_id": "9d9103e9-a4e6-4a28-b4a1-49e19839727f", + "1a981ce1-aa27-460c-a875-250d93d5cdbb": "1", + "6c6ac20d-cb40-4af9-b8f2-b5fc409467b6": "country == 'US' and currency == 'CAD'", + "d4cb5b78-e7a2-4214-8e0c-0f5499d51395": "" + } + ], + "inputs": [ + { + "id": "6c6ac20d-cb40-4af9-b8f2-b5fc409467b6", + "name": "Input" + } + ], + "outputs": [ + { + "id": "1a981ce1-aa27-460c-a875-250d93d5cdbb", + "name": "Fee", + "field": "fee" + }, + { + "id": "d4cb5b78-e7a2-4214-8e0c-0f5499d51395", + "name": "Percent", + "field": "percent" + } + ], + "hitPolicy": "first", + "passThrough": true + }, + "position": { + "x": 430, + "y": 120 + } + }, + { + "id": "bf832d94-d01d-4842-a60b-3db28375bccb", + "name": "fee set?", + "type": "switchNode", + "content": { + "hitPolicy": "first", + "statements": [ + { + "id": "94664df3-8233-4019-8381-affec373bffe", + "condition": "percent == null", + "isDefault": false + }, + { + "id": "718244e6-9471-4d20-af18-fce456d36f7e", + "condition": "", + "isDefault": true + } + ] + }, + "position": { + "x": 710, + "y": 120 + } + }, + { + "id": "09ecc83e-f89b-491b-9bfb-39311323ab94", + "name": "vars", + "type": "expressionNode", + "content": { + "expressions": [ + { + "id": "f7e02818-b2aa-4659-9095-4dbebda02a64", + "key": "cat_1", + "value": "3" + }, + { + "id": "7cebe9d1-5b10-44c3-8f74-36faafa1bde1", + "key": "cat_2", + "value": "4" + } + ], + "passThrough": true + }, + "position": { + "x": 985, + "y": 120 + } + }, + { + "id": "5e3c541a-7cf3-4765-b4c0-e20d07510379", + "name": "set percent", + "type": "decisionTableNode", + "content": { + "rules": [ + { + "_id": "34388171-596a-4159-bc8b-493a9858e992", + "da65cb09-9dd2-4a08-af3c-f59569f999e2": "currency == 'AUD'", + "c706d872-d774-461c-bd4e-0e26dc6f6f5f": "cat_1", + "bbedc722-3422-40f0-bbd1-15d407d9b348": "0.7" + }, + { + "_id": "dbe9b32c-73be-4d16-b67e-fcc087a38a86", + "da65cb09-9dd2-4a08-af3c-f59569f999e2": "currency == 'EUR'", + "c706d872-d774-461c-bd4e-0e26dc6f6f5f": "cat_2", + "bbedc722-3422-40f0-bbd1-15d407d9b348": "0.6" + } + ], + "inputs": [ + { + "id": "da65cb09-9dd2-4a08-af3c-f59569f999e2", + "name": "Input" + } + ], + "outputs": [ + { + "id": "c706d872-d774-461c-bd4e-0e26dc6f6f5f", + "name": "Percent", + "field": "percent" + }, + { + "id": "bbedc722-3422-40f0-bbd1-15d407d9b348", + "field": "fee", + "name": "Fee" + } + ], + "hitPolicy": "first", + "passThrough": true + }, + "position": { + "x": 1260, + "y": 120 + } + }, + { + "id": "1a8c6564-c566-44ff-b07a-bce27d7f5a02", + "name": "response", + "type": "outputNode", + "position": { + "x": 985, + "y": 225 + } + } + ], + "edges": [ + { + "id": "540ab38b-c0b1-4147-8826-c764766738a5", + "type": "edge", + "sourceId": "3a801d19-b611-4163-bfdd-f247a689a275", + "targetId": "bf832d94-d01d-4842-a60b-3db28375bccb" + }, + { + "id": "47000d98-4e3c-4276-a479-6febfe4740cb", + "type": "edge", + "sourceId": "e0f3612a-6019-4647-a682-4f3019b67fce", + "targetId": "3a801d19-b611-4163-bfdd-f247a689a275" + }, + { + "id": "b577da62-7398-46d0-a8ce-8ca23f08e20a", + "type": "edge", + "sourceId": "bf832d94-d01d-4842-a60b-3db28375bccb", + "targetId": "09ecc83e-f89b-491b-9bfb-39311323ab94", + "sourceHandle": "94664df3-8233-4019-8381-affec373bffe" + }, + { + "id": "2032b991-a0d0-4ba3-859d-7e1a805c7d8b", + "type": "edge", + "sourceId": "09ecc83e-f89b-491b-9bfb-39311323ab94", + "targetId": "5e3c541a-7cf3-4765-b4c0-e20d07510379" + }, + { + "id": "2361a49d-2c8d-4763-b638-057640b80dfb", + "type": "edge", + "sourceId": "bf832d94-d01d-4842-a60b-3db28375bccb", + "targetId": "1a8c6564-c566-44ff-b07a-bce27d7f5a02", + "sourceHandle": "718244e6-9471-4d20-af18-fce456d36f7e" + } + ] +} \ No newline at end of file