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:
stefan-gorules
2024-04-03 16:09:51 +02:00
committed by GitHub
parent 685a0345f5
commit daecf901e6
65 changed files with 1729 additions and 235 deletions
@@ -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,
})
}
}
+14 -4
View File
@@ -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,
})?
+43 -15
View File
@@ -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 -2
View File
@@ -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;
+1 -1
View File
@@ -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,
+27 -3
View File
@@ -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);