resolve snapshot tests

This commit is contained in:
Stefan
2025-09-03 00:12:28 +02:00
parent 042eb232e6
commit f1dff9be48
29 changed files with 425 additions and 383 deletions
+17 -24
View File
@@ -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<OnceLock<tokio::runtime::Runtime>>,
}
impl From<DecisionContent> for Decision {
@@ -28,7 +27,6 @@ impl From<DecisionContent> 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<Arc<DecisionContent>> 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<OnceLock<tokio::runtime::Runtime>>) -> 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<DecisionGraphResponse, Box<EvaluationError>> {
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<DecisionGraphResponse, Box<EvaluationError>> {
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<Value, Value> {
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
+78 -88
View File
@@ -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<DecisionGraphResponse, Box<EvaluationError>> {
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::<InputNodeData, InputNodeTrace>::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::<OutputNodeData, OutputNodeTrace>::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::<FunctionNodeData, FunctionNodeTrace>::from_base(
base_ctx,
content.clone(),
);
FunctionNodeHandler.handle(ctx)
handle_node(base_ctx, content.clone(), FunctionNodeHandler).await
}
DecisionNodeKind::DecisionNode { content } => {
let ctx = NodeContext::<DecisionNodeData, DecisionNodeTrace>::from_base(
base_ctx,
content.clone(),
);
DecisionNodeHandler.handle(ctx)
handle_node(base_ctx, content.clone(), DecisionNodeHandler::default()).await
}
DecisionNodeKind::DecisionTableNode { content } => {
let ctx =
NodeContext::<DecisionTableNodeData, DecisionTableNodeTrace>::from_base(
base_ctx,
content.clone(),
);
DecisionTableNodeHandler.handle(ctx)
handle_node(base_ctx, content.clone(), DecisionTableNodeHandler).await
}
DecisionNodeKind::ExpressionNode { content } => {
let ctx = NodeContext::<ExpressionNodeData, ExpressionNodeTrace>::from_base(
base_ctx,
content.clone(),
);
ExpressionNodeHandler.handle(ctx)
handle_node(base_ctx, content.clone(), ExpressionNodeHandler).await
}
DecisionNodeKind::CustomNode { content } => {
let ctx = NodeContext::<CustomNodeData, CustomNodeTrace>::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<HashMap<String, DecisionGraphTrace>>) -> Option<Variable> {
trace.as_ref().map(|s| {
s.values().for_each(|v| {
v.input.dot_remove("$nodes");
v.output.dot_remove("$nodes");
});
async fn handle_node<NodeData, TraceData, NodeHandlerType>(
base_ctx: NodeContextBase,
content: NodeData,
handler: NodeHandlerType,
) -> NodeResult
where
TraceData: TraceDataType,
NodeData: NodeDataType,
NodeHandlerType: NodeHandler<NodeData = NodeData, TraceData = TraceData>,
{
let ctx = NodeContext::<NodeData, TraceData>::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
}
+16 -18
View File
@@ -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<I>(&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<SwitchStatementTraceRow> = 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())
+28 -15
View File
@@ -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<OnceLock<tokio::runtime::Runtime>>,
}
#[derive(Debug, Default)]
#[derive(Debug)]
pub struct EvaluationOptions {
pub trace: Option<bool>,
pub max_depth: Option<u8>,
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<u8>,
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<F, O>(mut self, loader: F) -> Self
where
F: Fn(String) -> O + Sync + Send + Debug + 'static,
F: Fn(String) -> O + Sync + Send + 'static,
O: Future<Output = LoaderResponse> + 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<K>(
@@ -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
+2 -5
View File
@@ -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))?;
+11 -3
View File
@@ -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<F>
where
F: Sync + Send,
@@ -12,6 +11,15 @@ where
closure: F,
}
impl<T> Debug for ClosureLoader<T>
where
T: Sync + Send,
{
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "ClosureLoader")
}
}
impl<F, O> ClosureLoader<F>
where
F: Fn(String) -> O + Sync + Send,
@@ -24,7 +32,7 @@ where
impl<F, O> DecisionLoader for ClosureLoader<F>
where
F: Fn(String) -> O + Sync + Send + Debug,
F: Fn(String) -> O + Sync + Send,
O: Future<Output = LoaderResponse> + Send,
{
fn load<'a>(&'a self, key: &'a str) -> Pin<Box<dyn Future<Output = LoaderResponse> + 'a>> {
+32 -27
View File
@@ -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<NodeData, TraceData>
where
NodeData: NodeDataType,
@@ -27,6 +29,7 @@ where
pub trace: Option<RefCell<TraceData>>,
pub extensions: NodeHandlerExtensions,
pub iteration: u8,
pub config: NodeContextConfig,
}
impl<NodeData, TraceData> NodeContext<NodeData, TraceData>
@@ -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<Error>(&self, error: Error) -> NodeResult
where
Error: Into<Box<dyn std::error::Error>>,
@@ -78,29 +78,14 @@ where
Error: Into<Box<dyn std::error::Error>>,
{
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<Fut>(&self, future: Fut) -> Result<Fut::Output, NodeError>
where
Fut: Future,
{
let tokio_runtime = self.extensions.tokio_runtime();
Ok(tokio_runtime.block_on(future))
}
pub fn try_block_on<Fut, Output>(&self, future: Fut) -> Result<Output, NodeError>
where
Fut: Future<Output = Result<Output, NodeError>>,
{
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<str>,
pub name: Arc<str>,
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<Box<dyn std::error::Error>>,
{
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<ValidationError<'a>> 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,
}
}
}
+1 -1
View File
@@ -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(),
})
})
+7 -2
View File
@@ -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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> 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
}
}
+42 -14
View File
@@ -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<RefCell<Option<DecisionGraph>>>,
}
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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn after_transform_attributes(
&self,
_ctx: &NodeContext<Self::NodeData, Self::TraceData>,
) -> Result<(), NodeError> {
if let Some(graph) = self.decision_graph.borrow_mut().as_mut() {
graph.reset_graph();
};
Ok(())
}
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> 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();
+18 -13
View File
@@ -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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> 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<RowResult> {
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<Rc<str>, 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::<str>::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;
+16 -29
View File
@@ -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<T> NodeDataType for T where T: Clone + Debug + Serialize + for<'de> Deseria
pub trait TraceDataType: Clone + Debug + Default + ToVariable {}
impl<T> 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<Self::NodeData, Self::TraceData>,
ctx: &NodeContext<Self::NodeData, Self::TraceData>,
) -> Option<TransformAttributes> {
None
}
fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult;
}
#[allow(unused_variables)]
fn after_transform_attributes(
&self,
ctx: &NodeContext<Self::NodeData, Self::TraceData>,
) -> impl std::future::Future<Output = Result<(), NodeError>> {
Box::pin(async { Ok(()) })
}
pub struct NodeHandlers {
pub input: Arc<dyn NodeHandler<NodeData = InputNodeContent, TraceData = InputNodeTrace>>,
pub output: Arc<dyn NodeHandler<NodeData = OutputNodeContent, TraceData = OutputNodeTrace>>,
pub function:
Arc<dyn NodeHandler<NodeData = FunctionNodeContent, TraceData = FunctionNodeTrace>>,
pub expression:
Box<dyn NodeHandler<NodeData = ExpressionNodeContent, TraceData = ExpressionNodeContent>>,
}
pub enum NodeHandlerKind {
Input(),
Output(Box<dyn NodeHandler<NodeData = OutputNodeContent, TraceData = ()>>),
Function(Box<dyn NodeHandler<NodeData = FunctionNodeContent, TraceData = ()>>),
Expression(Box<dyn NodeHandler<NodeData = ExpressionNodeContent, TraceData = ()>>),
DecisionTable(Box<dyn NodeHandler<NodeData = DecisionNodeContent, TraceData = ()>>),
Decision(Box<dyn NodeHandler<NodeData = DecisionNodeContent, TraceData = ()>>),
Custom(Box<dyn NodeHandler<NodeData = CustomNodeContent, TraceData = ()>>),
fn handle(
&self,
ctx: NodeContext<Self::NodeData, Self::TraceData>,
) -> impl std::future::Future<Output = NodeResult>;
}
+3 -3
View File
@@ -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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> 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;
}
+14 -27
View File
@@ -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<OnceCell<Function>>,
pub(crate) tokio_runtime: Arc<OnceLock<tokio::runtime::Runtime>>,
pub(crate) function_runtime: Arc<tokio::sync::OnceCell<Function>>,
pub(crate) validator_cache: Arc<OnceCell<ValidatorCache>>,
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 {
+8 -5
View File
@@ -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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult {
match &ctx.node {
FunctionNodeContent::Version1(source) => {
let v1_context = NodeContext::<Arc<str>, 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::<FunctionContent, FunctionV2Trace> {
@@ -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
}
}
}
+3 -2
View File
@@ -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<str>;
type TraceData = FunctionV1Trace;
fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> 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);
@@ -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<Self> {
let module_loader = ModuleLoader::new();
+35 -45
View File
@@ -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<Function>,
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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> 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)
}
}
}
}
@@ -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<Box<dyn Future<Output = FunctionResult> + '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<Object<'js>>| {
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));
+2 -1
View File
@@ -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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult {
if let Some(json_schema) = &ctx.node.schema {
let input_json = ctx.input.to_value();
ctx.validate(json_schema, &input_json)?;
+3 -2
View File
@@ -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};
+2 -1
View File
@@ -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<Self::NodeData, Self::TraceData>) -> NodeResult {
async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult {
if let Some(json_schema) = &ctx.node.schema {
let input_json = ctx.input.to_value();
ctx.validate(json_schema, &input_json)?;
+1 -2
View File
@@ -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<NodeResponse, NodeError>;
#[derive(Debug, Error)]
pub struct NodeError {
pub node_id: Option<Arc<str>>,
pub node_id: Arc<str>,
pub trace: Option<Variable>,
pub source: Box<dyn std::error::Error>,
}
@@ -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<F, Fut>(&self, ctx: NodeContextBase, evaluate: F) -> NodeResult
where
F: Fn(Variable) -> Fut,
F: Fn(Variable, bool) -> Fut,
Fut: Future<Output = NodeResult>;
}
impl TransformAttributesExecution for TransformAttributes {
async fn run_with<F, Fut>(&self, ctx: NodeContextBase, evaluate: F) -> NodeResult
where
F: Fn(Variable) -> Fut,
F: Fn(Variable, bool) -> Fut,
Fut: Future<Output = NodeResult>,
{
let input = match &self.input_field {
@@ -54,7 +54,7 @@ impl TransformAttributesExecution for TransformAttributes {
let mut trace_data: Option<Variable> = 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);
}
+5 -4
View File
@@ -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"),
+20 -19
View File
@@ -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",