mirror of
https://github.com/gorules/zen.git
synced 2026-10-05 08:02:28 +00:00
feat: passthrough nodes
This commit is contained in:
@@ -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"] }
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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}""#))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 })
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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();
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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))),
|
||||
|
||||
Reference in New Issue
Block a user