diff --git a/core/engine/src/decision.rs b/core/engine/src/decision.rs index 28e91254..c3443d73 100644 --- a/core/engine/src/decision.rs +++ b/core/engine/src/decision.rs @@ -8,7 +8,7 @@ use crate::nodes::NodeHandlerExtensions; use crate::{DecisionGraphValidationError, EvaluationError}; use serde_json::Value; use std::cell::OnceCell; -use std::sync::{Arc, OnceLock}; +use std::sync::Arc; use zen_expression::variable::Variable; /// Represents a JDM decision which can be evaluated @@ -18,7 +18,6 @@ pub struct Decision { loader: DynamicLoader, adapter: DynamicCustomNode, validator_cache: ValidatorCache, - tokio_runtime: Arc>, } impl From for Decision { @@ -28,7 +27,6 @@ impl From for Decision { loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), validator_cache: ValidatorCache::default(), - tokio_runtime: Arc::new(OnceLock::new()), } } } @@ -40,7 +38,6 @@ impl From> for Decision { loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), validator_cache: ValidatorCache::default(), - tokio_runtime: Arc::new(OnceLock::new()), } } } @@ -51,61 +48,57 @@ impl Decision { self } - 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 fn evaluate( + pub async fn evaluate( &self, context: Variable, ) -> Result> { - self.evaluate_with_opts(context, Default::default()) + self.evaluate_with_opts(context, Default::default()).await } /// Evaluates a decision using in-memory reference with advanced options - pub fn evaluate_with_opts( + pub async fn evaluate_with_opts( &self, context: Variable, 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(), + max_depth: options.max_depth, + trace: options.trace, iteration: 0, 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)?; + let response = decision_graph.evaluate(context).await?; Ok(response) } - pub fn evaluate_serialized( + pub async 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, - }, - ); + let response = self + .evaluate_with_opts( + context, + EvaluationOptions { + trace: options.trace != EvaluationTraceKind::None, + max_depth: options.max_depth, + }, + ) + .await; match response { Ok(ok) => Ok(ok diff --git a/core/engine/src/decision_graph/graph.rs b/core/engine/src/decision_graph/graph.rs index 0137c7e9..ad9780f8 100644 --- a/core/engine/src/decision_graph/graph.rs +++ b/core/engine/src/decision_graph/graph.rs @@ -1,25 +1,23 @@ use crate::decision_graph::walker::{GraphWalker, NodeData, StableDiDecisionGraph}; use crate::engine::EvaluationTraceKind; 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::custom::CustomNodeHandler; +use crate::nodes::decision::DecisionNodeHandler; +use crate::nodes::decision_table::DecisionTableNodeHandler; +use crate::nodes::expression::ExpressionNodeHandler; +use crate::nodes::function::FunctionNodeHandler; +use crate::nodes::input::InputNodeHandler; +use crate::nodes::output::OutputNodeHandler; +use crate::nodes::transform_attributes::TransformAttributesExecution; use crate::nodes::{ - NodeContext, NodeContextBase, NodeHandler, NodeHandlerExtensions, NodeResponse, + NodeContext, NodeContextBase, NodeContextConfig, NodeDataType, NodeHandler, + NodeHandlerExtensions, NodeResponse, NodeResult, TraceDataType, }; use crate::EvaluationError; use ahash::{HashMap, HashMapExt}; 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; @@ -27,6 +25,7 @@ use std::time::Instant; use thiserror::Error; use zen_expression::variable::{ToVariable, Variable}; +#[derive(Debug)] pub struct DecisionGraph { initial_graph: StableDiDecisionGraph, graph: StableDiDecisionGraph, @@ -105,14 +104,13 @@ impl DecisionGraph { .count() } - pub fn evaluate( + pub async fn evaluate( &mut self, context: Variable, ) -> Result> { let root_start = Instant::now(); self.validate()?; - if self.iteration >= self.max_depth { return Err(Box::new(EvaluationError::DepthLimitExceeded)); } @@ -136,103 +134,71 @@ impl DecisionGraph { 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); + let (input, input_trace) = walker.incoming_node_data(&self.graph, nid, true); let mut base_ctx = NodeContextBase { id: node.id.clone(), name: node.name.clone(), - input: incoming_data.clone(), + input, extensions: self.extensions.clone(), iteration: self.iteration, - trace: self.trace, + config: NodeContextConfig { + max_depth: self.max_depth, + trace: self.trace, + ..Default::default() + }, }; let node_execution = match &node.kind { DecisionNodeKind::InputNode { content } => { base_ctx.input = context.clone(); - let ctx = NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - InputNodeHandler.handle(ctx) + handle_node(base_ctx, content.clone(), InputNodeHandler).await } DecisionNodeKind::OutputNode { content } => { terminate = true; - let ctx = NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - OutputNodeHandler.handle(ctx) - } - DecisionNodeKind::SwitchNode { .. } => { - let input_data = walker.incoming_node_data(&self.graph, nid, false); - - // walker.set_node_data(nid, input_data); - Ok(NodeResponse { - output: Variable::Null, - trace_data: None, - }) + handle_node(base_ctx, content.clone(), OutputNodeHandler).await } + DecisionNodeKind::SwitchNode { .. } => Ok(NodeResponse { + output: input_trace.clone(), + trace_data: None, + }), DecisionNodeKind::FunctionNode { content } => { - let ctx = NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - FunctionNodeHandler.handle(ctx) + handle_node(base_ctx, content.clone(), FunctionNodeHandler).await } DecisionNodeKind::DecisionNode { content } => { - let ctx = NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - DecisionNodeHandler.handle(ctx) + handle_node(base_ctx, content.clone(), DecisionNodeHandler::default()).await } DecisionNodeKind::DecisionTableNode { content } => { - let ctx = - NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - DecisionTableNodeHandler.handle(ctx) + handle_node(base_ctx, content.clone(), DecisionTableNodeHandler).await } DecisionNodeKind::ExpressionNode { content } => { - let ctx = NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - ExpressionNodeHandler.handle(ctx) + handle_node(base_ctx, content.clone(), ExpressionNodeHandler).await } DecisionNodeKind::CustomNode { content } => { - let ctx = NodeContext::::from_base( - base_ctx, - content.clone(), - ); - - CustomNodeHandler.handle(ctx) + handle_node(base_ctx, content.clone(), CustomNodeHandler).await } }; if let Some(nt) = &mut node_traces { + let input_trace = match &node.kind { + DecisionNodeKind::InputNode { .. } => Variable::Null, + _ => input_trace, + }; + let trace = match &node_execution { Ok(ok) => DecisionGraphTrace { id: node.id.clone(), name: node.name.clone(), - input: incoming_data, + input: input_trace, order: nt.len() as u32, - output: ok.output.clone(), + output: (!terminate).then(|| ok.output.clone()).unwrap_or_default(), 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, + input: input_trace, order: nt.len() as u32, output: Variable::Null, trace_data: err.trace.clone(), @@ -240,18 +206,24 @@ impl DecisionGraph { }, }; - nt.insert(node.id.clone(), trace); + if !matches!(node.kind, DecisionNodeKind::SwitchNode { .. }) { + nt.insert(node.id.clone(), trace); + } } + let output = node_execution?.output; + output.dot_remove("$nodes"); + walker.set_node_data( nid, NodeData { name: Rc::from(node.name.deref()), - data: node_execution?.output, + data: output, }, ); - if terminate { + // Terminate once Output node is reached + if matches!(node.kind, DecisionNodeKind::OutputNode { .. }) { break; } } @@ -349,19 +321,37 @@ pub struct DecisionGraphTrace { pub order: u32, } -pub(crate) fn error_trace(trace: &Option>) -> Option { - trace.as_ref().map(|s| { - s.values().for_each(|v| { - v.input.dot_remove("$nodes"); - v.output.dot_remove("$nodes"); - }); +async fn handle_node( + base_ctx: NodeContextBase, + content: NodeData, + handler: NodeHandlerType, +) -> NodeResult +where + TraceData: TraceDataType, + NodeData: NodeDataType, + NodeHandlerType: NodeHandler, +{ + let ctx = NodeContext::::from_base(base_ctx.clone(), content); + if let Some(transform_attributes) = handler.transform_attributes(&ctx) { + return transform_attributes + .run_with(base_ctx, move |input, has_more| { + let handler = handler.clone(); + let mut new_ctx = ctx.clone(); + new_ctx.input = input; - s.to_variable() - }) -} + async move { + match has_more { + false => handler.handle(new_ctx).await, + true => { + let result = handler.handle(new_ctx.clone()).await; + handler.after_transform_attributes(&new_ctx).await?; + result + } + } + } + }) + .await; + } -fn create_validator_cache_key(content: &Value) -> u64 { - let mut hasher = DefaultHasher::new(); - content.hash(&mut hasher); - hasher.finish() + handler.handle(ctx).await } diff --git a/core/engine/src/decision_graph/walker.rs b/core/engine/src/decision_graph/walker.rs index b26c9afb..38c03602 100644 --- a/core/engine/src/decision_graph/walker.rs +++ b/core/engine/src/decision_graph/walker.rs @@ -3,7 +3,7 @@ use fixedbitset::FixedBitSet; use petgraph::data::DataMap; use petgraph::matrix_graph::Zero; use petgraph::prelude::{EdgeIndex, NodeIndex, StableDiGraph}; -use petgraph::visit::{EdgeRef, IntoEdgesDirected, IntoNodeIdentifiers, VisitMap, Visitable}; +use petgraph::visit::{EdgeRef, IntoNodeIdentifiers, VisitMap, Visitable}; use petgraph::{Incoming, Outgoing}; use std::ops::Deref; use std::rc::Rc; @@ -92,7 +92,7 @@ impl GraphWalker { }) } - pub fn get_all_node_data(&self, g: &StableDiDecisionGraph) -> Variable { + pub fn get_all_node_data(&self) -> Variable { let node_values = self .node_data .iter() @@ -111,19 +111,19 @@ impl GraphWalker { g: &StableDiDecisionGraph, node_id: NodeIndex, with_nodes: bool, - ) -> Variable { - let value = self - .merge_node_data(g.neighbors_directed(node_id, Incoming)) - .depth_clone(1); + ) -> (Variable, Variable) { + let value = self.merge_node_data(g.neighbors_directed(node_id, Incoming)); - 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(); - object.insert(Rc::from("$nodes"), self.get_all_node_data(g)); + if self.nodes_in_context && with_nodes { + if let Some(object_ref) = value.as_object() { + let mut new_object = object_ref.borrow().clone(); + new_object.insert(Rc::from("$nodes"), self.get_all_node_data()); + + return (Variable::from_object(new_object), value); } } - value + (value.depth_clone(1), value) } pub fn merge_node_data(&self, iter: I) -> Variable @@ -163,8 +163,8 @@ impl GraphWalker { 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 mut isolate = Isolate::with_environment(input_data.clone()); + let (input, input_trace) = self.incoming_node_data(g, nid, true); + let mut isolate = Isolate::with_environment(input); let mut statement_iter = content.statements.iter(); let valid_statements: Vec = match content.hit_policy { @@ -181,14 +181,12 @@ impl GraphWalker { .collect(), }; - input_data.dot_remove("$nodes"); - if let Some(on_trace) = &mut on_trace { on_trace(DecisionGraphTrace { id: decision_node.id.clone(), name: decision_node.name.clone(), - input: input_data.shallow_clone(), - output: input_data.shallow_clone(), + input: input_trace.shallow_clone(), + output: input_trace, order: 0, performance: Some(Arc::from(format!("{:.1?}", start.elapsed()))), trace_data: Some( @@ -207,7 +205,7 @@ impl GraphWalker { edge.weight().source_handle.as_ref().map_or(true, |handle| { !valid_statements .iter() - .any(|s| s.id.deref() == handle.as_str()) + .any(|s| s.id.deref() == handle.deref()) }) }) .map(|edge| edge.id()) diff --git a/core/engine/src/engine.rs b/core/engine/src/engine.rs index ae9191e3..9df7947e 100644 --- a/core/engine/src/engine.rs +++ b/core/engine/src/engine.rs @@ -1,15 +1,13 @@ use crate::decision::Decision; use crate::decision_graph::graph::DecisionGraphResponse; -use crate::loader::{ - ClosureLoader, DecisionLoader, DynamicLoader, LoaderResponse, LoaderResult, NoopLoader, -}; +use crate::loader::{ClosureLoader, 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, OnceLock}; +use std::sync::Arc; use strum::{EnumString, IntoStaticStr}; use zen_expression::variable::Variable; @@ -18,19 +16,36 @@ use zen_expression::variable::Variable; pub struct DecisionEngine { loader: DynamicLoader, adapter: DynamicCustomNode, - runtime: Arc>, } -#[derive(Debug, Default)] +#[derive(Debug)] pub struct EvaluationOptions { - pub trace: Option, - pub max_depth: Option, + pub trace: bool, + pub max_depth: u8, } -#[derive(Debug, Default)] +impl Default for EvaluationOptions { + fn default() -> Self { + Self { + trace: false, + max_depth: 10, + } + } +} + +#[derive(Debug)] pub struct EvaluationSerializedOptions { pub trace: EvaluationTraceKind, - pub max_depth: Option, + pub max_depth: u8, +} + +impl Default for EvaluationSerializedOptions { + fn default() -> Self { + Self { + trace: EvaluationTraceKind::None, + max_depth: 10, + } + } } #[derive(Debug, Default, PartialEq, Eq, EnumString, IntoStaticStr)] @@ -67,7 +82,6 @@ impl Default for DecisionEngine { Self { loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), - runtime: Arc::new(OnceLock::new()), } } } @@ -93,7 +107,7 @@ impl DecisionEngine { pub fn with_closure_loader(mut self, loader: F) -> Self where - F: Fn(String) -> O + Sync + Send + Debug + 'static, + F: Fn(String) -> O + Sync + Send + 'static, O: Future + Send, { self.loader = Arc::new(ClosureLoader::new(loader)); @@ -125,7 +139,7 @@ impl DecisionEngine { { let content = self.loader.load(key.as_ref()).await?; let decision = self.create_decision(content); - decision.evaluate_with_opts(context, options) + decision.evaluate_with_opts(context, options).await } pub async fn evaluate_serialized( @@ -144,7 +158,7 @@ impl DecisionEngine { .map_err(|err| Value::String(err.to_string()))?; let decision = self.create_decision(content); - decision.evaluate_serialized(context, options) + decision.evaluate_serialized(context, options).await } /// Creates a decision from DecisionContent, exists for easier binding creation @@ -152,7 +166,6 @@ 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 0999f288..05a1e59e 100644 --- a/core/engine/src/error.rs +++ b/core/engine/src/error.rs @@ -12,7 +12,7 @@ pub enum EvaluationError { #[error("Loader error")] LoaderError(LoaderError), - #[error("Node error")] + #[error("{0}")] NodeError(NodeError), #[error("Depth limit exceeded")] @@ -43,10 +43,7 @@ impl EvaluationError { EvaluationError::NodeError(err) => { map.serialize_entry("type", "NodeError")?; map.serialize_entry("source", &err.source.to_string())?; - - if let Some(node_id) = &err.node_id { - map.serialize_entry("nodeId", node_id)?; - } + map.serialize_entry("nodeId", &err.node_id)?; if let Some(trace) = &err.trace { map.serialize_entry("trace", &mode.serialize_trace(trace))?; diff --git a/core/engine/src/loader/closure.rs b/core/engine/src/loader/closure.rs index 651261a7..ebf7ff9d 100644 --- a/core/engine/src/loader/closure.rs +++ b/core/engine/src/loader/closure.rs @@ -1,10 +1,9 @@ use crate::loader::{DecisionLoader, LoaderResponse}; -use std::fmt::Debug; +use std::fmt::{Debug, Formatter}; use std::future::Future; use std::pin::Pin; /// Loads decisions using an async closure -#[derive(Debug)] pub struct ClosureLoader where F: Sync + Send, @@ -12,6 +11,15 @@ where closure: F, } +impl Debug for ClosureLoader +where + T: Sync + Send, +{ + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "ClosureLoader") + } +} + impl ClosureLoader where F: Fn(String) -> O + Sync + Send, @@ -24,7 +32,7 @@ where impl DecisionLoader for ClosureLoader where - F: Fn(String) -> O + Sync + Send + Debug, + F: Fn(String) -> O + Sync + Send, O: Future + 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 89e87d1e..ee19b64f 100644 --- a/core/engine/src/nodes/context.rs +++ b/core/engine/src/nodes/context.rs @@ -3,18 +3,20 @@ use crate::nodes::extensions::NodeHandlerExtensions; use crate::nodes::function::v2::function::Function; use crate::nodes::result::{NodeResponse, NodeResult}; use crate::nodes::NodeError; +use crate::ZEN_CONFIG; 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::hash::Hasher; +use std::sync::atomic::Ordering; use std::sync::Arc; use thiserror::Error; use zen_types::variable::Variable; +#[derive(Clone)] pub struct NodeContext where NodeData: NodeDataType, @@ -27,6 +29,7 @@ where pub trace: Option>, pub extensions: NodeHandlerExtensions, pub iteration: u8, + pub config: NodeContextConfig, } impl NodeContext @@ -41,8 +44,9 @@ where input: base.input, extensions: base.extensions, iteration: base.iteration, - trace: base.trace.then(|| Default::default()), + trace: base.config.trace.then(|| Default::default()), node: data, + config: base.config, } } @@ -55,10 +59,6 @@ where } } - pub fn has_trace(&self) -> bool { - self.trace.is_some() - } - pub fn error(&self, error: Error) -> NodeResult where Error: Into>, @@ -78,29 +78,14 @@ where Error: Into>, { NodeError { - node_id: Some(self.id.clone()), + node_id: self.id.clone(), trace: self.trace.as_ref().map(|v| (*v.borrow()).to_variable()), source: error.into(), } } - pub fn block_on(&self, future: Fut) -> Result - where - Fut: Future, - { - let tokio_runtime = self.extensions.tokio_runtime(); - Ok(tokio_runtime.block_on(future)) - } - - pub fn try_block_on(&self, future: Fut) -> Result - where - Fut: Future>, - { - self.block_on(future)? - } - - pub(crate) fn function_runtime(&self) -> Result<&Function, NodeError> { - Ok(self.extensions.function_runtime()) + pub(crate) async fn function_runtime(&self) -> Result<&Function, NodeError> { + self.extensions.function_runtime().await.node_context(self) } pub fn validate(&self, schema: &Value, value: &Value) -> Result<(), NodeError> { @@ -189,13 +174,14 @@ where } } +#[derive(Clone)] pub struct NodeContextBase { pub id: Arc, pub name: Arc, pub input: Variable, pub iteration: u8, pub extensions: NodeHandlerExtensions, - pub trace: bool, + pub config: NodeContextConfig, } impl NodeContextBase { @@ -218,7 +204,7 @@ impl NodeContextBase { Error: Into>, { NodeError { - node_id: Some(self.id.clone()), + node_id: self.id.clone(), source: error.into(), trace: None, } @@ -237,7 +223,7 @@ where input: value.input, extensions: value.extensions, iteration: value.iteration, - trace: value.trace.is_some(), + config: value.config, } } } @@ -298,3 +284,22 @@ impl<'a> From> for ValidationErrorJson { } } } + +#[derive(Clone)] +pub struct NodeContextConfig { + pub trace: bool, + pub nodes_in_context: bool, + pub max_depth: u8, + pub function_timeout_millis: u64, +} + +impl Default for NodeContextConfig { + fn default() -> Self { + Self { + trace: false, + nodes_in_context: ZEN_CONFIG.nodes_in_context.load(Ordering::Relaxed), + function_timeout_millis: ZEN_CONFIG.function_timeout_millis.load(Ordering::Relaxed), + max_depth: 5, + } + } +} diff --git a/core/engine/src/nodes/custom/adapter.rs b/core/engine/src/nodes/custom/adapter.rs index 34417453..513922d7 100644 --- a/core/engine/src/nodes/custom/adapter.rs +++ b/core/engine/src/nodes/custom/adapter.rs @@ -27,7 +27,7 @@ impl CustomNodeAdapter for NoopCustomNode { Box::pin(async move { Err(NodeError { trace: None, - node_id: Some(request.node.id.clone()), + node_id: request.node.id.clone(), source: "Custom node handler not provided".to_string().into(), }) }) diff --git a/core/engine/src/nodes/custom/mod.rs b/core/engine/src/nodes/custom/mod.rs index 2af0a5ff..f875538e 100644 --- a/core/engine/src/nodes/custom/mod.rs +++ b/core/engine/src/nodes/custom/mod.rs @@ -9,7 +9,9 @@ pub use adapter::{ mod adapter; +#[derive(Debug, Clone)] pub struct CustomNodeHandler; + pub type CustomNodeData = CustomNodeContent; pub type CustomNodeTrace = Variable; @@ -17,7 +19,7 @@ impl NodeHandler for CustomNodeHandler { type NodeData = CustomNodeData; type TraceData = CustomNodeTrace; - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn handle(&self, ctx: NodeContext) -> NodeResult { let custom_node_request = CustomNodeRequest { input: ctx.input.clone(), node: CustomDecisionNode { @@ -28,6 +30,9 @@ impl NodeHandler for CustomNodeHandler { }, }; - ctx.block_on(ctx.extensions.custom_node().handle(custom_node_request))? + ctx.extensions + .custom_node() + .handle(custom_node_request) + .await } } diff --git a/core/engine/src/nodes/decision/mod.rs b/core/engine/src/nodes/decision/mod.rs index ff4f024e..66875dcd 100644 --- a/core/engine/src/nodes/decision/mod.rs +++ b/core/engine/src/nodes/decision/mod.rs @@ -1,11 +1,16 @@ use crate::decision_graph::graph::{DecisionGraph, DecisionGraphConfig}; -use crate::nodes::{NodeContext, NodeContextExt, NodeHandler, NodeResult}; +use crate::nodes::{NodeContext, NodeContextExt, NodeError, NodeHandler, NodeResult}; use crate::EvaluationError; +use std::cell::RefCell; use std::ops::Deref; +use std::rc::Rc; use zen_types::decision::{DecisionNodeContent, TransformAttributes}; use zen_types::variable::{ToVariable, Variable}; -pub struct DecisionNodeHandler; +#[derive(Debug, Clone, Default)] +pub struct DecisionNodeHandler { + decision_graph: Rc>>, +} pub type DecisionNodeData = DecisionNodeContent; pub type DecisionNodeTrace = Variable; @@ -21,21 +26,44 @@ impl NodeHandler for DecisionNodeHandler { Some(ctx.node.transform_attributes.clone()) } - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn after_transform_attributes( + &self, + _ctx: &NodeContext, + ) -> Result<(), NodeError> { + if let Some(graph) = self.decision_graph.borrow_mut().as_mut() { + graph.reset_graph(); + }; + + Ok(()) + } + + async 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 sub_decision = 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)?; + let mut decision_graph_ref = self.decision_graph.borrow_mut(); + let decision_graph = match decision_graph_ref.as_mut() { + Some(dg) => dg, + None => { + let dg = DecisionGraph::try_new(DecisionGraphConfig { + content: sub_decision, + extensions: ctx.extensions.clone(), + trace: ctx.config.trace, + iteration: ctx.iteration + 1, + max_depth: ctx.config.max_depth, + }) + .node_context(&ctx)?; - match decision_graph.evaluate(ctx.input.clone()) { + *decision_graph_ref = Some(dg); + match decision_graph_ref.as_mut() { + Some(dg) => dg, + None => return ctx.error("Failed to initialize decision graph".to_string()), + } + } + }; + + let evaluate_result = Box::pin(decision_graph.evaluate(ctx.input.clone())).await; + match evaluate_result { Ok(result) => { ctx.trace(|trace| { *trace = result.trace.to_variable(); diff --git a/core/engine/src/nodes/decision_table/mod.rs b/core/engine/src/nodes/decision_table/mod.rs index 2aea2236..18b46661 100644 --- a/core/engine/src/nodes/decision_table/mod.rs +++ b/core/engine/src/nodes/decision_table/mod.rs @@ -1,6 +1,6 @@ use crate::nodes::definition::NodeHandler; use crate::nodes::result::NodeResult; -use crate::nodes::NodeContext; +use crate::nodes::{NodeContext, NodeResponse}; use ahash::HashMap; use serde::Serialize; use std::ops::Deref; @@ -10,7 +10,7 @@ use zen_expression::variable::ToVariable; use zen_expression::Isolate; use zen_types::decision::{DecisionTableContent, DecisionTableHitPolicy, TransformAttributes}; use zen_types::variable::Variable; - +#[derive(Debug, Clone)] pub struct DecisionTableNodeHandler; pub type DecisionTableNodeData = DecisionTableContent; @@ -28,7 +28,7 @@ impl NodeHandler for DecisionTableNodeHandler { Some(ctx.node.transform_attributes.clone()) } - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn handle(&self, ctx: NodeContext) -> NodeResult { match ctx.node.hit_policy { DecisionTableHitPolicy::First => self.handle_first_hit(ctx), DecisionTableHitPolicy::Collect => self.handle_collect(ctx), @@ -39,6 +39,7 @@ impl NodeHandler for DecisionTableNodeHandler { impl DecisionTableNodeHandler { fn handle_first_hit(&self, ctx: DecisionTableContext) -> NodeResult { let mut isolate = Isolate::new(); + isolate.set_environment(ctx.input.depth_clone(1)); for (index, rule) in ctx.node.rules.iter().enumerate() { if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate) { @@ -63,13 +64,17 @@ impl DecisionTableNodeHandler { } } - ctx.success(Variable::Null) + Ok(NodeResponse { + output: Variable::Null, + trace_data: None, + }) } fn handle_collect(&self, ctx: DecisionTableContext) -> NodeResult { let mut isolate = Isolate::new(); let mut outputs = Vec::new(); let mut traces = Vec::new(); + isolate.set_environment(ctx.input.depth_clone(1)); for (index, rule) in ctx.node.rules.iter().enumerate() { if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate) { @@ -107,9 +112,9 @@ impl DecisionTableNodeHandler { isolate: &mut Isolate<'a>, ) -> Option { let content = &ctx.node; - for input in &content.inputs { + for input in content.inputs.iter() { let rule_value = rule.get(&input.id)?; - if rule_value.trim().is_empty() { + if rule_value.is_empty() { continue; } @@ -129,19 +134,19 @@ impl DecisionTableNodeHandler { } } - let mut outputs: HashMap, Variable> = Default::default(); - for output in &content.outputs { + let outputs = Variable::empty_object(); + for output in content.outputs.iter() { let rule_value = rule.get(&output.id)?; - if rule_value.trim().is_empty() { + if rule_value.is_empty() { continue; } let res = isolate.run_standard(rule_value).ok()?; - outputs.insert(Rc::from(&*output.field), res); + outputs.dot_insert(output.field.deref(), res); } - if !ctx.has_trace() { - return Some(RowResult::Output(outputs.to_variable())); + if !ctx.config.trace { + return Some(RowResult::Output(outputs)); } let id_str = Rc::::from("_id"); @@ -160,7 +165,7 @@ impl DecisionTableNodeHandler { expressions.insert(description_str.clone(), Rc::from(description.deref())); } - for input in &content.inputs { + for input in content.inputs.iter() { let rule_value = rule.get(input.id.deref())?; let Some(input_field) = &input.field else { continue; diff --git a/core/engine/src/nodes/definition.rs b/core/engine/src/nodes/definition.rs index 219e7024..ee660fa6 100644 --- a/core/engine/src/nodes/definition.rs +++ b/core/engine/src/nodes/definition.rs @@ -1,15 +1,9 @@ use crate::nodes::context::NodeContext; -use crate::nodes::function::FunctionNodeTrace; -use crate::nodes::input::InputNodeTrace; -use crate::nodes::output::OutputNodeTrace; use crate::nodes::result::NodeResult; +use crate::nodes::NodeError; use serde::{Deserialize, Serialize}; use std::fmt::Debug; -use std::sync::Arc; -use zen_types::decision::{ - CustomNodeContent, DecisionNodeContent, ExpressionNodeContent, FunctionNodeContent, - InputNodeContent, OutputNodeContent, TransformAttributes, -}; +use zen_types::decision::TransformAttributes; use zen_types::variable::ToVariable; pub trait NodeDataType: Clone + Debug + Serialize + for<'de> Deserialize<'de> {} @@ -18,35 +12,28 @@ impl NodeDataType for T where T: Clone + Debug + Serialize + for<'de> Deseria pub trait TraceDataType: Clone + Debug + Default + ToVariable {} impl TraceDataType for T where T: Clone + Debug + Default + ToVariable {} -pub trait NodeHandler { +pub trait NodeHandler: Clone { type NodeData: NodeDataType; type TraceData: TraceDataType; + #[allow(unused_variables)] fn transform_attributes( &self, - _ctx: &NodeContext, + ctx: &NodeContext, ) -> Option { None } - fn handle(&self, ctx: NodeContext) -> NodeResult; -} + #[allow(unused_variables)] + fn after_transform_attributes( + &self, + ctx: &NodeContext, + ) -> impl std::future::Future> { + Box::pin(async { Ok(()) }) + } -pub struct NodeHandlers { - pub input: Arc>, - pub output: Arc>, - pub function: - Arc>, - pub expression: - Box>, -} - -pub enum NodeHandlerKind { - Input(), - Output(Box>), - Function(Box>), - Expression(Box>), - DecisionTable(Box>), - Decision(Box>), - Custom(Box>), + fn handle( + &self, + ctx: NodeContext, + ) -> impl std::future::Future; } diff --git a/core/engine/src/nodes/expression/mod.rs b/core/engine/src/nodes/expression/mod.rs index a4f59b8c..7c6d2d75 100644 --- a/core/engine/src/nodes/expression/mod.rs +++ b/core/engine/src/nodes/expression/mod.rs @@ -9,6 +9,7 @@ use zen_expression::variable::{ToVariable, Variable}; use zen_expression::Isolate; use zen_types::decision::TransformAttributes; +#[derive(Debug, Clone)] pub struct ExpressionNodeHandler; pub type ExpressionNodeData = ExpressionNodeContent; @@ -25,13 +26,12 @@ impl NodeHandler for ExpressionNodeHandler { Some(ctx.node.transform_attributes.clone()) } - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn handle(&self, ctx: NodeContext) -> NodeResult { let result = Variable::empty_object(); let mut isolate = Isolate::new(); - isolate.set_environment(ctx.input.depth_clone(1)); - for expression in &ctx.node.expressions { + for expression in ctx.node.expressions.iter() { if expression.key.is_empty() || expression.value.is_empty() { continue; } diff --git a/core/engine/src/nodes/extensions.rs b/core/engine/src/nodes/extensions.rs index 3c865ac0..d271ae1d 100644 --- a/core/engine/src/nodes/extensions.rs +++ b/core/engine/src/nodes/extensions.rs @@ -4,14 +4,14 @@ 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 crate::nodes::validator_cache::ValidatorCache; +use anyhow::Context; use std::cell::OnceCell; -use std::sync::{Arc, OnceLock}; +use std::sync::Arc; /// This is created on every graph evaluation -#[derive(Clone)] +#[derive(Debug, Clone)] pub struct NodeHandlerExtensions { - pub(crate) function_runtime: Arc>, - pub(crate) tokio_runtime: Arc>, + pub(crate) function_runtime: Arc>, pub(crate) validator_cache: Arc>, pub(crate) loader: DynamicLoader, pub(crate) custom_node: DynamicCustomNode, @@ -20,7 +20,6 @@ pub struct NodeHandlerExtensions { impl Default for NodeHandlerExtensions { fn default() -> Self { Self { - tokio_runtime: Default::default(), function_runtime: Default::default(), validator_cache: Default::default(), @@ -31,33 +30,21 @@ impl Default for NodeHandlerExtensions { } impl NodeHandlerExtensions { - pub fn tokio_runtime(&self) -> &tokio::runtime::Runtime { - self.tokio_runtime.get_or_init(|| { - println!("Creating tokio runtime"); - - tokio::runtime::Builder::new_multi_thread() - .worker_threads(2) - .enable_all() - .build() - .expect("Failed to build tokio runtime") - }) - } - - pub fn function_runtime(&self) -> &Function { - self.function_runtime.get_or_init(|| { - let tokio_runtime = self.tokio_runtime(); - - tokio_runtime - .block_on(Function::create(FunctionConfig { + pub async fn function_runtime(&self) -> anyhow::Result<&Function> { + self.function_runtime + .get_or_try_init(|| { + Function::create(FunctionConfig { listeners: Some(vec![ Box::new(ConsoleListener), Box::new(ZenListener { - extensions: self.clone(), + loader: self.loader.clone(), + custom_node: self.custom_node.clone(), }), ]), - })) - .expect("Failed to create async function") - }) + }) + }) + .await + .context("Failed to create function") } pub fn validator_cache(&self) -> &ValidatorCache { diff --git a/core/engine/src/nodes/function/mod.rs b/core/engine/src/nodes/function/mod.rs index 4e3c947b..014eda44 100644 --- a/core/engine/src/nodes/function/mod.rs +++ b/core/engine/src/nodes/function/mod.rs @@ -10,6 +10,7 @@ use std::sync::Arc; use zen_types::decision::{FunctionContent, FunctionNodeContent}; use zen_types::variable::Variable; +#[derive(Debug, Clone)] pub struct FunctionNodeHandler; pub type FunctionNodeData = FunctionNodeContent; @@ -20,7 +21,7 @@ impl NodeHandler for FunctionNodeHandler { type NodeData = FunctionNodeData; type TraceData = FunctionNodeTrace; - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn handle(&self, ctx: NodeContext) -> NodeResult { match &ctx.node { FunctionNodeContent::Version1(source) => { let v1_context = NodeContext::, FunctionV1Trace> { @@ -28,12 +29,13 @@ impl NodeHandler for FunctionNodeHandler { name: ctx.name.clone(), input: ctx.input.clone(), extensions: ctx.extensions.clone(), + trace: ctx.config.trace.then(|| Default::default()), iteration: ctx.iteration, + config: ctx.config, node: source.clone(), - trace: None, }; - FunctionV1NodeHandler.handle(v1_context) + FunctionV1NodeHandler.handle(v1_context).await } FunctionNodeContent::Version2(content) => { let v2_context = NodeContext:: { @@ -41,12 +43,13 @@ impl NodeHandler for FunctionNodeHandler { name: ctx.name.clone(), input: ctx.input.clone(), extensions: ctx.extensions.clone(), + trace: ctx.config.trace.then(|| Default::default()), iteration: ctx.iteration, + config: ctx.config, node: content.clone(), - trace: None, }; - FunctionV2NodeHandler.handle(v2_context) + FunctionV2NodeHandler.handle(v2_context).await } } } diff --git a/core/engine/src/nodes/function/v1/mod.rs b/core/engine/src/nodes/function/v1/mod.rs index b68e7b8c..421690cf 100644 --- a/core/engine/src/nodes/function/v1/mod.rs +++ b/core/engine/src/nodes/function/v1/mod.rs @@ -13,6 +13,7 @@ use zen_expression::variable::ToVariable; pub(crate) mod runtime; mod script; +#[derive(Debug, Clone)] pub struct FunctionV1NodeHandler; const MAX_DURATION: Duration = Duration::from_millis(500); @@ -21,7 +22,7 @@ impl NodeHandler for FunctionV1NodeHandler { type NodeData = Arc; type TraceData = FunctionV1Trace; - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn handle(&self, ctx: NodeContext) -> NodeResult { let start = Instant::now(); let runtime = create_runtime().node_context_message(&ctx, "Failed to create JS Runtime")?; let interrupt_handler = Box::new(move || start.elapsed() > MAX_DURATION); @@ -29,7 +30,7 @@ impl NodeHandler for FunctionV1NodeHandler { runtime.set_interrupt_handler(Some(interrupt_handler)); let mut script = Script::new(runtime.clone()); - let result_response = ctx.block_on(script.call(ctx.node.deref(), &ctx.input))?; + let result_response = script.call(ctx.node.deref(), &ctx.input).await; runtime.set_interrupt_handler(None); diff --git a/core/engine/src/nodes/function/v2/function.rs b/core/engine/src/nodes/function/v2/function.rs index e3d31804..269f2f0e 100644 --- a/core/engine/src/nodes/function/v2/function.rs +++ b/core/engine/src/nodes/function/v2/function.rs @@ -1,3 +1,4 @@ +use std::fmt::{Debug, Formatter}; use std::hash::{DefaultHasher, Hash, Hasher}; use std::sync::Arc; @@ -22,6 +23,12 @@ pub struct Function { module_loader: ModuleLoader, } +impl Debug for Function { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "Function") + } +} + impl Function { pub async fn create<'js>(config: FunctionConfig) -> FunctionResult { let module_loader = ModuleLoader::new(); diff --git a/core/engine/src/nodes/function/v2/mod.rs b/core/engine/src/nodes/function/v2/mod.rs index d01f39b1..631f5bdc 100644 --- a/core/engine/src/nodes/function/v2/mod.rs +++ b/core/engine/src/nodes/function/v2/mod.rs @@ -1,5 +1,4 @@ use std::ops::Deref; -use std::rc::Rc; use std::time::Duration; use crate::nodes::definition::NodeHandler; @@ -21,72 +20,63 @@ pub(crate) mod listener; pub(crate) mod module; pub(crate) mod serde; -pub struct FunctionHandler { - function: Rc, - trace: bool, - iteration: u8, - max_depth: u8, - max_duration: Duration, -} - +#[derive(Debug, Clone)] pub struct FunctionV2NodeHandler; impl NodeHandler for FunctionV2NodeHandler { type NodeData = FunctionContent; type TraceData = FunctionV2Trace; - fn handle(&self, ctx: NodeContext) -> NodeResult { + async fn handle(&self, ctx: NodeContext) -> NodeResult { let start = std::time::Instant::now(); // TODO: Smart node omit - let function = ctx.function_runtime()?; + let function = ctx.function_runtime().await?; let module_name = function.suggest_module_name(ctx.id.deref(), ctx.node.source.deref()); // TODO: Add duration from configuration let max_duration = Duration::from_millis(500); let interrupt_handler = Box::new(move || start.elapsed() > max_duration); - ctx.try_block_on(async { - function - .runtime() - .set_interrupt_handler(Some(interrupt_handler)) - .await; + function + .runtime() + .set_interrupt_handler(Some(interrupt_handler)) + .await; - self.attach_globals(function).await.node_context(&ctx)?; + self.attach_globals(function).await.node_context(&ctx)?; - function - .register_module(&module_name, ctx.node.source.deref()) - .await - .node_context(&ctx)?; + function + .register_module(&module_name, ctx.node.source.deref()) + .await + .node_context(&ctx)?; - let response_result = function - .call_handler(&module_name, JsValue(ctx.input.clone())) - .await; + let response_result = function + .call_handler(&module_name, JsValue(ctx.input.clone())) + .await; - match response_result { - Ok(response) => { - function.runtime().set_interrupt_handler(None).await; - ctx.trace(|t| { - t.log = response.logs.clone(); - }); + match response_result { + Ok(response) => { + function.runtime().set_interrupt_handler(None).await; + ctx.trace(|t| { + t.log = response.logs.clone(); + }); - ctx.success(response.data) - } - Err(e) => { - let log = function.extract_logs().await; - ctx.trace(|t| { - t.log = log; - t.log.push(Log { - lines: vec![json!(e.to_string()).to_string()], - ms_since_run: start.elapsed().as_millis() as usize, - }); - }); - - ctx.error(e) - } + ctx.success(response.data) } - }) + Err(e) => { + let log = function.extract_logs().await; + ctx.trace(|t| { + t.log = log; + t.log.push(Log { + lines: vec![json!(e.to_string()).to_string()], + ms_since_run: start.elapsed().as_millis() as usize, + }); + }); + + ctx.error(e) + } + } } } diff --git a/core/engine/src/nodes/function/v2/module/zen.rs b/core/engine/src/nodes/function/v2/module/zen.rs index ed1bed32..565b2292 100644 --- a/core/engine/src/nodes/function/v2/module/zen.rs +++ b/core/engine/src/nodes/function/v2/module/zen.rs @@ -1,7 +1,6 @@ -use std::future::Future; -use std::pin::Pin; - use crate::decision_graph::graph::{DecisionGraph, DecisionGraphConfig}; +use crate::loader::DynamicLoader; +use crate::nodes::custom::DynamicCustomNode; use crate::nodes::function::v2::error::{FunctionResult, ResultExt}; use crate::nodes::function::v2::listener::{RuntimeEvent, RuntimeListener}; use crate::nodes::function::v2::module::export_default; @@ -10,9 +9,12 @@ use crate::nodes::NodeHandlerExtensions; use rquickjs::module::{Declarations, Exports, ModuleDef}; use rquickjs::prelude::{Async, Func, Opt}; use rquickjs::{CatchResultExt, Ctx, Function, Object}; +use std::future::Future; +use std::pin::Pin; pub(crate) struct ZenListener { - pub extensions: NodeHandlerExtensions, + pub loader: DynamicLoader, + pub custom_node: DynamicCustomNode, } impl RuntimeListener for ZenListener { @@ -21,7 +23,9 @@ impl RuntimeListener for ZenListener { ctx: Ctx<'js>, event: RuntimeEvent, ) -> Pin + 'js>> { - let extensions = self.extensions.clone(); + let loader = self.loader.clone(); + let custom_node = self.custom_node.clone(); + Box::pin(async move { if event != RuntimeEvent::Startup { return Ok(()); @@ -35,7 +39,8 @@ impl RuntimeListener for ZenListener { key: String, context: JsValue, opts: Opt>| { - let extensions = extensions.clone(); + let loader = loader.clone(); + let custom_node = custom_node.clone(); async move { let config: Object = ctx.globals().get("config").or_throw(&ctx)?; @@ -47,18 +52,22 @@ impl RuntimeListener for ZenListener { .map(|opt| opt.get::<_, bool>("trace").unwrap_or_default()) .unwrap_or_default(); - let load_result = extensions.loader().load(key.as_str()).await; + 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, max_depth, iteration: iteration + 1, trace, - extensions: extensions.clone(), + extensions: NodeHandlerExtensions { + loader: loader.clone(), + custom_node: custom_node.clone(), + ..Default::default() + }, }) .or_throw(&ctx)?; - let response = sub_tree.evaluate(context.0).or_throw(&ctx)?; + let response = sub_tree.evaluate(context.0).await.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 11b82f6d..11b87cf2 100644 --- a/core/engine/src/nodes/input/mod.rs +++ b/core/engine/src/nodes/input/mod.rs @@ -4,6 +4,7 @@ use crate::nodes::NodeContext; use zen_types::decision::InputNodeContent; use zen_types::variable::Variable; +#[derive(Debug, Clone)] pub struct InputNodeHandler; pub type InputNodeData = InputNodeContent; @@ -13,7 +14,7 @@ impl NodeHandler for InputNodeHandler { type NodeData = InputNodeData; type TraceData = InputNodeTrace; - fn handle(&self, ctx: NodeContext) -> NodeResult { + async 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)?; diff --git a/core/engine/src/nodes/mod.rs b/core/engine/src/nodes/mod.rs index 2b3cdd07..7202f4a4 100644 --- a/core/engine/src/nodes/mod.rs +++ b/core/engine/src/nodes/mod.rs @@ -12,7 +12,8 @@ mod result; pub(crate) mod transform_attributes; pub(crate) mod validator_cache; -pub use context::{NodeContext, NodeContextBase, NodeContextExt}; -pub use definition::{NodeHandler, NodeHandlerKind}; +pub use context::{NodeContext, NodeContextBase, NodeContextConfig, NodeContextExt}; +pub use definition::NodeHandler; +pub(crate) use definition::{NodeDataType, TraceDataType}; 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 16b5d910..47ce3c75 100644 --- a/core/engine/src/nodes/output/mod.rs +++ b/core/engine/src/nodes/output/mod.rs @@ -4,6 +4,7 @@ use crate::nodes::NodeContext; use zen_types::decision::OutputNodeContent; use zen_types::variable::Variable; +#[derive(Debug, Clone)] pub struct OutputNodeHandler; pub type OutputNodeData = OutputNodeContent; @@ -13,7 +14,7 @@ impl NodeHandler for OutputNodeHandler { type NodeData = OutputNodeData; type TraceData = OutputNodeTrace; - fn handle(&self, ctx: NodeContext) -> NodeResult { + async 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)?; diff --git a/core/engine/src/nodes/result.rs b/core/engine/src/nodes/result.rs index 8b1b4da9..e84f56ce 100644 --- a/core/engine/src/nodes/result.rs +++ b/core/engine/src/nodes/result.rs @@ -1,7 +1,6 @@ use crate::model::DecisionNode; use serde::{Deserialize, Serialize}; use std::fmt::{Display, Formatter}; -use std::rc::Rc; use std::sync::Arc; use thiserror::Error; use zen_expression::variable::Variable; @@ -24,7 +23,7 @@ pub type NodeResult = Result; #[derive(Debug, Error)] pub struct NodeError { - pub node_id: Option>, + pub node_id: Arc, pub trace: Option, pub source: Box, } diff --git a/core/engine/src/nodes/transform_attributes.rs b/core/engine/src/nodes/transform_attributes.rs index 8818cace..4f94366a 100644 --- a/core/engine/src/nodes/transform_attributes.rs +++ b/core/engine/src/nodes/transform_attributes.rs @@ -5,17 +5,17 @@ use std::future::Future; use std::ops::Deref; use zen_expression::{Isolate, Variable}; -pub trait TransformAttributesExecution { +pub(crate) trait TransformAttributesExecution { async fn run_with(&self, ctx: NodeContextBase, evaluate: F) -> NodeResult where - F: Fn(Variable) -> Fut, + F: Fn(Variable, bool) -> Fut, Fut: Future; } impl TransformAttributesExecution for TransformAttributes { async fn run_with(&self, ctx: NodeContextBase, evaluate: F) -> NodeResult where - F: Fn(Variable) -> Fut, + F: Fn(Variable, bool) -> Fut, Fut: Future, { let input = match &self.input_field { @@ -54,7 +54,7 @@ impl TransformAttributesExecution for TransformAttributes { let mut trace_data: Option = None; let mut output = match self.execution_mode { TransformExecutionMode::Single => { - let response = evaluate(input).await?; + let response = evaluate(input, false).await?; if let Some(td) = response.trace_data { trace_data.replace(td); } @@ -70,8 +70,9 @@ impl TransformAttributesExecution for TransformAttributes { let mut output_array = Vec::with_capacity(input_array.len()); let mut trace_datum = Vec::with_capacity(input_array.len()); - for input in input_array.iter() { - let mut response = evaluate(input.clone()).await?; + for (index, input) in input_array.iter().enumerate() { + let has_more = index < input_array.len() - 1; + let mut response = evaluate(input.clone(), has_more).await?; if let Some(td) = response.trace_data { trace_datum.push(td); } diff --git a/core/engine/tests/decision.rs b/core/engine/tests/decision.rs index a487fb6b..2f78ce5d 100644 --- a/core/engine/tests/decision.rs +++ b/core/engine/tests/decision.rs @@ -3,7 +3,8 @@ use serde_json::json; use std::ops::Deref; use std::sync::Arc; use tokio::runtime::Builder; -use zen_engine::{Decision, DecisionGraphValidationError, EvaluationError, NodeError}; +use zen_engine::nodes::NodeError; +use zen_engine::{Decision, DecisionGraphValidationError, EvaluationError}; mod support; @@ -28,10 +29,10 @@ async fn decision_from_content_recursive() { let context = json!({}); let result = decision.evaluate(context.clone().into()).await; match result.unwrap_err().deref() { - EvaluationError::NodeError(NodeError::Node { + EvaluationError::NodeError(NodeError { node_id, source, .. }) => { - assert_eq!(node_id, "0b8dcf6b-fc04-47cb-bf82-bda764e6c09b"); + assert_eq!(node_id.deref(), "0b8dcf6b-fc04-47cb-bf82-bda764e6c09b"); assert!(source.to_string().contains("Loader failed")); } _ => assert!(false, "Depth limit not exceeded"), @@ -40,7 +41,7 @@ async fn decision_from_content_recursive() { let with_loader = decision.with_loader(Arc::new(create_fs_loader())); let new_result = with_loader.evaluate(context.clone().into()).await; match new_result.unwrap_err().deref() { - EvaluationError::NodeError(NodeError::Node { source, .. }) => { + EvaluationError::NodeError(NodeError { source, .. }) => { assert_eq!(source.to_string(), "Depth limit exceeded") } _ => assert!(false, "Depth limit not exceeded"), diff --git a/core/engine/tests/engine.rs b/core/engine/tests/engine.rs index 31592e7e..2e4f7e02 100644 --- a/core/engine/tests/engine.rs +++ b/core/engine/tests/engine.rs @@ -1,5 +1,4 @@ use crate::support::{create_fs_loader, load_raw_test_data, load_test_data, test_data_root}; -use chrono::{TimeZone, Utc}; use serde::Deserialize; use serde_json::json; use std::fs; @@ -10,8 +9,9 @@ use std::sync::Arc; use tokio::runtime::Builder; use zen_engine::loader::{LoaderError, MemoryLoader}; use zen_engine::model::{DecisionContent, DecisionNode, DecisionNodeKind, FunctionNodeContent}; +use zen_engine::nodes::NodeError; +use zen_engine::Variable; use zen_engine::{DecisionEngine, EvaluationError, EvaluationOptions}; -use zen_engine::{NodeError, Variable}; mod support; @@ -41,7 +41,7 @@ async fn engine_memory_loader() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_filesystem_loader() { - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); let table = engine .evaluate("table.json", json!({ "input": 12 }).into()) .await; @@ -92,7 +92,7 @@ fn engine_noop_loader() { #[test] fn engine_get_decision() { let rt = Builder::new_current_thread().build().unwrap(); - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); assert!(rt.block_on(engine.get_decision("table.json")).is_ok()); assert!(rt.block_on(engine.get_decision("any.json")).is_err()); @@ -107,16 +107,16 @@ fn engine_create_decision() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_errors() { - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); let infinite_fn = engine .evaluate("infinite-function.json", json!({}).into()) .await; match infinite_fn.unwrap_err().deref() { - EvaluationError::NodeError(NodeError::Node { + EvaluationError::NodeError(NodeError { node_id, source, .. }) => { - assert_eq!(node_id, "e0fd96d0-44dc-4f0e-b825-06e56b442d78"); + assert_eq!(node_id.deref(), "e0fd96d0-44dc-4f0e-b825-06e56b442d78"); assert!(source.to_string().contains("interrupted")); } _ => assert!(false, "Wrong error type"), @@ -126,7 +126,8 @@ async fn engine_errors() { .evaluate("recursive-table1.json", json!({}).into()) .await; match recursive.unwrap_err().deref() { - EvaluationError::NodeError(NodeError::Node { source, .. }) => { + EvaluationError::NodeError(NodeError { source, .. }) => { + println!("{:?}", source); assert_eq!(source.to_string(), "Depth limit exceeded") } _ => assert!(false, "Depth limit not exceeded"), @@ -136,15 +137,15 @@ async fn engine_errors() { #[test] fn engine_with_trace() { let rt = Builder::new_current_thread().build().unwrap(); - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); let table_r = rt.block_on(engine.evaluate("table.json", json!({ "input": 12 }).into())); let table_opt_r = rt.block_on(engine.evaluate_with_opts( "table.json", json!({ "input": 12 }).into(), EvaluationOptions { - trace: Some(true), - max_depth: None, + trace: true, + ..Default::default() }, )); @@ -174,7 +175,7 @@ async fn engine_function_imports() { .map(|node| match &node.kind { DecisionNodeKind::FunctionNode { .. } => { let new_kind = DecisionNodeKind::FunctionNode { - content: FunctionNodeContent::Version1(replace_data.clone()), + content: FunctionNodeContent::Version1(Arc::from(replace_data.as_str())), }; Arc::new(DecisionNode { @@ -213,7 +214,7 @@ async fn engine_function_imports() { #[tokio::test] async fn engine_switch_node() { - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); let switch_node_r = engine .evaluate("switch-node.json", json!({ "color": "yellow" }).into()) @@ -318,8 +319,8 @@ async fn engine_snapshot_tests() { .evaluate_with_opts( input.clone(), EvaluationOptions { - trace: Some(true), - max_depth: None, + trace: true, + ..Default::default() }, ) .await @@ -336,7 +337,7 @@ async fn engine_snapshot_tests() { #[tokio::test] #[cfg_attr(miri, ignore)] async fn engine_function_v2() { - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); for _ in 0..1_000 { let function_opt_r = engine @@ -344,8 +345,8 @@ async fn engine_function_v2() { "function-v2.json", json!({ "input": 12 }).into(), EvaluationOptions { - trace: Some(true), - max_depth: None, + trace: true, + ..Default::default() }, ) .await; @@ -365,7 +366,7 @@ async fn engine_function_v2() { #[tokio::test] async fn test_validation() { - let engine = DecisionEngine::default().with_loader(create_fs_loader().into()); + let engine = DecisionEngine::default().with_loader(Arc::new(create_fs_loader())); let context_valid = json!({ "color": "red", diff --git a/core/expression/src/isolate.rs b/core/expression/src/isolate.rs index 8412f817..6d98760e 100644 --- a/core/expression/src/isolate.rs +++ b/core/expression/src/isolate.rs @@ -1,8 +1,6 @@ -use ahash::AHasher; +use ahash::HashMap; use serde::ser::SerializeMap; use serde::{Serialize, Serializer}; -use std::collections::HashMap; -use std::hash::BuildHasherDefault; use std::rc::Rc; use std::sync::Arc; use thiserror::Error; @@ -16,8 +14,6 @@ use crate::variable::Variable; use crate::vm::{VMError, VM}; use crate::{Expression, ExpressionKind}; -type ADefHasher = BuildHasherDefault; - /// Isolate is a component that encapsulates an isolated environment for executing expressions. /// /// Rerunning the Isolate allows for efficient memory reuse through an arena allocator. @@ -32,7 +28,7 @@ pub struct Isolate<'arena> { bump: UnsafeArena<'arena>, environment: Option, - references: HashMap, + references: HashMap, } impl<'a> Isolate<'a> { diff --git a/core/macros/src/to_variable.rs b/core/macros/src/to_variable.rs index 005c5f66..bd676355 100644 --- a/core/macros/src/to_variable.rs +++ b/core/macros/src/to_variable.rs @@ -97,12 +97,6 @@ fn generate_enum_body( .filter(|variant| !variant.attrs.skip_serializing()) .collect(); - // Check if the enum is untagged - let is_untagged = matches!( - container.attrs.tag(), - serde_derive_internals::attr::TagType::None - ); - match container.attrs.tag() { TagType::None => { let variant_arms = active_variants diff --git a/core/types/src/decision/mod.rs b/core/types/src/decision/mod.rs index f4138b24..8d81f29e 100644 --- a/core/types/src/decision/mod.rs +++ b/core/types/src/decision/mod.rs @@ -17,7 +17,7 @@ pub struct DecisionEdge { pub id: Arc, pub source_id: Arc, pub target_id: Arc, - pub source_handle: Option, + pub source_handle: Option>, } #[derive(Clone, Debug, Deserialize, Serialize)] @@ -109,9 +109,10 @@ pub struct DecisionNodeContent { #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] #[serde(rename_all = "camelCase")] pub struct DecisionTableContent { - pub rules: Vec, Arc>>, - pub inputs: Vec, - pub outputs: Vec, + #[serde(deserialize_with = "deserialize_trim_rules")] + pub rules: Arc, Arc>>>, + pub inputs: Arc>, + pub outputs: Arc>, pub hit_policy: DecisionTableHitPolicy, #[serde(flatten)] pub transform_attributes: TransformAttributes, @@ -144,7 +145,7 @@ pub struct DecisionTableOutputField { #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] #[serde(rename_all = "camelCase")] pub struct ExpressionNodeContent { - pub expressions: Vec, + pub expressions: Arc>, #[serde(flatten)] pub transform_attributes: TransformAttributes, } @@ -162,7 +163,7 @@ pub struct Expression { pub struct SwitchNodeContent { #[serde(default)] pub hit_policy: SwitchStatementHitPolicy, - pub statements: Vec, + pub statements: Arc>, } #[derive(Clone, Debug, PartialEq, Deserialize, Serialize)] @@ -239,3 +240,23 @@ where serde_json::from_str(data.as_ref()).map_err(serde::de::Error::custom)?, )) } + +fn deserialize_trim_rules<'de, D>( + deserializer: D, +) -> Result, Arc>>>, D::Error> +where + D: Deserializer<'de>, +{ + let rules: Vec, Arc>> = Vec::deserialize(deserializer)?; + + let filtered_rules: Vec, Arc>> = rules + .into_iter() + .map(|rule| { + rule.into_iter() + .map(|(k, v)| (k, Arc::from(v.trim()))) + .collect() + }) + .collect(); + + Ok(Arc::new(filtered_rules)) +}