Skip to main content

zen_engine/nodes/decision_table/
mod.rs

1use crate::nodes::definition::NodeHandler;
2use crate::nodes::result::NodeResult;
3use crate::nodes::{NodeContext, NodeResponse};
4use ahash::HashMap;
5use serde::Serialize;
6use std::ops::Deref;
7use std::rc::Rc;
8use std::sync::Arc;
9use zen_expression::variable::ToVariable;
10use zen_expression::Isolate;
11use zen_types::decision::{DecisionTableContent, DecisionTableHitPolicy, TransformAttributes};
12use zen_types::variable::Variable;
13#[derive(Debug, Clone)]
14pub struct DecisionTableNodeHandler;
15
16pub type DecisionTableNodeData = DecisionTableContent;
17
18type DecisionTableContext = NodeContext<DecisionTableNodeData, DecisionTableNodeTrace>;
19
20impl NodeHandler for DecisionTableNodeHandler {
21    type NodeData = DecisionTableNodeData;
22    type TraceData = DecisionTableNodeTrace;
23
24    fn transform_attributes(
25        &self,
26        ctx: &NodeContext<Self::NodeData, Self::TraceData>,
27    ) -> Option<TransformAttributes> {
28        Some(ctx.node.transform_attributes.clone())
29    }
30
31    async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult {
32        match ctx.node.hit_policy {
33            DecisionTableHitPolicy::First => self.handle_first_hit(ctx),
34            DecisionTableHitPolicy::Collect => self.handle_collect(ctx),
35        }
36    }
37}
38
39impl DecisionTableNodeHandler {
40    fn handle_first_hit(&self, ctx: DecisionTableContext) -> NodeResult {
41        let mut isolate = Isolate::with_environment(ctx.input.depth_clone(1))
42            .with_cache(ctx.extensions.compiled_cache.clone());
43
44        if !ctx.config.trace {
45            for rule in ctx.node.rules.iter() {
46                if let Some(RowResult::Output(output)) = self.evaluate_row(&ctx, rule, &mut isolate)
47                {
48                    return ctx.success(output);
49                }
50            }
51            return Ok(NodeResponse {
52                output: Variable::Null,
53                trace_data: None,
54            });
55        }
56
57        let hit = ctx.node.rules.iter().enumerate().find_map(|(index, rule)| {
58            match self.evaluate_row(&ctx, rule, &mut isolate)? {
59                RowResult::WithTrace {
60                    output,
61                    reference_map,
62                    rule,
63                } => Some((index, output, reference_map, rule)),
64                RowResult::Output(output) => {
65                    Some((index, output, Default::default(), Default::default()))
66                }
67            }
68        });
69
70        match hit {
71            Some((index, output, reference_map, rule)) => {
72                ctx.trace(|t| {
73                    *t = DecisionTableNodeTrace::FirstHit(DecisionTableRowTrace {
74                        reference_map,
75                        index,
76                        rule,
77                    })
78                });
79                ctx.success(output)
80            }
81            None => Ok(NodeResponse {
82                output: Variable::Null,
83                trace_data: None,
84            }),
85        }
86    }
87
88    fn handle_collect(&self, ctx: DecisionTableContext) -> NodeResult {
89        let mut outputs = Vec::new();
90        let mut traces = Vec::new();
91        let mut isolate = Isolate::with_environment(ctx.input.depth_clone(1))
92            .with_cache(ctx.extensions.compiled_cache.clone());
93
94        for (index, rule) in ctx.node.rules.iter().enumerate() {
95            if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate) {
96                match result {
97                    RowResult::Output(output) => {
98                        outputs.push(output);
99                    }
100                    RowResult::WithTrace {
101                        output,
102                        reference_map,
103                        rule,
104                    } => {
105                        outputs.push(output);
106                        traces.push(DecisionTableRowTrace {
107                            index,
108                            rule,
109                            reference_map,
110                        });
111                    }
112                }
113            }
114        }
115
116        ctx.trace(|t| {
117            *t = DecisionTableNodeTrace::Collect(traces);
118        });
119
120        ctx.success(Variable::from_array(outputs))
121    }
122
123    pub(crate) fn cell_passes(
124        rule: &HashMap<Arc<str>, Arc<str>>,
125        input: &zen_types::decision::DecisionTableInputField,
126        isolate: &mut Isolate,
127    ) -> bool {
128        let Some(rule_value) = rule.get(&input.id) else {
129            return true;
130        };
131        if rule_value.is_empty() {
132            return true;
133        }
134        match &input.field {
135            None => isolate
136                .run_standard(rule_value)
137                .ok()
138                .and_then(|result| result.as_bool())
139                .unwrap_or(false),
140            Some(field) => {
141                if isolate.set_reference(field).is_err() {
142                    return false;
143                }
144                isolate.run_unary(rule_value).unwrap_or(false)
145            }
146        }
147    }
148
149    fn evaluate_row<'a>(
150        &self,
151        ctx: &'a DecisionTableContext,
152        rule: &'a HashMap<Arc<str>, Arc<str>>,
153        isolate: &mut Isolate,
154    ) -> Option<RowResult> {
155        let content = &ctx.node;
156        for input in content.inputs.iter() {
157            let Some(rule_value) = rule.get(&input.id) else {
158                continue;
159            };
160            if rule_value.is_empty() {
161                continue;
162            }
163
164            match &input.field {
165                None => {
166                    let result = isolate.run_standard(rule_value).ok()?;
167                    if !result.as_bool().unwrap_or(false) {
168                        return None;
169                    }
170                }
171                Some(field) => {
172                    isolate.set_reference(&field).ok()?;
173                    if !isolate.run_unary(rule_value).ok()? {
174                        return None;
175                    }
176                }
177            }
178        }
179
180        let outputs = Variable::empty_object();
181        for output in content.outputs.iter() {
182            let Some(rule_value) = rule.get(&output.id) else {
183                continue;
184            };
185            if rule_value.is_empty() {
186                continue;
187            }
188
189            let res = isolate.run_standard(rule_value).ok()?;
190            outputs.dot_insert(output.field.deref(), res);
191        }
192
193        if !ctx.config.trace {
194            return Some(RowResult::Output(outputs));
195        }
196
197        let id_str = Rc::<str>::from("_id");
198        let description_str = Rc::<str>::from("_description");
199
200        let rule_id = match rule.get(id_str.as_ref()) {
201            Some(rid) => Rc::<str>::from(rid.deref()),
202            None => Rc::from(""),
203        };
204
205        let mut expressions: HashMap<Rc<str>, Rc<str>> = Default::default();
206        let mut reference_map: HashMap<Rc<str>, Variable> = Default::default();
207
208        expressions.insert(id_str.clone(), rule_id.clone());
209        if let Some(description) = rule.get(description_str.as_ref()) {
210            expressions.insert(description_str.clone(), Rc::from(description.deref()));
211        }
212
213        for input in content.inputs.iter() {
214            let Some(rule_value) = rule.get(input.id.deref()) else {
215                continue;
216            };
217            let Some(input_field) = &input.field else {
218                continue;
219            };
220
221            if let Some(reference) = isolate.get_reference(input_field.deref()) {
222                reference_map.insert(Rc::from(input_field.deref()), reference);
223            } else if let Some(reference) = isolate.run_standard(input_field.deref()).ok() {
224                reference_map.insert(Rc::from(input_field.deref()), reference);
225            }
226
227            let input_identifier = format!("{input_field}[{}]", &input.id);
228            expressions.insert(
229                Rc::from(input_identifier.as_str()),
230                Rc::from(rule_value.deref()),
231            );
232        }
233
234        Some(RowResult::WithTrace {
235            output: outputs.to_variable(),
236            reference_map,
237            rule: expressions,
238        })
239    }
240}
241
242enum RowResult {
243    Output(Variable),
244    WithTrace {
245        output: Variable,
246        reference_map: HashMap<Rc<str>, Variable>,
247        rule: HashMap<Rc<str>, Rc<str>>,
248    },
249}
250
251#[derive(Debug, Clone, Serialize, ToVariable)]
252pub struct DecisionTableRowTrace {
253    index: usize,
254    reference_map: HashMap<Rc<str>, Variable>,
255    rule: HashMap<Rc<str>, Rc<str>>,
256}
257
258#[derive(Debug, Clone, Serialize, ToVariable)]
259#[serde(untagged)]
260pub enum DecisionTableNodeTrace {
261    FirstHit(DecisionTableRowTrace),
262    Collect(Vec<DecisionTableRowTrace>),
263}
264
265impl Default for DecisionTableNodeTrace {
266    fn default() -> Self {
267        DecisionTableNodeTrace::Collect(Default::default())
268    }
269}