From 042eb232e67530ff66ff06a7991666e695633dfe Mon Sep 17 00:00:00 2001 From: Stefan Date: Mon, 1 Sep 2025 20:10:00 +0200 Subject: [PATCH] implement remaining nodes --- Cargo.toml | 2 +- core/engine/src/decision.rs | 86 ++-- core/engine/src/decision_graph/graph.rs | 456 +++++------------- core/engine/src/decision_graph/mod.rs | 2 +- .../{traversal.rs => walker.rs} | 136 +++--- core/engine/src/engine.rs | 50 +- core/engine/src/error.rs | 43 +- core/engine/src/lib.rs | 4 +- core/engine/src/loader/mod.rs | 4 +- core/engine/src/nodes/context.rs | 84 +++- core/engine/src/nodes/custom/adapter.rs | 21 +- core/engine/src/nodes/custom/mod.rs | 13 +- core/engine/src/nodes/decision/mod.rs | 54 +++ core/engine/src/nodes/decision_table/mod.rs | 40 +- core/engine/src/nodes/expression/mod.rs | 11 +- core/engine/src/nodes/extensions.rs | 92 ++-- core/engine/src/nodes/function/mod.rs | 10 +- .../src/nodes/function/v2/module/zen.rs | 22 +- core/engine/src/nodes/input/mod.rs | 14 +- core/engine/src/nodes/mod.rs | 26 +- core/engine/src/nodes/output/mod.rs | 12 +- core/engine/src/nodes/result.rs | 2 +- core/engine/src/nodes/validator_cache.rs | 23 +- core/types/src/decision/mod.rs | 22 +- core/types/src/variable/impls.rs | 15 + 25 files changed, 570 insertions(+), 674 deletions(-) rename core/engine/src/decision_graph/{traversal.rs => walker.rs} (73%) diff --git a/Cargo.toml b/Cargo.toml index 0b5b525f..331f4d3e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,7 +12,7 @@ chrono = "0.4" criterion = "0.5" fastrand = "2" humantime = "2" -tokio = "1" +tokio = { version = "1", features = ["rt-multi-thread"] } tokio-util = "0.7" once_cell = "1" petgraph = "0.8" diff --git a/core/engine/src/decision.rs b/core/engine/src/decision.rs index 33ec9e5f..28e91254 100644 --- a/core/engine/src/decision.rs +++ b/core/engine/src/decision.rs @@ -1,25 +1,24 @@ use crate::decision_graph::graph::{DecisionGraph, DecisionGraphConfig, DecisionGraphResponse}; use crate::engine::{EvaluationOptions, EvaluationSerializedOptions, EvaluationTraceKind}; -use crate::handler::custom_node_adapter::{CustomNodeAdapter, NoopCustomNode}; -use crate::loader::{CachedLoader, DecisionLoader, NoopLoader}; +use crate::loader::{DynamicLoader, NoopLoader}; use crate::model::DecisionContent; +use crate::nodes::custom::{DynamicCustomNode, NoopCustomNode}; use crate::nodes::validator_cache::ValidatorCache; +use crate::nodes::NodeHandlerExtensions; use crate::{DecisionGraphValidationError, EvaluationError}; use serde_json::Value; -use std::sync::Arc; +use std::cell::OnceCell; +use std::sync::{Arc, OnceLock}; use zen_expression::variable::Variable; -type DynamicLoader = Arc; -type DynamicCustomNode = Arc; - /// Represents a JDM decision which can be evaluated #[derive(Debug, Clone)] pub struct Decision { content: Arc, loader: DynamicLoader, adapter: DynamicCustomNode, - validator_cache: ValidatorCache, + tokio_runtime: Arc>, } impl From for Decision { @@ -28,8 +27,8 @@ impl From for Decision { content: value.into(), loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), - - validator_cache: Default::default(), + validator_cache: ValidatorCache::default(), + tokio_runtime: Arc::new(OnceLock::new()), } } } @@ -40,41 +39,38 @@ impl From> for Decision { content: value, loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), - - validator_cache: Default::default(), + validator_cache: ValidatorCache::default(), + tokio_runtime: Arc::new(OnceLock::new()), } } } impl Decision { - pub fn with_loader(self, loader: DynamicLoader) -> Self { - Decision { - loader, - adapter: self.adapter, - content: self.content, - validator_cache: self.validator_cache, - } + pub fn with_loader(mut self, loader: DynamicLoader) -> Self { + self.loader = loader; + self } - pub fn with_adapter(self, adapter: DynamicCustomNode) -> Self { - Decision { - adapter, - loader: self.loader, - content: self.content, - validator_cache: self.validator_cache, - } + pub fn with_runtime(mut self, runtime: Arc>) -> Self { + self.tokio_runtime = runtime; + self + } + + pub fn with_adapter(mut self, adapter: DynamicCustomNode) -> Self { + self.adapter = adapter; + self } /// Evaluates a decision using an in-memory reference stored in struct - pub async fn evaluate( + pub fn evaluate( &self, context: Variable, ) -> Result> { - self.evaluate_with_opts(context, Default::default()).await + self.evaluate_with_opts(context, Default::default()) } /// Evaluates a decision using in-memory reference with advanced options - pub async fn evaluate_with_opts( + pub fn evaluate_with_opts( &self, context: Variable, options: EvaluationOptions, @@ -83,31 +79,33 @@ impl Decision { 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, - validator_cache: Some(self.validator_cache.clone()), + extensions: NodeHandlerExtensions { + loader: self.loader.clone(), + custom_node: self.adapter.clone(), + validator_cache: Arc::new(OnceCell::from(self.validator_cache.clone())), + tokio_runtime: self.tokio_runtime.clone(), + ..Default::default() + }, })?; - let response = decision_graph.evaluate(context).await?; + let response = decision_graph.evaluate(context)?; Ok(response) } - pub async fn evaluate_serialized( + pub fn evaluate_serialized( &self, context: Variable, options: EvaluationSerializedOptions, ) -> Result { - let response = self - .evaluate_with_opts( - context, - EvaluationOptions { - trace: Some(options.trace != EvaluationTraceKind::None), - max_depth: options.max_depth, - }, - ) - .await; + let response = self.evaluate_with_opts( + context, + EvaluationOptions { + trace: Some(options.trace != EvaluationTraceKind::None), + max_depth: options.max_depth, + }, + ); match response { Ok(ok) => Ok(ok @@ -124,10 +122,8 @@ impl Decision { content: self.content.clone(), max_depth: 1, trace: false, - loader: Arc::new(CachedLoader::from(self.loader.clone())), - adapter: self.adapter.clone(), iteration: 0, - validator_cache: Some(self.validator_cache.clone()), + extensions: Default::default(), })?; decision_graph.validate() diff --git a/core/engine/src/decision_graph/graph.rs b/core/engine/src/decision_graph/graph.rs index c9d646a3..0137c7e9 100644 --- a/core/engine/src/decision_graph/graph.rs +++ b/core/engine/src/decision_graph/graph.rs @@ -1,48 +1,51 @@ +use crate::decision_graph::walker::{GraphWalker, NodeData, StableDiDecisionGraph}; use crate::engine::EvaluationTraceKind; -use crate::loader::DecisionLoader; -use crate::model::{DecisionContent, DecisionNodeKind, FunctionNodeContent}; -use crate::nodes::result::NodeRequest; -use crate::nodes::validator_cache::ValidatorCache; -use crate::{EvaluationError, NodeError}; +use crate::model::{DecisionContent, DecisionNodeKind}; +use crate::nodes::custom::{CustomNodeData, CustomNodeHandler, CustomNodeTrace}; +use crate::nodes::decision::{DecisionNodeData, DecisionNodeHandler, DecisionNodeTrace}; +use crate::nodes::decision_table::{ + DecisionTableNodeData, DecisionTableNodeHandler, DecisionTableNodeTrace, +}; +use crate::nodes::expression::{ExpressionNodeData, ExpressionNodeHandler, ExpressionNodeTrace}; +use crate::nodes::function::{FunctionNodeData, FunctionNodeHandler, FunctionNodeTrace}; +use crate::nodes::input::{InputNodeData, InputNodeHandler, InputNodeTrace}; +use crate::nodes::output::{OutputNodeData, OutputNodeHandler, OutputNodeTrace}; +use crate::nodes::{ + NodeContext, NodeContextBase, NodeHandler, NodeHandlerExtensions, NodeResponse, +}; +use crate::EvaluationError; use ahash::{HashMap, HashMapExt}; -use anyhow::anyhow; use petgraph::algo::is_cyclic_directed; use serde::ser::SerializeMap; use serde::{Deserialize, Serialize, Serializer}; use serde_json::Value; use std::hash::{DefaultHasher, Hash, Hasher}; +use std::ops::Deref; use std::rc::Rc; use std::sync::Arc; use std::time::Instant; use thiserror::Error; use zen_expression::variable::{ToVariable, Variable}; -pub struct DecisionGraph { +pub struct DecisionGraph { initial_graph: StableDiDecisionGraph, graph: StableDiDecisionGraph, - adapter: Arc, - loader: Arc, trace: bool, max_depth: u8, iteration: u8, - runtime: Option>, - validator_cache: ValidatorCache, + extensions: NodeHandlerExtensions, } -pub struct DecisionGraphConfig { - pub loader: Arc, - pub adapter: Arc, +pub struct DecisionGraphConfig { pub content: Arc, pub trace: bool, pub iteration: u8, pub max_depth: u8, - pub validator_cache: Option, + pub extensions: NodeHandlerExtensions, } -impl DecisionGraph { - pub fn try_new( - config: DecisionGraphConfig, - ) -> Result { +impl DecisionGraph { + pub fn try_new(config: DecisionGraphConfig) -> Result { let content = config.content; let mut graph = StableDiDecisionGraph::new(); let mut index_map = HashMap::new(); @@ -71,45 +74,15 @@ impl DecisionGraph< graph, iteration: config.iteration, trace: config.trace, - loader: config.loader, - adapter: config.adapter, max_depth: config.max_depth, - validator_cache: config.validator_cache.unwrap_or_default(), - runtime: None, + extensions: config.extensions, }) } - pub(crate) fn with_function(mut self, runtime: Option>) -> Self { - self.runtime = runtime; - self - } - pub(crate) fn reset_graph(&mut self) { self.graph = self.initial_graph.clone(); } - async fn get_or_insert_function(&mut self) -> anyhow::Result> { - if let Some(function) = &self.runtime { - return Ok(function.clone()); - } - - let function = Function::create(FunctionConfig { - listeners: Some(vec![ - Box::new(ConsoleListener), - Box::new(ZenListener { - loader: self.loader.clone(), - adapter: self.adapter.clone(), - }), - ]), - }) - .await - .map_err(|err| anyhow!(err.to_string()))?; - let rc_function = Rc::new(function); - self.runtime.replace(rc_function.clone()); - - Ok(rc_function) - } - pub fn validate(&self) -> Result<(), DecisionGraphValidationError> { let input_count = self.input_node_count(); if input_count != 1 { @@ -132,24 +105,16 @@ impl DecisionGraph< .count() } - pub async fn evaluate( + pub fn evaluate( &mut self, context: Variable, - ) -> Result { + ) -> Result> { let root_start = Instant::now(); - self.validate().map_err(|e| NodeError::Node { - node_id: "".to_string(), - source: anyhow!(e).into(), - trace: None, - })?; + self.validate()?; if self.iteration >= self.max_depth { - return Err(NodeError::Node { - node_id: "".to_string(), - source: Box::new(NodeError::Display("Depth limit exceeded".to_string())), - trace: None, - }); + return Err(Box::new(EvaluationError::DepthLimitExceeded)); } let mut walker = GraphWalker::new(&self.graph); @@ -168,303 +133,126 @@ impl DecisionGraph< continue; } + let mut terminate = false; let node = (&self.graph[nid]).clone(); let start = Instant::now(); + let incoming_data = walker.incoming_node_data(&self.graph, nid, true); - macro_rules! trace { - ({ $($field:ident: $value:expr),* $(,)? }) => { - if let Some(nt) = &mut node_traces { - nt.insert( - node.id.clone(), - DecisionGraphTrace { - name: node.name.clone(), - id: node.id.clone(), - performance: Some(format!("{:.1?}", start.elapsed())), - order: nt.len() as u32, - $($field: $value,)* - } - ); - } - }; - } + let mut base_ctx = NodeContextBase { + id: node.id.clone(), + name: node.name.clone(), + input: incoming_data.clone(), + extensions: self.extensions.clone(), + iteration: self.iteration, + trace: self.trace, + }; - match &node.kind { + let node_execution = match &node.kind { DecisionNodeKind::InputNode { content } => { - trace!({ - input: Variable::Null, - output: context.clone(), - trace_data: None, - }); + base_ctx.input = context.clone(); + let ctx = NodeContext::::from_base( + base_ctx, + content.clone(), + ); - if let Some(json_schema) = content - .schema - .as_ref() - .map(|s| serde_json::from_str::(&s).ok()) - .flatten() - { - let validator_key = create_validator_cache_key(&json_schema); - let validator = self - .validator_cache - .get_or_insert(validator_key, &json_schema) - .await - .map_err(|e| NodeError::Node { - source: NodeError::from(e.to_string()).into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })?; - - let context_json = context.to_value(); - validator - .validate(&context_json) - .map_err(|e| NodeError::Node { - source: anyhow!(serde_json::to_value( - Box::::from(e) - ) - .unwrap_or_default()) - .into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })?; - } - - walker.set_node_data(nid, context.clone()); + InputNodeHandler.handle(ctx) } DecisionNodeKind::OutputNode { content } => { - let incoming_data = walker.incoming_node_data(&self.graph, nid, false); + terminate = true; + let ctx = NodeContext::::from_base( + base_ctx, + content.clone(), + ); - trace!({ - input: incoming_data.clone(), - output: Variable::Null, - trace_data: None, - }); - - if let Some(json_schema) = content - .schema - .as_ref() - .map(|s| serde_json::from_str::(&s).ok()) - .flatten() - { - let validator_key = create_validator_cache_key(&json_schema); - let validator = self - .validator_cache - .get_or_insert(validator_key, &json_schema) - .await - .map_err(|e| NodeError::Node { - source: NodeError::from(e.to_string()).into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })?; - - let incoming_data_json = incoming_data.to_value(); - validator - .validate(&incoming_data_json) - .map_err(|e| NodeError::Node { - source: NodeError::from(e.to_string()).into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })?; - } - - return Ok(DecisionGraphResponse { - result: incoming_data, - performance: format!("{:.1?}", root_start.elapsed()), - trace: node_traces, - }); + OutputNodeHandler.handle(ctx) } DecisionNodeKind::SwitchNode { .. } => { let input_data = walker.incoming_node_data(&self.graph, nid, false); - walker.set_node_data(nid, input_data); + // walker.set_node_data(nid, input_data); + Ok(NodeResponse { + output: Variable::Null, + trace_data: None, + }) } DecisionNodeKind::FunctionNode { content } => { - let function = - self.get_or_insert_function() - .await - .map_err(|e| NodeError::Node { - source: e.into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })?; + let ctx = NodeContext::::from_base( + base_ctx, + content.clone(), + ); - let node_request = NodeRequest { - node: node.clone(), - iteration: self.iteration, - input: walker.incoming_node_data(&self.graph, nid, true), - }; - let res = match content { - FunctionNodeContent::Version2(_) => FunctionHandler::new( - function, - self.trace, - self.iteration, - self.max_depth, - ) - .handle(node_request.clone()) - .await - .map_err(|e| { - if let NodeError::PartialTrace { trace, .. } = &e { - trace!({ - input: node_request.input.clone(), - output: Variable::Null, - trace_data: trace.clone(), - }); - } - - NodeError::Node { - source: e.into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - } - })?, - FunctionNodeContent::Version1(_) => { - let runtime = create_runtime().map_err(|e| NodeError::Node { - source: e.into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })?; - - function_v1::FunctionHandler::new(self.trace, runtime) - .handle(node_request.clone()) - .await - .map_err(|e| NodeError::Node { - source: e.into(), - node_id: node.id.clone(), - trace: error_trace(&node_traces), - })? - } - }; - - node_request.input.dot_remove("$nodes"); - res.output.dot_remove("$nodes"); - - trace!({ - input: node_request.input, - output: res.output.clone(), - trace_data: res.trace_data, - }); - walker.set_node_data(nid, res.output); + FunctionNodeHandler.handle(ctx) } - DecisionNodeKind::DecisionNode { .. } => { - let node_request = NodeRequest { - node: node.clone(), - iteration: self.iteration, - input: walker.incoming_node_data(&self.graph, nid, true), - }; + DecisionNodeKind::DecisionNode { content } => { + let ctx = NodeContext::::from_base( + base_ctx, + content.clone(), + ); - let res = DecisionHandler::new( - self.trace, - self.max_depth, - self.loader.clone(), - self.adapter.clone(), - self.runtime.clone(), - self.validator_cache.clone(), - ) - .handle(node_request.clone()) - .await - .map_err(|e| NodeError::Node { - source: e.into(), - node_id: node.id.to_string(), - trace: error_trace(&node_traces), - })?; - - node_request.input.dot_remove("$nodes"); - res.output.dot_remove("$nodes"); - - trace!({ - input: node_request.input, - output: res.output.clone(), - trace_data: res.trace_data, - }); - walker.set_node_data(nid, res.output); + DecisionNodeHandler.handle(ctx) } - DecisionNodeKind::DecisionTableNode { .. } => { - let node_request = NodeRequest { - node: node.clone(), - iteration: self.iteration, - input: walker.incoming_node_data(&self.graph, nid, true), - }; + DecisionNodeKind::DecisionTableNode { content } => { + let ctx = + NodeContext::::from_base( + base_ctx, + content.clone(), + ); - let res = DecisionTableHandler::new(self.trace) - .handle(node_request.clone()) - .await - .map_err(|e| NodeError::Node { - node_id: node.id.clone(), - source: e.into(), - trace: error_trace(&node_traces), - })?; - - node_request.input.dot_remove("$nodes"); - res.output.dot_remove("$nodes"); - - trace!({ - input: node_request.input, - output: res.output.clone(), - trace_data: res.trace_data, - }); - walker.set_node_data(nid, res.output); + DecisionTableNodeHandler.handle(ctx) } - DecisionNodeKind::ExpressionNode { .. } => { - let node_request = NodeRequest { - node: node.clone(), - iteration: self.iteration, - input: walker.incoming_node_data(&self.graph, nid, true), - }; + DecisionNodeKind::ExpressionNode { content } => { + let ctx = NodeContext::::from_base( + base_ctx, + content.clone(), + ); - let res = ExpressionHandler::new(self.trace) - .handle(node_request.clone()) - .await - .map_err(|e| { - if let NodeError::PartialTrace { trace, .. } = &e { - trace!({ - input: node_request.input.clone(), - output: Variable::Null, - trace_data: trace.clone(), - }); - } - - NodeError::Node { - node_id: node.id.clone(), - source: e.into(), - trace: error_trace(&node_traces), - } - })?; - - node_request.input.dot_remove("$nodes"); - res.output.dot_remove("$nodes"); - - trace!({ - input: node_request.input, - output: res.output.clone(), - trace_data: res.trace_data, - }); - walker.set_node_data(nid, res.output); + ExpressionNodeHandler.handle(ctx) } - DecisionNodeKind::CustomNode { .. } => { - let node_request = NodeRequest { - node: node.clone(), - iteration: self.iteration, - input: walker.incoming_node_data(&self.graph, nid, true), - }; + DecisionNodeKind::CustomNode { content } => { + let ctx = NodeContext::::from_base( + base_ctx, + content.clone(), + ); - let res = self - .adapter - .handle(CustomNodeRequest::try_from(node_request.clone()).unwrap()) - .await - .map_err(|e| NodeError::Node { - node_id: node.id.clone(), - source: e.into(), - trace: error_trace(&node_traces), - })?; - - node_request.input.dot_remove("$nodes"); - res.output.dot_remove("$nodes"); - - trace!({ - input: node_request.input, - output: res.output.clone(), - trace_data: res.trace_data, - }); - walker.set_node_data(nid, res.output); + CustomNodeHandler.handle(ctx) } + }; + + if let Some(nt) = &mut node_traces { + let trace = match &node_execution { + Ok(ok) => DecisionGraphTrace { + id: node.id.clone(), + name: node.name.clone(), + input: incoming_data, + order: nt.len() as u32, + output: ok.output.clone(), + trace_data: ok.trace_data.clone(), + performance: Some(Arc::from(format!("{:.1?}", start.elapsed()))), + }, + Err(err) => DecisionGraphTrace { + id: node.id.clone(), + name: node.name.clone(), + input: incoming_data, + order: nt.len() as u32, + output: Variable::Null, + trace_data: err.trace.clone(), + performance: Some(Arc::from(format!("{:.1?}", start.elapsed()))), + }, + }; + + nt.insert(node.id.clone(), trace); + } + + walker.set_node_data( + nid, + NodeData { + name: Rc::from(node.name.deref()), + data: node_execution?.output, + }, + ); + + if terminate { + break; } } diff --git a/core/engine/src/decision_graph/mod.rs b/core/engine/src/decision_graph/mod.rs index aec805f7..8758b834 100644 --- a/core/engine/src/decision_graph/mod.rs +++ b/core/engine/src/decision_graph/mod.rs @@ -1,2 +1,2 @@ pub mod graph; -mod traversal; +mod walker; diff --git a/core/engine/src/decision_graph/traversal.rs b/core/engine/src/decision_graph/walker.rs similarity index 73% rename from core/engine/src/decision_graph/traversal.rs rename to core/engine/src/decision_graph/walker.rs index e3af418c..b26c9afb 100644 --- a/core/engine/src/decision_graph/traversal.rs +++ b/core/engine/src/decision_graph/walker.rs @@ -3,9 +3,8 @@ use fixedbitset::FixedBitSet; use petgraph::data::DataMap; use petgraph::matrix_graph::Zero; use petgraph::prelude::{EdgeIndex, NodeIndex, StableDiGraph}; -use petgraph::visit::{EdgeRef, IntoNodeIdentifiers, VisitMap, Visitable}; +use petgraph::visit::{EdgeRef, IntoEdgesDirected, IntoNodeIdentifiers, VisitMap, Visitable}; use petgraph::{Incoming, Outgoing}; -use serde_json::json; use std::ops::Deref; use std::rc::Rc; use std::sync::atomic::Ordering; @@ -17,16 +16,21 @@ use crate::model::{ DecisionEdge, DecisionNode, DecisionNodeKind, SwitchStatement, SwitchStatementHitPolicy, }; use crate::DecisionGraphTrace; -use zen_expression::variable::Variable; +use zen_expression::variable::{ToVariable, Variable}; use zen_expression::Isolate; pub(crate) type StableDiDecisionGraph = StableDiGraph, Arc>; +pub(crate) struct NodeData { + pub name: Rc, + pub data: Variable, +} + pub(crate) struct GraphWalker { + iter: usize, + node_data: HashMap, ordered: FixedBitSet, to_visit: Vec, - node_data: HashMap, - iter: usize, visited_switch_nodes: Vec, nodes_in_context: bool, @@ -36,9 +40,9 @@ const ITER_MAX: usize = 1_000; impl GraphWalker { pub fn new(graph: &StableDiDecisionGraph) -> Self { - let mut topo = Self::empty(graph); - topo.extend_with_initials(graph); - topo + let mut walker = Self::empty(graph); + walker.extend_with_initials(graph); + walker } fn extend_with_initials(&mut self, g: &StableDiDecisionGraph) { @@ -71,7 +75,7 @@ impl GraphWalker { } pub fn get_node_data(&self, node_id: NodeIndex) -> Option { - self.node_data.get(&node_id).cloned() + Some(self.node_data.get(&node_id)?.data.clone()) } pub fn ending_variables(&self, g: &StableDiDecisionGraph) -> Variable { @@ -83,7 +87,7 @@ impl GraphWalker { .fold(Variable::empty_object(), |mut acc, curr| { match self.node_data.get(&curr) { None => acc, - Some(data) => acc.merge(data), + Some(nd) => acc.merge(&nd.data), } }) } @@ -92,16 +96,13 @@ impl GraphWalker { let node_values = self .node_data .iter() - .filter_map(|(idx, value)| { - let weight = g.node_weight(*idx)?; - Some((Rc::from(weight.name.deref()), value.clone())) - }) + .filter_map(|(_, nd)| Some((nd.name.clone(), nd.data.clone()))) .collect(); Variable::from_object(node_values) } - pub fn set_node_data(&mut self, node_id: NodeIndex, value: Variable) { + pub fn set_node_data(&mut self, node_id: NodeIndex, value: NodeData) { self.node_data.insert(node_id, value); } @@ -114,6 +115,7 @@ impl GraphWalker { let value = self .merge_node_data(g.neighbors_directed(node_id, Incoming)) .depth_clone(1); + if self.nodes_in_context { if let Some(object_ref) = with_nodes.then_some(value.as_object()).flatten() { let mut object = object_ref.borrow_mut(); @@ -128,11 +130,10 @@ impl GraphWalker { where I: Iterator, { - let default_map = Variable::empty_object(); - iter.fold(Variable::empty_object(), |mut prev, curr| { - let data = self.node_data.get(&curr).unwrap_or(&default_map); - prev.merge_clone(data) - }) + iter.filter_map(|nid| self.node_data.get(&nid)) + .fold(Variable::empty_object(), |mut prev, nd| { + prev.merge_clone(&nd.data) + }) } pub fn next( @@ -146,7 +147,6 @@ 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)?.clone(); if self.ordered.is_visited(&nid) { continue; } @@ -160,41 +160,27 @@ impl GraphWalker { self.ordered.visit(nid); + let decision_node = g.node_weight(nid)?.clone(); if let DecisionNodeKind::SwitchNode { content } = &decision_node.kind { if !self.visited_switch_nodes.contains(&nid) { let input_data = self.incoming_node_data(g, nid, true); - - let env = input_data.depth_clone(1); - env.dot_insert("$", input_data.depth_clone(1)); - - let mut isolate = Isolate::with_environment(env); + let mut isolate = Isolate::with_environment(input_data.clone()); let mut statement_iter = content.statements.iter(); - let valid_statements: Vec<&SwitchStatement> = match content.hit_policy { + let valid_statements: Vec = match content.hit_policy { SwitchStatementHitPolicy::First => statement_iter .find(|&s| switch_statement_evaluate(&mut isolate, &s)) .into_iter() + .cloned() + .map(SwitchStatementTraceRow::from) .collect(), SwitchStatementHitPolicy::Collect => statement_iter .filter(|&s| switch_statement_evaluate(&mut isolate, &s)) + .cloned() + .map(SwitchStatementTraceRow::from) .collect(), }; - let valid_statements_trace = Variable::from_array( - valid_statements - .iter() - .map(|&statement| { - let v = Variable::empty_object(); - v.dot_insert( - "id", - Variable::String(Rc::from(statement.id.deref())), - ); - - v - }) - .collect(), - ); - input_data.dot_remove("$nodes"); if let Some(on_trace) = &mut on_trace { @@ -204,9 +190,12 @@ impl GraphWalker { input: input_data.shallow_clone(), output: input_data.shallow_clone(), order: 0, - performance: Some(format!("{:.1?}", start.elapsed())), + performance: Some(Arc::from(format!("{:.1?}", start.elapsed()))), trace_data: Some( - json!({ "statements": valid_statements_trace }).into(), + SwitchStatementTrace { + statements: valid_statements.clone(), + } + .to_variable(), ), }); } @@ -282,37 +271,38 @@ fn remove_edge_recursive(g: &mut StableDiDecisionGraph, edge_id: EdgeIndex) { g.remove_edge(edge_id); - // Remove dead branches from target - let target_incoming_count = g.edges_directed(target_nid, Incoming).count(); - if target_incoming_count.is_zero() { - let edge_ids: Vec = g - .edges_directed(target_nid, Outgoing) - .map(|edge| edge.id()) - .collect(); + for (nid, direction) in [(target_nid, Incoming), (source_nid, Outgoing)] { + let count = g.edges_directed(nid, direction).count(); + if count.is_zero() { + let edge_ids: Vec = g + .edges_directed(nid, direction.opposite()) + .map(|edge| edge.id()) + .collect(); - edge_ids.iter().for_each(|edge_id| { - remove_edge_recursive(g, edge_id.clone()); - }); + edge_ids.iter().for_each(|&edge_id| { + remove_edge_recursive(g, edge_id); + }); - if g.edges(target_nid).count().is_zero() { - g.remove_node(target_nid); - } - } - - // Remove dead branches from source - let source_outgoing_count = g.edges_directed(source_nid, Outgoing).count(); - if source_outgoing_count.is_zero() { - let edge_ids: Vec = g - .edges_directed(source_nid, Incoming) - .map(|edge| edge.id()) - .collect(); - - edge_ids.iter().for_each(|edge_id| { - remove_edge_recursive(g, edge_id.clone()); - }); - - if g.edges(source_nid).count().is_zero() { - g.remove_node(source_nid); + if g.edges(nid).count().is_zero() { + g.remove_node(nid); + } } } } + +#[derive(ToVariable)] +struct SwitchStatementTrace { + statements: Vec, +} + +#[derive(ToVariable, Clone)] +#[serde(rename_all = "camelCase")] +struct SwitchStatementTraceRow { + pub id: Arc, +} + +impl From for SwitchStatementTraceRow { + fn from(value: SwitchStatement) -> Self { + Self { id: value.id } + } +} diff --git a/core/engine/src/engine.rs b/core/engine/src/engine.rs index a11b0edb..ae9191e3 100644 --- a/core/engine/src/engine.rs +++ b/core/engine/src/engine.rs @@ -1,25 +1,24 @@ use crate::decision::Decision; use crate::decision_graph::graph::DecisionGraphResponse; -use crate::handler::custom_node_adapter::{CustomNodeAdapter, NoopCustomNode}; use crate::loader::{ ClosureLoader, DecisionLoader, DynamicLoader, LoaderResponse, LoaderResult, NoopLoader, }; use crate::model::DecisionContent; +use crate::nodes::custom::{DynamicCustomNode, NoopCustomNode}; use crate::EvaluationError; use serde_json::Value; use std::fmt::Debug; use std::future::Future; -use std::sync::Arc; +use std::sync::{Arc, OnceLock}; use strum::{EnumString, IntoStaticStr}; use zen_expression::variable::Variable; -type DynamicCustomNode = Arc; - /// Structure used for generating and evaluating JDM decisions #[derive(Debug, Clone)] pub struct DecisionEngine { loader: DynamicLoader, adapter: DynamicCustomNode, + runtime: Arc>, } #[derive(Debug, Default)] @@ -68,41 +67,37 @@ impl Default for DecisionEngine { Self { loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), + runtime: Arc::new(OnceLock::new()), } } } impl DecisionEngine { pub fn new(loader: DynamicLoader, adapter: DynamicCustomNode) -> Self { - Self { loader, adapter } - } - - pub fn with_adapter(self, adapter: DynamicCustomNode) -> Self - where - CustomNode: CustomNodeAdapter, - { - DecisionEngine { - loader: self.loader, - adapter, - } - } - - pub fn with_loader(self, loader: DynamicLoader) -> Self { - DecisionEngine { + Self { loader, - adapter: self.adapter, + adapter, + ..Default::default() } } - pub fn with_closure_loader(self, loader: F) -> Self + pub fn with_adapter(mut self, adapter: DynamicCustomNode) -> Self { + self.adapter = adapter; + self + } + + pub fn with_loader(mut self, loader: DynamicLoader) -> Self { + self.loader = loader; + self + } + + pub fn with_closure_loader(mut self, loader: F) -> Self where F: Fn(String) -> O + Sync + Send + Debug + 'static, O: Future + Send, { - DecisionEngine { - loader: Arc::new(ClosureLoader::new(loader)), - adapter: self.adapter, - } + self.loader = Arc::new(ClosureLoader::new(loader)); + self } /// Evaluates a decision through loader using a key @@ -130,7 +125,7 @@ impl DecisionEngine { { let content = self.loader.load(key.as_ref()).await?; let decision = self.create_decision(content); - decision.evaluate_with_opts(context, options).await + decision.evaluate_with_opts(context, options) } pub async fn evaluate_serialized( @@ -149,7 +144,7 @@ impl DecisionEngine { .map_err(|err| Value::String(err.to_string()))?; let decision = self.create_decision(content); - decision.evaluate_serialized(context, options).await + decision.evaluate_serialized(context, options) } /// Creates a decision from DecisionContent, exists for easier binding creation @@ -157,6 +152,7 @@ impl DecisionEngine { Decision::from(content) .with_loader(self.loader.clone()) .with_adapter(self.adapter.clone()) + .with_runtime(self.runtime.clone()) } /// Retrieves a decision based on the loader diff --git a/core/engine/src/error.rs b/core/engine/src/error.rs index 0dfdee79..0999f288 100644 --- a/core/engine/src/error.rs +++ b/core/engine/src/error.rs @@ -1,12 +1,10 @@ use crate::decision_graph::graph::DecisionGraphValidationError; use crate::engine::EvaluationTraceKind; use crate::loader::LoaderError; -pub use crate::nodes::result::NodeError; -use jsonschema::{ErrorIterator, ValidationError}; +use crate::nodes::NodeError; use serde::ser::SerializeMap; use serde::{Serialize, Serializer}; -use serde_json::{Map, Value}; -use std::iter::once; +use serde_json::Value; use thiserror::Error; #[derive(Debug, Error)] @@ -106,40 +104,3 @@ impl From for Box { Box::new(EvaluationError::InvalidGraph(error.into())) } } - -#[derive(Serialize)] -#[serde(rename_all = "camelCase")] -struct ValidationErrorJson { - path: String, - message: String, -} - -impl<'a> From> for ValidationErrorJson { - fn from(value: ValidationError<'a>) -> Self { - ValidationErrorJson { - path: value.instance_path.to_string(), - message: format!("{}", value), - } - } -} - -impl<'a> From> for Box { - fn from(error_iter: ErrorIterator<'a>) -> Self { - let errors: Vec = error_iter.into_iter().map(From::from).collect(); - - let mut json_map = Map::new(); - json_map.insert( - "errors".to_string(), - serde_json::to_value(errors).unwrap_or_default(), - ); - - Box::new(EvaluationError::Validation(Value::Object(json_map))) - } -} - -impl<'a> From> for Box { - fn from(value: ValidationError<'a>) -> Self { - let iterator: ErrorIterator<'a> = Box::new(once(value)); - Box::::from(iterator) - } -} diff --git a/core/engine/src/lib.rs b/core/engine/src/lib.rs index 51b219bb..980c9ebf 100644 --- a/core/engine/src/lib.rs +++ b/core/engine/src/lib.rs @@ -127,10 +127,9 @@ mod decision; mod decision_graph; mod engine; pub mod error; -pub mod handler; pub mod loader; pub mod model; -mod nodes; +pub mod nodes; pub use config::ZEN_CONFIG; pub use decision::Decision; @@ -141,5 +140,4 @@ pub use engine::{ DecisionEngine, EvaluationOptions, EvaluationSerializedOptions, EvaluationTraceKind, }; pub use error::EvaluationError; -pub use nodes::result::NodeError; pub use zen_expression::Variable; diff --git a/core/engine/src/loader/mod.rs b/core/engine/src/loader/mod.rs index 48b5e2a2..129fb982 100644 --- a/core/engine/src/loader/mod.rs +++ b/core/engine/src/loader/mod.rs @@ -19,13 +19,13 @@ mod filesystem; mod memory; mod noop; -pub type DynamicLoader = Arc; +pub type DynamicLoader = Arc; pub type LoaderResult = Result; pub type LoaderResponse = LoaderResult>; /// Trait used for implementing a loader for decisions -pub trait DecisionLoader: Debug { +pub trait DecisionLoader: Debug + Send { fn load<'a>(&'a self, key: &'a str) -> Pin + 'a>>; } diff --git a/core/engine/src/nodes/context.rs b/core/engine/src/nodes/context.rs index 656458a3..89e87d1e 100644 --- a/core/engine/src/nodes/context.rs +++ b/core/engine/src/nodes/context.rs @@ -2,11 +2,17 @@ use crate::nodes::definition::{NodeDataType, TraceDataType}; use crate::nodes::extensions::NodeHandlerExtensions; use crate::nodes::function::v2::function::Function; use crate::nodes::result::{NodeResponse, NodeResult}; -use crate::NodeError; +use crate::nodes::NodeError; +use ahash::AHasher; +use jsonschema::ValidationError; +use serde::Serialize; +use serde_json::Value; use std::cell::RefCell; +use std::fmt::{Display, Formatter}; use std::future::Future; -use std::rc::Rc; +use std::hash::Hasher; use std::sync::Arc; +use thiserror::Error; use zen_types::variable::Variable; pub struct NodeContext @@ -20,6 +26,7 @@ where pub input: Variable, pub trace: Option>, pub extensions: NodeHandlerExtensions, + pub iteration: u8, } impl NodeContext @@ -27,6 +34,18 @@ where NodeData: NodeDataType, TraceData: TraceDataType, { + pub fn from_base(base: NodeContextBase, data: NodeData) -> Self { + Self { + id: base.id, + name: base.name, + input: base.input, + extensions: base.extensions, + iteration: base.iteration, + trace: base.trace.then(|| Default::default()), + node: data, + } + } + pub fn trace(&self, mutator: Function) where Function: FnOnce(&mut TraceData), @@ -69,7 +88,7 @@ where where Fut: Future, { - let tokio_runtime = self.extensions.tokio_runtime().node_context(self)?; + let tokio_runtime = self.extensions.tokio_runtime(); Ok(tokio_runtime.block_on(future)) } @@ -81,7 +100,30 @@ where } pub(crate) fn function_runtime(&self) -> Result<&Function, NodeError> { - self.extensions.function_runtime().node_context(self) + Ok(self.extensions.function_runtime()) + } + + pub fn validate(&self, schema: &Value, value: &Value) -> Result<(), NodeError> { + let validator_cache = self.extensions.validator_cache(); + let hash = self.hash_node(); + + let validator = validator_cache + .get_or_insert(hash, schema) + .node_context(self)?; + + validator + .validate(value) + .map_err(|err| ValidationErrorJson::from(err)) + .node_context(self)?; + + Ok(()) + } + + fn hash_node(&self) -> u64 { + let mut hasher = AHasher::default(); + hasher.write(self.id.as_bytes()); + hasher.write(self.name.as_bytes()); + hasher.finish() } } @@ -151,6 +193,9 @@ pub struct NodeContextBase { pub id: Arc, pub name: Arc, pub input: Variable, + pub iteration: u8, + pub extensions: NodeHandlerExtensions, + pub trace: bool, } impl NodeContextBase { @@ -187,9 +232,12 @@ where { fn from(value: NodeContext) -> Self { Self { - id: value.id.clone(), - name: value.name.clone(), - input: value.input.clone(), + id: value.id, + name: value.name, + input: value.input, + extensions: value.extensions, + iteration: value.iteration, + trace: value.trace.is_some(), } } } @@ -228,3 +276,25 @@ impl NodeContextExt for Option { self.ok_or_else(|| ctx.make_error(f("None"))) } } + +#[derive(Debug, Serialize, Error)] +#[serde(rename_all = "camelCase")] +struct ValidationErrorJson { + path: String, + message: String, +} + +impl Display for ValidationErrorJson { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "{}: {}", self.path, self.message) + } +} + +impl<'a> From> for ValidationErrorJson { + fn from(value: ValidationError<'a>) -> Self { + ValidationErrorJson { + path: value.instance_path.to_string(), + message: format!("{}", value), + } + } +} diff --git a/core/engine/src/nodes/custom/adapter.rs b/core/engine/src/nodes/custom/adapter.rs index 32475ec2..34417453 100644 --- a/core/engine/src/nodes/custom/adapter.rs +++ b/core/engine/src/nodes/custom/adapter.rs @@ -1,17 +1,15 @@ -use crate::nodes::result::{NodeError, NodeRequest, NodeResult}; +use crate::nodes::result::{NodeError, NodeResult}; use json_dotpath::DotPaths; use serde::Serialize; use serde_json::Value; use std::fmt::Debug; use std::future::Future; -use std::ops::Deref; use std::pin::Pin; -use std::rc::Rc; use std::sync::Arc; use zen_expression::variable::Variable; use zen_tmpl::TemplateRenderError; -pub trait CustomNodeAdapter: Debug { +pub trait CustomNodeAdapter: Debug + Send { fn handle( &self, request: CustomNodeRequest, @@ -29,7 +27,7 @@ impl CustomNodeAdapter for NoopCustomNode { Box::pin(async move { Err(NodeError { trace: None, - node_id: Some(Rc::from(request.node.id.deref())), + node_id: Some(request.node.id.clone()), source: "Custom node handler not provided".to_string().into(), }) }) @@ -43,17 +41,6 @@ pub struct CustomNodeRequest { pub node: CustomDecisionNode, } -impl TryFrom for CustomNodeRequest { - type Error = (); - - fn try_from(value: NodeRequest) -> Result { - Ok(Self { - input: value.input.clone(), - node: value.node.deref().try_into()?, - }) - } -} - impl CustomNodeRequest { pub fn get_field(&self, path: &str) -> Result, TemplateRenderError> { let Some(selected_value) = self.get_field_raw(path) else { @@ -82,4 +69,4 @@ pub struct CustomDecisionNode { pub config: Arc, } -pub type DynamicCustomNode = Arc; +pub type DynamicCustomNode = Arc; diff --git a/core/engine/src/nodes/custom/mod.rs b/core/engine/src/nodes/custom/mod.rs index e336a9bb..2af0a5ff 100644 --- a/core/engine/src/nodes/custom/mod.rs +++ b/core/engine/src/nodes/custom/mod.rs @@ -1,16 +1,21 @@ -use crate::nodes::custom::adapter::{CustomDecisionNode, CustomNodeRequest}; use crate::nodes::result::NodeResult; use crate::nodes::{NodeContext, NodeHandler}; use zen_types::decision::CustomNodeContent; use zen_types::variable::Variable; +pub use adapter::{ + CustomDecisionNode, CustomNodeAdapter, CustomNodeRequest, DynamicCustomNode, NoopCustomNode, +}; + mod adapter; pub struct CustomNodeHandler; +pub type CustomNodeData = CustomNodeContent; +pub type CustomNodeTrace = Variable; impl NodeHandler for CustomNodeHandler { - type NodeData = CustomNodeContent; - type TraceData = Variable; + type NodeData = CustomNodeData; + type TraceData = CustomNodeTrace; fn handle(&self, ctx: NodeContext) -> NodeResult { let custom_node_request = CustomNodeRequest { @@ -26,5 +31,3 @@ impl NodeHandler for CustomNodeHandler { ctx.block_on(ctx.extensions.custom_node().handle(custom_node_request))? } } - -pub use adapter::DynamicCustomNode; diff --git a/core/engine/src/nodes/decision/mod.rs b/core/engine/src/nodes/decision/mod.rs index 279f1e55..ff4f024e 100644 --- a/core/engine/src/nodes/decision/mod.rs +++ b/core/engine/src/nodes/decision/mod.rs @@ -1 +1,55 @@ +use crate::decision_graph::graph::{DecisionGraph, DecisionGraphConfig}; +use crate::nodes::{NodeContext, NodeContextExt, NodeHandler, NodeResult}; +use crate::EvaluationError; +use std::ops::Deref; +use zen_types::decision::{DecisionNodeContent, TransformAttributes}; +use zen_types::variable::{ToVariable, Variable}; + pub struct DecisionNodeHandler; + +pub type DecisionNodeData = DecisionNodeContent; +pub type DecisionNodeTrace = Variable; + +impl NodeHandler for DecisionNodeHandler { + type NodeData = DecisionNodeData; + type TraceData = DecisionNodeTrace; + + fn transform_attributes( + &self, + ctx: &NodeContext, + ) -> Option { + Some(ctx.node.transform_attributes.clone()) + } + + fn handle(&self, ctx: NodeContext) -> NodeResult { + let loader = ctx.extensions.loader(); + let sub_decision = + ctx.try_block_on(async { loader.load(ctx.node.key.deref()).await.node_context(&ctx) })?; + + let mut decision_graph = DecisionGraph::try_new(DecisionGraphConfig { + content: sub_decision, + extensions: ctx.extensions.clone(), + trace: ctx.has_trace(), + iteration: ctx.iteration, + max_depth: 10, + }) + .node_context(&ctx)?; + + match decision_graph.evaluate(ctx.input.clone()) { + Ok(result) => { + ctx.trace(|trace| { + *trace = result.trace.to_variable(); + }); + + ctx.success(result.result) + } + Err(err) => { + if let EvaluationError::NodeError(node_error) = err.deref() { + ctx.trace(|trace| *trace = node_error.trace.to_variable()); + } + + ctx.error(err.to_string()) + } + } + } +} diff --git a/core/engine/src/nodes/decision_table/mod.rs b/core/engine/src/nodes/decision_table/mod.rs index 596e4d74..2aea2236 100644 --- a/core/engine/src/nodes/decision_table/mod.rs +++ b/core/engine/src/nodes/decision_table/mod.rs @@ -8,15 +8,18 @@ use std::rc::Rc; use std::sync::Arc; use zen_expression::variable::ToVariable; use zen_expression::Isolate; -use zen_types::decision::{DecisionTableContent, TransformAttributes}; +use zen_types::decision::{DecisionTableContent, DecisionTableHitPolicy, TransformAttributes}; use zen_types::variable::Variable; -pub struct DecisionTableHandler; -type DecisionTableContext = NodeContext; +pub struct DecisionTableNodeHandler; -impl NodeHandler for DecisionTableHandler { - type NodeData = DecisionTableContent; - type TraceData = DecisionTableTrace; +pub type DecisionTableNodeData = DecisionTableContent; + +type DecisionTableContext = NodeContext; + +impl NodeHandler for DecisionTableNodeHandler { + type NodeData = DecisionTableNodeData; + type TraceData = DecisionTableNodeTrace; fn transform_attributes( &self, @@ -26,12 +29,15 @@ impl NodeHandler for DecisionTableHandler { } fn handle(&self, ctx: NodeContext) -> NodeResult { - todo!() + match ctx.node.hit_policy { + DecisionTableHitPolicy::First => self.handle_first_hit(ctx), + DecisionTableHitPolicy::Collect => self.handle_collect(ctx), + } } } -impl DecisionTableHandler { - fn handle_first_hit(&mut self, ctx: DecisionTableContext) -> NodeResult { +impl DecisionTableNodeHandler { + fn handle_first_hit(&self, ctx: DecisionTableContext) -> NodeResult { let mut isolate = Isolate::new(); for (index, rule) in ctx.node.rules.iter().enumerate() { @@ -44,7 +50,7 @@ impl DecisionTableHandler { rule, } => { ctx.trace(|t| { - *t = DecisionTableTrace::FirstHit(DecisionTableRowTrace { + *t = DecisionTableNodeTrace::FirstHit(DecisionTableRowTrace { reference_map, index, rule, @@ -60,7 +66,7 @@ impl DecisionTableHandler { ctx.success(Variable::Null) } - fn handle_collect(&mut self, ctx: DecisionTableContext) -> NodeResult { + fn handle_collect(&self, ctx: DecisionTableContext) -> NodeResult { let mut isolate = Isolate::new(); let mut outputs = Vec::new(); let mut traces = Vec::new(); @@ -88,14 +94,14 @@ impl DecisionTableHandler { } ctx.trace(|t| { - *t = DecisionTableTrace::Collect(traces); + *t = DecisionTableNodeTrace::Collect(traces); }); ctx.success(Variable::from_array(outputs)) } fn evaluate_row<'a>( - &mut self, + &self, ctx: &'a DecisionTableContext, rule: &'a HashMap, Arc>, isolate: &mut Isolate<'a>, @@ -191,7 +197,7 @@ enum RowResult { } #[derive(Debug, Clone, Serialize, ToVariable)] -struct DecisionTableRowTrace { +pub struct DecisionTableRowTrace { index: usize, reference_map: HashMap, Variable>, rule: HashMap, Rc>, @@ -199,13 +205,13 @@ struct DecisionTableRowTrace { #[derive(Debug, Clone, Serialize, ToVariable)] #[serde(untagged)] -enum DecisionTableTrace { +pub enum DecisionTableNodeTrace { FirstHit(DecisionTableRowTrace), Collect(Vec), } -impl Default for DecisionTableTrace { +impl Default for DecisionTableNodeTrace { fn default() -> Self { - DecisionTableTrace::Collect(Default::default()) + DecisionTableNodeTrace::Collect(Default::default()) } } diff --git a/core/engine/src/nodes/expression/mod.rs b/core/engine/src/nodes/expression/mod.rs index 665aa33c..a4f59b8c 100644 --- a/core/engine/src/nodes/expression/mod.rs +++ b/core/engine/src/nodes/expression/mod.rs @@ -11,9 +11,12 @@ use zen_types::decision::TransformAttributes; pub struct ExpressionNodeHandler; +pub type ExpressionNodeData = ExpressionNodeContent; +pub type ExpressionNodeTrace = HashMap, ExpressionNodeTraceItem>; + impl NodeHandler for ExpressionNodeHandler { - type NodeData = ExpressionNodeContent; - type TraceData = HashMap, ExpressionTrace>; + type NodeData = ExpressionNodeData; + type TraceData = ExpressionNodeTrace; fn transform_attributes( &self, @@ -41,7 +44,7 @@ impl NodeHandler for ExpressionNodeHandler { ctx.trace(|trace| { trace.insert( Rc::from(&*expression.key), - ExpressionTrace { + ExpressionNodeTraceItem { result: value.clone(), }, ); @@ -64,6 +67,6 @@ impl NodeHandler for ExpressionNodeHandler { } #[derive(Debug, Clone, ToVariable)] -pub struct ExpressionTrace { +pub struct ExpressionNodeTraceItem { result: Variable, } diff --git a/core/engine/src/nodes/extensions.rs b/core/engine/src/nodes/extensions.rs index d4bb8d96..3c865ac0 100644 --- a/core/engine/src/nodes/extensions.rs +++ b/core/engine/src/nodes/extensions.rs @@ -1,66 +1,72 @@ -use crate::loader::DynamicLoader; -use crate::nodes::custom::DynamicCustomNode; +use crate::loader::{DynamicLoader, NoopLoader}; +use crate::nodes::custom::{DynamicCustomNode, NoopCustomNode}; use crate::nodes::function::v2::function::{Function, FunctionConfig}; use crate::nodes::function::v2::module::console::ConsoleListener; use crate::nodes::function::v2::module::zen::ZenListener; -use anyhow::Context; +use crate::nodes::validator_cache::ValidatorCache; use std::cell::OnceCell; -use std::sync::Arc; +use std::sync::{Arc, OnceLock}; /// This is created on every graph evaluation #[derive(Clone)] pub struct NodeHandlerExtensions { - tokio_runtime: Arc>, - function_runtime: Arc>, + pub(crate) function_runtime: Arc>, + pub(crate) tokio_runtime: Arc>, + pub(crate) validator_cache: Arc>, + pub(crate) loader: DynamicLoader, + pub(crate) custom_node: DynamicCustomNode, +} - loader: DynamicLoader, - custom_node_adapter: DynamicCustomNode, +impl Default for NodeHandlerExtensions { + fn default() -> Self { + Self { + tokio_runtime: Default::default(), + function_runtime: Default::default(), + validator_cache: Default::default(), + + loader: Arc::new(NoopLoader::default()), + custom_node: Arc::new(NoopCustomNode::default()), + } + } } impl NodeHandlerExtensions { - pub fn tokio_runtime(&self) -> anyhow::Result<&tokio::runtime::Runtime> { - if let Some(tokio_runtime) = self.tokio_runtime.get() { - return Ok(tokio_runtime); - } + pub fn tokio_runtime(&self) -> &tokio::runtime::Runtime { + self.tokio_runtime.get_or_init(|| { + println!("Creating tokio runtime"); - let runtime = tokio::runtime::Builder::new_multi_thread() - .worker_threads(2) - .enable_all() - .build() - .context("Failed to build tokio runtime")?; - - let _ = self.tokio_runtime.set(runtime); - self.tokio_runtime - .get() - .context("Tokio runtime is not initialized") + tokio::runtime::Builder::new_multi_thread() + .worker_threads(2) + .enable_all() + .build() + .expect("Failed to build tokio runtime") + }) } - pub fn function_runtime(&self) -> anyhow::Result<&Function> { - if let Some(function_runtime) = self.function_runtime.get() { - return Ok(function_runtime); - } + pub fn function_runtime(&self) -> &Function { + self.function_runtime.get_or_init(|| { + let tokio_runtime = self.tokio_runtime(); - let tokio_runtime = self.tokio_runtime()?; - let function_runtime = tokio_runtime - .block_on(Function::create(FunctionConfig { - listeners: Some(vec![ - Box::new(ConsoleListener), - Box::new(ZenListener { - loader: self.decision_loader.clone(), - adapter: self.custom_node_adapter.clone(), - }), - ]), - })) - .context("Failed to create async function")?; + tokio_runtime + .block_on(Function::create(FunctionConfig { + listeners: Some(vec![ + Box::new(ConsoleListener), + Box::new(ZenListener { + extensions: self.clone(), + }), + ]), + })) + .expect("Failed to create async function") + }) + } - let _ = self.function_runtime.set(function_runtime); - self.function_runtime - .get() - .context("Tokio runtime is not initialized") + pub fn validator_cache(&self) -> &ValidatorCache { + self.validator_cache + .get_or_init(|| ValidatorCache::default()) } pub fn custom_node(&self) -> &DynamicCustomNode { - &self.custom_node_adapter + &self.custom_node } pub fn loader(&self) -> &DynamicLoader { diff --git a/core/engine/src/nodes/function/mod.rs b/core/engine/src/nodes/function/mod.rs index 0937549d..4e3c947b 100644 --- a/core/engine/src/nodes/function/mod.rs +++ b/core/engine/src/nodes/function/mod.rs @@ -12,8 +12,12 @@ use zen_types::variable::Variable; pub struct FunctionNodeHandler; +pub type FunctionNodeData = FunctionNodeContent; + +pub type FunctionNodeTrace = Variable; + impl NodeHandler for FunctionNodeHandler { - type NodeData = FunctionNodeContent; + type NodeData = FunctionNodeData; type TraceData = FunctionNodeTrace; fn handle(&self, ctx: NodeContext) -> NodeResult { @@ -24,6 +28,7 @@ impl NodeHandler for FunctionNodeHandler { name: ctx.name.clone(), input: ctx.input.clone(), extensions: ctx.extensions.clone(), + iteration: ctx.iteration, node: source.clone(), trace: None, }; @@ -36,6 +41,7 @@ impl NodeHandler for FunctionNodeHandler { name: ctx.name.clone(), input: ctx.input.clone(), extensions: ctx.extensions.clone(), + iteration: ctx.iteration, node: content.clone(), trace: None, }; @@ -45,5 +51,3 @@ impl NodeHandler for FunctionNodeHandler { } } } - -pub type FunctionNodeTrace = Variable; diff --git a/core/engine/src/nodes/function/v2/module/zen.rs b/core/engine/src/nodes/function/v2/module/zen.rs index 77dd37dd..ed1bed32 100644 --- a/core/engine/src/nodes/function/v2/module/zen.rs +++ b/core/engine/src/nodes/function/v2/module/zen.rs @@ -1,21 +1,18 @@ use std::future::Future; use std::pin::Pin; -use std::sync::Arc; use crate::decision_graph::graph::{DecisionGraph, DecisionGraphConfig}; -use crate::handler::custom_node_adapter::DynamicCustomNode; -use crate::loader::{DecisionLoader, DynamicLoader}; use crate::nodes::function::v2::error::{FunctionResult, ResultExt}; use crate::nodes::function::v2::listener::{RuntimeEvent, RuntimeListener}; use crate::nodes::function::v2::module::export_default; use crate::nodes::function::v2::serde::JsValue; +use crate::nodes::NodeHandlerExtensions; use rquickjs::module::{Declarations, Exports, ModuleDef}; use rquickjs::prelude::{Async, Func, Opt}; use rquickjs::{CatchResultExt, Ctx, Function, Object}; pub(crate) struct ZenListener { - pub loader: Arc, - pub adapter: Arc, + pub extensions: NodeHandlerExtensions, } impl RuntimeListener for ZenListener { @@ -24,9 +21,7 @@ impl RuntimeListener for ZenListener { ctx: Ctx<'js>, event: RuntimeEvent, ) -> Pin + 'js>> { - let loader = self.loader.clone(); - let adapter = self.adapter.clone(); - + let extensions = self.extensions.clone(); Box::pin(async move { if event != RuntimeEvent::Startup { return Ok(()); @@ -40,8 +35,7 @@ impl RuntimeListener for ZenListener { key: String, context: JsValue, opts: Opt>| { - let loader = loader.clone(); - let adapter = adapter.clone(); + let extensions = extensions.clone(); async move { let config: Object = ctx.globals().get("config").or_throw(&ctx)?; @@ -53,20 +47,18 @@ impl RuntimeListener for ZenListener { .map(|opt| opt.get::<_, bool>("trace").unwrap_or_default()) .unwrap_or_default(); - let load_result = loader.load(key.as_str()).await; + let load_result = extensions.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, max_depth, - loader, - adapter, iteration: iteration + 1, trace, - validator_cache: None, + extensions: extensions.clone(), }) .or_throw(&ctx)?; - let response = sub_tree.evaluate(context.0).await.or_throw(&ctx)?; + let response = sub_tree.evaluate(context.0).or_throw(&ctx)?; let k = serde_json::to_value(response).or_throw(&ctx)?.into(); return rquickjs::Result::Ok(JsValue(k)); diff --git a/core/engine/src/nodes/input/mod.rs b/core/engine/src/nodes/input/mod.rs index 8ca2ecae..11b82f6d 100644 --- a/core/engine/src/nodes/input/mod.rs +++ b/core/engine/src/nodes/input/mod.rs @@ -6,13 +6,19 @@ use zen_types::variable::Variable; pub struct InputNodeHandler; +pub type InputNodeData = InputNodeContent; +pub type InputNodeTrace = Variable; + impl NodeHandler for InputNodeHandler { - type NodeData = InputNodeContent; - type TraceData = Variable; + type NodeData = InputNodeData; + type TraceData = InputNodeTrace; fn handle(&self, ctx: NodeContext) -> NodeResult { + if let Some(json_schema) = &ctx.node.schema { + let input_json = ctx.input.to_value(); + ctx.validate(json_schema, &input_json)?; + }; + ctx.success(ctx.input.clone()) } } - -pub type InputNodeTrace = Variable; diff --git a/core/engine/src/nodes/mod.rs b/core/engine/src/nodes/mod.rs index 8e2b2164..2b3cdd07 100644 --- a/core/engine/src/nodes/mod.rs +++ b/core/engine/src/nodes/mod.rs @@ -1,16 +1,18 @@ mod context; -mod custom; -mod decision; -mod decision_table; -pub(crate) mod definition; -mod expression; +pub mod custom; +pub mod decision; +pub mod decision_table; +mod definition; +pub mod expression; mod extensions; -pub(crate) mod function; -mod input; -mod output; -pub mod result; -mod transform_attributes; -pub mod validator_cache; +pub mod function; +pub mod input; +pub mod output; +mod result; +pub(crate) mod transform_attributes; +pub(crate) mod validator_cache; pub use context::{NodeContext, NodeContextBase, NodeContextExt}; -pub use definition::NodeHandler; +pub use definition::{NodeHandler, NodeHandlerKind}; +pub use extensions::NodeHandlerExtensions; +pub use result::{NodeError, NodeRequest, NodeResponse, NodeResult}; diff --git a/core/engine/src/nodes/output/mod.rs b/core/engine/src/nodes/output/mod.rs index 1fca47d1..16b5d910 100644 --- a/core/engine/src/nodes/output/mod.rs +++ b/core/engine/src/nodes/output/mod.rs @@ -6,13 +6,19 @@ use zen_types::variable::Variable; pub struct OutputNodeHandler; +pub type OutputNodeData = OutputNodeContent; +pub type OutputNodeTrace = Variable; + impl NodeHandler for OutputNodeHandler { - type NodeData = OutputNodeContent; + type NodeData = OutputNodeData; type TraceData = OutputNodeTrace; fn handle(&self, ctx: NodeContext) -> NodeResult { + if let Some(json_schema) = &ctx.node.schema { + let input_json = ctx.input.to_value(); + ctx.validate(json_schema, &input_json)?; + }; + ctx.success(ctx.input.clone()) } } - -pub type OutputNodeTrace = Variable; diff --git a/core/engine/src/nodes/result.rs b/core/engine/src/nodes/result.rs index 0672fbbc..8b1b4da9 100644 --- a/core/engine/src/nodes/result.rs +++ b/core/engine/src/nodes/result.rs @@ -24,7 +24,7 @@ pub type NodeResult = Result; #[derive(Debug, Error)] pub struct NodeError { - pub node_id: Option>, + pub node_id: Option>, pub trace: Option, pub source: Box, } diff --git a/core/engine/src/nodes/validator_cache.rs b/core/engine/src/nodes/validator_cache.rs index cce0335e..bb7687ed 100644 --- a/core/engine/src/nodes/validator_cache.rs +++ b/core/engine/src/nodes/validator_cache.rs @@ -1,9 +1,8 @@ -use crate::EvaluationError; use ahash::HashMap; +use anyhow::Context; use jsonschema::Validator; use serde_json::Value; -use std::sync::Arc; -use tokio::sync::RwLock; +use std::sync::{Arc, RwLock}; #[derive(Clone, Default, Debug)] pub struct ValidatorCache { @@ -11,21 +10,21 @@ pub struct ValidatorCache { } impl ValidatorCache { - pub async fn get(&self, key: u64) -> Option> { - let read = self.inner.read().await; + pub fn get(&self, key: u64) -> Option> { + let read = self.inner.read().ok()?; read.get(&key).cloned() } - pub async fn get_or_insert( - &self, - key: u64, - schema: &Value, - ) -> Result, Box> { - if let Some(v) = self.get(key).await { + pub fn get_or_insert(&self, key: u64, schema: &Value) -> anyhow::Result> { + if let Some(v) = self.get(key) { return Ok(v); } - let mut w_shared = self.inner.write().await; + let mut w_shared = self + .inner + .write() + .ok() + .context("Failed to acquire lock on validator cache")?; let validator = Arc::new(jsonschema::draft7::new(&schema)?); w_shared.insert(key, validator.clone()); diff --git a/core/types/src/decision/mod.rs b/core/types/src/decision/mod.rs index 58e23e82..f4138b24 100644 --- a/core/types/src/decision/mod.rs +++ b/core/types/src/decision/mod.rs @@ -71,15 +71,15 @@ pub enum DecisionNodeKind { #[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Default)] #[serde(rename_all = "camelCase")] pub struct InputNodeContent { - #[serde(default, deserialize_with = "empty_string_is_none")] - pub schema: Option>, + #[serde(default, deserialize_with = "empty_value_string_is_none")] + pub schema: Option>, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize, Default)] #[serde(rename_all = "camelCase")] pub struct OutputNodeContent { - #[serde(default, deserialize_with = "empty_string_is_none")] - pub schema: Option>, + #[serde(default, deserialize_with = "empty_value_string_is_none")] + pub schema: Option>, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -225,3 +225,17 @@ where StringOrNull::Null => Ok(None), } } + +fn empty_value_string_is_none<'de, D>(deserializer: D) -> Result>, D::Error> +where + D: Deserializer<'de>, +{ + let s = empty_string_is_none(deserializer)?; + let Some(data) = s else { + return Ok(None); + }; + + Ok(Some( + serde_json::from_str(data.as_ref()).map_err(serde::de::Error::custom)?, + )) +} diff --git a/core/types/src/variable/impls.rs b/core/types/src/variable/impls.rs index 3b6508b3..e37be1cc 100644 --- a/core/types/src/variable/impls.rs +++ b/core/types/src/variable/impls.rs @@ -3,6 +3,7 @@ use rust_decimal::Decimal; use rust_decimal::prelude::FromPrimitive; use serde_json::Value; use std::collections::HashMap; +use std::ops::Deref; use std::rc::Rc; use std::sync::Arc; @@ -97,6 +98,20 @@ where } } +impl ToVariable for HashMap, V, S> +where + V: ToVariable, + S: std::hash::BuildHasher, +{ + fn to_variable(&self) -> Variable { + Variable::from_object( + self.iter() + .map(|(k, v)| (Rc::::from(k.deref()), v.to_variable())) + .collect(), + ) + } +} + impl ToVariable for HashMap where V: ToVariable,