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        let has_collect_columns = ctx.node.outputs.iter().any(|output| output.write_path().1);
39        match ctx.node.hit_policy {
40            DecisionTableHitPolicy::First if has_collect_columns => {
41                self.handle_first_hit_collect(ctx)
42            }
43            DecisionTableHitPolicy::First => self.handle_first_hit(ctx),
44            DecisionTableHitPolicy::Collect => self.handle_collect(ctx),
45        }
46    }
47}
48
49impl DecisionTableNodeHandler {
50    fn handle_first_hit(&self, ctx: DecisionTableContext) -> NodeResult {
51        let mut isolate = ctx.isolate();
52
53        if !ctx.config.trace {
54            let index = Self::table_index(&ctx);
55            let candidates =
56                index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate));
57            let pruner = candidates.as_ref().and(index);
58            for (row_idx, rule) in ctx.node.rules.iter().enumerate() {
59                if candidates.as_ref().is_some_and(|c| !c.contains(row_idx)) {
60                    continue;
61                }
62                let pruned = pruner.map(|ix| (ix, row_idx));
63                if let Some(RowResult::Output(output)) =
64                    self.evaluate_row(&ctx, rule, &mut isolate, pruned)
65                {
66                    return ctx.success(output);
67                }
68            }
69            return Ok(NodeResponse {
70                output: Variable::Null,
71                trace_data: None,
72            });
73        }
74
75        let hit = ctx.node.rules.iter().enumerate().find_map(|(index, rule)| {
76            match self.evaluate_row(&ctx, rule, &mut isolate, None)? {
77                RowResult::WithTrace {
78                    output,
79                    reference_map,
80                    rule,
81                } => Some((index, output, reference_map, rule)),
82                RowResult::Output(output) => {
83                    Some((index, output, Default::default(), Default::default()))
84                }
85            }
86        });
87
88        match hit {
89            Some((index, output, reference_map, rule)) => {
90                ctx.trace(|t| {
91                    *t = DecisionTableNodeTrace::FirstHit(DecisionTableRowTrace {
92                        reference_map,
93                        index,
94                        rule,
95                    })
96                });
97                ctx.success(output)
98            }
99            None => Ok(NodeResponse {
100                output: Variable::Null,
101                trace_data: None,
102            }),
103        }
104    }
105
106    fn handle_collect(&self, ctx: DecisionTableContext) -> NodeResult {
107        let mut outputs = Vec::new();
108        let mut traces = Vec::new();
109        let mut isolate = ctx.isolate();
110
111        let table_index = (!ctx.config.trace)
112            .then(|| Self::table_index(&ctx))
113            .flatten();
114        let candidates =
115            table_index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate));
116        let pruner = candidates.as_ref().and(table_index);
117
118        for (index, rule) in ctx.node.rules.iter().enumerate() {
119            if candidates.as_ref().is_some_and(|c| !c.contains(index)) {
120                continue;
121            }
122            let pruned = pruner.map(|ix| (ix, index));
123            if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate, pruned) {
124                match result {
125                    RowResult::Output(output) => {
126                        outputs.push(output);
127                    }
128                    RowResult::WithTrace {
129                        output,
130                        reference_map,
131                        rule,
132                    } => {
133                        outputs.push(output);
134                        traces.push(DecisionTableRowTrace {
135                            index,
136                            rule,
137                            reference_map,
138                        });
139                    }
140                }
141            }
142        }
143
144        ctx.trace(|t| {
145            *t = DecisionTableNodeTrace::Collect(traces);
146        });
147
148        ctx.success(Variable::from_array(outputs))
149    }
150
151    pub(crate) fn cell_passes(
152        rule: &HashMap<Arc<str>, Arc<str>>,
153        input: &zen_types::decision::DecisionTableInputField,
154        isolate: &mut Isolate,
155    ) -> bool {
156        let Some(rule_value) = rule.get(&input.id) else {
157            return true;
158        };
159        if rule_value.is_empty() {
160            return true;
161        }
162        match &input.field {
163            None => isolate
164                .run_standard(rule_value)
165                .ok()
166                .and_then(|result| result.as_bool())
167                .unwrap_or(false),
168            Some(field) => {
169                if isolate.set_reference(field).is_err() {
170                    return false;
171                }
172                isolate.run_unary(rule_value).unwrap_or(false)
173            }
174        }
175    }
176
177    fn table_index(ctx: &DecisionTableContext) -> Option<&TableIndex> {
178        ctx.extensions.dt_indexes.as_ref()?.get(&ctx.id)
179    }
180
181    fn candidate_rows(
182        index: &TableIndex,
183        inputs: &[DecisionTableInputField],
184        isolate: &mut Isolate,
185    ) -> Option<FixedBitSet> {
186        let mut acc: Option<FixedBitSet> = None;
187        for (col_idx, column) in index.columns.iter().enumerate() {
188            let Some(column) = column else {
189                continue;
190            };
191            let Some(field) = inputs[col_idx].field.as_ref().filter(|f| !f.is_empty()) else {
192                continue;
193            };
194            isolate.set_reference(field).ok()?;
195            let value = isolate.get_reference(field)?;
196            if matches!(value, Variable::Dynamic(_)) {
197                return None;
198            }
199            let hit = column.rows_for(&value);
200            match &mut acc {
201                None => {
202                    let mut first = column.fallback.clone();
203                    if let Some(hit) = hit {
204                        first.union_with(hit);
205                    }
206                    acc = Some(first);
207                }
208                Some(acc) => {
209                    let fallback = column.fallback.as_slice();
210                    let hit = hit.map(FixedBitSet::as_slice).unwrap_or_default();
211                    for (i, word) in acc.as_mut_slice().iter_mut().enumerate() {
212                        let f = fallback.get(i).copied().unwrap_or(0);
213                        let h = hit.get(i).copied().unwrap_or(0);
214                        *word &= f | h;
215                    }
216                }
217            }
218        }
219        acc
220    }
221
222    fn handle_first_hit_collect(&self, ctx: DecisionTableContext) -> NodeResult {
223        let mut isolate = ctx.isolate();
224
225        let table_index = (!ctx.config.trace)
226            .then(|| Self::table_index(&ctx))
227            .flatten();
228        let candidates =
229            table_index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate));
230        let pruner = candidates.as_ref().and(table_index);
231
232        let mut scalars: Option<Variable> = None;
233        let mut collected: Vec<Vec<Variable>> = vec![Vec::new(); ctx.node.outputs.len()];
234        let mut matched = false;
235        let mut traces = Vec::new();
236
237        for (row_idx, rule) in ctx.node.rules.iter().enumerate() {
238            if candidates.as_ref().is_some_and(|c| !c.contains(row_idx)) {
239                continue;
240            }
241            let pruned = pruner.map(|ix| (ix, row_idx));
242            if !Self::row_matches(&ctx, rule, &mut isolate, pruned) {
243                continue;
244            }
245
246            let Some((row_scalars, row_collects)) =
247                Self::evaluate_row_cells(&ctx, rule, &mut isolate, scalars.is_none())
248            else {
249                continue;
250            };
251            matched = true;
252            if let Some(row_scalars) = row_scalars {
253                scalars = Some(row_scalars);
254            }
255            for (column_idx, value) in row_collects {
256                collected[column_idx].push(value);
257            }
258
259            if ctx.config.trace {
260                let (reference_map, rule_trace) = Self::row_trace_parts(&ctx, rule, &mut isolate);
261                traces.push(DecisionTableRowTrace {
262                    index: row_idx,
263                    reference_map,
264                    rule: rule_trace,
265                });
266            }
267        }
268
269        if !matched {
270            return Ok(NodeResponse {
271                output: Variable::Null,
272                trace_data: None,
273            });
274        }
275
276        let output = scalars.unwrap_or_else(Variable::empty_object);
277        for (column_idx, column) in ctx.node.outputs.iter().enumerate() {
278            let (path, collect) = column.write_path();
279            if !collect || path.is_empty() {
280                continue;
281            }
282            let values = std::mem::take(&mut collected[column_idx]);
283            output.dot_insert(path, Variable::from_array(values));
284        }
285
286        ctx.trace(|t| {
287            *t = DecisionTableNodeTrace::Collect(traces);
288        });
289        ctx.success(output)
290    }
291
292    fn row_matches(
293        ctx: &DecisionTableContext,
294        rule: &HashMap<Arc<str>, Arc<str>>,
295        isolate: &mut Isolate,
296        pruned: Option<(&TableIndex, usize)>,
297    ) -> bool {
298        for (col_idx, input) in ctx.node.inputs.iter().enumerate() {
299            if pruned.is_some_and(|(ix, row_idx)| ix.decides(col_idx, row_idx)) {
300                continue;
301            }
302            let Some(rule_value) = rule.get(&input.id) else {
303                continue;
304            };
305            if rule_value.is_empty() {
306                continue;
307            }
308
309            let passed = match &input.field {
310                None => isolate
311                    .run_standard(rule_value)
312                    .ok()
313                    .and_then(|result| result.as_bool())
314                    .unwrap_or(false),
315                Some(field) => {
316                    isolate.set_reference(field).is_ok()
317                        && isolate.run_unary(rule_value).unwrap_or(false)
318                }
319            };
320            if !passed {
321                return false;
322            }
323        }
324        true
325    }
326
327    fn evaluate_row_cells(
328        ctx: &DecisionTableContext,
329        rule: &HashMap<Arc<str>, Arc<str>>,
330        isolate: &mut Isolate,
331        include_scalars: bool,
332    ) -> Option<(Option<Variable>, Vec<(usize, Variable)>)> {
333        let scalars = include_scalars.then(Variable::empty_object);
334        let mut collects = Vec::new();
335        for (column_idx, output) in ctx.node.outputs.iter().enumerate() {
336            let (path, collect) = output.write_path();
337            if path.is_empty() || (!collect && !include_scalars) {
338                continue;
339            }
340            let Some(rule_value) = rule.get(&output.id) else {
341                continue;
342            };
343            if rule_value.is_empty() {
344                continue;
345            }
346
347            let value = isolate.run_standard(rule_value).ok()?.deep_clone();
348            if collect {
349                collects.push((column_idx, value));
350            } else if let Some(scalars) = &scalars {
351                scalars.dot_insert(path, value);
352            }
353        }
354        Some((scalars, collects))
355    }
356
357    fn row_trace_parts(
358        ctx: &DecisionTableContext,
359        rule: &HashMap<Arc<str>, Arc<str>>,
360        isolate: &mut Isolate,
361    ) -> (HashMap<Rc<str>, Variable>, HashMap<Rc<str>, Rc<str>>) {
362        let id_str = Rc::<str>::from("_id");
363        let description_str = Rc::<str>::from("_description");
364
365        let rule_id = match rule.get(id_str.as_ref()) {
366            Some(rid) => Rc::<str>::from(rid.deref()),
367            None => Rc::from(""),
368        };
369
370        let mut expressions: HashMap<Rc<str>, Rc<str>> = Default::default();
371        let mut reference_map: HashMap<Rc<str>, Variable> = Default::default();
372
373        expressions.insert(id_str.clone(), rule_id.clone());
374        if let Some(description) = rule.get(description_str.as_ref()) {
375            expressions.insert(description_str.clone(), Rc::from(description.deref()));
376        }
377
378        for input in ctx.node.inputs.iter() {
379            let Some(rule_value) = rule.get(input.id.deref()) else {
380                continue;
381            };
382            let Some(input_field) = &input.field else {
383                continue;
384            };
385
386            if let Some(reference) = isolate.get_reference(input_field.deref()) {
387                reference_map.insert(Rc::from(input_field.deref()), reference);
388            } else if let Some(reference) = isolate.run_standard(input_field.deref()).ok() {
389                reference_map.insert(Rc::from(input_field.deref()), reference);
390            }
391
392            let input_identifier = format!("{input_field}[{}]", &input.id);
393            expressions.insert(
394                Rc::from(input_identifier.as_str()),
395                Rc::from(rule_value.deref()),
396            );
397        }
398
399        (reference_map, expressions)
400    }
401
402    fn evaluate_row<'a>(
403        &self,
404        ctx: &'a DecisionTableContext,
405        rule: &'a HashMap<Arc<str>, Arc<str>>,
406        isolate: &mut Isolate,
407        pruned: Option<(&TableIndex, usize)>,
408    ) -> Option<RowResult> {
409        if !Self::row_matches(ctx, rule, isolate, pruned) {
410            return None;
411        }
412
413        let outputs = Variable::empty_object();
414        for output in ctx.node.outputs.iter() {
415            let (path, _) = output.write_path();
416            if path.is_empty() {
417                continue;
418            }
419            let Some(rule_value) = rule.get(&output.id) else {
420                continue;
421            };
422            if rule_value.is_empty() {
423                continue;
424            }
425
426            let res = isolate.run_standard(rule_value).ok()?;
427            outputs.dot_insert(path, res.deep_clone());
428        }
429
430        if !ctx.config.trace {
431            return Some(RowResult::Output(outputs));
432        }
433
434        let (reference_map, expressions) = Self::row_trace_parts(ctx, rule, isolate);
435        Some(RowResult::WithTrace {
436            output: outputs.to_variable(),
437            reference_map,
438            rule: expressions,
439        })
440    }
441}
442
443enum RowResult {
444    Output(Variable),
445    WithTrace {
446        output: Variable,
447        reference_map: HashMap<Rc<str>, Variable>,
448        rule: HashMap<Rc<str>, Rc<str>>,
449    },
450}
451
452#[derive(Debug, Clone, Serialize, ToVariable)]
453pub struct DecisionTableRowTrace {
454    index: usize,
455    reference_map: HashMap<Rc<str>, Variable>,
456    rule: HashMap<Rc<str>, Rc<str>>,
457}
458
459#[derive(Debug, Clone, Serialize, ToVariable)]
460#[serde(untagged)]
461pub enum DecisionTableNodeTrace {
462    FirstHit(DecisionTableRowTrace),
463    Collect(Vec<DecisionTableRowTrace>),
464}
465
466impl Default for DecisionTableNodeTrace {
467    fn default() -> Self {
468        DecisionTableNodeTrace::Collect(Default::default())
469    }
470}