mirror of
https://github.com/gorules/zen.git
synced 2026-10-06 16:02:19 +00:00
feat: custom node (#138)
* feat: custom node * add zen template, expose $nodes and $root * add custom handler to go and nodejs * add support for python, improve bindings, add error to zen template * fix: correct binding exports for nodejs and python * fix benchmark * improve rust api, trim template in zen templates * update cargo action * update expression version * compile action * fix action format
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
use crate::handler::node::{NodeRequest, NodeResult};
|
||||
use crate::model::{DecisionNode, DecisionNodeKind};
|
||||
use anyhow::anyhow;
|
||||
use json_dotpath::DotPaths;
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
use zen_template::TemplateRenderError;
|
||||
|
||||
pub trait CustomNodeAdapter {
|
||||
fn handle(
|
||||
&self,
|
||||
request: CustomNodeRequest<'_>,
|
||||
) -> impl std::future::Future<Output = NodeResult> + Send;
|
||||
}
|
||||
|
||||
#[derive(Default, Debug)]
|
||||
pub struct NoopCustomNode;
|
||||
|
||||
impl CustomNodeAdapter for NoopCustomNode {
|
||||
async fn handle(&self, _: CustomNodeRequest<'_>) -> NodeResult {
|
||||
Err(anyhow!("Custom node handler not provided"))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct CustomNodeRequest<'a> {
|
||||
pub input: &'a Value,
|
||||
pub node: CustomDecisionNode<'a>,
|
||||
}
|
||||
|
||||
impl<'a> TryFrom<&'a NodeRequest<'a>> for CustomNodeRequest<'a> {
|
||||
type Error = ();
|
||||
|
||||
fn try_from(value: &'a NodeRequest<'a>) -> Result<Self, Self::Error> {
|
||||
Ok(Self {
|
||||
input: &value.input,
|
||||
node: value.node.try_into()?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> CustomNodeRequest<'a> {
|
||||
pub fn get_field(&self, path: &str) -> Result<Option<Value>, TemplateRenderError> {
|
||||
let Some(selected_value) = self.get_field_raw(path) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Value::String(template) = selected_value else {
|
||||
return Ok(Some(selected_value));
|
||||
};
|
||||
|
||||
let template_value = zen_template::render(template.as_str(), &self.input)?;
|
||||
Ok(Some(template_value))
|
||||
}
|
||||
|
||||
fn get_field_raw(&self, path: &str) -> Option<Value> {
|
||||
self.node.config.dot_get(path).ok().flatten()
|
||||
}
|
||||
}
|
||||
|
||||
#[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,
|
||||
}
|
||||
|
||||
impl<'a> TryFrom<&'a DecisionNode> for CustomDecisionNode<'a> {
|
||||
type Error = ();
|
||||
|
||||
fn try_from(value: &'a 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,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::handler::custom_node_adapter::CustomNodeAdapter;
|
||||
use crate::handler::graph::{DecisionGraph, DecisionGraphConfig};
|
||||
use crate::handler::node::{NodeRequest, NodeResponse, NodeResult};
|
||||
use crate::loader::DecisionLoader;
|
||||
@@ -8,18 +9,26 @@ use rquickjs::Runtime;
|
||||
use std::ops::Deref;
|
||||
use std::sync::Arc;
|
||||
|
||||
pub struct DecisionHandler<T: DecisionLoader> {
|
||||
pub struct DecisionHandler<L: DecisionLoader, A: CustomNodeAdapter> {
|
||||
trace: bool,
|
||||
loader: Arc<T>,
|
||||
loader: Arc<L>,
|
||||
adapter: Arc<A>,
|
||||
max_depth: u8,
|
||||
js_runtime: Option<Runtime>,
|
||||
}
|
||||
|
||||
impl<T: DecisionLoader> DecisionHandler<T> {
|
||||
pub fn new(trace: bool, max_depth: u8, loader: Arc<T>, js_runtime: Option<Runtime>) -> Self {
|
||||
impl<L: DecisionLoader, A: CustomNodeAdapter> DecisionHandler<L, A> {
|
||||
pub fn new(
|
||||
trace: bool,
|
||||
max_depth: u8,
|
||||
loader: Arc<L>,
|
||||
adapter: Arc<A>,
|
||||
js_runtime: Option<Runtime>,
|
||||
) -> Self {
|
||||
Self {
|
||||
trace,
|
||||
loader,
|
||||
adapter,
|
||||
max_depth,
|
||||
js_runtime,
|
||||
}
|
||||
@@ -37,6 +46,7 @@ impl<T: DecisionLoader> DecisionHandler<T> {
|
||||
content: sub_decision.deref(),
|
||||
max_depth: self.max_depth,
|
||||
loader: self.loader.clone(),
|
||||
adapter: self.adapter.clone(),
|
||||
iteration: request.iteration + 1,
|
||||
trace: self.trace,
|
||||
})?
|
||||
|
||||
@@ -10,6 +10,7 @@ use serde::{Deserialize, Serialize, Serializer};
|
||||
use serde_json::Value;
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::handler::custom_node_adapter::{CustomNodeAdapter, CustomNodeRequest};
|
||||
use crate::handler::decision::DecisionHandler;
|
||||
use crate::handler::expression::ExpressionHandler;
|
||||
use crate::handler::function::runtime::create_runtime;
|
||||
@@ -21,8 +22,9 @@ use crate::loader::DecisionLoader;
|
||||
use crate::model::{DecisionContent, DecisionNodeKind};
|
||||
use crate::{EvaluationError, NodeError};
|
||||
|
||||
pub struct DecisionGraph<'a, L: DecisionLoader> {
|
||||
pub struct DecisionGraph<'a, L: DecisionLoader, A: CustomNodeAdapter> {
|
||||
graph: StableDiDecisionGraph<'a>,
|
||||
adapter: Arc<A>,
|
||||
loader: Arc<L>,
|
||||
trace: bool,
|
||||
max_depth: u8,
|
||||
@@ -30,17 +32,18 @@ pub struct DecisionGraph<'a, L: DecisionLoader> {
|
||||
runtime: Option<Runtime>,
|
||||
}
|
||||
|
||||
pub struct DecisionGraphConfig<'a, T: DecisionLoader> {
|
||||
pub loader: Arc<T>,
|
||||
pub struct DecisionGraphConfig<'a, L: DecisionLoader, A: CustomNodeAdapter> {
|
||||
pub loader: Arc<L>,
|
||||
pub adapter: Arc<A>,
|
||||
pub content: &'a DecisionContent,
|
||||
pub trace: bool,
|
||||
pub iteration: u8,
|
||||
pub max_depth: u8,
|
||||
}
|
||||
|
||||
impl<'a, L: DecisionLoader> DecisionGraph<'a, L> {
|
||||
impl<'a, L: DecisionLoader, A: CustomNodeAdapter> DecisionGraph<'a, L, A> {
|
||||
pub fn try_new(
|
||||
config: DecisionGraphConfig<'a, L>,
|
||||
config: DecisionGraphConfig<'a, L, A>,
|
||||
) -> Result<Self, DecisionGraphValidationError> {
|
||||
let content = config.content;
|
||||
let mut graph = StableDiDecisionGraph::new();
|
||||
@@ -69,7 +72,8 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> {
|
||||
graph,
|
||||
iteration: config.iteration,
|
||||
trace: config.trace,
|
||||
loader: config.loader.clone(),
|
||||
loader: config.loader,
|
||||
adapter: config.adapter,
|
||||
max_depth: config.max_depth,
|
||||
runtime: None,
|
||||
})
|
||||
@@ -154,13 +158,13 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> {
|
||||
};
|
||||
}
|
||||
|
||||
let node_request = NodeRequest {
|
||||
let mut node_request = NodeRequest {
|
||||
node,
|
||||
iteration: self.iteration,
|
||||
input: walker.incoming_node_data(&self.graph, nid),
|
||||
input: walker.incoming_node_data(&self.graph, nid, true),
|
||||
};
|
||||
|
||||
match node.kind {
|
||||
match &node.kind {
|
||||
DecisionNodeKind::InputNode => {
|
||||
walker.set_node_data(nid, context.clone());
|
||||
trace!({
|
||||
@@ -181,6 +185,9 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> {
|
||||
performance: None,
|
||||
trace_data: None,
|
||||
});
|
||||
if let Some(obj) = node_request.input.as_object_mut() {
|
||||
obj.remove("$nodes");
|
||||
}
|
||||
|
||||
return Ok(DecisionGraphResponse {
|
||||
result: node_request.input,
|
||||
@@ -228,6 +235,7 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> {
|
||||
self.trace,
|
||||
self.max_depth,
|
||||
self.loader.clone(),
|
||||
self.adapter.clone(),
|
||||
self.runtime.clone(),
|
||||
)
|
||||
.handle(&node_request)
|
||||
@@ -285,6 +293,26 @@ impl<'a, L: DecisionLoader> DecisionGraph<'a, L> {
|
||||
trace_data: res.trace_data,
|
||||
});
|
||||
}
|
||||
DecisionNodeKind::CustomNode { .. } => {
|
||||
let res = self
|
||||
.adapter
|
||||
.handle(CustomNodeRequest::try_from(&node_request).unwrap())
|
||||
.await
|
||||
.map_err(|e| NodeError {
|
||||
node_id: node.id.clone(),
|
||||
source: e.into(),
|
||||
})?;
|
||||
|
||||
walker.set_node_data(nid, res.output.clone());
|
||||
trace!({
|
||||
input: node_request.input,
|
||||
output: res.output,
|
||||
name: node.name.clone(),
|
||||
id: node.id.clone(),
|
||||
performance: Some(format!("{:?}", start.elapsed())),
|
||||
trace_data: res.trace_data,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -351,10 +379,10 @@ pub struct DecisionGraphResponse {
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct DecisionGraphTrace {
|
||||
input: Value,
|
||||
output: Value,
|
||||
name: String,
|
||||
id: String,
|
||||
performance: Option<String>,
|
||||
trace_data: Option<Value>,
|
||||
pub input: Value,
|
||||
pub output: Value,
|
||||
pub name: String,
|
||||
pub id: String,
|
||||
pub performance: Option<String>,
|
||||
pub trace_data: Option<Value>,
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ pub mod expression;
|
||||
pub mod function;
|
||||
pub mod table;
|
||||
|
||||
pub(crate) mod graph;
|
||||
pub(crate) mod node;
|
||||
pub mod custom_node_adapter;
|
||||
pub mod graph;
|
||||
pub mod node;
|
||||
pub(crate) mod traversal;
|
||||
|
||||
@@ -11,7 +11,7 @@ pub struct NodeResponse {
|
||||
pub trace_data: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct NodeRequest<'a> {
|
||||
pub input: Value,
|
||||
pub iteration: u8,
|
||||
|
||||
@@ -62,12 +62,36 @@ impl GraphWalker {
|
||||
self.node_data.get(&node_id)
|
||||
}
|
||||
|
||||
pub fn get_all_node_data(&self, g: &StableDiDecisionGraph) -> Value {
|
||||
let node_values: Map<String, Value> = self
|
||||
.node_data
|
||||
.iter()
|
||||
.map(|(idx, value)| {
|
||||
let weight = g.node_weight(*idx).unwrap();
|
||||
(weight.name.clone(), value.clone())
|
||||
})
|
||||
.collect();
|
||||
|
||||
Value::Object(node_values)
|
||||
}
|
||||
|
||||
pub fn set_node_data(&mut self, node_id: NodeIndex, value: Value) {
|
||||
self.node_data.insert(node_id, value);
|
||||
}
|
||||
|
||||
pub fn incoming_node_data(&self, g: &StableDiDecisionGraph, node_id: NodeIndex) -> Value {
|
||||
self.merge_node_data(g.neighbors_directed(node_id, Incoming))
|
||||
pub fn incoming_node_data(
|
||||
&self,
|
||||
g: &StableDiDecisionGraph,
|
||||
node_id: NodeIndex,
|
||||
with_nodes: bool,
|
||||
) -> Value {
|
||||
let mut value = self.merge_node_data(g.neighbors_directed(node_id, Incoming));
|
||||
|
||||
if let Some(object) = with_nodes.then_some(value.as_object_mut()).flatten() {
|
||||
object.insert("$nodes".to_string(), self.get_all_node_data(g));
|
||||
}
|
||||
|
||||
value
|
||||
}
|
||||
|
||||
pub fn merge_node_data<I>(&self, iter: I) -> Value
|
||||
@@ -98,7 +122,7 @@ impl GraphWalker {
|
||||
self.ordered.visit(nid);
|
||||
|
||||
if let DecisionNodeKind::SwitchNode { content } = &decision_node.kind {
|
||||
let mut input_data = self.incoming_node_data(g, nid);
|
||||
let mut input_data = self.incoming_node_data(g, nid, true);
|
||||
let input_context = json!({ "$": &input_data });
|
||||
merge_json(&mut input_data, &input_context, true);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user