From 9dcb694c8c1a9dc44c8870771fb36d90f6403878 Mon Sep 17 00:00:00 2001 From: stefan-gorules <127550877+stefan-gorules@users.noreply.github.com> Date: Tue, 2 May 2023 18:07:06 +0200 Subject: [PATCH] refactor: cleanup decision tree (#32) * refactor: cleanup decision tree; * remove unused error; add graph tests; * add test and fix for depth limit --- core/engine/src/decision.rs | 12 +- core/engine/src/engine.rs | 26 +- core/engine/src/error.rs | 3 - core/engine/src/handler/decision.rs | 15 +- core/engine/src/handler/function/mod.rs | 2 +- core/engine/src/handler/function/script.rs | 20 +- core/engine/src/handler/graph.rs | 310 +++++++++++++++++ core/engine/src/handler/mod.rs | 2 +- core/engine/src/handler/table/mod.rs | 7 +- core/engine/src/handler/table/zen.rs | 20 +- core/engine/src/handler/tree.rs | 385 --------------------- test-data/recursive-table1.json | 49 +++ test-data/recursive-table2.json | 49 +++ 13 files changed, 466 insertions(+), 434 deletions(-) create mode 100644 core/engine/src/handler/graph.rs delete mode 100644 core/engine/src/handler/tree.rs create mode 100644 test-data/recursive-table1.json create mode 100644 test-data/recursive-table2.json diff --git a/core/engine/src/decision.rs b/core/engine/src/decision.rs index 254dab5f..a149848c 100644 --- a/core/engine/src/decision.rs +++ b/core/engine/src/decision.rs @@ -1,5 +1,5 @@ use crate::engine::EvaluationOptions; -use crate::handler::tree::{GraphResponse, GraphTree, GraphTreeConfig}; +use crate::handler::graph::{DecisionGraph, DecisionGraphConfig, DecisionGraphResponse}; use crate::loader::{DecisionLoader, NoopLoader}; use crate::model::DecisionContent; use crate::EvaluationError; @@ -49,7 +49,10 @@ where } /// Evaluates a decision using an in-memory reference stored in struct - pub async fn evaluate(&self, context: &Value) -> Result> { + pub async fn evaluate( + &self, + context: &Value, + ) -> Result> { self.evaluate_with_opts(context, Default::default()).await } @@ -58,8 +61,8 @@ where &self, context: &Value, options: EvaluationOptions, - ) -> Result> { - let tree = GraphTree::new(GraphTreeConfig { + ) -> Result> { + let tree = DecisionGraph::new(DecisionGraphConfig { max_depth: options.max_depth.unwrap_or(5), trace: options.trace.unwrap_or_default(), loader: self.loader.clone(), @@ -67,7 +70,6 @@ where content: &self.content, }); - tree.connect()?; Ok(tree.evaluate(context).await?) } } diff --git a/core/engine/src/engine.rs b/core/engine/src/engine.rs index f3491a0b..34116714 100644 --- a/core/engine/src/engine.rs +++ b/core/engine/src/engine.rs @@ -1,11 +1,11 @@ use crate::decision::Decision; -use crate::handler::tree::GraphResponse; use crate::loader::{ClosureLoader, DecisionLoader, LoaderResponse, LoaderResult, NoopLoader}; use crate::model::DecisionContent; use serde_json::Value; use std::future::Future; +use crate::handler::graph::DecisionGraphResponse; use crate::EvaluationError; use std::sync::Arc; @@ -63,7 +63,7 @@ impl DecisionEngine { &self, key: K, context: &Value, - ) -> Result> + ) -> Result> where K: AsRef, { @@ -77,7 +77,7 @@ impl DecisionEngine { key: K, context: &Value, options: EvaluationOptions, - ) -> Result> + ) -> Result> where K: AsRef, { @@ -191,4 +191,24 @@ mod tests { _ => assert!(false, "Wrong error type"), } } + + #[test] + fn it_terminates_when_depth_limit_exceeded() { + let cargo_root = Path::new(env!("CARGO_MANIFEST_DIR")); + let test_data_root = cargo_root.join("../../").join("test-data"); + let fs_loader = FilesystemLoader::new(FilesystemLoaderOptions { + keep_in_memory: true, + root: test_data_root.to_str().unwrap(), + }); + + let graph = DecisionEngine::new(fs_loader); + let recursive = tokio_test::block_on(graph.evaluate("recursive-table1.json", &json!({}))); + + match recursive.unwrap_err().deref() { + EvaluationError::NodeError(e) => { + assert_eq!(e.source.to_string(), "Depth limit exceeded") + } + _ => assert!(false, "Depth limit not exceeded"), + } + } } diff --git a/core/engine/src/error.rs b/core/engine/src/error.rs index f7aa71b3..31e4f2b6 100644 --- a/core/engine/src/error.rs +++ b/core/engine/src/error.rs @@ -12,9 +12,6 @@ pub enum EvaluationError { #[error("Depth limit exceeded")] DepthLimitExceeded, - - #[error("Node not found {0}")] - NodeConnectError(String), } impl From for Box { diff --git a/core/engine/src/handler/decision.rs b/core/engine/src/handler/decision.rs index 1052b52f..32e67cd0 100644 --- a/core/engine/src/handler/decision.rs +++ b/core/engine/src/handler/decision.rs @@ -1,5 +1,5 @@ +use crate::handler::graph::{DecisionGraph, DecisionGraphConfig}; use crate::handler::node::{NodeRequest, NodeResponse, NodeResult}; -use crate::handler::tree::{GraphTree, GraphTreeConfig}; use crate::loader::DecisionLoader; use crate::model::DecisionNodeKind; use anyhow::{anyhow, Context}; @@ -30,7 +30,7 @@ impl DecisionHandler { }?; let sub_decision = self.loader.load(&content.key).await?; - let sub_tree = GraphTree::new(GraphTreeConfig { + let sub_tree = DecisionGraph::new(DecisionGraphConfig { content: sub_decision.deref(), max_depth: self.max_depth, loader: self.loader.clone(), @@ -38,7 +38,6 @@ impl DecisionHandler { trace: self.trace, }); - sub_tree.connect()?; let result = sub_tree .evaluate(&request.input) .await @@ -46,12 +45,10 @@ impl DecisionHandler { Ok(NodeResponse { output: result.result, - trace_data: match self.trace { - true => { - Some(serde_json::to_value(result.trace).context("Failed to parse trace data")?) - } - false => None, - }, + trace_data: self + .trace + .then(|| serde_json::to_value(result.trace).context("Failed to parse trace data")) + .transpose()?, }) } } diff --git a/core/engine/src/handler/function/mod.rs b/core/engine/src/handler/function/mod.rs index b200d1bc..1df79b6e 100644 --- a/core/engine/src/handler/function/mod.rs +++ b/core/engine/src/handler/function/mod.rs @@ -10,7 +10,7 @@ use crate::model::DecisionNodeKind; mod script; mod vm; -pub async fn evaluate(source: &str, args: &Value) -> anyhow::Result { +async fn evaluate(source: &str, args: &Value) -> anyhow::Result { let mut script = Script::new().with_timeout(Duration::from_millis(50)); script.call(source, args).await } diff --git a/core/engine/src/handler/function/script.rs b/core/engine/src/handler/function/script.rs index 849d760c..7125b9c2 100644 --- a/core/engine/src/handler/function/script.rs +++ b/core/engine/src/handler/function/script.rs @@ -66,22 +66,17 @@ impl Script { let Some(src_script) = v8::Script::compile(tc_scope, src, None) else { let exception = tc_scope.exception().context("Failed to load script")?; - return Err(anyhow!(exception.to_rust_string_lossy(tc_scope).to_string())); + return Err(anyhow!(exception.to_rust_string_lossy(tc_scope))); }; - match src_script.run(tc_scope) { - Some(..) => {} - None => { - let exception = tc_scope.exception().unwrap(); - return Err(anyhow!(exception - .to_rust_string_lossy(tc_scope) - .to_string())); - } + if let None = src_script.run(tc_scope) { + let exception = tc_scope.exception().unwrap(); + return Err(anyhow!(exception.to_rust_string_lossy(tc_scope))); } let Some(js_script) = v8::Script::compile(tc_scope, js_src, None) else { let exception = tc_scope.exception().context("Failed to load script")?; - return Err(anyhow!(exception.to_rust_string_lossy(tc_scope).to_string())); + return Err(anyhow!(exception.to_rust_string_lossy(tc_scope))); }; let Some(result) = js_script.run(tc_scope) else { @@ -90,10 +85,9 @@ impl Script { } let exception = tc_scope.exception().context("Failed to run loaded script")?; - return Err(anyhow!(exception.to_rust_string_lossy(tc_scope).to_string())); + return Err(anyhow!(exception.to_rust_string_lossy(tc_scope))); }; - let res: EvaluateResponse = serde_v8::from_v8(tc_scope, result).unwrap(); - Ok(res) + serde_v8::from_v8(tc_scope, result).context("Failed to parse function result") } } diff --git a/core/engine/src/handler/graph.rs b/core/engine/src/handler/graph.rs new file mode 100644 index 00000000..99989439 --- /dev/null +++ b/core/engine/src/handler/graph.rs @@ -0,0 +1,310 @@ +use crate::loader::DecisionLoader; +use crate::model::{DecisionContent, DecisionNode, DecisionNodeKind}; +use std::collections::HashMap; + +use crate::handler::decision::DecisionHandler; +use crate::handler::function::FunctionHandler; +use crate::handler::node::NodeRequest; +use crate::handler::table::zen::DecisionTableHandler; + +use crate::{EvaluationError, NodeError}; +use anyhow::anyhow; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; +use std::sync::Arc; +use std::time::Instant; + +pub struct DecisionGraph<'a, T: DecisionLoader> { + nodes: Vec>, + loader: Arc, + trace: bool, + max_depth: u8, + iteration: u8, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct DecisionGraphResponse { + pub performance: String, + pub result: Value, + #[serde(skip_serializing_if = "Option::is_none")] + pub trace: Option>, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct DecisionGraphTrace { + input: Value, + output: Value, + name: String, + id: String, + performance: Option, + trace_data: Option, +} + +pub struct DecisionGraphConfig<'a, T: DecisionLoader> { + pub loader: Arc, + pub content: &'a DecisionContent, + pub trace: bool, + pub iteration: u8, + pub max_depth: u8, +} + +impl<'a, T: DecisionLoader> DecisionGraph<'a, T> { + pub fn new(config: DecisionGraphConfig<'a, T>) -> Self { + let nodes = config + .content + .nodes + .iter() + .map(|node| { + let parents: Vec<&'a str> = config + .content + .edges + .iter() + .filter(|edge| edge.target_id == node.id) + .map(|edge| edge.source_id.as_str()) + .collect(); + + DecisionGraphNode { parents, node } + }) + .collect(); + + Self { + nodes, + max_depth: config.max_depth, + iteration: config.iteration, + trace: config.trace, + loader: config.loader, + } + } + + pub async fn evaluate(&self, state: &Value) -> Result { + if self.iteration >= self.max_depth { + return Err(NodeError { + node_id: "".to_string(), + source: anyhow!(EvaluationError::DepthLimitExceeded), + }); + } + + let root_start = Instant::now(); + let mut node_data = HashMap::<&str, Value>::default(); + let mut node_traces = self.trace.then(|| HashMap::default()); + + for graph_node in &self.nodes { + let node = graph_node.node; + let start = Instant::now(); + + macro_rules! trace { + ($data: tt) => { + if let Some(nt) = &mut node_traces { + nt.insert(node.id.clone(), DecisionGraphTrace $data); + }; + }; + } + + match node.kind { + DecisionNodeKind::InputNode => { + node_data.insert(&node.id, state.clone()); + trace!({ + input: Value::Null, + output: Value::Null, + name: node.name.clone(), + id: node.id.clone(), + performance: None, + trace_data: None, + }); + } + DecisionNodeKind::OutputNode => { + trace!({ + input: Value::Null, + output: Value::Null, + name: node.name.clone(), + id: node.id.clone(), + performance: None, + trace_data: None, + }); + + return Ok(DecisionGraphResponse { + result: graph_node.parent_data(&node_data)?, + performance: format!("{:?}", root_start.elapsed()), + trace: node_traces, + }); + } + DecisionNodeKind::FunctionNode { .. } => { + let input = graph_node.parent_data(&node_data)?; + let req = NodeRequest { + node, + iteration: self.iteration, + input, + }; + + let res = FunctionHandler::new(self.trace) + .handle(&req) + .await + .map_err(|e| NodeError { + source: e.into(), + node_id: node.id.clone(), + })?; + + node_data.insert(&node.id, res.output.clone()); + trace!({ + input: req.input, + output: res.output, + name: node.name.clone(), + id: node.id.clone(), + performance: Some(format!("{:?}", start.elapsed())), + trace_data: res.trace_data, + }); + } + DecisionNodeKind::DecisionNode { .. } => { + let input = graph_node.parent_data(&node_data)?; + + let req = NodeRequest { + node, + iteration: self.iteration, + input, + }; + + let res = DecisionHandler::new(self.trace, self.max_depth, self.loader.clone()) + .handle(&req) + .await + .map_err(|e| NodeError { + source: e.into(), + node_id: node.id.to_string(), + })?; + + node_data.insert(&node.id, res.output.clone()); + trace!({ + input: req.input, + output: res.output, + name: node.name.clone(), + id: node.id.clone(), + performance: Some(format!("{:?}", start.elapsed())), + trace_data: res.trace_data, + }); + } + DecisionNodeKind::DecisionTableNode { .. } => { + let input = graph_node.parent_data(&node_data)?; + + let req = NodeRequest { + node, + iteration: self.iteration, + input, + }; + + let res = DecisionTableHandler::new(self.trace) + .handle(&req) + .await + .map_err(|e| NodeError { + node_id: node.id.clone(), + source: e.into(), + })?; + + node_data.insert(&node.id, res.output.clone()); + trace!({ + input: req.input, + output: res.output, + name: node.name.clone(), + id: node.id.clone(), + performance: Some(format!("{:?}", start.elapsed())), + trace_data: res.trace_data, + }); + } + } + } + + Err(NodeError { + node_id: "".to_string(), + source: anyhow!("Graph did not halt. Missing output node."), + }) + } +} + +struct DecisionGraphNode<'a> { + parents: Vec<&'a str>, + node: &'a DecisionNode, +} + +impl<'a> DecisionGraphNode<'a> { + pub fn parent_data(&self, node_data: &HashMap<&str, Value>) -> Result { + let mut object = Value::Object(Map::new()); + + for pid in &self.parents { + let data = node_data.get(pid).ok_or_else(|| NodeError { + node_id: self.node.id.clone(), + source: anyhow!("Failed to parse node data"), + })?; + + merge_json(&mut object, data, true); + } + + Ok(object) + } +} + +fn merge_json(doc: &mut Value, patch: &Value, top_level: bool) { + if !patch.is_object() && !patch.is_array() && top_level { + return; + } + + if doc.is_object() && patch.is_object() { + let map = doc.as_object_mut().unwrap(); + for (key, value) in patch.as_object().unwrap() { + if value.is_null() { + map.remove(key.as_str()); + } else { + merge_json(map.entry(key.as_str()).or_insert(Value::Null), value, false); + } + } + } else if doc.is_array() && patch.is_array() { + let arr = doc.as_array_mut().unwrap(); + arr.extend(patch.as_array().unwrap().clone()); + } else { + *doc = patch.clone(); + } +} + +#[cfg(test)] +mod tests { + use crate::handler::graph::{DecisionGraph, DecisionGraphConfig}; + use crate::loader::MemoryLoader; + use serde_json::json; + use std::sync::Arc; + + #[test] + fn decision_table() { + let content = + &serde_json::from_str(include_str!("../../../../test-data/table.json")).unwrap(); + let tree = DecisionGraph::new(DecisionGraphConfig { + max_depth: 5, + trace: false, + iteration: 0, + content, + loader: Arc::new(MemoryLoader::default()), + }); + + let result = + tokio_test::block_on(async { tree.evaluate(&json!({ "input": 15 })).await.unwrap() }); + + assert_eq!(result.result, json!({ "output": 10 })); + } + + #[test] + #[cfg_attr(miri, ignore)] + fn function() { + let content = + &serde_json::from_str(include_str!("../../../../test-data/function.json")).unwrap(); + let tree = DecisionGraph::new(DecisionGraphConfig { + max_depth: 5, + trace: false, + iteration: 0, + content, + loader: Arc::new(MemoryLoader::default()), + }); + + let result = + tokio_test::block_on(async { tree.evaluate(&json!({ "input": 15 })).await.unwrap() }); + + assert_eq!(result.result, json!({ "output": 30 })); + } +} diff --git a/core/engine/src/handler/mod.rs b/core/engine/src/handler/mod.rs index 3bdf81ca..c3a661c8 100644 --- a/core/engine/src/handler/mod.rs +++ b/core/engine/src/handler/mod.rs @@ -2,5 +2,5 @@ pub mod decision; pub mod function; pub mod table; +pub(crate) mod graph; pub(crate) mod node; -pub(crate) mod tree; diff --git a/core/engine/src/handler/table/mod.rs b/core/engine/src/handler/table/mod.rs index 1e3dd837..b04334b9 100644 --- a/core/engine/src/handler/table/mod.rs +++ b/core/engine/src/handler/table/mod.rs @@ -48,20 +48,19 @@ impl RowOutput { let mut result: BTreeMap = BTreeMap::new(); for inner_map in map { for (key, value) in inner_map { - let rk = RowKey::from(key); match value { // Unexpected, as we've filtered out all objects in prior step Value::Object(_) => return Err(RowOutputError::FailedToParse), Value::Array(arr) => { - let maybe_exist = result.get_mut(&rk).map(|a| a.as_array_mut()).flatten(); + let maybe_exist = result.get_mut(&key).map(|a| a.as_array_mut()).flatten(); if let Some(exist) = maybe_exist { exist.extend_from_slice(&arr); } else { - result.insert(rk, Value::Array(arr)); + result.insert(key, Value::Array(arr)); } } _ => { - result.insert(rk, value); + result.insert(key, value); } } } diff --git a/core/engine/src/handler/table/zen.rs b/core/engine/src/handler/table/zen.rs index e300ea1b..c473d124 100644 --- a/core/engine/src/handler/table/zen.rs +++ b/core/engine/src/handler/table/zen.rs @@ -51,12 +51,12 @@ impl<'a> DecisionTableHandler<'a> { if let Some(result) = self.evaluate_row(&content, i).await { return Ok(NodeResponse { output: result.output.to_json().await?, - trace_data: match self.trace { - true => Some( - serde_json::to_value(result).context("Failed to parse trace data")?, - ), - false => None, - }, + trace_data: self + .trace + .then(|| { + serde_json::to_value(&result).context("Failed to parse trace data") + }) + .transpose()?, }); } } @@ -82,10 +82,10 @@ impl<'a> DecisionTableHandler<'a> { Ok(NodeResponse { output: serde_json::to_value(&outputs).context("Failed to parse table row output")?, - trace_data: match self.trace { - true => Some(serde_json::to_value(&results).context("Failed to parse trace data")?), - false => None, - }, + trace_data: self + .trace + .then(|| serde_json::to_value(&results).context("Failed to parse trace data")) + .transpose()?, }) } diff --git a/core/engine/src/handler/tree.rs b/core/engine/src/handler/tree.rs deleted file mode 100644 index e0ae4002..00000000 --- a/core/engine/src/handler/tree.rs +++ /dev/null @@ -1,385 +0,0 @@ -use crate::handler::decision::DecisionHandler; -use crate::handler::function::FunctionHandler; -use crate::handler::node::{NodeError, NodeRequest}; -use crate::handler::table::zen::DecisionTableHandler; -use crate::loader::DecisionLoader; -use crate::model::{DecisionContent, DecisionNode, DecisionNodeKind}; -use crate::EvaluationError; -use anyhow::anyhow; -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; -use std::cell::RefCell; -use std::collections::HashMap; -use std::ops::Deref; -use std::sync::Arc; -use std::time::Instant; - -#[derive(Debug)] -pub struct GraphTreeNode<'a> { - pub parents: Vec<&'a str>, - pub node: &'a DecisionNode, - pub data: Value, - pub trace: Value, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GraphTrace { - input: Value, - output: Value, - name: String, - id: String, - performance: Option, - trace_data: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct GraphResponse { - pub performance: String, - pub result: Value, - #[serde(skip_serializing_if = "Option::is_none")] - pub trace: Option>, -} - -impl<'a> From<&'a DecisionNode> for GraphTreeNode<'a> { - fn from(node: &'a DecisionNode) -> Self { - Self { - node, - parents: Default::default(), - data: Value::Null, - trace: Value::Null, - } - } -} - -pub struct GraphTree<'a, T: DecisionLoader> { - nodes: RefCell>>>, - node_ids: RefCell>, - node_data: RefCell>, - node_trace: RefCell>, - content: &'a DecisionContent, - trace: bool, - max_depth: u8, - loader: Arc, - pub iteration: u8, -} - -pub struct GraphTreeConfig<'a, T: DecisionLoader> { - pub loader: Arc, - pub content: &'a DecisionContent, - pub trace: bool, - pub iteration: u8, - pub max_depth: u8, -} - -impl<'a, T: DecisionLoader> GraphTree<'a, T> { - pub fn new(config: GraphTreeConfig<'a, T>) -> Self { - Self { - nodes: Default::default(), - node_ids: Default::default(), - node_data: Default::default(), - node_trace: Default::default(), - max_depth: config.max_depth, - iteration: config.iteration, - trace: config.trace, - content: config.content, - loader: config.loader, - } - } - - pub fn connect(&self) -> Result<(), EvaluationError> { - if self.iteration >= self.max_depth { - return Err(EvaluationError::DepthLimitExceeded); - } - - let mut node_ids = self.node_ids.borrow_mut(); - let mut nodes = self.nodes.borrow_mut(); - - self.content.nodes.iter().for_each(|node| { - let key = node.id.as_str(); - node_ids.push(key); - nodes.insert(key, RefCell::new(GraphTreeNode::from(node))); - }); - - self.content - .edges - .iter() - .try_for_each::<_, Result<(), EvaluationError>>(|edge| { - let source_ref = nodes - .get(edge.source_id.as_str()) - .ok_or_else(|| EvaluationError::NodeConnectError(edge.source_id.to_string()))?; - - let target_ref = nodes - .get(edge.target_id.as_str()) - .ok_or_else(|| EvaluationError::NodeConnectError(edge.target_id.to_string()))?; - - let source = source_ref.borrow(); - let mut target = target_ref.borrow_mut(); - - target.parents.push(source.node.id.as_str()); - - Ok(()) - })?; - - Ok(()) - } - - fn merge(doc: &mut Value, patch: &Value, top_level: bool) { - if !patch.is_object() && !patch.is_array() && top_level { - return; - } - - if doc.is_object() && patch.is_object() { - let map = doc.as_object_mut().unwrap(); - for (key, value) in patch.as_object().unwrap() { - if value.is_null() { - map.remove(key.as_str()); - } else { - Self::merge(map.entry(key.as_str()).or_insert(Value::Null), value, false); - } - } - } else if doc.is_array() && patch.is_array() { - let arr = doc.as_array_mut().unwrap(); - arr.extend(patch.as_array().unwrap().clone()); - } else { - *doc = patch.clone(); - } - } - - fn parent_data( - &self, - node: &GraphTreeNode, - node_data: &HashMap<&'a str, Value>, - ) -> anyhow::Result { - let mut object = Value::Object(Map::new()); - - for pid in &node.parents { - let data = node_data - .get(pid) - .ok_or_else(|| anyhow!("Failed to parse node data"))?; - - Self::merge(&mut object, data, true); - } - - Ok(object) - } - - fn trace(&self, key: String, trace: GraphTrace) { - if !self.trace { - return; - } - - let mut node_trace = self.node_trace.borrow_mut(); - node_trace.insert(key, trace); - } - - pub async fn evaluate(&self, state: &Value) -> Result { - let initial_start = Instant::now(); - let node_ids = self.node_ids.borrow(); - let nodes = self.nodes.borrow(); - let mut node_data = self.node_data.borrow_mut(); - - for id in node_ids.deref() { - let start = Instant::now(); - let node_ref = nodes.get(id).ok_or_else(|| NodeError { - source: anyhow!("Failed to parse a reference"), - node_id: id.to_string(), - })?; - let node = node_ref.borrow(); - - match &node.node.kind { - DecisionNodeKind::InputNode => { - node_data.insert(id, state.clone()); - self.trace( - id.to_string(), - GraphTrace { - input: Value::Null, - output: Value::Null, - name: node.node.name.clone(), - id: node.node.id.clone(), - performance: None, - trace_data: None, - }, - ); - } - - DecisionNodeKind::OutputNode => { - self.trace( - id.to_string(), - GraphTrace { - input: Value::Null, - output: Value::Null, - name: node.node.name.clone(), - id: node.node.id.clone(), - performance: None, - trace_data: None, - }, - ); - - return Ok(GraphResponse { - result: self.parent_data(&node, &node_data).map_err(|e| NodeError { - source: e, - node_id: id.to_string(), - })?, - performance: format!("{:?}", initial_start.elapsed()), - trace: match self.trace { - true => Some(self.node_trace.borrow().clone()), - false => None, - }, - }); - } - DecisionNodeKind::FunctionNode { .. } => { - let input = self.parent_data(&node, &node_data).map_err(|e| NodeError { - node_id: id.to_string(), - source: e, - })?; - let req = NodeRequest { - node: node.node, - iteration: self.iteration, - input, - }; - - let res = FunctionHandler::new(self.trace) - .handle(&req) - .await - .map_err(|e| NodeError { - source: e.into(), - node_id: node.node.id.clone(), - })?; - - node_data.insert(id, res.output.clone()); - self.trace( - id.to_string(), - GraphTrace { - input: req.input, - output: res.output, - name: node.node.name.clone(), - id: node.node.id.clone(), - performance: Some(format!("{:?}", start.elapsed())), - trace_data: res.trace_data, - }, - ); - } - DecisionNodeKind::DecisionNode { .. } => { - let input = self.parent_data(&node, &node_data).map_err(|e| NodeError { - source: e, - node_id: id.to_string(), - })?; - - let req = NodeRequest { - node: node.node, - iteration: self.iteration, - input, - }; - - let res = DecisionHandler::new(self.trace, self.max_depth, self.loader.clone()) - .handle(&req) - .await - .map_err(|e| NodeError { - source: e.into(), - node_id: id.to_string(), - })?; - - node_data.insert(id, res.output.clone()); - self.trace( - id.to_string(), - GraphTrace { - input: req.input, - output: res.output, - name: node.node.name.clone(), - id: node.node.id.clone(), - performance: Some(format!("{:?}", start.elapsed())), - trace_data: res.trace_data, - }, - ); - } - DecisionNodeKind::DecisionTableNode { .. } => { - let input = self.parent_data(&node, &node_data).map_err(|e| NodeError { - source: e, - node_id: id.to_string(), - })?; - - let req = NodeRequest { - node: node.node, - iteration: self.iteration, - input, - }; - - let res = DecisionTableHandler::new(self.trace) - .handle(&req) - .await - .map_err(|e| NodeError { - node_id: id.to_string(), - source: e.into(), - })?; - - node_data.insert(id, res.output.clone()); - self.trace( - id.to_string(), - GraphTrace { - input: req.input, - output: res.output, - name: node.node.name.clone(), - id: node.node.id.clone(), - performance: Some(format!("{:?}", start.elapsed())), - trace_data: res.trace_data, - }, - ); - } - }; - } - - Err(NodeError { - node_id: "".to_string(), - source: anyhow!("Graph did not halt. Missing output node."), - }) - } -} - -#[cfg(test)] -mod tests { - use crate::handler::tree::{GraphTree, GraphTreeConfig}; - use crate::loader::MemoryLoader; - use serde_json::json; - use std::sync::Arc; - - #[test] - fn decision_table() { - let content = - &serde_json::from_str(include_str!("../../../../test-data/table.json")).unwrap(); - let tree = GraphTree::new(GraphTreeConfig { - max_depth: 5, - trace: false, - iteration: 0, - content, - loader: Arc::new(MemoryLoader::default()), - }); - - tree.connect().unwrap(); - let result = - tokio_test::block_on(async { tree.evaluate(&json!({ "input": 15 })).await.unwrap() }); - - assert_eq!(result.result, json!({ "output": 10 })); - } - - #[test] - #[cfg_attr(miri, ignore)] - fn function() { - let content = - &serde_json::from_str(include_str!("../../../../test-data/function.json")).unwrap(); - let tree = GraphTree::new(GraphTreeConfig { - max_depth: 5, - trace: false, - iteration: 0, - content, - loader: Arc::new(MemoryLoader::default()), - }); - - tree.connect().unwrap(); - let result = - tokio_test::block_on(async { tree.evaluate(&json!({ "input": 15 })).await.unwrap() }); - - assert_eq!(result.result, json!({ "output": 30 })); - } -} diff --git a/test-data/recursive-table1.json b/test-data/recursive-table1.json new file mode 100644 index 00000000..996beb6f --- /dev/null +++ b/test-data/recursive-table1.json @@ -0,0 +1,49 @@ +{ + "contentType": "application/vnd.gorules.decision", + "edges": [ + { + "id": "0f5ca374-811e-44da-b882-0d5566f43b65", + "type": "edge", + "sourceId": "341e36a6-be77-44e1-99a5-d7c7ff1b7aba", + "targetId": "0b8dcf6b-fc04-47cb-bf82-bda764e6c09b" + }, + { + "id": "f07bc0ac-05f1-43ce-b942-9ba7e0bca430", + "type": "edge", + "sourceId": "0b8dcf6b-fc04-47cb-bf82-bda764e6c09b", + "targetId": "513e554b-ecf7-42f8-97fc-caca1849930c" + } + ], + "nodes": [ + { + "id": "341e36a6-be77-44e1-99a5-d7c7ff1b7aba", + "name": "Request", + "type": "inputNode", + "position": { + "x": 40, + "y": 240 + } + }, + { + "id": "0b8dcf6b-fc04-47cb-bf82-bda764e6c09b", + "name": "inf2", + "type": "decisionNode", + "content": { + "key": "recursive-table2.json" + }, + "position": { + "x": 370, + "y": 240 + } + }, + { + "id": "513e554b-ecf7-42f8-97fc-caca1849930c", + "name": "Response", + "type": "outputNode", + "position": { + "x": 710, + "y": 240 + } + } + ] +} \ No newline at end of file diff --git a/test-data/recursive-table2.json b/test-data/recursive-table2.json new file mode 100644 index 00000000..5280966a --- /dev/null +++ b/test-data/recursive-table2.json @@ -0,0 +1,49 @@ +{ + "contentType": "application/vnd.gorules.decision", + "edges": [ + { + "id": "47fd6cf6-fc8b-41ac-af3a-b59022b51911", + "type": "edge", + "sourceId": "137d8749-f625-4b50-a3c2-e0ec96f6a8bf", + "targetId": "c7e1277c-4f1d-4073-a132-e00b832d0061" + }, + { + "id": "5999c500-fc6c-4177-842f-e5ff1cefc8a1", + "type": "edge", + "sourceId": "c7e1277c-4f1d-4073-a132-e00b832d0061", + "targetId": "8e5f573b-6eab-45e9-927d-65ea5da6d5f8" + } + ], + "nodes": [ + { + "id": "137d8749-f625-4b50-a3c2-e0ec96f6a8bf", + "name": "Request", + "type": "inputNode", + "position": { + "x": 130, + "y": 210 + } + }, + { + "id": "c7e1277c-4f1d-4073-a132-e00b832d0061", + "name": "inf1", + "type": "decisionNode", + "content": { + "key": "recursive-table1.json" + }, + "position": { + "x": 430, + "y": 210 + } + }, + { + "id": "8e5f573b-6eab-45e9-927d-65ea5da6d5f8", + "name": "Response", + "type": "outputNode", + "position": { + "x": 740, + "y": 210 + } + } + ] +} \ No newline at end of file