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 fixedbitset::FixedBitSet;
6use index::TableIndex;
7use serde::Serialize;
8use std::ops::Deref;
9use std::rc::Rc;
10use std::sync::Arc;
11use zen_expression::variable::ToVariable;
12use zen_expression::Isolate;
13use zen_types::decision::{
14    DecisionTableContent, DecisionTableHitPolicy, DecisionTableInputField, TransformAttributes,
15};
16use zen_types::variable::Variable;
17pub(crate) mod index;
18
19#[derive(Debug, Clone)]
20pub struct DecisionTableNodeHandler;
21
22pub type DecisionTableNodeData = DecisionTableContent;
23
24type DecisionTableContext = NodeContext<DecisionTableNodeData, DecisionTableNodeTrace>;
25
26impl NodeHandler for DecisionTableNodeHandler {
27    type NodeData = DecisionTableNodeData;
28    type TraceData = DecisionTableNodeTrace;
29
30    fn transform_attributes(
31        &self,
32        ctx: &NodeContext<Self::NodeData, Self::TraceData>,
33    ) -> Option<TransformAttributes> {
34        Some(ctx.node.transform_attributes.clone())
35    }
36
37    async fn handle(&self, ctx: NodeContext<Self::NodeData, Self::TraceData>) -> NodeResult {
38        match ctx.node.hit_policy {
39            DecisionTableHitPolicy::First => self.handle_first_hit(ctx),
40            DecisionTableHitPolicy::Collect => self.handle_collect(ctx),
41        }
42    }
43}
44
45impl DecisionTableNodeHandler {
46    fn handle_first_hit(&self, ctx: DecisionTableContext) -> NodeResult {
47        let mut isolate = ctx.isolate();
48
49        if !ctx.config.trace {
50            let index = Self::table_index(&ctx);
51            let candidates =
52                index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate));
53            let pruner = candidates.as_ref().and(index);
54            for (row_idx, rule) in ctx.node.rules.iter().enumerate() {
55                if candidates.as_ref().is_some_and(|c| !c.contains(row_idx)) {
56                    continue;
57                }
58                let pruned = pruner.map(|ix| (ix, row_idx));
59                if let Some(RowResult::Output(output)) =
60                    self.evaluate_row(&ctx, rule, &mut isolate, pruned)
61                {
62                    return ctx.success(output);
63                }
64            }
65            return Ok(NodeResponse {
66                output: Variable::Null,
67                trace_data: None,
68            });
69        }
70
71        let hit = ctx.node.rules.iter().enumerate().find_map(|(index, rule)| {
72            match self.evaluate_row(&ctx, rule, &mut isolate, None)? {
73                RowResult::WithTrace {
74                    output,
75                    reference_map,
76                    rule,
77                } => Some((index, output, reference_map, rule)),
78                RowResult::Output(output) => {
79                    Some((index, output, Default::default(), Default::default()))
80                }
81            }
82        });
83
84        match hit {
85            Some((index, output, reference_map, rule)) => {
86                ctx.trace(|t| {
87                    *t = DecisionTableNodeTrace::FirstHit(DecisionTableRowTrace {
88                        reference_map,
89                        index,
90                        rule,
91                    })
92                });
93                ctx.success(output)
94            }
95            None => Ok(NodeResponse {
96                output: Variable::Null,
97                trace_data: None,
98            }),
99        }
100    }
101
102    fn handle_collect(&self, ctx: DecisionTableContext) -> NodeResult {
103        let mut outputs = Vec::new();
104        let mut traces = Vec::new();
105        let mut isolate = ctx.isolate();
106
107        let table_index = (!ctx.config.trace)
108            .then(|| Self::table_index(&ctx))
109            .flatten();
110        let candidates =
111            table_index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate));
112        let pruner = candidates.as_ref().and(table_index);
113
114        for (index, rule) in ctx.node.rules.iter().enumerate() {
115            if candidates.as_ref().is_some_and(|c| !c.contains(index)) {
116                continue;
117            }
118            let pruned = pruner.map(|ix| (ix, index));
119            if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate, pruned) {
120                match result {
121                    RowResult::Output(output) => {
122                        outputs.push(output);
123                    }
124                    RowResult::WithTrace {
125                        output,
126                        reference_map,
127                        rule,
128                    } => {
129                        outputs.push(output);
130                        traces.push(DecisionTableRowTrace {
131                            index,
132                            rule,
133                            reference_map,
134                        });
135                    }
136                }
137            }
138        }
139
140        ctx.trace(|t| {
141            *t = DecisionTableNodeTrace::Collect(traces);
142        });
143
144        ctx.success(Variable::from_array(outputs))
145    }
146
147    pub(crate) fn cell_passes(
148        rule: &HashMap<Arc<str>, Arc<str>>,
149        input: &zen_types::decision::DecisionTableInputField,
150        isolate: &mut Isolate,
151    ) -> bool {
152        let Some(rule_value) = rule.get(&input.id) else {
153            return true;
154        };
155        if rule_value.is_empty() {
156            return true;
157        }
158        match &input.field {
159            None => isolate
160                .run_standard(rule_value)
161                .ok()
162                .and_then(|result| result.as_bool())
163                .unwrap_or(false),
164            Some(field) => {
165                if isolate.set_reference(field).is_err() {
166                    return false;
167                }
168                isolate.run_unary(rule_value).unwrap_or(false)
169            }
170        }
171    }
172
173    fn table_index(ctx: &DecisionTableContext) -> Option<&TableIndex> {
174        ctx.extensions.dt_indexes.as_ref()?.get(&ctx.id)
175    }
176
177    fn candidate_rows(
178        index: &TableIndex,
179        inputs: &[DecisionTableInputField],
180        isolate: &mut Isolate,
181    ) -> Option<FixedBitSet> {
182        let mut acc: Option<FixedBitSet> = None;
183        for (col_idx, column) in index.columns.iter().enumerate() {
184            let Some(column) = column else {
185                continue;
186            };
187            let Some(field) = inputs[col_idx].field.as_ref().filter(|f| !f.is_empty()) else {
188                continue;
189            };
190            isolate.set_reference(field).ok()?;
191            let value = isolate.get_reference(field)?;
192            if matches!(value, Variable::Dynamic(_)) {
193                return None;
194            }
195            let hit = column.rows_for(&value);
196            match &mut acc {
197                None => {
198                    let mut first = column.fallback.clone();
199                    if let Some(hit) = hit {
200                        first.union_with(hit);
201                    }
202                    acc = Some(first);
203                }
204                Some(acc) => {
205                    let fallback = column.fallback.as_slice();
206                    let hit = hit.map(FixedBitSet::as_slice).unwrap_or_default();
207                    for (i, word) in acc.as_mut_slice().iter_mut().enumerate() {
208                        let f = fallback.get(i).copied().unwrap_or(0);
209                        let h = hit.get(i).copied().unwrap_or(0);
210                        *word &= f | h;
211                    }
212                }
213            }
214        }
215        acc
216    }
217
218    fn evaluate_row<'a>(
219        &self,
220        ctx: &'a DecisionTableContext,
221        rule: &'a HashMap<Arc<str>, Arc<str>>,
222        isolate: &mut Isolate,
223        pruned: Option<(&TableIndex, usize)>,
224    ) -> Option<RowResult> {
225        let content = &ctx.node;
226        for (col_idx, input) in content.inputs.iter().enumerate() {
227            if pruned.is_some_and(|(ix, row_idx)| ix.decides(col_idx, row_idx)) {
228                continue;
229            }
230            let Some(rule_value) = rule.get(&input.id) else {
231                continue;
232            };
233            if rule_value.is_empty() {
234                continue;
235            }
236
237            match &input.field {
238                None => {
239                    let result = isolate.run_standard(rule_value).ok()?;
240                    if !result.as_bool().unwrap_or(false) {
241                        return None;
242                    }
243                }
244                Some(field) => {
245                    isolate.set_reference(&field).ok()?;
246                    if !isolate.run_unary(rule_value).ok()? {
247                        return None;
248                    }
249                }
250            }
251        }
252
253        let outputs = Variable::empty_object();
254        for output in content.outputs.iter() {
255            let Some(rule_value) = rule.get(&output.id) else {
256                continue;
257            };
258            if rule_value.is_empty() {
259                continue;
260            }
261
262            let res = isolate.run_standard(rule_value).ok()?;
263            outputs.dot_insert(output.field.deref(), res);
264        }
265
266        if !ctx.config.trace {
267            return Some(RowResult::Output(outputs));
268        }
269
270        let id_str = Rc::<str>::from("_id");
271        let description_str = Rc::<str>::from("_description");
272
273        let rule_id = match rule.get(id_str.as_ref()) {
274            Some(rid) => Rc::<str>::from(rid.deref()),
275            None => Rc::from(""),
276        };
277
278        let mut expressions: HashMap<Rc<str>, Rc<str>> = Default::default();
279        let mut reference_map: HashMap<Rc<str>, Variable> = Default::default();
280
281        expressions.insert(id_str.clone(), rule_id.clone());
282        if let Some(description) = rule.get(description_str.as_ref()) {
283            expressions.insert(description_str.clone(), Rc::from(description.deref()));
284        }
285
286        for input in content.inputs.iter() {
287            let Some(rule_value) = rule.get(input.id.deref()) else {
288                continue;
289            };
290            let Some(input_field) = &input.field else {
291                continue;
292            };
293
294            if let Some(reference) = isolate.get_reference(input_field.deref()) {
295                reference_map.insert(Rc::from(input_field.deref()), reference);
296            } else if let Some(reference) = isolate.run_standard(input_field.deref()).ok() {
297                reference_map.insert(Rc::from(input_field.deref()), reference);
298            }
299
300            let input_identifier = format!("{input_field}[{}]", &input.id);
301            expressions.insert(
302                Rc::from(input_identifier.as_str()),
303                Rc::from(rule_value.deref()),
304            );
305        }
306
307        Some(RowResult::WithTrace {
308            output: outputs.to_variable(),
309            reference_map,
310            rule: expressions,
311        })
312    }
313}
314
315enum RowResult {
316    Output(Variable),
317    WithTrace {
318        output: Variable,
319        reference_map: HashMap<Rc<str>, Variable>,
320        rule: HashMap<Rc<str>, Rc<str>>,
321    },
322}
323
324#[derive(Debug, Clone, Serialize, ToVariable)]
325pub struct DecisionTableRowTrace {
326    index: usize,
327    reference_map: HashMap<Rc<str>, Variable>,
328    rule: HashMap<Rc<str>, Rc<str>>,
329}
330
331#[derive(Debug, Clone, Serialize, ToVariable)]
332#[serde(untagged)]
333pub enum DecisionTableNodeTrace {
334    FirstHit(DecisionTableRowTrace),
335    Collect(Vec<DecisionTableRowTrace>),
336}
337
338impl Default for DecisionTableNodeTrace {
339    fn default() -> Self {
340        DecisionTableNodeTrace::Collect(Default::default())
341    }
342}