feat: passthrough nodes

This commit is contained in:
Stefan
2024-10-18 20:09:58 +02:00
parent 8a7e72bf53
commit 302e82afbd
20 changed files with 325 additions and 189 deletions
+1 -1
View File
@@ -17,7 +17,7 @@ thiserror = { workspace = true }
bincode = { workspace = true, optional = true }
petgraph = { workspace = true }
serde_json = { workspace = true, features = ["arbitrary_precision"] }
serde = { workspace = true, features = ["derive"] }
serde = { workspace = true, features = ["derive", "rc"] }
once_cell = { workspace = true }
json_dotpath = { workspace = true }
rust_decimal = { workspace = true, features = ["maths-nopanic"] }
+2 -2
View File
@@ -82,12 +82,12 @@ where
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(),
loader: Arc::new(CachedLoader::from(self.loader.clone())),
adapter: self.adapter.clone(),
iteration: 0,
content: &self.content,
})?;
Ok(decision_graph.evaluate(context).await?)
@@ -95,12 +95,12 @@ where
pub fn validate(&self) -> Result<(), DecisionGraphValidationError> {
let decision_graph = DecisionGraph::try_new(DecisionGraphConfig {
content: self.content.clone(),
max_depth: 1,
trace: false,
loader: Arc::new(CachedLoader::from(self.loader.clone())),
adapter: self.adapter.clone(),
iteration: 0,
content: &self.content,
})?;
decision_graph.validate()
+21 -22
View File
@@ -4,44 +4,43 @@ use anyhow::anyhow;
use json_dotpath::DotPaths;
use serde::Serialize;
use serde_json::Value;
use std::ops::Deref;
use std::sync::Arc;
use zen_expression::variable::Variable;
use zen_tmpl::TemplateRenderError;
pub trait CustomNodeAdapter {
fn handle(
&self,
request: CustomNodeRequest<'_>,
) -> impl std::future::Future<Output = NodeResult>;
fn handle(&self, request: CustomNodeRequest) -> impl std::future::Future<Output = NodeResult>;
}
#[derive(Default, Debug)]
pub struct NoopCustomNode;
impl CustomNodeAdapter for NoopCustomNode {
async fn handle(&self, _: CustomNodeRequest<'_>) -> NodeResult {
async fn handle(&self, _: CustomNodeRequest) -> NodeResult {
Err(anyhow!("Custom node handler not provided"))
}
}
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CustomNodeRequest<'a> {
pub struct CustomNodeRequest {
pub input: Variable,
pub node: CustomDecisionNode<'a>,
pub node: CustomDecisionNode,
}
impl<'a> TryFrom<&'a NodeRequest<'a>> for CustomNodeRequest<'a> {
impl TryFrom<NodeRequest> for CustomNodeRequest {
type Error = ();
fn try_from(value: &'a NodeRequest<'a>) -> Result<Self, Self::Error> {
fn try_from(value: NodeRequest) -> Result<Self, Self::Error> {
Ok(Self {
input: value.input.clone(),
node: value.node.try_into()?,
node: value.node.deref().try_into()?,
})
}
}
impl<'a> CustomNodeRequest<'a> {
impl CustomNodeRequest {
pub fn get_field(&self, path: &str) -> Result<Option<Variable>, TemplateRenderError> {
let Some(selected_value) = self.get_field_raw(path) else {
return Ok(None);
@@ -62,26 +61,26 @@ impl<'a> CustomNodeRequest<'a> {
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
pub struct CustomDecisionNode<'a> {
pub id: &'a str,
pub name: &'a str,
pub kind: &'a str,
pub config: &'a Value,
pub struct CustomDecisionNode {
pub id: String,
pub name: String,
pub kind: String,
pub config: Arc<Value>,
}
impl<'a> TryFrom<&'a DecisionNode> for CustomDecisionNode<'a> {
impl TryFrom<&DecisionNode> for CustomDecisionNode {
type Error = ();
fn try_from(value: &'a DecisionNode) -> Result<Self, Self::Error> {
fn try_from(value: &DecisionNode) -> Result<Self, Self::Error> {
let DecisionNodeKind::CustomNode { content } = &value.kind else {
return Err(());
};
Ok(Self {
id: &value.id,
name: &value.name,
kind: &content.kind,
config: &content.config,
id: value.id.clone(),
name: value.name.clone(),
kind: content.kind.clone(),
config: content.config.clone(),
})
}
}
+23 -68
View File
@@ -3,15 +3,14 @@ use crate::handler::function::function::Function;
use crate::handler::graph::{DecisionGraph, DecisionGraphConfig};
use crate::handler::node::{NodeRequest, NodeResponse, NodeResult};
use crate::loader::DecisionLoader;
use crate::model::{DecisionNodeKind, TransformExecutionMode};
use anyhow::{anyhow, Context};
use serde_json::Value;
use crate::model::DecisionNodeKind;
use anyhow::anyhow;
use std::future::Future;
use std::ops::Deref;
use std::pin::Pin;
use std::rc::Rc;
use std::sync::Arc;
use zen_expression::{Isolate, Variable};
use tokio::sync::Mutex;
use zen_expression::Isolate;
pub struct DecisionHandler<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> {
trace: bool,
@@ -40,7 +39,7 @@ impl<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionHandle
pub fn handle<'s, 'arg, 'recursion>(
&'s self,
request: &'arg NodeRequest<'_>,
request: NodeRequest,
) -> Pin<Box<dyn Future<Output = NodeResult> + 'recursion>>
where
's: 'recursion,
@@ -52,11 +51,9 @@ impl<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionHandle
_ => Err(anyhow!("Unexpected node type")),
}?;
let mut isolate = Isolate::new();
let sub_decision = self.loader.load(&content.key).await?;
let mut sub_tree = DecisionGraph::try_new(DecisionGraphConfig {
content: sub_decision.deref(),
let sub_tree = DecisionGraph::try_new(DecisionGraphConfig {
content: sub_decision,
max_depth: self.max_depth,
loader: self.loader.clone(),
adapter: self.adapter.clone(),
@@ -65,69 +62,27 @@ impl<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionHandle
})?
.with_function(self.js_function.clone());
let input_data = match &content.transform_attributes.input_field {
None => request.input.clone(),
Some(input_field) => {
isolate.set_environment(request.input.clone());
isolate.run_standard(input_field.as_str())?
}
};
let sub_tree_mutex = Arc::new(Mutex::new(sub_tree));
let mut trace_data: Option<Value> = None;
let mut output_data = match &content.transform_attributes.execution_mode {
TransformExecutionMode::Single => {
let response = sub_tree
.evaluate(request.input.clone())
.await
.map_err(|e| e.source)?;
content
.transform_attributes
.run_with(request.input, |input| {
let sub_tree_mutex = sub_tree_mutex.clone();
if self.trace {
trace_data.replace(
serde_json::to_value(response.trace)
.context("Failed to serialize trace")?,
);
}
async move {
let mut sub_tree_ref = sub_tree_mutex.lock().await;
response.result
}
TransformExecutionMode::Loop => {
let input_array_ref = input_data.as_array().context("Expected an array")?;
let input_array = input_array_ref.borrow();
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 response = sub_tree
.evaluate(input.clone())
sub_tree_ref
.evaluate(input)
.await
.map_err(|e| e.source)?;
output_array.push(response.result);
trace_datum.push(response.trace);
.map(|r| NodeResponse {
output: r.result,
trace_data: serde_json::to_value(r.trace).ok(),
})
.map_err(|e| e.source)
}
if self.trace {
trace_data.replace(
serde_json::to_value(trace_datum)
.context("Failed to parse trace data")?,
);
}
Variable::from_array(output_array)
}
};
if let Some(output_path) = &content.transform_attributes.output_path {
let new_output_data = Variable::empty_object();
new_output_data.dot_insert(output_path.as_str(), output_data);
output_data = new_output_data;
}
Ok(NodeResponse {
output: output_data,
trace_data,
})
})
.await
})
}
}
+48 -17
View File
@@ -1,15 +1,16 @@
use crate::handler::node::{NodeRequest, NodeResponse, NodeResult};
use crate::model::DecisionNodeKind;
use crate::model::{DecisionNodeKind, ExpressionNodeContent};
use ahash::{HashMap, HashMapExt};
use std::sync::Arc;
use anyhow::{anyhow, Context};
use serde::Serialize;
use tokio::sync::Mutex;
use zen_expression::variable::Variable;
use zen_expression::Isolate;
pub struct ExpressionHandler<'a> {
pub struct ExpressionHandler {
trace: bool,
isolate: Isolate<'a>,
}
#[derive(Debug, Serialize)]
@@ -17,26 +18,62 @@ struct ExpressionTrace {
result: String,
}
impl<'a> ExpressionHandler<'a> {
impl ExpressionHandler {
pub fn new(trace: bool) -> Self {
Self {
trace,
isolate: Isolate::new(),
}
Self { trace }
}
pub async fn handle(&mut self, request: &'a NodeRequest<'_>) -> NodeResult {
pub async fn handle(&mut self, request: NodeRequest) -> NodeResult {
let content = match &request.node.kind {
DecisionNodeKind::ExpressionNode { content } => Ok(content),
_ => Err(anyhow!("Unexpected node type")),
}?;
let inner_handler_mutex = Arc::new(Mutex::new(ExpressionHandlerInner::new(self.trace)));
content
.transform_attributes
.run_with(request.input, |input| {
let inner_handler_mutex = inner_handler_mutex.clone();
async move {
let mut inner_handler_ref = inner_handler_mutex.lock().await;
inner_handler_ref.handle(input, content).await
}
})
.await
}
}
struct ExpressionHandlerInner<'a> {
isolate: Isolate<'a>,
trace: bool,
}
impl<'a> ExpressionHandlerInner<'a> {
pub fn new(trace: bool) -> Self {
Self {
isolate: Isolate::new(),
trace,
}
}
async fn handle(&mut self, input: Variable, content: &'a ExpressionNodeContent) -> NodeResult {
let result = Variable::empty_object();
let mut trace_map = self.trace.then(|| HashMap::<&str, ExpressionTrace>::new());
self.isolate.set_environment(request.input.depth_clone(1));
self.isolate.set_environment(input.depth_clone(1));
for expression in &content.expressions {
let value = self.evaluate_expression(&expression.value)?;
if expression.key.is_empty() || expression.value.is_empty() {
continue;
}
let value = self
.isolate
.run_standard(&expression.value)
.with_context(|| {
format!(r#"Failed to evaluate expression: "{}""#, &expression.value)
})?;
if let Some(tmap) = &mut trace_map {
tmap.insert(
&expression.key,
@@ -66,10 +103,4 @@ impl<'a> ExpressionHandler<'a> {
.context("Failed to serialize trace data")?,
})
}
fn evaluate_expression(&mut self, expression: &'a str) -> anyhow::Result<Variable> {
self.isolate
.run_standard(expression)
.with_context(|| format!(r#"Failed to evaluate expression: "{expression}""#))
}
}
+1 -1
View File
@@ -43,7 +43,7 @@ impl FunctionHandler {
}
}
pub async fn handle(&self, request: &NodeRequest<'_>) -> NodeResult {
pub async fn handle(&self, request: NodeRequest) -> NodeResult {
let content = match &request.node.kind {
DecisionNodeKind::FunctionNode { content } => match content {
FunctionNodeContent::Version2(content) => Ok(content),
@@ -58,7 +58,7 @@ impl<Loader: DecisionLoader + 'static, Adapter: CustomNodeAdapter + 'static> Run
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,
content: decision_content,
max_depth,
loader,
adapter,
+1 -1
View File
@@ -22,7 +22,7 @@ impl FunctionHandler {
Self { trace, runtime }
}
pub async fn handle(&self, request: &NodeRequest<'_>) -> NodeResult {
pub async fn handle(&self, request: NodeRequest) -> NodeResult {
let content = match &request.node.kind {
DecisionNodeKind::FunctionNode { content } => match content {
FunctionNodeContent::Version1(content) => Ok(content),
+38 -34
View File
@@ -25,8 +25,8 @@ use std::time::Instant;
use thiserror::Error;
use zen_expression::variable::Variable;
pub struct DecisionGraph<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> {
graph: StableDiDecisionGraph<'a>,
pub struct DecisionGraph<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> {
graph: StableDiDecisionGraph,
adapter: Arc<A>,
loader: Arc<L>,
trace: bool,
@@ -35,18 +35,18 @@ pub struct DecisionGraph<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter +
runtime: Option<Rc<Function>>,
}
pub struct DecisionGraphConfig<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> {
pub struct DecisionGraphConfig<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> {
pub loader: Arc<L>,
pub adapter: Arc<A>,
pub content: &'a DecisionContent,
pub content: Arc<DecisionContent>,
pub trace: bool,
pub iteration: u8,
pub max_depth: u8,
}
impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGraph<'a, L, A> {
impl<L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGraph<L, A> {
pub fn try_new(
config: DecisionGraphConfig<'a, L, A>,
config: DecisionGraphConfig<L, A>,
) -> Result<Self, DecisionGraphValidationError> {
let content = config.content;
let mut graph = StableDiDecisionGraph::new();
@@ -54,7 +54,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
for node in &content.nodes {
let node_id = node.id.clone();
let node_index = graph.add_node(node);
let node_index = graph.add_node(node.clone());
index_map.insert(node_id, node_index);
}
@@ -68,7 +68,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
DecisionGraphValidationError::MissingNode(edge.target_id.to_string())
})?;
graph.add_edge(source_index.clone(), target_index.clone(), edge);
graph.add_edge(source_index.clone(), target_index.clone(), edge.clone());
}
Ok(Self {
@@ -117,13 +117,6 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
));
}
let output_count = self.node_kind_count(DecisionNodeKind::OutputNode);
if output_count < 1 {
return Err(DecisionGraphValidationError::InvalidOutputCount(
output_count as u32,
));
}
if is_cyclic_directed(&self.graph) {
return Err(DecisionGraphValidationError::CyclicGraph);
}
@@ -171,7 +164,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
continue;
}
let node = self.graph[nid];
let node = (&self.graph[nid]).clone();
let start = Instant::now();
macro_rules! trace {
@@ -222,7 +215,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
})?;
let node_request = NodeRequest {
node,
node: node.clone(),
iteration: self.iteration,
input: walker.incoming_node_data(&self.graph, nid, true),
};
@@ -233,7 +226,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
self.iteration,
self.max_depth,
)
.handle(&node_request)
.handle(node_request.clone())
.await
.map_err(|e| NodeError {
source: e.into(),
@@ -246,7 +239,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
})?;
function_v1::FunctionHandler::new(self.trace, runtime)
.handle(&node_request)
.handle(node_request.clone())
.await
.map_err(|e| NodeError {
source: e.into(),
@@ -270,7 +263,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
}
DecisionNodeKind::DecisionNode { .. } => {
let node_request = NodeRequest {
node,
node: node.clone(),
iteration: self.iteration,
input: walker.incoming_node_data(&self.graph, nid, true),
};
@@ -282,7 +275,7 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
self.adapter.clone(),
self.runtime.clone(),
)
.handle(&node_request)
.handle(node_request.clone())
.await
.map_err(|e| NodeError {
source: e.into(),
@@ -302,15 +295,15 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
});
walker.set_node_data(nid, res.output);
}
DecisionNodeKind::DecisionTableNode { .. } => {
DecisionNodeKind::DecisionTableNode { content } => {
let node_request = NodeRequest {
node,
node: node.clone(),
iteration: self.iteration,
input: walker.incoming_node_data(&self.graph, nid, true),
};
let res = DecisionTableHandler::new(self.trace)
.handle(&node_request)
let mut res = DecisionTableHandler::new(self.trace)
.handle(node_request.clone())
.await
.map_err(|e| NodeError {
node_id: node.id.clone(),
@@ -320,6 +313,11 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
node_request.input.dot_remove("$nodes");
res.output.dot_remove("$nodes");
if content.pass_through {
let mut base = node_request.input.clone();
res.output = base.merge(&res.output)
}
trace!({
input: node_request.input,
output: res.output.clone(),
@@ -330,15 +328,15 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
});
walker.set_node_data(nid, res.output);
}
DecisionNodeKind::ExpressionNode { .. } => {
DecisionNodeKind::ExpressionNode { content } => {
let node_request = NodeRequest {
node,
node: node.clone(),
iteration: self.iteration,
input: walker.incoming_node_data(&self.graph, nid, true),
};
let res = ExpressionHandler::new(self.trace)
.handle(&node_request)
let mut res = ExpressionHandler::new(self.trace)
.handle(node_request.clone())
.await
.map_err(|e| NodeError {
node_id: node.id.clone(),
@@ -348,6 +346,11 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
node_request.input.dot_remove("$nodes");
res.output.dot_remove("$nodes");
if content.pass_through {
let mut base = node_request.input.clone();
res.output = base.merge(&res.output)
}
trace!({
input: node_request.input,
output: res.output.clone(),
@@ -360,14 +363,14 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
}
DecisionNodeKind::CustomNode { .. } => {
let node_request = NodeRequest {
node,
node: node.clone(),
iteration: self.iteration,
input: walker.incoming_node_data(&self.graph, nid, true),
};
let res = self
.adapter
.handle(CustomNodeRequest::try_from(&node_request).unwrap())
.handle(CustomNodeRequest::try_from(node_request.clone()).unwrap())
.await
.map_err(|e| NodeError {
node_id: node.id.clone(),
@@ -390,9 +393,10 @@ impl<'a, L: DecisionLoader + 'static, A: CustomNodeAdapter + 'static> DecisionGr
}
}
Err(NodeError {
node_id: "".to_string(),
source: anyhow!("Graph did not halt. Missing output node."),
Ok(DecisionGraphResponse {
result: walker.ending_variables(&self.graph),
performance: format!("{:?}", root_start.elapsed()),
trace: node_traces,
})
}
}
+4 -3
View File
@@ -2,6 +2,7 @@ use crate::model::DecisionNode;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::fmt::{Display, Formatter};
use std::sync::Arc;
use thiserror::Error;
use zen_expression::variable::Variable;
@@ -12,11 +13,11 @@ pub struct NodeResponse {
pub trace_data: Option<Value>,
}
#[derive(Debug, Serialize)]
pub struct NodeRequest<'a> {
#[derive(Debug, Serialize, Clone)]
pub struct NodeRequest {
pub input: Variable,
pub iteration: u8,
pub node: &'a DecisionNode,
pub node: Arc<DecisionNode>,
}
#[derive(Error, Debug)]
+40 -12
View File
@@ -1,10 +1,12 @@
use ahash::HashMap;
use anyhow::{anyhow, Context};
use std::sync::Arc;
use crate::handler::node::{NodeRequest, NodeResponse, NodeResult};
use crate::handler::table::{RowOutput, RowOutputKind};
use crate::model::{DecisionNodeKind, DecisionTableContent, DecisionTableHitPolicy};
use serde::Serialize;
use tokio::sync::Mutex;
use zen_expression::variable::Variable;
use zen_expression::Isolate;
@@ -18,12 +20,32 @@ struct RowResult {
}
#[derive(Debug)]
pub struct DecisionTableHandler<'a> {
pub struct DecisionTableHandler {
trace: bool,
}
impl DecisionTableHandler {
pub fn new(trace: bool) -> Self {
Self { trace }
}
pub async fn handle(&mut self, request: NodeRequest) -> NodeResult {
let content = match &request.node.kind {
DecisionNodeKind::DecisionTableNode { content } => Ok(content),
_ => Err(anyhow!("Unexpected node type")),
}?;
let inner_handler = DecisionTableHandlerInner::new(self.trace);
inner_handler.handle(request.input, content).await
}
}
struct DecisionTableHandlerInner<'a> {
isolate: Isolate<'a>,
trace: bool,
}
impl<'a> DecisionTableHandler<'a> {
impl<'a> DecisionTableHandlerInner<'a> {
pub fn new(trace: bool) -> Self {
Self {
isolate: Isolate::new(),
@@ -31,18 +53,24 @@ impl<'a> DecisionTableHandler<'a> {
}
}
pub async fn handle(&mut self, request: &'a NodeRequest<'_>) -> NodeResult {
let content = match &request.node.kind {
DecisionNodeKind::DecisionTableNode { content } => Ok(content),
_ => Err(anyhow!("Unexpected node type")),
}?;
pub async fn handle(self, input: Variable, content: &'a DecisionTableContent) -> NodeResult {
let self_mutex = Arc::new(Mutex::new(self));
self.isolate.set_environment(request.input.depth_clone(1));
content
.transform_attributes
.run_with(input, |input| {
let self_mutex = self_mutex.clone();
async move {
let mut self_ref = self_mutex.lock().await;
match &content.hit_policy {
DecisionTableHitPolicy::First => self.handle_first_hit(&content).await,
DecisionTableHitPolicy::Collect => self.handle_collect(&content).await,
}
self_ref.isolate.set_environment(input);
match &content.hit_policy {
DecisionTableHitPolicy::First => self_ref.handle_first_hit(&content).await,
DecisionTableHitPolicy::Collect => self_ref.handle_collect(&content).await,
}
}
})
.await
}
async fn handle_first_hit(&mut self, content: &'a DecisionTableContent) -> NodeResult {
+17 -2
View File
@@ -8,6 +8,7 @@ use petgraph::{Incoming, Outgoing};
use serde_json::json;
use std::rc::Rc;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::Instant;
use crate::config::ZEN_CONFIG;
@@ -18,7 +19,7 @@ use crate::DecisionGraphTrace;
use zen_expression::variable::Variable;
use zen_expression::Isolate;
pub(crate) type StableDiDecisionGraph<'a> = StableDiGraph<&'a DecisionNode, &'a DecisionEdge>;
pub(crate) type StableDiDecisionGraph = StableDiGraph<Arc<DecisionNode>, Arc<DecisionEdge>>;
pub(crate) struct GraphWalker {
ordered: FixedBitSet,
@@ -72,6 +73,20 @@ impl GraphWalker {
self.node_data.get(&node_id).cloned()
}
pub fn ending_variables(&self, g: &StableDiDecisionGraph) -> Variable {
g.node_indices()
.filter(|nid| {
self.ordered.is_visited(nid)
&& g.neighbors_directed(*nid, Outgoing).count().is_zero()
})
.fold(Variable::empty_object(), |mut acc, curr| {
match self.node_data.get(&curr) {
None => acc,
Some(data) => acc.merge(data),
}
})
}
pub fn get_all_node_data(&self, g: &StableDiDecisionGraph) -> Variable {
let node_values = self
.node_data
@@ -130,7 +145,7 @@ 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)?;
let decision_node = g.node_weight(nid)?.clone();
if self.ordered.is_visited(&nid) {
continue;
}
+1
View File
@@ -130,6 +130,7 @@ pub mod handler;
pub mod loader;
#[path = "model/mod.rs"]
pub mod model;
mod util;
pub use config::ZEN_CONFIG;
pub use decision::Decision;
+34 -4
View File
@@ -1,14 +1,15 @@
use ahash::HashMap;
use serde::{Deserialize, Serialize};
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::Value;
use std::sync::Arc;
/// JDM Decision model
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))]
#[serde(rename_all = "camelCase")]
pub struct DecisionContent {
pub nodes: Vec<DecisionNode>,
pub edges: Vec<DecisionEdge>,
pub nodes: Vec<Arc<DecisionNode>>,
pub edges: Vec<Arc<DecisionEdge>>,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
@@ -86,6 +87,10 @@ pub struct DecisionTableContent {
pub inputs: Vec<DecisionTableInputField>,
pub outputs: Vec<DecisionTableOutputField>,
pub hit_policy: DecisionTableHitPolicy,
#[serde(default)]
pub pass_through: bool,
#[serde(flatten)]
pub transform_attributes: TransformAttributes,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
@@ -102,6 +107,7 @@ pub enum DecisionTableHitPolicy {
pub struct DecisionTableInputField {
pub id: String,
pub name: String,
#[serde(default, deserialize_with = "empty_string_is_none")]
pub field: Option<String>,
}
@@ -119,6 +125,10 @@ pub struct DecisionTableOutputField {
#[serde(rename_all = "camelCase")]
pub struct ExpressionNodeContent {
pub expressions: Vec<Expression>,
#[serde(default)]
pub pass_through: bool,
#[serde(flatten)]
pub transform_attributes: TransformAttributes,
}
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
@@ -160,7 +170,9 @@ pub enum SwitchStatementHitPolicy {
#[cfg_attr(feature = "bincode", derive(bincode::Encode, bincode::Decode))]
#[serde(rename_all = "camelCase")]
pub struct TransformAttributes {
#[serde(default, deserialize_with = "empty_string_is_none")]
pub input_field: Option<String>,
#[serde(default, deserialize_with = "empty_string_is_none")]
pub output_path: Option<String>,
#[serde(default)]
pub execution_mode: TransformExecutionMode,
@@ -179,7 +191,7 @@ pub enum TransformExecutionMode {
#[serde(rename_all = "camelCase")]
pub struct CustomNodeContent {
pub kind: String,
pub config: Value,
pub config: Arc<Value>,
}
#[cfg(feature = "bincode")]
@@ -225,3 +237,21 @@ impl<'__de> ::bincode::BorrowDecode<'__de> for CustomNodeContent {
Ok(Self { kind, config })
}
}
fn empty_string_is_none<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
where
D: Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(untagged)]
enum StringOrNull {
String(String),
Null,
}
match StringOrNull::deserialize(deserializer)? {
StringOrNull::String(s) if s.trim().is_empty() => Ok(None),
StringOrNull::String(s) => Ok(Some(s)),
StringOrNull::Null => Ok(None),
}
}
+1
View File
@@ -0,0 +1 @@
mod transform_attribute;
@@ -0,0 +1,62 @@
use crate::handler::node::{NodeResponse, NodeResult};
use crate::model::{TransformAttributes, TransformExecutionMode};
use anyhow::Context;
use serde_json::Value;
use std::future::Future;
use zen_expression::{Isolate, Variable};
impl TransformAttributes {
pub(crate) async fn run_with<F, Fut>(&self, input: Variable, evaluate: F) -> NodeResult
where
F: Fn(Variable) -> Fut,
Fut: Future<Output = NodeResult>,
{
let input = match &self.input_field {
None => input.clone(),
Some(input_field) => {
let mut isolate = Isolate::new();
isolate.set_environment(input.clone());
isolate.run_standard(input_field.as_str())?
}
};
let mut trace_data: Option<Value> = None;
let mut output = match self.execution_mode {
TransformExecutionMode::Single => {
let response = evaluate(input).await?;
if let Some(td) = response.trace_data {
trace_data.replace(td);
}
response.output
}
TransformExecutionMode::Loop => {
let input_array_ref = input.as_array().context("Expected an array")?;
let input_array = input_array_ref.borrow();
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 response = evaluate(input.clone()).await?;
output_array.push(response.output);
if let Some(td) = response.trace_data {
trace_datum.push(td);
}
}
trace_data.replace(Value::Array(trace_datum));
Variable::from_array(output_array)
}
};
if let Some(output_path) = &self.output_path {
let new_output = Variable::empty_object();
new_output.dot_insert(output_path.as_str(), output);
output = new_output;
}
Ok(NodeResponse { output, trace_data })
}
}
-7
View File
@@ -87,11 +87,4 @@ fn decision_validation() {
missing_input_error,
DecisionGraphValidationError::InvalidInputCount(_)
));
let missing_output_decision = Decision::from(load_test_data("error-missing-output.json"));
let missing_output_error = missing_output_decision.validate().unwrap_err();
assert!(matches!(
missing_output_error,
DecisionGraphValidationError::InvalidOutputCount(_)
));
}
+24 -12
View File
@@ -8,7 +8,7 @@ use std::path::Path;
use std::sync::Arc;
use tokio::runtime::Builder;
use zen_engine::loader::{LoaderError, MemoryLoader};
use zen_engine::model::{DecisionContent, DecisionNodeKind, FunctionNodeContent};
use zen_engine::model::{DecisionContent, DecisionNode, DecisionNodeKind, FunctionNodeContent};
use zen_engine::Variable;
use zen_engine::{DecisionEngine, EvaluationError, EvaluationOptions};
@@ -158,24 +158,36 @@ fn engine_with_trace() {
#[tokio::test]
#[cfg_attr(miri, ignore)]
async fn engine_function_imports() {
let mut function_content = load_test_data("function.json");
let function_content = load_test_data("function.json");
let imports_js_path = Path::new("js").join("imports.js");
let mut replace_buffer = load_raw_test_data(imports_js_path.to_str().unwrap());
let mut replace_data = String::new();
replace_buffer.read_to_string(&mut replace_data).unwrap();
function_content.nodes.iter_mut().for_each(|node| {
if let DecisionNodeKind::FunctionNode { content, .. } = &mut node.kind {
match content {
FunctionNodeContent::Version1(content) => {
let _ = std::mem::replace(content, replace_data.clone());
}
_ => {}
}
}
});
let new_nodes = function_content
.nodes
.into_iter()
.map(|node| match &node.kind {
DecisionNodeKind::FunctionNode { .. } => {
let new_kind = DecisionNodeKind::FunctionNode {
content: FunctionNodeContent::Version1(replace_data.clone()),
};
Arc::new(DecisionNode {
id: node.id.clone(),
name: node.name.clone(),
kind: new_kind,
})
}
_ => node,
})
.collect::<Vec<_>>();
let function_content = DecisionContent {
edges: function_content.edges,
nodes: new_nodes,
};
let decision = DecisionEngine::default().create_decision(function_content.into());
let response = decision.evaluate(json!({}).into()).await.unwrap();
+5 -1
View File
@@ -128,7 +128,11 @@ impl<'arena, 'token_ref, Flavor> Parser<'arena, 'token_ref, Flavor> {
}
pub(crate) fn prev_token_end(&self) -> u32 {
match self.tokens.get(self.position() - 1) {
let Some(pos) = self.position().checked_sub(1) else {
return self.token_start();
};
match self.tokens.get(pos) {
None => self.token_start(),
Some(t) => t.span.1,
}
+1 -1
View File
@@ -62,7 +62,7 @@ impl VariableType {
(VariableType::Bool, VariableType::Bool) => true,
(VariableType::String, VariableType::String) => true,
(VariableType::Number, VariableType::Number) => true,
(VariableType::Array(a1), VariableType::Array(a2)) => a1 == a2,
(VariableType::Array(a1), VariableType::Array(a2)) => a1.satisfies(a2),
(VariableType::Object(o1), VariableType::Object(o2)) => o1
.iter()
.all(|(k, v)| o2.get(k).is_some_and(|tv| v.satisfies(tv))),