zen_engine/nodes/decision_table/
mod.rs1use 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}