mirror of
https://github.com/gorules/zen.git
synced 2026-10-04 16:02:18 +00:00
implement remaining nodes
This commit is contained in:
+1
-1
@@ -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
@@ -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()
|
||||
|
||||
@@ -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,2 +1,2 @@
|
||||
pub mod graph;
|
||||
mod traversal;
|
||||
mod walker;
|
||||
|
||||
+63
-73
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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>>;
|
||||
}
|
||||
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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>;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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));
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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>,
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
|
||||
|
||||
@@ -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)?,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user