From 302e82afbd0ef313abf7f495d2a447e1b7dd27db Mon Sep 17 00:00:00 2001 From: Stefan Date: Fri, 18 Oct 2024 20:09:58 +0200 Subject: [PATCH] feat: passthrough nodes --- core/engine/Cargo.toml | 2 +- core/engine/src/decision.rs | 4 +- .../engine/src/handler/custom_node_adapter.rs | 43 +++++---- core/engine/src/handler/decision.rs | 91 +++++-------------- core/engine/src/handler/expression/mod.rs | 65 +++++++++---- core/engine/src/handler/function/mod.rs | 2 +- .../engine/src/handler/function/module/zen.rs | 2 +- core/engine/src/handler/function_v1/mod.rs | 2 +- core/engine/src/handler/graph.rs | 72 ++++++++------- core/engine/src/handler/node.rs | 7 +- core/engine/src/handler/table/zen.rs | 52 ++++++++--- core/engine/src/handler/traversal.rs | 19 +++- core/engine/src/lib.rs | 1 + core/engine/src/model/mod.rs | 38 +++++++- core/engine/src/util/mod.rs | 1 + core/engine/src/util/transform_attribute.rs | 62 +++++++++++++ core/engine/tests/decision.rs | 7 -- core/engine/tests/engine.rs | 36 +++++--- core/expression/src/parser/parser.rs | 6 +- core/expression/src/variable/types/util.rs | 2 +- 20 files changed, 325 insertions(+), 189 deletions(-) create mode 100644 core/engine/src/util/mod.rs create mode 100644 core/engine/src/util/transform_attribute.rs 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..38d5285c 100644 --- a/core/engine/src/handler/decision.rs +++ b/core/engine/src/handler/decision.rs @@ -3,15 +3,14 @@ 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; +use zen_expression::Isolate; pub struct DecisionHandler { trace: bool, @@ -40,7 +39,7 @@ impl DecisionHandle pub fn handle<'s, 'arg, 'recursion>( &'s self, - request: &'arg NodeRequest<'_>, + request: NodeRequest, ) -> Pin + 'recursion>> where 's: 'recursion, @@ -52,11 +51,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 +62,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..aa5727ff 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(), @@ -302,15 +295,15 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr }); walker.set_node_data(nid, res.output); } - DecisionNodeKind::DecisionTableNode { .. } => { + DecisionNodeKind::DecisionTableNode { content } => { 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) + let mut res = DecisionTableHandler::new(self.trace) + .handle(node_request.clone()) .await .map_err(|e| NodeError { node_id: node.id.clone(), @@ -320,6 +313,11 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr node_request.input.dot_remove("$nodes"); res.output.dot_remove("$nodes"); + if content.pass_through { + let mut base = node_request.input.clone(); + res.output = base.merge(&res.output) + } + trace!({ input: node_request.input, output: res.output.clone(), @@ -330,15 +328,15 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr }); walker.set_node_data(nid, res.output); } - DecisionNodeKind::ExpressionNode { .. } => { + DecisionNodeKind::ExpressionNode { content } => { 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) + let mut res = ExpressionHandler::new(self.trace) + .handle(node_request.clone()) .await .map_err(|e| NodeError { node_id: node.id.clone(), @@ -348,6 +346,11 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr node_request.input.dot_remove("$nodes"); res.output.dot_remove("$nodes"); + if content.pass_through { + let mut base = node_request.input.clone(); + res.output = base.merge(&res.output) + } + trace!({ input: node_request.input, output: res.output.clone(), @@ -360,14 +363,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 +393,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..fb1602da 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,32 @@ 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, 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 +53,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..211433b6 100644 --- a/core/engine/src/handler/traversal.rs +++ b/core/engine/src/handler/traversal.rs @@ -8,6 +8,7 @@ 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, @@ -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,7 +145,7 @@ 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; } 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..6c571262 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,10 @@ pub struct DecisionTableContent { pub inputs: Vec, pub outputs: Vec, pub hit_policy: DecisionTableHitPolicy, + #[serde(default)] + pub pass_through: bool, + #[serde(flatten)] + pub transform_attributes: TransformAttributes, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -102,6 +107,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 +125,10 @@ pub struct DecisionTableOutputField { #[serde(rename_all = "camelCase")] pub struct ExpressionNodeContent { pub expressions: Vec, + #[serde(default)] + pub pass_through: bool, + #[serde(flatten)] + pub transform_attributes: TransformAttributes, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -160,7 +170,9 @@ 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, @@ -179,7 +191,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 +237,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..e9930e20 --- /dev/null +++ b/core/engine/src/util/transform_attribute.rs @@ -0,0 +1,62 @@ +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, input: Variable, evaluate: F) -> NodeResult + where + F: Fn(Variable) -> Fut, + Fut: Future, + { + let input = match &self.input_field { + None => input.clone(), + Some(input_field) => { + let mut isolate = Isolate::new(); + isolate.set_environment(input.clone()); + isolate.run_standard(input_field.as_str())? + } + }; + + 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 + } + 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 response = evaluate(input.clone()).await?; + + output_array.push(response.output); + if let Some(td) = response.trace_data { + trace_datum.push(td); + } + } + + 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; + } + + 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/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/types/util.rs b/core/expression/src/variable/types/util.rs index ce1d341e..2ba586ef 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))),