diff --git a/bindings/nodejs/src/content.rs b/bindings/nodejs/src/content.rs index 6e0039e3..75fd2f30 100644 --- a/bindings/nodejs/src/content.rs +++ b/bindings/nodejs/src/content.rs @@ -24,7 +24,7 @@ impl ZenDecisionContent { } }; if let DecisionContent::Graph(g) = &mut decision_content { - g.compile(); + Arc::make_mut(g).compile(); } Ok(Self { diff --git a/bindings/python/src/content.rs b/bindings/python/src/content.rs index 0b08f0b6..9d6516fd 100644 --- a/bindings/python/src/content.rs +++ b/bindings/python/src/content.rs @@ -17,7 +17,7 @@ impl PyZenDecisionContent { let mut content: DecisionContent = serde_json::from_str(data).context("Failed to parse JSON")?; if let DecisionContent::Graph(g) = &mut content { - g.compile(); + Arc::make_mut(g).compile(); } Ok(Self(Arc::new(content))) } diff --git a/core/engine/src/decision.rs b/core/engine/src/decision.rs index 38cc17ec..ed0a4af6 100644 --- a/core/engine/src/decision.rs +++ b/core/engine/src/decision.rs @@ -4,7 +4,6 @@ use crate::loader::{DynamicLoader, NoopLoader}; use crate::model::GraphContent; use crate::nodes::custom::{DynamicCustomNode, NoopCustomNode}; use crate::nodes::function::http_handler::DynamicHttpHandler; -use crate::nodes::validator_cache::ValidatorCache; use crate::nodes::NodeHandlerExtensions; use crate::{DecisionGraphValidationError, EvaluationError}; use serde_json::Value; @@ -19,7 +18,6 @@ pub struct Decision { loader: DynamicLoader, adapter: DynamicCustomNode, http_handler: DynamicHttpHandler, - validator_cache: ValidatorCache, } impl From for Decision { @@ -29,7 +27,6 @@ impl From for Decision { loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), http_handler: None, - validator_cache: ValidatorCache::default(), } } } @@ -41,7 +38,6 @@ impl From> for Decision { loader: Arc::new(NoopLoader::default()), adapter: Arc::new(NoopCustomNode::default()), http_handler: None, - validator_cache: ValidatorCache::default(), } } } @@ -86,8 +82,9 @@ impl Decision { custom_node: self.adapter.clone(), http_handler: self.http_handler.clone(), compiled_cache: self.content.compiled_cache.clone(), + dt_indexes: self.content.dt_indexes.clone(), stripped_functions: self.content.stripped_functions.clone(), - validator_cache: Arc::new(OnceCell::from(self.validator_cache.clone())), + validator_cache: Arc::new(OnceCell::from(self.content.validator_cache.clone())), ..Default::default() }, })?; diff --git a/core/engine/src/engine.rs b/core/engine/src/engine.rs index 9d85018a..bfc19405 100644 --- a/core/engine/src/engine.rs +++ b/core/engine/src/engine.rs @@ -306,9 +306,9 @@ impl DecisionEngine { fn decision_from_graph_arc(&self, content: Arc) -> Decision { let graph: Arc = match Arc::try_unwrap(content) { - Ok(DecisionContent::Graph(g)) => Arc::new(g), + Ok(DecisionContent::Graph(g)) => g, Err(arc) => match arc.as_ref() { - DecisionContent::Graph(g) => Arc::new(g.clone()), + DecisionContent::Graph(g) => g.clone(), DecisionContent::Policy(_) => { panic!("decision_from_graph_arc called with Policy variant") } diff --git a/core/engine/src/loader/cached.rs b/core/engine/src/loader/cached.rs index 70f82f4c..1542ad6f 100644 --- a/core/engine/src/loader/cached.rs +++ b/core/engine/src/loader/cached.rs @@ -29,9 +29,9 @@ fn compiled(content: Arc) -> Arc { return content; } - let mut owned = graph.clone(); + let mut owned = (**graph).clone(); owned.compile(); - Arc::new(DecisionContent::Graph(owned)) + Arc::new(DecisionContent::Graph(Arc::new(owned))) } async fn prepared(loader: &DynamicLoader, content: Arc) -> Arc { @@ -42,10 +42,10 @@ async fn prepared(loader: &DynamicLoader, content: Arc) -> Arc< return content; } - let mut owned = graph.clone(); + let mut owned = (**graph).clone(); owned.compile(); let _ = owned.resolve_schemas(loader).await; - Arc::new(DecisionContent::Graph(owned)) + Arc::new(DecisionContent::Graph(Arc::new(owned))) } impl DecisionLoader for CachedLoader { diff --git a/core/engine/src/model/decision_content.rs b/core/engine/src/model/decision_content.rs index 76fbac62..b920ca2f 100644 --- a/core/engine/src/model/decision_content.rs +++ b/core/engine/src/model/decision_content.rs @@ -1,6 +1,8 @@ use crate::decision_graph::schema_dict; use crate::loader::DynamicLoader; +use crate::nodes::decision_table::index::TableIndex; use crate::nodes::function::v2::strip::TypeStripper; +use crate::nodes::validator_cache::ValidatorCache; use crate::policy::PolicyDocument; use ahash::{HashMap, HashMapExt}; use serde::{Deserialize, Deserializer, Serialize}; @@ -11,7 +13,7 @@ use zen_types::decision::{DecisionEdge, DecisionNode, DecisionNodeKind, Function #[derive(Clone, Debug, Serialize)] #[serde(untagged)] pub enum DecisionContent { - Graph(GraphContent), + Graph(Arc), Policy(PolicyContent), } @@ -28,7 +30,8 @@ impl<'de> Deserialize<'de> for DecisionContent { let content = if is_policy { serde_path_to_error::deserialize::<_, PolicyContent>(value).map(Self::Policy) } else { - serde_path_to_error::deserialize::<_, GraphContent>(value).map(Self::Graph) + serde_path_to_error::deserialize::<_, GraphContent>(value) + .map(|graph| Self::Graph(Arc::new(graph))) }; content.map_err(serde::de::Error::custom) @@ -37,7 +40,7 @@ impl<'de> Deserialize<'de> for DecisionContent { impl Default for DecisionContent { fn default() -> Self { - Self::Graph(GraphContent::default()) + Self::Graph(Arc::new(GraphContent::default())) } } @@ -64,19 +67,21 @@ impl DecisionContent { } pub fn into_graph_arc(self: Arc) -> Option> { - match Arc::try_unwrap(self) { - Ok(Self::Graph(g)) => Some(Arc::new(g)), - Ok(Self::Policy(_)) => None, - Err(arc) => match arc.as_ref() { - Self::Graph(g) => Some(Arc::new(g.clone())), - Self::Policy(_) => None, - }, + match self.as_ref() { + Self::Graph(g) => Some(g.clone()), + Self::Policy(_) => None, } } } impl From for DecisionContent { fn from(value: GraphContent) -> Self { + Self::Graph(Arc::new(value)) + } +} + +impl From> for DecisionContent { + fn from(value: Arc) -> Self { Self::Graph(value) } } @@ -110,6 +115,12 @@ pub struct GraphContent { #[serde(skip)] pub resolved_schemas: Option, (Arc, u64)>>>, + + #[serde(skip)] + pub(crate) validator_cache: ValidatorCache, + + #[serde(skip)] + pub(crate) dt_indexes: Option, TableIndex>>>, } #[derive(Clone, Debug, Deserialize, Serialize)] @@ -119,6 +130,7 @@ pub struct PolicyContent(pub Arc); impl GraphContent { pub fn compile(&mut self) { self.compile_functions(); + self.build_dt_indexes(); if self.compiled_cache.is_some() { return; } @@ -190,6 +202,26 @@ impl GraphContent { self.compiled_cache.replace(Arc::new(cache)); } + fn build_dt_indexes(&mut self) { + if self.dt_indexes.is_some() { + return; + } + + let indexes: HashMap, TableIndex> = self + .nodes + .iter() + .filter_map(|node| match &node.kind { + DecisionNodeKind::DecisionTableNode { content } => { + TableIndex::build(&content.inputs, &content.rules) + .map(|index| (node.id.clone(), index)) + } + _ => None, + }) + .collect(); + + self.dt_indexes = Some(Arc::new(indexes)); + } + pub async fn resolve_schemas(&mut self, loader: &DynamicLoader) -> Result<(), String> { if self.resolved_schemas.is_some() { return Ok(()); diff --git a/core/engine/src/nodes/decision/mod.rs b/core/engine/src/nodes/decision/mod.rs index dccdf04f..026a7c10 100644 --- a/core/engine/src/nodes/decision/mod.rs +++ b/core/engine/src/nodes/decision/mod.rs @@ -63,6 +63,9 @@ impl NodeHandler for DecisionNodeHandler { let mut extensions = ctx.extensions.clone(); extensions.compiled_cache = sub_graph.compiled_cache.clone(); + extensions.dt_indexes = sub_graph.dt_indexes.clone(); + extensions.validator_cache = + std::sync::Arc::new(std::cell::OnceCell::from(sub_graph.validator_cache.clone())); let dg = DecisionGraph::try_new(DecisionGraphConfig { content: sub_graph, diff --git a/core/engine/src/nodes/decision_table/index.rs b/core/engine/src/nodes/decision_table/index.rs new file mode 100644 index 00000000..9d3b4e9a --- /dev/null +++ b/core/engine/src/nodes/decision_table/index.rs @@ -0,0 +1,130 @@ +use ahash::HashMap; +use fixedbitset::FixedBitSet; +use rust_decimal::Decimal; +use std::sync::Arc; +use zen_expression::intellisense::{ArmTest, IntelliSense}; +use zen_types::decision::DecisionTableInputField; +use zen_types::variable::Variable; + +pub(crate) const MIN_INDEX_ROWS: usize = 8; + +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct TableIndex { + pub(crate) columns: Vec>, +} + +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct ColumnIndex { + strings: HashMap, FixedBitSet>, + numbers: HashMap, + bools: HashMap, + captured: FixedBitSet, + pub(crate) fallback: FixedBitSet, +} + +impl TableIndex { + pub(crate) fn build( + inputs: &[DecisionTableInputField], + rules: &[HashMap, Arc>], + ) -> Option { + let rows = rules.len(); + if rows < MIN_INDEX_ROWS { + return None; + } + let mut intellisense = IntelliSense::new(); + let columns: Vec> = inputs + .iter() + .map(|col| ColumnIndex::build(col, rules, rows, &mut intellisense)) + .collect(); + columns + .iter() + .any(Option::is_some) + .then_some(TableIndex { columns }) + } + + pub(crate) fn decides(&self, col_idx: usize, row_idx: usize) -> bool { + self.columns + .get(col_idx) + .and_then(Option::as_ref) + .is_some_and(|c| c.captured.contains(row_idx)) + } +} + +impl ColumnIndex { + fn build( + col: &DecisionTableInputField, + rules: &[HashMap, Arc>], + rows: usize, + intellisense: &mut IntelliSense, + ) -> Option { + if col.field.as_deref().is_none_or(|f| f.is_empty()) { + return None; + } + let mut strings: HashMap, FixedBitSet> = HashMap::default(); + let mut numbers: HashMap = HashMap::default(); + let mut bools: HashMap = HashMap::default(); + let mut captured = FixedBitSet::with_capacity(rows); + let mut fallback = FixedBitSet::with_capacity(rows); + + for (row_idx, rule) in rules.iter().enumerate() { + let Some(cell) = rule.get(&col.id).filter(|c| !c.is_empty()) else { + fallback.insert(row_idx); + continue; + }; + match intellisense.cell_test(cell) { + ArmTest::Enum { values, .. } => { + for value in values { + strings + .entry(Arc::from(value.as_ref())) + .or_insert_with(|| FixedBitSet::with_capacity(rows)) + .insert(row_idx); + } + captured.insert(row_idx); + } + ArmTest::Bool { values, .. } => { + for value in values { + bools + .entry(value) + .or_insert_with(|| FixedBitSet::with_capacity(rows)) + .insert(row_idx); + } + captured.insert(row_idx); + } + ArmTest::Number { cover, .. } => match cover.points() { + Some(points) => { + for point in points { + numbers + .entry(point.normalize()) + .or_insert_with(|| FixedBitSet::with_capacity(rows)) + .insert(row_idx); + } + captured.insert(row_idx); + } + None => { + fallback.insert(row_idx); + } + }, + ArmTest::Default | ArmTest::Unrecognized => { + fallback.insert(row_idx); + } + } + } + + (captured.count_ones(..) > 0).then_some(ColumnIndex { + strings, + numbers, + bools, + captured, + fallback, + }) + } + + pub(crate) fn rows_for(&self, value: &Variable) -> Option<&FixedBitSet> { + match value { + Variable::String(s) => self.strings.get(s.as_str()), + Variable::Number(n) => self.numbers.get(&n.normalize()), + Variable::Bool(b) => self.bools.get(b), + _ => None, + } + } +} diff --git a/core/engine/src/nodes/decision_table/mod.rs b/core/engine/src/nodes/decision_table/mod.rs index 2d0d330f..86b8475b 100644 --- a/core/engine/src/nodes/decision_table/mod.rs +++ b/core/engine/src/nodes/decision_table/mod.rs @@ -2,14 +2,20 @@ use crate::nodes::definition::NodeHandler; use crate::nodes::result::NodeResult; use crate::nodes::{NodeContext, NodeResponse}; use ahash::HashMap; +use fixedbitset::FixedBitSet; +use index::TableIndex; use serde::Serialize; use std::ops::Deref; use std::rc::Rc; use std::sync::Arc; use zen_expression::variable::ToVariable; use zen_expression::Isolate; -use zen_types::decision::{DecisionTableContent, DecisionTableHitPolicy, TransformAttributes}; +use zen_types::decision::{ + DecisionTableContent, DecisionTableHitPolicy, DecisionTableInputField, TransformAttributes, +}; use zen_types::variable::Variable; +pub(crate) mod index; + #[derive(Debug, Clone)] pub struct DecisionTableNodeHandler; @@ -41,8 +47,17 @@ impl DecisionTableNodeHandler { let mut isolate = ctx.isolate(); if !ctx.config.trace { - for rule in ctx.node.rules.iter() { - if let Some(RowResult::Output(output)) = self.evaluate_row(&ctx, rule, &mut isolate) + let index = Self::table_index(&ctx); + let candidates = + index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate)); + let pruner = candidates.as_ref().and(index); + for (row_idx, rule) in ctx.node.rules.iter().enumerate() { + if candidates.as_ref().is_some_and(|c| !c.contains(row_idx)) { + continue; + } + let pruned = pruner.map(|ix| (ix, row_idx)); + if let Some(RowResult::Output(output)) = + self.evaluate_row(&ctx, rule, &mut isolate, pruned) { return ctx.success(output); } @@ -54,7 +69,7 @@ impl DecisionTableNodeHandler { } let hit = ctx.node.rules.iter().enumerate().find_map(|(index, rule)| { - match self.evaluate_row(&ctx, rule, &mut isolate)? { + match self.evaluate_row(&ctx, rule, &mut isolate, None)? { RowResult::WithTrace { output, reference_map, @@ -89,8 +104,19 @@ impl DecisionTableNodeHandler { let mut traces = Vec::new(); let mut isolate = ctx.isolate(); + let table_index = (!ctx.config.trace) + .then(|| Self::table_index(&ctx)) + .flatten(); + let candidates = + table_index.and_then(|ix| Self::candidate_rows(ix, &ctx.node.inputs, &mut isolate)); + let pruner = candidates.as_ref().and(table_index); + for (index, rule) in ctx.node.rules.iter().enumerate() { - if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate) { + if candidates.as_ref().is_some_and(|c| !c.contains(index)) { + continue; + } + let pruned = pruner.map(|ix| (ix, index)); + if let Some(result) = self.evaluate_row(&ctx, rule, &mut isolate, pruned) { match result { RowResult::Output(output) => { outputs.push(output); @@ -144,14 +170,63 @@ impl DecisionTableNodeHandler { } } + fn table_index(ctx: &DecisionTableContext) -> Option<&TableIndex> { + ctx.extensions.dt_indexes.as_ref()?.get(&ctx.id) + } + + fn candidate_rows( + index: &TableIndex, + inputs: &[DecisionTableInputField], + isolate: &mut Isolate, + ) -> Option { + let mut acc: Option = None; + for (col_idx, column) in index.columns.iter().enumerate() { + let Some(column) = column else { + continue; + }; + let Some(field) = inputs[col_idx].field.as_ref().filter(|f| !f.is_empty()) else { + continue; + }; + isolate.set_reference(field).ok()?; + let value = isolate.get_reference(field)?; + if matches!(value, Variable::Dynamic(_)) { + return None; + } + let hit = column.rows_for(&value); + match &mut acc { + None => { + let mut first = column.fallback.clone(); + if let Some(hit) = hit { + first.union_with(hit); + } + acc = Some(first); + } + Some(acc) => { + let fallback = column.fallback.as_slice(); + let hit = hit.map(FixedBitSet::as_slice).unwrap_or_default(); + for (i, word) in acc.as_mut_slice().iter_mut().enumerate() { + let f = fallback.get(i).copied().unwrap_or(0); + let h = hit.get(i).copied().unwrap_or(0); + *word &= f | h; + } + } + } + } + acc + } + fn evaluate_row<'a>( &self, ctx: &'a DecisionTableContext, rule: &'a HashMap, Arc>, isolate: &mut Isolate, + pruned: Option<(&TableIndex, usize)>, ) -> Option { let content = &ctx.node; - for input in content.inputs.iter() { + for (col_idx, input) in content.inputs.iter().enumerate() { + if pruned.is_some_and(|(ix, row_idx)| ix.decides(col_idx, row_idx)) { + continue; + } let Some(rule_value) = rule.get(&input.id) else { continue; }; diff --git a/core/engine/src/nodes/extensions.rs b/core/engine/src/nodes/extensions.rs index 2a9784f7..4ee46328 100644 --- a/core/engine/src/nodes/extensions.rs +++ b/core/engine/src/nodes/extensions.rs @@ -1,5 +1,6 @@ use crate::loader::{DynamicLoader, NoopLoader}; use crate::nodes::custom::{DynamicCustomNode, NoopCustomNode}; +use crate::nodes::decision_table::index::TableIndex; use crate::nodes::function::http_handler::DynamicHttpHandler; use crate::nodes::function::v2::function::{Function, FunctionConfig}; use crate::nodes::function::v2::module::console::ConsoleListener; @@ -21,6 +22,7 @@ pub struct NodeHandlerExtensions { pub(crate) http_handler: DynamicHttpHandler, pub(crate) compiled_cache: Option>, pub(crate) stripped_functions: Option, Arc>>>, + pub(crate) dt_indexes: Option, TableIndex>>>, } impl Default for NodeHandlerExtensions { @@ -33,6 +35,7 @@ impl Default for NodeHandlerExtensions { custom_node: Arc::new(NoopCustomNode::default()), compiled_cache: None, stripped_functions: None, + dt_indexes: None, http_handler: None, } } diff --git a/core/engine/src/nodes/validator_cache.rs b/core/engine/src/nodes/validator_cache.rs index c9e047a8..2b061882 100644 --- a/core/engine/src/nodes/validator_cache.rs +++ b/core/engine/src/nodes/validator_cache.rs @@ -10,6 +10,12 @@ pub struct ValidatorCache { inner: Arc>>>>, } +impl PartialEq for ValidatorCache { + fn eq(&self, _: &Self) -> bool { + true + } +} + impl ValidatorCache { pub fn get(&self, key: u64) -> Option>> { let read = self.inner.read().ok()?; diff --git a/core/engine/src/policy/blocks/decision_table.rs b/core/engine/src/policy/blocks/decision_table.rs index a74bc6c6..8b3506c4 100644 --- a/core/engine/src/policy/blocks/decision_table.rs +++ b/core/engine/src/policy/blocks/decision_table.rs @@ -2,7 +2,6 @@ use std::sync::{Arc, OnceLock}; use ahash::{HashMap, HashSet}; use fixedbitset::FixedBitSet; -use rust_decimal::Decimal; use serde::{Deserialize, Serialize}; use zen_expression::intellisense::{ArmTest, IntelliSense, NumberCover}; use zen_expression::variable::{Variable, VariableType}; @@ -26,6 +25,7 @@ use super::{ Block, BlockKind, BlockReadPlan, CellReads, ConditionalReads, ExpressionLocation, ParseContext, ReadFlattenFn, WriteSite, WriteTarget, }; +use crate::nodes::decision_table::index::TableIndex; pub(crate) struct TableSelection { pub(crate) matched_rows: Vec, @@ -101,127 +101,6 @@ pub struct DecisionTableIr { index: OnceLock>, } -const MIN_INDEX_ROWS: usize = 8; - -#[derive(Debug, Clone)] -struct TableIndex { - columns: Vec>, -} - -#[derive(Debug, Clone)] -struct ColumnIndex { - strings: HashMap, FixedBitSet>, - numbers: HashMap, - bools: HashMap, - captured: FixedBitSet, - fallback: FixedBitSet, -} - -impl TableIndex { - fn build(table: &DecisionTableIr) -> Option { - let rows = table.rules.len(); - if rows < MIN_INDEX_ROWS { - return None; - } - let mut intellisense = IntelliSense::new(); - let columns: Vec> = table - .inputs - .iter() - .map(|col| ColumnIndex::build(table, col, rows, &mut intellisense)) - .collect(); - columns - .iter() - .any(Option::is_some) - .then_some(TableIndex { columns }) - } - - fn decides(&self, col_idx: usize, row_idx: usize) -> bool { - self.columns - .get(col_idx) - .and_then(Option::as_ref) - .is_some_and(|c| c.captured.contains(row_idx)) - } -} - -impl ColumnIndex { - fn build( - table: &DecisionTableIr, - col: &DecisionTableInputField, - rows: usize, - intellisense: &mut IntelliSense, - ) -> Option { - if col.field.as_deref().is_none_or(|f| f.is_empty()) { - return None; - } - let mut strings: HashMap, FixedBitSet> = HashMap::default(); - let mut numbers: HashMap = HashMap::default(); - let mut bools: HashMap = HashMap::default(); - let mut captured = FixedBitSet::with_capacity(rows); - let mut fallback = FixedBitSet::with_capacity(rows); - - for (row_idx, rule) in table.rules.iter().enumerate() { - let Some(cell) = rule.get(&col.id).filter(|c| !c.is_empty()) else { - fallback.insert(row_idx); - continue; - }; - match intellisense.cell_test(cell) { - ArmTest::Enum { values, .. } => { - for value in values { - strings - .entry(Arc::from(value.as_ref())) - .or_insert_with(|| FixedBitSet::with_capacity(rows)) - .insert(row_idx); - } - captured.insert(row_idx); - } - ArmTest::Bool { values, .. } => { - for value in values { - bools - .entry(value) - .or_insert_with(|| FixedBitSet::with_capacity(rows)) - .insert(row_idx); - } - captured.insert(row_idx); - } - ArmTest::Number { cover, .. } => match cover.points() { - Some(points) => { - for point in points { - numbers - .entry(point.normalize()) - .or_insert_with(|| FixedBitSet::with_capacity(rows)) - .insert(row_idx); - } - captured.insert(row_idx); - } - None => { - fallback.insert(row_idx); - } - }, - ArmTest::Default | ArmTest::Unrecognized => { - fallback.insert(row_idx); - } - } - } - - (captured.count_ones(..) > 0).then_some(ColumnIndex { - strings, - numbers, - bools, - captured, - fallback, - }) - } - - fn rows_for(&self, value: &Variable) -> Option<&FixedBitSet> { - match value { - Variable::String(s) => self.strings.get(s.as_str()), - Variable::Number(n) => self.numbers.get(&n.normalize()), - Variable::Bool(b) => self.bools.get(b), - _ => None, - } - } -} - #[derive(Debug, Clone)] pub struct OutputColumn { pub id: Arc, @@ -1173,7 +1052,9 @@ impl DecisionTableIr { } fn table_index(&self) -> Option<&TableIndex> { - self.index.get_or_init(|| TableIndex::build(self)).as_ref() + self.index + .get_or_init(|| TableIndex::build(&self.inputs, &self.rules)) + .as_ref() } fn candidate_rows( diff --git a/core/engine/src/policy/runtime.rs b/core/engine/src/policy/runtime.rs index 945f745a..cbc62be7 100644 --- a/core/engine/src/policy/runtime.rs +++ b/core/engine/src/policy/runtime.rs @@ -144,21 +144,19 @@ impl CompiledSet { workspace.set_policy_arc(key.clone(), policy.0.clone()); policy_keys.push(key.clone()); } - DecisionContent::Graph(graph) => { - match Decision::from(Arc::new(graph.clone())).validate() { - Err(error) => failures.push(CompileFailure { - key: key.clone(), - kind: "graph", - diagnostics: Vec::new(), - error: Some(error.to_string()), - }), - Ok(()) => { - let mut compiled = graph.clone(); - compiled.compile(); - entries.insert(key.clone(), CompiledEntry::Graph(Arc::new(compiled))); - } + DecisionContent::Graph(graph) => match Decision::from(graph.clone()).validate() { + Err(error) => failures.push(CompileFailure { + key: key.clone(), + kind: "graph", + diagnostics: Vec::new(), + error: Some(error.to_string()), + }), + Ok(()) => { + let mut compiled = graph.clone(); + Arc::make_mut(&mut compiled).compile(); + entries.insert(key.clone(), CompiledEntry::Graph(compiled)); } - } + }, } } diff --git a/core/engine/tests/engine.rs b/core/engine/tests/engine.rs index e9d29104..bceae4c5 100644 --- a/core/engine/tests/engine.rs +++ b/core/engine/tests/engine.rs @@ -209,13 +209,11 @@ async fn engine_function_imports() { }) .collect::>(); - let function_content = GraphContent { - edges: function_content.edges, - nodes: new_nodes, - imports: Vec::new(), - compiled_cache: None, - stripped_functions: None, - resolved_schemas: None, + let function_content = { + let mut content = GraphContent::default(); + content.edges = function_content.edges; + content.nodes = new_nodes; + content }; let decision = DecisionEngine::default() .create_decision(Arc::new(function_content.into())) @@ -498,3 +496,89 @@ async fn test_nodes_reference() { }) ); } + +#[tokio::test] +async fn decision_table_index_matches_linear_semantics() { + let table_node = json!({ + "id": "dt-node", + "type": "decisionTableNode", + "name": "dt", + "content": { + "hitPolicy": "first", + "inputs": [ + {"id": "c1", "name": "Tier", "field": "tier"}, + {"id": "c2", "name": "Amount", "field": "amount"}, + {"id": "c3", "name": "Active", "field": "active"} + ], + "outputs": [{"id": "o1", "name": "Rate", "field": "rate"}], + "rules": [ + {"_id": "r1", "c1": "'gold'", "c2": "100", "c3": "true", "o1": "1"}, + {"_id": "r2", "c1": "'gold'", "c2": "> 500", "c3": "", "o1": "2"}, + {"_id": "r3", "c1": "'silver'", "c2": "[100..200]", "c3": "", "o1": "3"}, + {"_id": "r4", "c1": "'silver', 'bronze'", "c2": "", "c3": "false", "o1": "4"}, + {"_id": "r5", "c1": "", "c2": "42", "c3": "", "o1": "5"}, + {"_id": "r6", "c1": "'gold'", "c2": "", "c3": "", "o1": "6"}, + {"_id": "r7", "c1": "'bronze'", "c2": "7, 8, 9", "c3": "true", "o1": "7"}, + {"_id": "r8", "c1": "", "c2": "", "c3": "", "o1": "8"}, + {"_id": "r9", "c1": "'platinum'", "c2": "1000", "c3": "true", "o1": "9"} + ] + } + }); + let graph = json!({ + "nodes": [ + {"id": "in", "type": "inputNode", "name": "request"}, + table_node, + {"id": "out", "type": "outputNode", "name": "response"} + ], + "edges": [ + {"id": "e1", "sourceId": "in", "targetId": "dt-node"}, + {"id": "e2", "sourceId": "dt-node", "targetId": "out"} + ] + }); + + let probes = [ + json!({"tier": "gold", "amount": 100, "active": true}), + json!({"tier": "gold", "amount": 600, "active": false}), + json!({"tier": "silver", "amount": 150, "active": true}), + json!({"tier": "silver", "amount": 50, "active": false}), + json!({"tier": "bronze", "amount": 8, "active": true}), + json!({"tier": "unknown", "amount": 42, "active": false}), + json!({"tier": "unknown", "amount": 0, "active": false}), + json!({"tier": "platinum", "amount": 1000, "active": true}), + json!({"tier": 5, "amount": "x", "active": null}), + ]; + + for hit_policy in ["first", "collect"] { + let mut graph = graph.clone(); + graph["nodes"][1]["content"]["hitPolicy"] = json!(hit_policy); + let mut content: GraphContent = serde_json::from_value(graph).unwrap(); + content.compile(); + let decision = DecisionEngine::default() + .create_decision(Arc::new(content.into())) + .unwrap(); + + for probe in &probes { + let indexed = decision + .evaluate(Variable::from(probe)) + .await + .unwrap() + .result; + let linear = decision + .evaluate_with_opts( + Variable::from(probe), + EvaluationOptions { + trace: true, + max_depth: 5, + }, + ) + .await + .unwrap() + .result; + assert_eq!( + serde_json::to_value(&indexed).unwrap(), + serde_json::to_value(&linear).unwrap(), + "hit_policy={hit_policy} probe={probe}" + ); + } + } +} diff --git a/core/engine/tests/support/mod.rs b/core/engine/tests/support/mod.rs index 85a6a552..ec652b7c 100644 --- a/core/engine/tests/support/mod.rs +++ b/core/engine/tests/support/mod.rs @@ -26,7 +26,7 @@ pub fn load_raw_test_data(key: &str) -> BufReader { pub fn load_test_data(key: &str) -> GraphContent { let content: DecisionContent = serde_json::from_reader(load_raw_test_data(key)).unwrap(); match content { - DecisionContent::Graph(g) => g, + DecisionContent::Graph(g) => (*g).clone(), DecisionContent::Policy(_) => { panic!("expected graph test fixture, got policy: {key}") }