implement remaining nodes

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