From 067ae5544df552751bbf3cd12cffda556ef493e9 Mon Sep 17 00:00:00 2001 From: Ivan Miletic Date: Thu, 1 Oct 2026 21:12:34 +0200 Subject: [PATCH] fix: cleanup --- core/engine/src/analysis/mod.rs | 1 + core/engine/src/analysis/nullable.rs | 490 +++++++++--------- core/engine/src/analysis/proof.rs | 277 ++++++++++ core/engine/src/analysis/table/cell.rs | 198 ++++--- core/engine/src/analysis/table/index.rs | 7 +- core/engine/src/analysis/table/merge.rs | 11 +- core/engine/src/analysis/table/partition.rs | 12 +- core/engine/src/analysis/table/print.rs | 71 ++- core/engine/src/analysis/table/value_set.rs | 65 ++- core/engine/src/analysis/table/witness.rs | 8 +- core/engine/src/policy/blocks/context.rs | 51 +- .../src/policy/blocks/decision_table.rs | 33 +- core/engine/src/policy/blocks/mod.rs | 2 +- core/engine/src/policy/linter/mod.rs | 33 +- .../policy/linter/redundant_parentheses.rs | 126 +---- core/engine/src/policy/queries/dependency.rs | 6 +- core/engine/src/policy/queries/diagnostics.rs | 107 +++- core/engine/src/workspace/graph/analysis.rs | 33 +- core/engine/tests/table_fix_regressions.rs | 433 ++++++++++++++++ core/engine/tests/table_verification.rs | 267 ++++++++++ 20 files changed, 1690 insertions(+), 541 deletions(-) create mode 100644 core/engine/src/analysis/proof.rs create mode 100644 core/engine/tests/table_fix_regressions.rs diff --git a/core/engine/src/analysis/mod.rs b/core/engine/src/analysis/mod.rs index 5bc55d36..fa0951bd 100644 --- a/core/engine/src/analysis/mod.rs +++ b/core/engine/src/analysis/mod.rs @@ -1,2 +1,3 @@ pub(crate) mod nullable; +pub(crate) mod proof; pub(crate) mod table; diff --git a/core/engine/src/analysis/nullable.rs b/core/engine/src/analysis/nullable.rs index 0a4776a6..0fa9d2b0 100644 --- a/core/engine/src/analysis/nullable.rs +++ b/core/engine/src/analysis/nullable.rs @@ -1,22 +1,16 @@ use std::cell::RefCell; +use ahash::{HashMap, HashMapExt, HashSet}; use zen_expression::intellisense::IntelliSense; use zen_expression::lexer::{ArithmeticOperator, ComparisonOperator, LogicalOperator, Operator}; use zen_expression::parser::Node; -use crate::policy::linter::{AstOps, RedundantParentheses}; +use crate::analysis::proof::{FixEdit, FixProof}; +use crate::policy::linter::AstOps; use crate::workspace::types::{Diagnostic, DiagnosticCode, Span}; pub(crate) struct NullableOperand; -struct FallbackEdit { - span: Span, - range: Span, - kept: String, - dropped: String, - keep_left: bool, -} - #[derive(Default)] struct Fallback { operands: Option<(Span, Span)>, @@ -29,252 +23,228 @@ struct Found { path: Option, } +struct Candidate { + idx: usize, + span: Span, + kept: Span, + dropped: Span, + keep_left: bool, + wrapper: Option, +} + +impl Candidate { + fn edit(&self, outer: Span) -> FixEdit { + FixEdit::unwrap(outer, self.kept, Some((self.span, self.keep_left))) + } +} + impl NullableOperand { pub(crate) fn annotate( - diagnostic: &mut Diagnostic, - is: &mut IntelliSense, - source: &str, - unary: bool, - ) { - if diagnostic.code == DiagnosticCode::RedundantNullish { - Self::fallback_fix(diagnostic, is, source, unary); - return; - } - if diagnostic.code != DiagnosticCode::TypeMismatch { - return; - } - let Some(span) = diagnostic.location.span else { - return; - }; - let Some((operator, left, right)) = Self::parse_message(&diagnostic.message) else { - return; - }; - let (left_nullable, right_nullable) = (left.ends_with('?'), right.ends_with('?')); - if left_nullable == right_nullable - || left.trim_end_matches('?') != "number" - || right.trim_end_matches('?') != "number" - { - return; - } - let Some(found) = Self::locate(is, source, unary, span, left_nullable) else { - return; - }; - if let Some(path) = found.path.filter(|p| !p.starts_with('$')) { - diagnostic.args.insert("nullablePath", path); - } - let defaultable = match operator.as_str() { - "+" | "-" | "*" | ">" | "<" | ">=" | "<=" => true, - "/" | "%" => found.left, - _ => false, - }; - if !defaultable { - return; - } - let operand: String = source - .chars() - .skip(found.operand.0 as usize) - .take((found.operand.1 - found.operand.0) as usize) - .collect(); - let replacement = format!("({operand} ?? 0)"); - let prefix: String = source.chars().take(found.operand.0 as usize).collect(); - let suffix: String = source.chars().skip(found.operand.1 as usize).collect(); - diagnostic - .args - .insert("fixSource", format!("{prefix}{replacement}{suffix}")); - diagnostic.args.insert("fixOriginal", source.to_string()); - diagnostic.args.insert("fixOperand", operand); - } - - fn fallback_fix(diagnostic: &mut Diagnostic, is: &mut IntelliSense, source: &str, unary: bool) { - let Some(edit) = Self::fallback_edit(diagnostic, is, source, unary) else { - return; - }; - diagnostic.args.insert("fixOriginal", source.to_string()); - diagnostic.args.insert( - "fixSource", - Self::splice(source, &[(edit.range, edit.kept.clone())]), - ); - diagnostic.args.insert( - "fixKeep", - if edit.keep_left { "left" } else { "right" }.to_string(), - ); - diagnostic.args.insert( - "fixFallback", - if edit.keep_left { - edit.dropped - } else { - edit.kept - }, - ); - } - - pub(crate) fn fallback_all( diagnostics: &mut [Diagnostic], is: &mut IntelliSense, source: &str, unary: bool, ) { - let mut edits: Vec<(usize, FallbackEdit)> = diagnostics - .iter() - .enumerate() - .filter(|(_, d)| { - d.code == DiagnosticCode::RedundantNullish && d.args.contains_key("fixSource") - }) - .filter_map(|(idx, d)| Some((idx, Self::fallback_edit(d, is, source, unary)?))) - .collect(); - if edits.len() < 2 { - return; - } - edits.sort_by_key(|(_, edit)| edit.range.0); - if edits - .windows(2) - .any(|pair| pair[0].1.range.1 > pair[1].1.range.0) - { - return; - } - let replacements: Vec<(Span, String)> = edits - .iter() - .map(|(_, edit)| (edit.range, edit.kept.clone())) - .collect(); - let combined = Self::splice(source, &replacements); - let targets: Vec<(Span, bool)> = edits - .iter() - .map(|(_, edit)| (edit.span, edit.keep_left)) - .collect(); - let expected = is.with_ast(source, unary, |root, metadata| { - let swaps: RefCell> = RefCell::new(Vec::new()); - root.walk(|node| { - let Node::Binary { - left, - operator: Operator::Logical(LogicalOperator::NullishCoalescing), - right, - } = node - else { - return; - }; - let Some(span) = AstOps::span(metadata, node) else { - return; - }; - if let Some((_, keep_left)) = targets.iter().find(|(target, _)| *target == span) { - let kept = if *keep_left { *left } else { *right }; - swaps - .borrow_mut() - .push((format!("{node:?}"), format!("{kept:?}"))); - } - }); - let swaps = swaps.into_inner(); - if swaps.len() != targets.len() { - return None; - } - let mut debug = format!("{root:?}"); - for (from, to) in swaps { - debug = debug.replace(&from, &to); - } - Some(RedundantParentheses::tree_shape(&debug)) - }); - let actual = is.with_ast(&combined, unary, |root, _| { - RedundantParentheses::tree_shape(&format!("{root:?}")) - }); - match (expected.flatten(), actual) { - (Some(expected), Some(actual)) if expected == actual => {} - _ => return, - } - for (idx, _) in edits { - diagnostics[idx].args.insert("fixAll", combined.clone()); - } + Self::default_operands(diagnostics, is, source, unary); + Self::fallbacks(diagnostics, is, source, unary); } - fn splice(source: &str, replacements: &[(Span, String)]) -> String { - let chars: Vec = source.chars().collect(); - let mut out = String::with_capacity(source.len()); - let mut cursor = 0usize; - let mut sorted: Vec<&(Span, String)> = replacements.iter().collect(); - sorted.sort_by_key(|(range, _)| range.0); - for (range, with) in sorted { - let (start, end) = (range.0 as usize, range.1 as usize); - out.extend(&chars[cursor.min(chars.len())..start.min(chars.len())]); - out.push_str(with); - cursor = end; - } - out.extend(&chars[cursor.min(chars.len())..]); - out - } - - fn fallback_edit( - diagnostic: &Diagnostic, + fn default_operands( + diagnostics: &mut [Diagnostic], is: &mut IntelliSense, source: &str, unary: bool, - ) -> Option { - let span = diagnostic.location.span?; - let keep_left = if diagnostic.message.contains("is never null") { - true - } else if diagnostic.message.contains("is always null") { - false - } else { - return None; - }; - let located = is.with_ast(source, unary, |root, metadata| { - let found: RefCell = RefCell::new(Fallback::default()); - root.walk(|node| match node { - Node::Binary { - left, - operator: Operator::Logical(LogicalOperator::NullishCoalescing), - right, - } if AstOps::span(metadata, node) == Some(span) => { - let (kept, dropped) = if keep_left { - (*left, *right) - } else { - (*right, *left) - }; - if let (Some(kept), Some(dropped)) = ( - AstOps::span(metadata, kept), - AstOps::span(metadata, dropped), - ) { - found.borrow_mut().operands = Some((kept, dropped)); + ) { + let requests: Vec<(usize, Span, String, bool)> = diagnostics + .iter() + .enumerate() + .filter(|(_, d)| d.code == DiagnosticCode::TypeMismatch) + .filter_map(|(idx, d)| { + let span = d.location.span?; + let (operator, left, right) = Self::parse_message(&d.message)?; + let (left_nullable, right_nullable) = (left.ends_with('?'), right.ends_with('?')); + (left_nullable != right_nullable + && left.trim_end_matches('?') == "number" + && right.trim_end_matches('?') == "number") + .then_some((idx, span, operator, left_nullable)) + }) + .collect(); + if requests.is_empty() { + return; + } + let spans: HashSet = requests.iter().map(|(_, span, _, _)| *span).collect(); + let operands = Self::locate(is, source, unary, &spans); + for (idx, span, operator, left_nullable) in requests { + let Some((left, right)) = operands.get(&span) else { + continue; + }; + let found = if left_nullable { left } else { right }; + let diagnostic = &mut diagnostics[idx]; + if let Some(path) = found.path.as_ref().filter(|p| !p.starts_with('$')) { + diagnostic.args.insert("nullablePath", path.clone()); + } + let defaultable = match operator.as_str() { + "+" | "-" | "*" | ">" | "<" | ">=" | "<=" => true, + "/" | "%" => found.left, + _ => false, + }; + if !defaultable { + continue; + } + let Some(operand) = AstOps::text(source, found.operand) else { + continue; + }; + let replacement = format!("({operand} ?? 0)"); + let Some(fixed) = AstOps::splice(source, &[(found.operand, replacement.as_str())]) + else { + continue; + }; + diagnostic.args.insert("fixSource", fixed); + diagnostic.args.insert("fixOriginal", source.to_string()); + diagnostic.args.insert("fixOperand", operand.to_string()); + } + } + + fn fallbacks(diagnostics: &mut [Diagnostic], is: &mut IntelliSense, source: &str, unary: bool) { + let candidates = Self::candidates(diagnostics, is, source, unary); + if candidates.is_empty() { + return; + } + let preferred: Vec = candidates + .iter() + .map(|candidate| candidate.edit(candidate.wrapper.unwrap_or(candidate.span))) + .collect(); + let mut edits: Vec> = FixProof::proven(is, source, unary, &preferred) + .into_iter() + .zip(preferred) + .map(|(proven, edit)| proven.then_some(edit)) + .collect(); + let retry: Vec = (0..candidates.len()) + .filter(|&i| edits[i].is_none() && candidates[i].wrapper.is_some()) + .collect(); + let alternatives: Vec = retry + .iter() + .map(|&i| candidates[i].edit(candidates[i].span)) + .collect(); + for ((i, proven), edit) in retry + .into_iter() + .zip(FixProof::proven(is, source, unary, &alternatives)) + .zip(alternatives) + { + if proven { + edits[i] = Some(edit); + } + } + let accepted: Vec<&FixEdit> = edits.iter().flatten().collect(); + let fix_all = (accepted.len() > 1) + .then(|| FixProof::holds(is, source, unary, &accepted)) + .flatten(); + for (candidate, edit) in candidates.iter().zip(&edits) { + let Some(fixed) = edit.as_ref().and_then(|edit| edit.apply(source)) else { + continue; + }; + let (Some(kept), Some(dropped)) = ( + AstOps::text(source, candidate.kept), + AstOps::text(source, candidate.dropped), + ) else { + continue; + }; + let args = &mut diagnostics[candidate.idx].args; + args.insert("fixOriginal", source.to_string()); + args.insert("fixSource", fixed); + args.insert( + "fixKeep", + if candidate.keep_left { "left" } else { "right" }.to_string(), + ); + args.insert( + "fixFallback", + if candidate.keep_left { dropped } else { kept }.to_string(), + ); + if let Some(all) = &fix_all { + args.insert("fixAll", all.clone()); + } + } + } + + fn candidates( + diagnostics: &[Diagnostic], + is: &mut IntelliSense, + source: &str, + unary: bool, + ) -> Vec { + let targets: Vec<(usize, Span, bool)> = diagnostics + .iter() + .enumerate() + .filter(|(_, d)| d.code == DiagnosticCode::RedundantNullish) + .filter_map(|(idx, d)| { + let keep_left = if d.message.contains("is never null") { + true + } else if d.message.contains("is always null") { + false + } else { + return None; + }; + Some((idx, d.location.span?, keep_left)) + }) + .collect(); + if targets.is_empty() { + return Vec::new(); + } + let spans: HashSet = targets.iter().map(|(_, span, _)| *span).collect(); + let located = is + .with_ast(source, unary, |root, metadata| { + let found: RefCell> = RefCell::new(HashMap::new()); + root.walk(|node| match node { + Node::Binary { + left, + operator: Operator::Logical(LogicalOperator::NullishCoalescing), + right, + } => { + let Some(span) = + AstOps::span(metadata, node).filter(|span| spans.contains(span)) + else { + return; + }; + if let (Some(left), Some(right)) = + (AstOps::span(metadata, left), AstOps::span(metadata, right)) + { + found.borrow_mut().entry(span).or_default().operands = + Some((left, right)); + } } - } - Node::Parenthesized(inner) if AstOps::span(metadata, inner) == Some(span) => { - found.borrow_mut().wrapper = AstOps::span(metadata, node); - } - _ => {} - }); - found.into_inner() - })?; - let Fallback { - operands: Some((kept, dropped)), - wrapper, - } = located - else { - return None; - }; - let text = |range: Span| -> String { - source - .chars() - .skip(range.0 as usize) - .take((range.1 - range.0) as usize) - .collect() - }; - let kept_text = text(kept); - let mut shape = |candidate: &str| { - is.with_ast(candidate, unary, |root, _| { - RedundantParentheses::tree_shape(&format!("{root:?}")) + Node::Parenthesized(inner) => { + if let Some(span) = + AstOps::span(metadata, inner).filter(|span| spans.contains(span)) + { + found.borrow_mut().entry(span).or_default().wrapper = + AstOps::span(metadata, node); + } + } + _ => {} + }); + found.into_inner() }) - }; - let plain_shape = shape(&Self::splice(source, &[(span, kept_text.clone())]))?; - let range = wrapper - .filter(|wrapper| { - shape(&Self::splice(source, &[(*wrapper, kept_text.clone())])).as_deref() - == Some(plain_shape.as_str()) + .unwrap_or_default(); + targets + .into_iter() + .filter_map(|(idx, span, keep_left)| { + let fallback = located.get(&span)?; + let (left, right) = fallback.operands?; + let (kept, dropped) = if keep_left { + (left, right) + } else { + (right, left) + }; + Some(Candidate { + idx, + span, + kept, + dropped, + keep_left, + wrapper: fallback.wrapper, + }) }) - .unwrap_or(span); - Some(FallbackEdit { - span, - range, - kept: kept_text, - dropped: text(dropped), - keep_left, - }) + .collect() } fn parse_message(message: &str) -> Option<(String, String, String)> { @@ -292,11 +262,10 @@ impl NullableOperand { is: &mut IntelliSense, source: &str, unary: bool, - span: Span, - left_nullable: bool, - ) -> Option { + spans: &HashSet, + ) -> HashMap { is.with_ast(source, unary, |root, metadata| { - let found: RefCell> = RefCell::new(None); + let found: RefCell> = RefCell::new(HashMap::new()); root.walk(|node| { let Node::Binary { left, @@ -321,22 +290,37 @@ impl NullableOperand { | ComparisonOperator::GreaterThanOrEqual ) ); - if !numeric || AstOps::span(metadata, node) != Some(span) { + if !numeric { return; } - let operand = if left_nullable { *left } else { *right }; - let Some(operand_span) = AstOps::span(metadata, operand) else { + let Some(span) = AstOps::span(metadata, node).filter(|span| spans.contains(span)) + else { return; }; - found.replace(Some(Found { - operand: operand_span, - left: left_nullable, - path: Self::path(operand), - })); + let (Some(left_span), Some(right_span)) = + (AstOps::span(metadata, left), AstOps::span(metadata, right)) + else { + return; + }; + found.borrow_mut().insert( + span, + ( + Found { + operand: left_span, + left: true, + path: Self::path(left), + }, + Found { + operand: right_span, + left: false, + path: Self::path(right), + }, + ), + ); }); found.into_inner() }) - .flatten() + .unwrap_or_default() } fn path(node: &Node) -> Option { diff --git a/core/engine/src/analysis/proof.rs b/core/engine/src/analysis/proof.rs new file mode 100644 index 00000000..dbf69851 --- /dev/null +++ b/core/engine/src/analysis/proof.rs @@ -0,0 +1,277 @@ +use std::collections::BTreeMap; +use std::fmt::Write; + +use ahash::{HashMap, HashMapExt, HashSet}; +use zen_expression::intellisense::{AstMetadata, IntelliSense}; +use zen_expression::lexer::{LogicalOperator, Operator}; +use zen_expression::parser::Node; + +use crate::policy::linter::AstOps; +use crate::workspace::types::Span; + +#[derive(Clone)] +pub(crate) struct FixEdit { + pub(crate) deletions: Vec, + pub(crate) kept: Span, + pub(crate) swap: Option<(Span, bool)>, +} + +#[derive(Default)] +struct Layer { + members: Vec, + taken: BTreeMap, + targets: HashSet, + kept: HashSet, +} + +impl Layer { + fn admits(&self, edit: &FixEdit) -> bool { + let free = edit.deletions.iter().all(|(start, end)| { + self.taken + .range(..*end) + .next_back() + .is_none_or(|(_, taken_end)| taken_end <= start) + }); + free && edit.swap.is_none_or(|(target, _)| { + !self.kept.contains(&target) && !self.targets.contains(&edit.kept) + }) + } + + fn insert(&mut self, idx: usize, edit: &FixEdit) { + self.members.push(idx); + self.taken.extend(edit.deletions.iter().copied()); + if let Some((target, _)) = edit.swap { + self.targets.insert(target); + self.kept.insert(edit.kept); + } + } +} + +impl FixEdit { + pub(crate) fn unwrap(outer: Span, kept: Span, swap: Option<(Span, bool)>) -> Self { + Self { + deletions: [(outer.0, kept.0), (kept.1, outer.1)] + .into_iter() + .filter(|(start, end)| start < end) + .collect(), + kept, + swap, + } + } + + pub(crate) fn apply(&self, source: &str) -> Option { + FixProof::splice(source, &[self]) + } +} + +pub(crate) struct FixProof; + +impl FixProof { + pub(crate) fn proven( + is: &mut IntelliSense, + source: &str, + unary: bool, + edits: &[FixEdit], + ) -> Vec { + let mut proven = vec![false; edits.len()]; + let mut layers: Vec = Vec::new(); + for (idx, edit) in edits.iter().enumerate() { + match layers.iter_mut().find(|layer| layer.admits(edit)) { + Some(layer) => layer.insert(idx, edit), + None => { + let mut layer = Layer::default(); + layer.insert(idx, edit); + layers.push(layer); + } + } + } + for layer in layers { + Self::bisect(is, source, unary, edits, &layer.members, &mut proven); + } + proven + } + + pub(crate) fn holds( + is: &mut IntelliSense, + source: &str, + unary: bool, + edits: &[&FixEdit], + ) -> Option { + let fixed = Self::splice(source, edits)?; + let swaps: HashMap = edits.iter().filter_map(|edit| edit.swap).collect(); + let expected = is + .with_ast(source, unary, |root, metadata| { + Shape::of(root, metadata, &swaps) + }) + .flatten()?; + let actual = is + .with_ast(&fixed, unary, |root, metadata| { + Shape::of(root, metadata, &HashMap::new()) + }) + .flatten()?; + (expected == actual).then_some(fixed) + } + + fn bisect( + is: &mut IntelliSense, + source: &str, + unary: bool, + edits: &[FixEdit], + members: &[usize], + proven: &mut [bool], + ) { + if members.is_empty() { + return; + } + let batch: Vec<&FixEdit> = members.iter().map(|&idx| &edits[idx]).collect(); + if Self::holds(is, source, unary, &batch).is_some() { + members.iter().for_each(|&idx| proven[idx] = true); + return; + } + if members.len() == 1 { + return; + } + let (left, right) = members.split_at(members.len() / 2); + Self::bisect(is, source, unary, edits, left, proven); + Self::bisect(is, source, unary, edits, right, proven); + } + + fn splice(source: &str, edits: &[&FixEdit]) -> Option { + let deletions: Vec<(Span, &str)> = edits + .iter() + .flat_map(|edit| edit.deletions.iter().map(|span| (*span, ""))) + .collect(); + AstOps::splice(source, &deletions) + } +} + +struct Shape<'m> { + metadata: &'m AstMetadata, + swaps: &'m HashMap, + matched: usize, + out: String, +} + +impl<'m> Shape<'m> { + fn of( + root: &Node, + metadata: &'m AstMetadata, + swaps: &'m HashMap, + ) -> Option { + let mut shape = Shape { + metadata, + swaps, + matched: 0, + out: String::new(), + }; + shape.write(root); + (shape.matched == swaps.len()).then_some(shape.out) + } + + fn write(&mut self, node: &Node) { + if let Node::Binary { + left, + operator: Operator::Logical(LogicalOperator::NullishCoalescing), + right, + } = node + { + if let Some(keep_left) = AstOps::span(self.metadata, node) + .and_then(|span| self.swaps.get(&span)) + .copied() + { + self.matched += 1; + return self.write(if keep_left { left } else { right }); + } + } + match node { + Node::Parenthesized(inner) => return self.write(inner), + Node::Null + | Node::Bool(_) + | Node::Number(_) + | Node::String(_) + | Node::Pointer + | Node::Identifier(_) + | Node::Root => { + let _ = write!(self.out, "{node:?};"); + return; + } + _ => {} + } + self.out.push_str(node.into()); + let _ = match node { + Node::Closure { alias, .. } => write!(self.out, "{alias:?}"), + Node::Interval { + left_bracket, + right_bracket, + .. + } => write!(self.out, "{left_bracket:?}{right_bracket:?}"), + Node::Unary { operator, .. } | Node::Binary { operator, .. } => { + write!(self.out, "{operator:?}") + } + Node::FunctionCall { kind, .. } => write!(self.out, "{kind:?}"), + Node::MethodCall { kind, .. } => write!(self.out, "{kind:?}"), + Node::Error { error, .. } => write!(self.out, "{error:?}"), + _ => Ok(()), + }; + self.out.push('('); + match node { + Node::TemplateString(items) | Node::Array(items) => { + items.iter().for_each(|item| self.write(item)) + } + Node::Object(entries) => entries.iter().for_each(|(key, value)| { + self.write(key); + self.write(value); + }), + Node::Assignments { list, output } => { + list.iter().for_each(|(key, value)| { + self.write(key); + self.write(value); + }); + self.optional(*output); + } + Node::Closure { body, .. } => self.write(body), + Node::Member { node, property } => { + self.write(node); + self.write(property); + } + Node::Slice { node, from, to } => { + self.write(node); + self.optional(*from); + self.optional(*to); + } + Node::Interval { left, right, .. } | Node::Binary { left, right, .. } => { + self.write(left); + self.write(right); + } + Node::Conditional { + condition, + on_true, + on_false, + } => { + self.write(condition); + self.write(on_true); + self.write(on_false); + } + Node::Unary { node, .. } => self.write(node), + Node::FunctionCall { arguments, .. } => { + arguments.iter().for_each(|argument| self.write(argument)) + } + Node::MethodCall { + this, arguments, .. + } => { + self.write(this); + arguments.iter().for_each(|argument| self.write(argument)); + } + Node::Error { node, .. } => self.optional(*node), + _ => {} + } + self.out.push(')'); + } + + fn optional(&mut self, node: Option<&Node>) { + match node { + Some(node) => self.write(node), + None => self.out.push('_'), + } + } +} diff --git a/core/engine/src/analysis/table/cell.rs b/core/engine/src/analysis/table/cell.rs index 1f756018..fd6ac120 100644 --- a/core/engine/src/analysis/table/cell.rs +++ b/core/engine/src/analysis/table/cell.rs @@ -1,3 +1,4 @@ +use std::cell::Cell; use std::rc::Rc; use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}; use std::sync::Arc; @@ -45,24 +46,41 @@ impl CellConstraint { } else { None }; - (truth, format!("{node:?}")) + match truth { + Some(truth) => Ok(truth.t), + None => Err((Self::is_random(node), format!("{node:?}"))), + } }); match parsed { - Some((Some(truth), _)) => CellConstraint::Known(truth.t), - Some((None, key)) => CellConstraint::Opaque(Self::atom_key(trimmed, Some(key))), - None => CellConstraint::Opaque(Self::atom_key(trimmed, None)), + Some(Ok(set)) => CellConstraint::Known(set), + Some(Err((random, key))) => CellConstraint::Opaque(Self::atom_key(random, key)), + None => CellConstraint::Opaque(Self::atom_key(false, format!("src:{trimmed}"))), } } - fn atom_key(source: &str, ast: Option) -> Rc { - if source.contains("rand(") { + fn is_random(node: &Node) -> bool { + let random = Cell::new(false); + node.walk(|n| { + if let Node::FunctionCall { + kind: FunctionKind::Internal(InternalFunction::Rand), + .. + } = n + { + random.set(true); + } + }); + random.get() + } + + fn atom_key(random: bool, key: String) -> Rc { + if random { static UNIQUE: AtomicUsize = AtomicUsize::new(0); return Rc::from(format!( "unique:{}", UNIQUE.fetch_add(1, AtomicOrdering::Relaxed) )); } - Rc::from(ast.unwrap_or_else(|| format!("src:{source}"))) + Rc::from(key) } pub(crate) fn known_set(&self) -> Option { @@ -134,10 +152,79 @@ struct Truth { } impl Truth { - fn unwrap<'a, 'n>(node: &'a Node<'n>) -> &'a Node<'n> { - match node { - Node::Parenthesized(inner) => Self::unwrap(inner), - other => other, + fn unwrap<'a, 'n>(mut node: &'a Node<'n>) -> &'a Node<'n> { + while let Node::Parenthesized(inner) = node { + node = inner; + } + node + } + + fn chain<'a, 'n>(node: &'a Node<'n>, op: LogicalOperator) -> Vec<&'a Node<'n>> { + let mut operands = Vec::new(); + let mut pending = vec![node]; + while let Some(next) = pending.pop() { + match Self::unwrap(next) { + Node::Binary { + left, + operator: Operator::Logical(found), + right, + } if *found == op => { + pending.push(right); + pending.push(left); + } + other => operands.push(other), + } + } + operands + } + + fn all(node: &Node, cx: &Scope) -> Option { + let mut operands = Self::chain(node, LogicalOperator::And).into_iter(); + let mut acc = Self::of(operands.next()?, cx)?; + for operand in operands { + let b = Self::of(operand, cx)?; + acc = Truth { + f: acc.f.union(&acc.t.intersect(&b.f)), + t: acc.t.intersect(&b.t), + }; + } + Some(acc) + } + + fn any(node: &Node, cx: &Scope) -> Option { + let mut acc: Option = None; + let mut total: Vec = Vec::new(); + for operand in Self::chain(node, LogicalOperator::Or) { + let b = Self::of(operand, cx)?; + if !b.t.intersects(&b.f) && b.t.union(&b.f).is_all() { + total.push(b.t); + continue; + } + if !total.is_empty() { + acc = Some(Self::either(acc, Self::total(&total))); + total.clear(); + } + acc = Some(Self::either(acc, b)); + } + if !total.is_empty() { + acc = Some(Self::either(acc, Self::total(&total))); + } + acc + } + + fn total(sets: &[ValueSet]) -> Truth { + let t = ValueSet::union_all(sets); + let f = t.complement(); + Truth { t, f } + } + + fn either(acc: Option, b: Truth) -> Truth { + match acc { + None => b, + Some(a) => Truth { + t: a.t.union(&a.f.intersect(&b.t)), + f: a.f.intersect(&b.f), + }, } } @@ -155,29 +242,13 @@ impl Truth { f: inner.t, }), Node::Binary { - left, operator: Operator::Logical(LogicalOperator::And), - right, - } => { - let a = Self::of(left, cx)?; - let b = Self::of(right, cx)?; - Some(Truth { - t: a.t.intersect(&b.t), - f: a.f.union(&a.t.intersect(&b.f)), - }) - } + .. + } => Self::all(node, cx), Node::Binary { - left, operator: Operator::Logical(LogicalOperator::Or), - right, - } => { - let a = Self::of(left, cx)?; - let b = Self::of(right, cx)?; - Some(Truth { - t: a.t.union(&a.f.intersect(&b.t)), - f: a.f.intersect(&b.f), - }) - } + .. + } => Self::any(node, cx), Node::Binary { left, operator: Operator::Comparison(op), @@ -222,42 +293,34 @@ impl Truth { } } - fn subject<'a>(node: &Node<'a>) -> Option> { - match Self::unwrap(node) { - Node::FunctionCall { - kind: FunctionKind::Internal(InternalFunction::Bool), - arguments: [argument], - } => Self::subject(argument), - Node::Unary { - operator: Operator::Logical(LogicalOperator::Not), - node, - } => Self::subject(node), - Node::Binary { - left, - operator: Operator::Logical(LogicalOperator::And | LogicalOperator::Or), - .. - } => Self::subject(left), - Node::Binary { - left, - operator: Operator::Comparison(_), - right, - } => Self::path(left).or_else(|| Self::path(right)), - _ => None, + fn subject<'a>(mut node: &Node<'a>) -> Option> { + loop { + node = match Self::unwrap(node) { + Node::FunctionCall { + kind: FunctionKind::Internal(InternalFunction::Bool), + arguments: [argument], + } => argument, + Node::Unary { + operator: Operator::Logical(LogicalOperator::Not), + node, + } => node, + Node::Binary { + left, + operator: Operator::Logical(LogicalOperator::And | LogicalOperator::Or), + .. + } => left, + Node::Binary { + left, + operator: Operator::Comparison(_), + right, + } => return Self::path(left).or_else(|| Self::path(right)), + _ => return None, + }; } } fn conjuncts<'n, 'a>(node: &'n Node<'a>, out: &mut Vec<&'n Node<'a>>) { - match Self::unwrap(node) { - Node::Binary { - left, - operator: Operator::Logical(LogicalOperator::And), - right, - } => { - Self::conjuncts(left, out); - Self::conjuncts(right, out); - } - other => out.push(other), - } + out.extend(Self::chain(node, LogicalOperator::And)); } fn comparison(left: &Node, op: ComparisonOperator, right: &Node, cx: &Scope) -> Option { @@ -319,10 +382,11 @@ impl Truth { fn membership(right: &Node, cx: &Scope) -> Option { match Self::unwrap(right) { Node::Array(items) => { - let mut t = ValueSet::empty(); - for item in items.iter() { - t = t.union(&Self::literal(item, cx)?); - } + let literals = items + .iter() + .map(|item| Self::literal(item, cx)) + .collect::>>()?; + let t = ValueSet::union_all(&literals); let f = ValueSet::scalars().difference(&t); Some(Truth { t, f }) } diff --git a/core/engine/src/analysis/table/index.rs b/core/engine/src/analysis/table/index.rs index a6f2e1ec..983bf4ee 100644 --- a/core/engine/src/analysis/table/index.rs +++ b/core/engine/src/analysis/table/index.rs @@ -154,9 +154,10 @@ impl RowIndex { pub(super) fn inner_points(&self, col: usize, interval: &Interval) -> Vec { let points = &self.columns[col].points; match points.range(interval) { - Some((first, last)) => { - vec![points.representative(last), points.representative(first)] - } + Some((first, last)) => [points.representative(last), points.representative(first)] + .into_iter() + .flatten() + .collect(), None => Vec::new(), } } diff --git a/core/engine/src/analysis/table/merge.rs b/core/engine/src/analysis/table/merge.rs index feecab6b..7e0286f8 100644 --- a/core/engine/src/analysis/table/merge.rs +++ b/core/engine/src/analysis/table/merge.rs @@ -321,14 +321,15 @@ impl VerifyTable<'_> { } let fresh: BTreeSet> = b.iter().filter(|key| !a.contains(*key)).cloned().collect(); + if fresh.iter().any(|key| CellText::string(key).is_none()) { + return false; + } let mut moved = Self::region(next_row); moved[col] = ValueSet { strings: StringSet::Finite(fresh.clone()), ..ValueSet::empty() }; - if self.mode != HitMode::Collect - && !Self::clear_between(rows, keep, next, &moved, work, index) - { + if !Self::clear_between(rows, keep, next, &moved, work, index) { return false; } if let Some(row) = rows[keep].as_mut() { @@ -361,9 +362,7 @@ impl VerifyTable<'_> { .enumerate() .map(|(idx, set)| if idx == col { set.difference(&a) } else { set }) .collect(); - if self.mode != HitMode::Collect - && !Self::clear_between(rows, keep, next, &moved, work, index) - { + if !Self::clear_between(rows, keep, next, &moved, work, index) { return false; } let union = a.union(&b); diff --git a/core/engine/src/analysis/table/partition.rs b/core/engine/src/analysis/table/partition.rs index 06413028..667fcd50 100644 --- a/core/engine/src/analysis/table/partition.rs +++ b/core/engine/src/analysis/table/partition.rs @@ -46,14 +46,14 @@ impl Points { } } - pub(super) fn representative(&self, piece: usize) -> Decimal { + pub(super) fn representative(&self, piece: usize) -> Option { let points = &self.0; match (piece % 2, points.len()) { - (1, _) => points[piece / 2], - (_, 0) => Decimal::ZERO, - _ if piece == 0 => points[0] - Decimal::ONE, - _ if piece / 2 == points.len() => points[points.len() - 1] + Decimal::ONE, - _ => (points[piece / 2 - 1] + points[piece / 2]) / Decimal::TWO, + (1, _) => Some(points[piece / 2]), + (_, 0) => Some(Decimal::ZERO), + _ if piece == 0 => points[0].checked_sub(Decimal::ONE), + _ if piece / 2 == points.len() => points[points.len() - 1].checked_add(Decimal::ONE), + _ => Interval::midpoint(points[piece / 2 - 1], points[piece / 2]), } } diff --git a/core/engine/src/analysis/table/print.rs b/core/engine/src/analysis/table/print.rs index 234d8991..7bfb18ca 100644 --- a/core/engine/src/analysis/table/print.rs +++ b/core/engine/src/analysis/table/print.rs @@ -2,7 +2,7 @@ use rust_decimal::prelude::ToPrimitive; use rust_decimal::Decimal; use serde_json::Value; -use super::value_set::{decimal_json, Bound, Interval, StringSet, ValueSet}; +use super::value_set::{Bound, Interval, StringSet, ValueSet}; pub(crate) struct DateDay; @@ -81,13 +81,13 @@ impl DateDay { let day = Decimal::from(Self::DAY); for interval in set.numbers.intervals() { let candidates = match (interval.lo, interval.hi) { - (Bound::Inclusive(a), _) => vec![a, a + day], - (Bound::Exclusive(a), _) => vec![a + day], - (Bound::Unbounded, Bound::Inclusive(b)) => vec![b, b - day], - (Bound::Unbounded, Bound::Exclusive(b)) => vec![b - day], - (Bound::Unbounded, Bound::Unbounded) => vec![Decimal::ZERO], + (Bound::Inclusive(a), _) => vec![Some(a), a.checked_add(day)], + (Bound::Exclusive(a), _) => vec![a.checked_add(day)], + (Bound::Unbounded, Bound::Inclusive(b)) => vec![Some(b), b.checked_sub(day)], + (Bound::Unbounded, Bound::Exclusive(b)) => vec![b.checked_sub(day)], + (Bound::Unbounded, Bound::Unbounded) => vec![Some(Decimal::ZERO)], }; - let found = candidates.into_iter().find_map(|c| { + let found = candidates.into_iter().flatten().find_map(|c| { ValueSet::number(c) .is_subset(set) .then(|| Self::format(c)) @@ -107,20 +107,19 @@ impl CellText { pub(crate) fn brief(text: &str) -> String { const KEEP: usize = 6; let chars: Vec = text.chars().collect(); - let mut quoted = false; + let mut quote: Option = None; let mut depth = 0usize; let mut commas: Vec<(usize, usize)> = Vec::new(); - let mut i = 0; - while i < chars.len() { - match chars[i] { - '\\' if quoted => i += 1, - '"' => quoted = !quoted, - '[' | '(' if !quoted => depth += 1, - ']' | ')' if !quoted => depth = depth.saturating_sub(1), - ',' if !quoted && depth <= 1 => commas.push((i, depth)), + for (i, &c) in chars.iter().enumerate() { + match (quote, c) { + (Some(open), _) if c == open => quote = None, + (Some(_), _) => {} + (None, '"' | '\'' | '`') => quote = Some(c), + (None, '[' | '(') => depth += 1, + (None, ']' | ')') => depth = depth.saturating_sub(1), + (None, ',') if depth <= 1 => commas.push((i, depth)), _ => {} } - i += 1; } if commas.len() < KEEP + 2 { return text.to_string(); @@ -143,7 +142,7 @@ impl CellText { return None; } let positive = Self::positive(&wanted, dated); - let negative = Self::negative(&domain.difference(&wanted), dated); + let negative = Self::negative(&domain.difference(&wanted), wanted.other, dated); match (positive, negative) { (Some(p), Some(n)) if n.len() < p.len() => Some(n), (Some(p), _) => Some(p), @@ -173,7 +172,11 @@ impl CellText { } } match &set.strings { - StringSet::Finite(values) => tokens.extend(values.iter().map(|v| Self::string(v))), + StringSet::Finite(values) => { + for value in values { + tokens.push(Self::string(value)?); + } + } StringSet::CoFinite(_) => return None, } if set.bools & ValueSet::TRUE != 0 { @@ -188,7 +191,7 @@ impl CellText { (!tokens.is_empty()).then(|| tokens.join(", ")) } - fn negative(excluded: &ValueSet, dated: bool) -> Option { + fn negative(excluded: &ValueSet, other: bool, dated: bool) -> Option { if excluded.other || (dated && !excluded.strings.is_empty()) { return None; } @@ -202,7 +205,11 @@ impl CellText { } } match &excluded.strings { - StringSet::Finite(values) => points.extend(values.iter().map(|v| Self::string(v))), + StringSet::Finite(values) => { + for value in values { + points.push(Self::string(value)?); + } + } StringSet::CoFinite(_) => return None, } if excluded.bools & ValueSet::TRUE != 0 { @@ -217,6 +224,13 @@ impl CellText { match points.as_slice() { [] => None, [single] => Some(format!("!= {single}")), + _ if other => Some( + points + .iter() + .map(|point| format!("!= {point}")) + .collect::>() + .join(" and "), + ), _ => Some(format!("not in [{}]", points.join(", "))), } } @@ -260,16 +274,17 @@ impl CellText { fn number(d: Decimal, dated: bool) -> Option { if dated { - return DateDay::format(d).map(|text| Self::string(&text)); + return DateDay::format(d).and_then(|text| Self::string(&text)); } - Some(match decimal_json(d) { - Value::Number(n) => n.to_string(), - _ => d.normalize().to_string(), - }) + Some(d.normalize().to_string()) } - fn string(s: &str) -> String { - serde_json::to_string(s).unwrap_or_else(|_| format!("\"{s}\"")) + pub(crate) fn string(s: &str) -> Option { + match (s.contains('"'), s.contains('\'')) { + (false, _) => Some(format!("\"{s}\"")), + (true, false) => Some(format!("'{s}'")), + (true, true) => None, + } } } diff --git a/core/engine/src/analysis/table/value_set.rs b/core/engine/src/analysis/table/value_set.rs index 3f8598fd..4faf3066 100644 --- a/core/engine/src/analysis/table/value_set.rs +++ b/core/engine/src/analysis/table/value_set.rs @@ -124,22 +124,35 @@ impl Interval { above && below } - fn example(&self) -> Decimal { + pub(crate) fn midpoint(l: Decimal, h: Decimal) -> Option { + [ + l.checked_add(h).map(|sum| sum / Decimal::TWO), + (l / Decimal::TWO).checked_add(h / Decimal::TWO), + h.checked_sub(l) + .and_then(|width| l.checked_add(width / Decimal::TWO)), + ] + .into_iter() + .flatten() + .find(|m| l < *m && *m < h) + } + + fn example(&self) -> Option { let candidates = match (self.lo, self.hi) { - (Bound::Unbounded, Bound::Unbounded) => vec![Decimal::ZERO], - (Bound::Inclusive(l), _) => vec![l], - (Bound::Exclusive(l), Bound::Unbounded) => vec![l.floor() + Decimal::ONE], - (Bound::Unbounded, Bound::Inclusive(h)) => vec![h], - (Bound::Unbounded, Bound::Exclusive(h)) => vec![h.ceil() - Decimal::ONE], + (Bound::Unbounded, Bound::Unbounded) => vec![Some(Decimal::ZERO)], + (Bound::Inclusive(l), _) => vec![Some(l)], + (Bound::Exclusive(l), Bound::Unbounded) => vec![l.floor().checked_add(Decimal::ONE)], + (Bound::Unbounded, Bound::Inclusive(h)) => vec![Some(h)], + (Bound::Unbounded, Bound::Exclusive(h)) => vec![h.ceil().checked_sub(Decimal::ONE)], (Bound::Exclusive(l), hi) => { let h = hi.value().unwrap_or(l); - vec![l.floor() + Decimal::ONE, (l + h) / Decimal::TWO] + vec![ + l.floor().checked_add(Decimal::ONE), + Self::midpoint(l, h), + Some(h), + ] } }; - candidates - .into_iter() - .find(|c| self.contains(*c)) - .unwrap_or(Decimal::ZERO) + candidates.into_iter().flatten().find(|c| self.contains(*c)) } } @@ -263,7 +276,7 @@ impl NumberSet { } fn example(&self) -> Option { - self.intervals.first().map(Interval::example) + self.intervals.iter().find_map(Interval::example) } } @@ -471,6 +484,34 @@ impl ValueSet { } } + pub(crate) fn union_all(sets: &[ValueSet]) -> Self { + let mut intervals = Vec::new(); + let mut finite = BTreeSet::new(); + let mut cofinite: Option = None; + let mut out = Self::empty(); + for set in sets { + intervals.extend(set.numbers.intervals.iter().copied()); + match &set.strings { + StringSet::Finite(values) => finite.extend(values.iter().cloned()), + strings => { + cofinite = Some(match cofinite { + Some(acc) => acc.union(strings), + None => strings.clone(), + }) + } + } + out.bools |= set.bools; + out.null |= set.null; + out.other |= set.other; + } + out.numbers = NumberSet::from_intervals(intervals); + out.strings = match cofinite { + Some(acc) => acc.union(&StringSet::Finite(finite)), + None => StringSet::Finite(finite), + }; + out + } + pub(crate) fn intersect(&self, other: &Self) -> Self { Self { numbers: self.numbers.intersect(&other.numbers), diff --git a/core/engine/src/analysis/table/witness.rs b/core/engine/src/analysis/table/witness.rs index d8ebf9bb..29c7a6b9 100644 --- a/core/engine/src/analysis/table/witness.rs +++ b/core/engine/src/analysis/table/witness.rs @@ -4,7 +4,7 @@ use rust_decimal::Decimal; use super::cell::CellConstraint; use super::index::RowIndex; -use super::value_set::{Bound, StringSet, ValueSet}; +use super::value_set::{Bound, Interval, StringSet, ValueSet}; use super::verify::VerifyTable; #[derive(Clone)] @@ -48,10 +48,10 @@ impl Point { let fallback = match (interval.lo, interval.hi) { (_, Bound::Inclusive(h)) => Some(h), (Bound::Inclusive(l), _) => Some(l), - (Bound::Unbounded, Bound::Exclusive(h)) => Some(h - Decimal::ONE), - (Bound::Exclusive(l), Bound::Unbounded) => Some(l + Decimal::ONE), + (Bound::Unbounded, Bound::Exclusive(h)) => h.checked_sub(Decimal::ONE), + (Bound::Exclusive(l), Bound::Unbounded) => l.checked_add(Decimal::ONE), (Bound::Unbounded, Bound::Unbounded) => Some(Decimal::ZERO), - (Bound::Exclusive(l), Bound::Exclusive(h)) => Some((l + h) / Decimal::TWO), + (Bound::Exclusive(l), Bound::Exclusive(h)) => Interval::midpoint(l, h), }; if let Some(x) = fallback.filter(|x| set.numbers.contains(*x)) { out.push(Point::Number(x)); diff --git a/core/engine/src/policy/blocks/context.rs b/core/engine/src/policy/blocks/context.rs index 3e5d2f01..59e5675f 100644 --- a/core/engine/src/policy/blocks/context.rs +++ b/core/engine/src/policy/blocks/context.rs @@ -9,7 +9,7 @@ use zen_expression::{Isolate, IsolateError}; use super::property_read::ReadFlattener; use super::type_check::TypeCheck; use crate::analysis::nullable::NullableOperand; -use crate::analysis::table::VerifyTable; +use crate::analysis::table::{HitMode, VerifyInput, VerifyOutput, VerifyTable}; use crate::policy::ir::PropertyPath; use crate::policy::queries::dependency::{DataModelPaths, PathPrefix}; use crate::policy::queries::scope::VariableTypeScope; @@ -63,6 +63,14 @@ pub type SharedDictionaryTypes = Rc, VariableType>>; pub type SharedPoisonedPaths = Rc>>>; pub type SharedDeclaredPaths = Rc; +#[derive(Debug, Clone)] +pub struct TableCheck { + pub(crate) at: usize, + pub(crate) mode: HitMode, + pub(crate) inputs: Vec, + pub(crate) outputs: Vec, +} + pub struct AnalysisContext { scope: VariableType, policy_path: Arc, @@ -70,6 +78,7 @@ pub struct AnalysisContext { reads: Vec, writes: Vec, diagnostics: Vec, + table_checks: Vec, pass: AnalysisPass, intellisense: SharedIntelliSense, dictionary_types: SharedDictionaryTypes, @@ -97,6 +106,7 @@ impl AnalysisContext { reads: Vec::new(), writes: Vec::new(), diagnostics: Vec::new(), + table_checks: Vec::new(), pass, intellisense, dictionary_types, @@ -308,24 +318,13 @@ impl AnalysisContext { self.declared_paths.declares(path) } - pub(super) fn push_table_diagnostics( - &mut self, - table: &VerifyTable, - row_key: impl Fn(usize) -> Arc, - ) { - let policy_path = self.policy_path.clone(); - let block_id = self.block_id.clone(); - let diagnostics = table.diagnostics( - &mut self.intellisense.borrow_mut(), - row_key, - |expression_id| match expression_id { - Some(id) => { - DiagnosticLocation::expression(policy_path.clone(), block_id.clone(), id, None) - } - None => DiagnosticLocation::block(policy_path.clone(), block_id.clone()), - }, - ); - self.diagnostics.extend(diagnostics); + pub(super) fn defer_table_check(&mut self, table: VerifyTable) { + self.table_checks.push(TableCheck { + at: self.diagnostics.len(), + mode: table.mode, + inputs: table.inputs, + outputs: table.outputs, + }); } pub fn hint_with_target( @@ -414,6 +413,7 @@ impl AnalysisContext { reads: self.reads, writes: self.writes, diagnostics: self.diagnostics, + table_checks: self.table_checks, } } @@ -494,16 +494,10 @@ impl AnalysisContext { span: Some(diag.span), target: self.default_target.clone(), }; - let mut diagnostic = Diagnostic::from_expression(diag, location); - NullableOperand::annotate( - &mut diagnostic, - &mut self.intellisense.borrow_mut(), - source, - matches!(kind, ExpressionKind::Unary), - ); - self.diagnostics.push(diagnostic); + self.diagnostics + .push(Diagnostic::from_expression(diag, location)); } - NullableOperand::fallback_all( + NullableOperand::annotate( &mut self.diagnostics[first..], &mut self.intellisense.borrow_mut(), source, @@ -517,6 +511,7 @@ pub struct AnalysisSummary { pub reads: Vec, pub writes: Vec, pub diagnostics: Vec, + pub table_checks: Vec, } pub struct ExecutionContext<'a> { diff --git a/core/engine/src/policy/blocks/decision_table.rs b/core/engine/src/policy/blocks/decision_table.rs index a1dcb375..ffd490af 100644 --- a/core/engine/src/policy/blocks/decision_table.rs +++ b/core/engine/src/policy/blocks/decision_table.rs @@ -3,7 +3,7 @@ use std::sync::{Arc, OnceLock}; use ahash::{HashMap, HashSet}; use fixedbitset::FixedBitSet; use serde::{Deserialize, Serialize}; -use zen_expression::intellisense::{ArmTest, NumberCover}; +use zen_expression::intellisense::{ArmTest, IntelliSense, NumberCover}; use zen_expression::variable::{Variable, VariableType}; use zen_expression::Isolate; use zen_types::decision::{ @@ -16,12 +16,12 @@ use crate::analysis::table::{HitMode, TableColumn, VerifyTable}; use crate::policy::queries::scope::VariableTypeScope; use crate::workspace::types::{ BlockTrace, Cursor, CursorTarget, DecisionTableExtras, Diagnostic, DiagnosticArgs, - DiagnosticCode, ExpressionKind, + DiagnosticCode, DiagnosticLocation, ExpressionKind, }; use crate::policy::ArcStrTrim; -use super::context::{AnalysisContext, ExecutionContext, ExecutionError}; +use super::context::{AnalysisContext, ExecutionContext, ExecutionError, TableCheck}; use super::{ Block, BlockKind, BlockReadPlan, CellReads, ConditionalReads, ExpressionLocation, ParseContext, ReadFlattenFn, WriteSite, WriteTarget, @@ -478,7 +478,7 @@ impl DecisionTableIr { .collect(), rules: &self.rules, }; - cx.push_table_diagnostics(&table, |row| Self::row_key(&self.rules[row], row)); + cx.defer_table_check(table); } for col in &self.outputs { @@ -703,6 +703,31 @@ impl DecisionTableIr { } } + pub(crate) fn verify( + &self, + check: &TableCheck, + is: &mut IntelliSense, + policy_path: &Arc, + block_id: &Arc, + ) -> Vec { + let table = VerifyTable { + mode: check.mode, + inputs: check.inputs.clone(), + outputs: check.outputs.clone(), + rules: &self.rules, + }; + table.diagnostics( + is, + |row| Self::row_key(&self.rules[row], row), + |expression_id| match expression_id { + Some(id) => { + DiagnosticLocation::expression(policy_path.clone(), block_id.clone(), id, None) + } + None => DiagnosticLocation::block(policy_path.clone(), block_id.clone()), + }, + ) + } + fn row_key(rule: &HashMap, Arc>, row: usize) -> Arc { rule.get(ROW_ID_KEY) .cloned() diff --git a/core/engine/src/policy/blocks/mod.rs b/core/engine/src/policy/blocks/mod.rs index a14c8d39..e9d2fd99 100644 --- a/core/engine/src/policy/blocks/mod.rs +++ b/core/engine/src/policy/blocks/mod.rs @@ -29,7 +29,7 @@ pub(crate) use context::IntelliSenseSource; pub use context::{ AnalysisContext, AnalysisSummary, ExecutionContext, ExecutionError, ExpressionLocation, InstanceSource, PropertyRead, SharedDeclaredPaths, SharedDictionaryTypes, SharedIntelliSense, - SharedPoisonedPaths, WriteTarget, + SharedPoisonedPaths, TableCheck, WriteTarget, }; pub use decision_table::{DecisionTableDoc, DecisionTableIr, DeclaredType}; pub(crate) use decision_table::{DictionaryCandidate, TableSelection, ROW_ID_KEY}; diff --git a/core/engine/src/policy/linter/mod.rs b/core/engine/src/policy/linter/mod.rs index 916d57e6..e07c912f 100644 --- a/core/engine/src/policy/linter/mod.rs +++ b/core/engine/src/policy/linter/mod.rs @@ -117,10 +117,35 @@ impl AstOps { } fn chars_at(source: &str, span: Span) -> impl Iterator + '_ { - source - .chars() - .skip(span.0 as usize) - .take((span.1 as usize).saturating_sub(span.0 as usize)) + Self::text(source, span).unwrap_or_default().chars() + } + + pub(crate) fn text(source: &str, span: Span) -> Option<&str> { + source.get(span.0 as usize..span.1 as usize) + } + + pub(crate) fn splice(source: &str, edits: &[(Span, &str)]) -> Option { + let word = |c: char| c.is_alphanumeric() || matches!(c, '_' | '$' | '#'); + let mut sorted: Vec<&(Span, &str)> = edits.iter().collect(); + sorted.sort_by_key(|(range, _)| *range); + let mut out = String::with_capacity(source.len()); + let mut push = |piece: &str| { + if out.chars().next_back().is_some_and(word) && piece.chars().next().is_some_and(word) { + out.push(' '); + } + out.push_str(piece); + }; + let mut cursor = 0u32; + for (range, with) in sorted { + if range.0 < cursor || range.1 < range.0 { + return None; + } + push(Self::text(source, (cursor, range.0))?); + push(with); + cursor = range.1; + } + push(source.get(cursor as usize..)?); + Some(out) } } diff --git a/core/engine/src/policy/linter/redundant_parentheses.rs b/core/engine/src/policy/linter/redundant_parentheses.rs index fb19bcc2..17ae71ab 100644 --- a/core/engine/src/policy/linter/redundant_parentheses.rs +++ b/core/engine/src/policy/linter/redundant_parentheses.rs @@ -1,7 +1,8 @@ -use zen_expression::intellisense::AstMetadata; +use zen_expression::intellisense::{AstMetadata, IntelliSense}; use zen_expression::lexer::Operator; use zen_expression::parser::{Associativity, Node, ParserOperator}; +use crate::analysis::proof::{FixEdit, FixProof}; use crate::workspace::types::{ Diagnostic, DiagnosticArgs, DiagnosticCode, DiagnosticLocation, ExpressionKind, Span, }; @@ -203,28 +204,34 @@ impl RedundantParentheses { impl RedundantParentheses { pub(crate) fn fix_args( + is: &mut IntelliSense, source: &str, findings: &[(Option, Option)], - mut shape: impl FnMut(&str) -> Option, ) -> Vec { - let pairs: Vec> = findings + let edits: Vec> = findings .iter() - .map(|(outer, inner)| Some(((*outer)?, (*inner)?))) + .map(|(outer, inner)| Some(FixEdit::unwrap((*outer)?, (*inner)?, None))) .collect(); - let Some(expected) = shape(source) else { - return vec![DiagnosticArgs::new(); findings.len()]; - }; - let mut verified = - |fixed: String| (shape(&fixed).as_deref() == Some(expected.as_str())).then_some(fixed); - let all: Vec<(Span, Span)> = pairs.iter().flatten().copied().collect(); - let fix_all = (all.len() > 1) - .then(|| Self::strip(source, &all)) - .and_then(&mut verified); - pairs + let candidates: Vec = edits.iter().flatten().cloned().collect(); + let proven = FixProof::proven(is, source, false, &candidates); + let accepted: Vec<&FixEdit> = candidates .iter() - .map(|pair| { + .zip(&proven) + .filter_map(|(edit, proven)| proven.then_some(edit)) + .collect(); + let fix_all = (accepted.len() > 1) + .then(|| FixProof::holds(is, source, false, &accepted)) + .flatten(); + let mut proven = proven.into_iter(); + edits + .iter() + .map(|edit| { let mut args = DiagnosticArgs::new(); - if let Some(fixed) = pair.and_then(|pair| verified(Self::strip(source, &[pair]))) { + let fixed = edit + .as_ref() + .filter(|_| proven.next() == Some(true)) + .and_then(|edit| edit.apply(source)); + if let Some(fixed) = fixed { args.insert("fixOriginal", source.to_string()); args.insert("fixSource", fixed); if let Some(all) = &fix_all { @@ -235,83 +242,6 @@ impl RedundantParentheses { }) .collect() } - - pub(crate) fn tree_shape(debug: &str) -> String { - const WRAPPER: &str = "Parenthesized("; - let chars: Vec = debug.chars().collect(); - let wrapper: Vec = WRAPPER.chars().collect(); - let mut drop_close: Vec = Vec::new(); - let mut depth = 0usize; - let mut quoted = false; - let mut out = String::with_capacity(debug.len()); - let mut i = 0; - while i < chars.len() { - let c = chars[i]; - if quoted { - out.push(c); - if c == '\\' && i + 1 < chars.len() { - out.push(chars[i + 1]); - i += 2; - continue; - } - quoted = c != '"'; - i += 1; - continue; - } - if c == '"' { - quoted = true; - out.push(c); - } else if chars[i..].starts_with(&wrapper) { - depth += 1; - drop_close.push(depth); - i += wrapper.len(); - continue; - } else if c == '(' { - depth += 1; - out.push(c); - } else if c == ')' { - if drop_close.last() == Some(&depth) { - drop_close.pop(); - } else { - out.push(c); - } - depth = depth.saturating_sub(1); - } else { - out.push(c); - } - i += 1; - } - out - } - - fn strip(source: &str, pairs: &[(Span, Span)]) -> String { - let chars: Vec = source.chars().collect(); - let mut removed = vec![false; chars.len()]; - for (outer, inner) in pairs { - for idx in - (outer.0 as usize..inner.0 as usize).chain(inner.1 as usize..outer.1 as usize) - { - if let Some(slot) = removed.get_mut(idx) { - *slot = true; - } - } - } - let word = |c: char| c.is_alphanumeric() || matches!(c, '_' | '$' | '#'); - let mut out = String::with_capacity(source.len()); - let mut gap = false; - for (idx, c) in chars.iter().enumerate() { - if removed[idx] { - gap = true; - continue; - } - if gap && out.chars().last().is_some_and(word) && word(*c) { - out.push(' '); - } - gap = false; - out.push(*c); - } - out - } } impl LintRule for RedundantParentheses { @@ -326,11 +256,11 @@ impl LintRule for RedundantParentheses { RedundantParentheses::scan(root, metadata) }) .unwrap_or_default(); - let fixes = Self::fix_args(&expression.source, &findings, |source| { - cx.with_ast(source, expression.kind, |root, _| { - Self::tree_shape(&format!("{root:?}")) - }) - }); + let fixes = Self::fix_args( + &mut cx.db.intellisense().borrow_mut(), + &expression.source, + &findings, + ); for ((span, inner_span), args) in findings.into_iter().zip(fixes) { let message = match inner_span { Some(inner) => format!( diff --git a/core/engine/src/policy/queries/dependency.rs b/core/engine/src/policy/queries/dependency.rs index 36efe4b5..5fefec93 100644 --- a/core/engine/src/policy/queries/dependency.rs +++ b/core/engine/src/policy/queries/dependency.rs @@ -9,7 +9,7 @@ use zen_expression::variable::VariableType; use crate::policy::blocks::{ AnalysisContext, AnalysisSummary, Block, InstanceSource, PropertyRead, SharedDeclaredPaths, - SharedDictionaryTypes, SharedIntelliSense, SharedPoisonedPaths, WriteTarget, + SharedDictionaryTypes, SharedIntelliSense, SharedPoisonedPaths, TableCheck, WriteTarget, }; use crate::policy::ir::{DataModelIr, ParsedPolicy, PropertyPath, PropertyTypeIr}; use crate::policy::queries::path::{PathClassifier, PathRoot}; @@ -142,7 +142,9 @@ impl EnrichedState { #[derive(Debug, Clone)] pub struct RuleEnrichedAnalysis { pub policy_path: Arc, + pub block_id: Arc, pub diagnostics: Vec, + pub table_checks: Vec, } #[derive(Debug)] @@ -801,7 +803,9 @@ impl Snapshot { per_rule.push(RuleEnrichedAnalysis { policy_path: policy_path.clone(), + block_id: key.block_id.clone(), diagnostics: summary.diagnostics, + table_checks: summary.table_checks, }); } diff --git a/core/engine/src/policy/queries/diagnostics.rs b/core/engine/src/policy/queries/diagnostics.rs index a66bfd0e..cd032057 100644 --- a/core/engine/src/policy/queries/diagnostics.rs +++ b/core/engine/src/policy/queries/diagnostics.rs @@ -2,9 +2,10 @@ use std::sync::Arc; use ahash::{HashMap, HashMapExt, HashSet}; +use crate::policy::blocks::BlockKind; use crate::policy::ir::PropertyTypeIr; use crate::policy::linter::Linter; -use crate::policy::queries::dependency::WriteScope; +use crate::policy::queries::dependency::{RuleEnrichedAnalysis, WriteScope}; use crate::policy::queries::path::PathRoot; use crate::workspace::db::{Db, Unit}; use crate::workspace::types::{BlockRef, Diagnostic, DiagnosticCode, DiagnosticLocation}; @@ -50,13 +51,13 @@ impl Db { .filter(|d| d.is_in(path)) .cloned(), ); - out.extend( - enriched - .per_rule - .iter() - .filter(|rule| rule.policy_path == *path) - .flat_map(|rule| rule.diagnostics.iter().cloned()), - ); + for rule in enriched + .per_rule + .iter() + .filter(|rule| rule.policy_path == *path) + { + out.extend(self.rule_diagnostics(rule)); + } out.extend(self.import_diagnostics(path)); @@ -71,6 +72,34 @@ impl Db { out } + fn rule_diagnostics(&self, rule: &RuleEnrichedAnalysis) -> Vec { + if rule.table_checks.is_empty() { + return rule.diagnostics.clone(); + } + let block = self.block_ir(&BlockRef { + policy_path: rule.policy_path.clone(), + block_id: rule.block_id.clone(), + }); + let Some(BlockKind::DecisionTable(table)) = block.as_ref().map(|block| &block.kind) else { + return rule.diagnostics.clone(); + }; + let intellisense = self.intellisense(); + let mut out = Vec::with_capacity(rule.diagnostics.len()); + let mut cursor = 0; + for check in &rule.table_checks { + out.extend(rule.diagnostics[cursor..check.at].iter().cloned()); + out.extend(table.verify( + check, + &mut intellisense.borrow_mut(), + &rule.policy_path, + &rule.block_id, + )); + cursor = check.at; + } + out.extend(rule.diagnostics[cursor..].iter().cloned()); + out + } + fn locate_nullable_sources(&self, path: &Arc, out: &mut [Diagnostic]) { if !out.iter().any(|d| d.args.contains_key("nullablePath")) { return; @@ -749,3 +778,65 @@ impl Db { out } } + +#[cfg(test)] +mod tests { + use serde_json::json; + + use crate::workspace::types::{DiagnosticCode, Severity}; + use crate::workspace::Workspace; + + #[test] + fn evaluation_diagnostics_skip_table_verification() { + let policy = |value: &str| { + json!({ "blocks": [ + { "id": "dm", "type": "dataModel", "props": { "data": { + "name": "applicant", + "properties": [ + { "id": "p1", "name": "age", "type": "number", "array": false, "optional": false } + ] + } } }, + { "id": "dt", "type": "decisionTable", "props": { "data": { + "hitPolicy": "first", + "inputs": [ { "id": "i0", "name": "Age", "field": "applicant.age" } ], + "outputs": [ { "id": "o0", "name": "Rate", "field": "applicant.rate" } ], + "rules": [ + { "_id": "r1", "i0": "< 18", "o0": "1" }, + { "_id": "r2", "i0": "< 10", "o0": "2" } + ] + } } }, + { "id": "calc", "type": "expression", "props": { "data": { "key": "applicant.total", "value": value } } } + ] }) + }; + let table_codes = [ + DiagnosticCode::MissingCases, + DiagnosticCode::UnreachableRule, + ]; + for (value, errors) in [("applicant.age + 1", 0), ("applicant.missing > 50", 1)] { + let mut ws = Workspace::new(); + ws.set_policy("p", serde_json::from_value(policy(value)).expect("policy")); + let editor = ws.diagnostics("p"); + assert_eq!( + editor + .iter() + .filter(|d| table_codes.contains(&d.code)) + .count(), + 2, + "{editor:?}" + ); + let evaluation = ws.evaluation_diagnostics("p"); + assert!( + evaluation.iter().all(|d| !table_codes.contains(&d.code)), + "{evaluation:?}" + ); + assert_eq!( + evaluation + .iter() + .filter(|d| d.severity == Severity::Error) + .count(), + errors, + "{evaluation:?}" + ); + } + } +} diff --git a/core/engine/src/workspace/graph/analysis.rs b/core/engine/src/workspace/graph/analysis.rs index 886494c4..9516105b 100644 --- a/core/engine/src/workspace/graph/analysis.rs +++ b/core/engine/src/workspace/graph/analysis.rs @@ -1169,8 +1169,13 @@ impl<'a> GraphAnalyzer<'a> { return None; } let mut written: Vec> = Vec::new(); + let mut seen: HashSet> = HashSet::default(); for (pred, _) in incoming { - written.extend(after.get(pred).cloned().flatten()?); + for path in after.get(pred)?.as_ref()? { + if seen.insert(path.clone()) { + written.push(path.clone()); + } + } } Some(written) } @@ -1280,7 +1285,7 @@ impl<'a> GraphAnalyzer<'a> { content: &DecisionTableContent, field: Option<&str>, ) -> Option { - if content.transform_attributes.input_field.is_some() { + if !self.preserved_input(content, field) { return None; } let field = field?.trim(); @@ -1647,13 +1652,11 @@ impl<'a> GraphAnalyzer<'a> { RedundantParentheses::scan(root, metadata) }) .unwrap_or_default(); - let fixes = RedundantParentheses::fix_args(&site.source, &findings, |source| { - intellisense - .borrow_mut() - .with_ast(source, false, |root, _| { - RedundantParentheses::tree_shape(&format!("{root:?}")) - }) - }); + let fixes = RedundantParentheses::fix_args( + &mut intellisense.borrow_mut(), + &site.source, + &findings, + ); for ((span, inner_span), args) in findings.into_iter().zip(fixes) { let message = match inner_span { Some(inner) => format!( @@ -2076,16 +2079,10 @@ impl<'a> GraphAnalyzer<'a> { span: Some(diagnostic.span), target: target.clone(), }; - let mut diagnostic = Diagnostic::from_expression(diagnostic, location); - NullableOperand::annotate( - &mut diagnostic, - &mut intellisense.borrow_mut(), - source, - matches!(kind, ExpressionKind::Unary), - ); - self.diagnostics.push(diagnostic); + self.diagnostics + .push(Diagnostic::from_expression(diagnostic, location)); } - NullableOperand::fallback_all( + NullableOperand::annotate( &mut self.diagnostics[first..], &mut intellisense.borrow_mut(), source, diff --git a/core/engine/tests/table_fix_regressions.rs b/core/engine/tests/table_fix_regressions.rs new file mode 100644 index 00000000..c40bd026 --- /dev/null +++ b/core/engine/tests/table_fix_regressions.rs @@ -0,0 +1,433 @@ +use serde_json::{json, Value}; +use std::sync::Arc; +use std::time::{Duration, Instant}; +use zen_engine::loader::MemoryLoader; +use zen_engine::model::{DecisionContent, PolicyContent}; +use zen_engine::policy::{Diagnostic, DiagnosticCode, PolicyWorkspace, Workspace}; +use zen_engine::DecisionEngine; + +fn with_code(diagnostics: &[Diagnostic], code: DiagnosticCode) -> Vec { + diagnostics + .iter() + .filter(|d| d.code == code) + .cloned() + .collect() +} + +fn arg(diagnostic: &Diagnostic, key: &str) -> Option { + diagnostic.args.get(key).cloned() +} + +fn applicant_model() -> Value { + json!({ "id": "dm", "type": "dataModel", "props": { "data": { + "name": "applicant", + "properties": [ + { "id": "p1", "name": "age", "type": "number", "array": false, "optional": false }, + { "id": "p2", "name": "target", "type": "number", "array": false, "optional": true }, + { "id": "p3", "name": "vip", "type": "boolean", "array": false, "optional": false } + ] + } } }) +} + +fn policy_diagnostics(blocks: Vec) -> Vec { + let mut all = vec![applicant_model()]; + all.extend(blocks); + let mut ws = PolicyWorkspace::new(); + ws.set_policy( + "p", + serde_json::from_value(json!({ "blocks": all })).expect("policy"), + ); + ws.diagnostics("p") +} + +fn policy_expression(value: &str) -> Vec { + policy_diagnostics(vec![ + json!({ "id": "calc", "type": "expression", "props": { "data": { "key": "applicant.total", "value": value } } }), + ]) +} + +fn graph_expression(value: &str) -> Vec { + let schema = json!({ + "type": "object", + "properties": { "amount": { "type": "number" }, "target": { "type": "number" } }, + "required": ["amount"] + }); + let content: DecisionContent = serde_json::from_value(json!({ + "nodes": [ + { "id": "in", "name": "in", "type": "inputNode", "content": { "schema": schema.to_string() } }, + { "id": "calc", "name": "calc", "type": "expressionNode", "content": { + "expressions": [ { "id": "x", "key": "total", "value": value } ], + "passThrough": true + } }, + { "id": "out", "name": "out", "type": "outputNode", "content": {} } + ], + "edges": [ + { "id": "e1", "sourceId": "in", "targetId": "calc" }, + { "id": "e2", "sourceId": "calc", "targetId": "out" } + ] + })) + .expect("graph"); + let mut ws = Workspace::new(); + ws.set_document("g", content); + ws.diagnostics("g") +} + +fn cell_table(field: &str, cell: &str) -> Value { + json!({ "id": "dt", "type": "decisionTable", "props": { "data": { + "hitPolicy": "first", + "inputs": [ { "id": "i0", "name": "In", "field": field } ], + "outputs": [ { "id": "o0", "name": "Rate", "field": "applicant.rate" } ], + "rules": [ + { "_id": "r1", "i0": cell, "o0": "1" }, + { "_id": "r2", "i0": "", "o0": "2" } + ] + } } }) +} + +#[test] +fn unary_fallback_fix_keeps_boolean_semantics() { + for cell in [ + "applicant.vip ?? false", + "$ ?? false", + "(applicant.vip ?? true)", + ] { + let found = with_code( + &policy_diagnostics(vec![cell_table("applicant.vip", cell)]), + DiagnosticCode::RedundantNullish, + ); + assert_eq!(found.len(), 1, "{cell}: {found:?}"); + assert_eq!(arg(&found[0], "fixSource"), None, "{cell}: {found:?}"); + } + + let found = with_code( + &policy_diagnostics(vec![cell_table( + "applicant.age", + "(applicant.age ?? 0) + 1", + )]), + DiagnosticCode::RedundantNullish, + ); + assert_eq!(found.len(), 1, "{found:?}"); + assert_eq!( + arg(&found[0], "fixSource").as_deref(), + Some("applicant.age + 1") + ); + + let found = with_code( + &policy_diagnostics(vec![cell_table( + "applicant.age", + "(applicant.age ?? 0) + (applicant.age ?? 1)", + )]), + DiagnosticCode::RedundantNullish, + ); + assert_eq!(found.len(), 2, "{found:?}"); + for d in &found { + assert_eq!(arg(d, "fixSource"), None, "{d:?}"); + } +} + +#[test] +fn fixes_splice_by_byte_offsets() { + for text in ["café", "é🎉"] { + let source = format!("\"{text}\" != \"x\" and applicant.target > 0"); + let found = with_code(&policy_expression(&source), DiagnosticCode::TypeMismatch); + assert_eq!(found.len(), 1, "{found:?}"); + assert_eq!( + arg(&found[0], "fixSource"), + Some(format!( + "\"{text}\" != \"x\" and (applicant.target ?? 0) > 0" + )) + ); + assert_eq!( + arg(&found[0], "fixOperand").as_deref(), + Some("applicant.target") + ); + + let source = format!("\"{text}\" != \"x\" and (applicant.age ?? 0) > 1"); + let found = with_code( + &policy_expression(&source), + DiagnosticCode::RedundantNullish, + ); + assert_eq!(found.len(), 1, "{found:?}"); + assert_eq!( + arg(&found[0], "fixSource"), + Some(format!("\"{text}\" != \"x\" and applicant.age > 1")) + ); + assert_eq!(arg(&found[0], "fixFallback").as_deref(), Some("0")); + + let source = format!("\"{text}\" != \"x\" and (applicant.age) > 1"); + let found = with_code( + &policy_expression(&source), + DiagnosticCode::RedundantParentheses, + ); + assert_eq!(found.len(), 1, "{found:?}"); + assert_eq!( + arg(&found[0], "fixSource"), + Some(format!("\"{text}\" != \"x\" and applicant.age > 1")) + ); + assert_eq!( + found[0].message, + "unnecessary parentheses around 'applicant.age'" + ); + } +} + +#[test] +fn quick_fix_proofs_scale_with_expression_length() { + let source = vec!["(amount ?? 0) + (1)"; 1500].join(" + "); + let started = Instant::now(); + let diagnostics = graph_expression(&source); + let elapsed = started.elapsed(); + eprintln!("{} bytes in {elapsed:?}", source.len()); + + for (code, all, first) in [ + ( + DiagnosticCode::RedundantNullish, + vec!["amount + (1)"; 1500].join(" + "), + "amount + (1) + (amount ?? 0) + (1)", + ), + ( + DiagnosticCode::RedundantParentheses, + vec!["(amount ?? 0) + 1"; 1500].join(" + "), + "(amount ?? 0) + 1 + (amount ?? 0) + (1)", + ), + ] { + let found = with_code(&diagnostics, code); + assert_eq!(found.len(), 1500); + for d in &found { + assert_eq!(arg(d, "fixOriginal").as_deref(), Some(source.as_str())); + assert_eq!(arg(d, "fixAll").as_deref(), Some(all.as_str())); + } + let fixed = arg(&found[0], "fixSource").expect("fix"); + assert!(fixed.starts_with(first), "{}", &fixed[..60]); + assert_eq!( + fixed.len(), + source.len() - (source.len() - all.len()) / 1500 + ); + } + assert!(elapsed < Duration::from_secs(10), "{elapsed:?}"); + + let source = vec!["target * 2"; 2000].join(" + "); + let started = Instant::now(); + let found = with_code(&graph_expression(&source), DiagnosticCode::TypeMismatch); + let elapsed = started.elapsed(); + assert_eq!(found.len(), 2000); + assert!(found.iter().all(|d| d.args.contains_key("fixSource"))); + let fixed = arg(&found[0], "fixSource").expect("fix"); + assert!( + fixed.starts_with("(target ?? 0) * 2 + target * 2"), + "{}", + &fixed[..60] + ); + assert!(elapsed < Duration::from_secs(10), "{elapsed:?}"); +} + +#[test] +fn nested_fallbacks_are_proven_separately() { + let found = with_code( + &graph_expression("-(amount ?? 0 ?? 1)"), + DiagnosticCode::RedundantNullish, + ); + let mut fixes: Vec = found + .iter() + .map(|d| arg(d, "fixSource").expect("fix")) + .collect(); + fixes.sort(); + assert_eq!(fixes, vec!["-(amount ?? 0)", "-(amount ?? 1)"]); + for d in &found { + assert_eq!(arg(d, "fixAll").as_deref(), Some("-(amount)")); + } +} + +fn number_schema() -> String { + json!({ + "type": "object", + "properties": { + "applicant": { + "type": "object", + "properties": { "age": { "type": "number", "minimum": 0, "maximum": 120 } }, + "required": ["age"] + } + }, + "required": ["applicant"] + }) + .to_string() +} + +fn age_table() -> Value { + json!({ "id": "dt", "name": "dt", "type": "decisionTableNode", "content": { + "hitPolicy": "first", + "inputs": [ { "id": "i0", "name": "Age", "field": "applicant.age" } ], + "outputs": [ { "id": "o0", "name": "Rate", "field": "rate" } ], + "rules": [ { "_id": "r1", "i0": "<= 120", "o0": "1" } ] + } }) +} + +fn expression_node(id: &str, key: &str, value: &str) -> Value { + json!({ "id": id, "name": id, "type": "expressionNode", "content": { + "expressions": [ { "id": format!("{id}-x"), "key": key, "value": value } ], + "passThrough": true + } }) +} + +fn graph_diagnostics(nodes: Vec, edges: &[(&str, &str)]) -> Vec { + let mut all = vec![ + json!({ "id": "in", "name": "in", "type": "inputNode", "content": { "schema": number_schema() } }), + ]; + all.extend(nodes); + all.push(json!({ "id": "out", "name": "out", "type": "outputNode", "content": {} })); + let edges: Vec = edges + .iter() + .enumerate() + .map(|(i, (a, b))| json!({ "id": format!("e{i}"), "sourceId": a, "targetId": b, "sourceHandle": null })) + .collect(); + let mut ws = Workspace::new(); + ws.set_document( + "g", + serde_json::from_value(json!({ "nodes": all, "edges": edges })).expect("graph"), + ); + ws.diagnostics("g") +} + +#[test] +fn schema_ranges_ignore_rewritten_fields() { + let direct = graph_diagnostics(vec![age_table()], &[("in", "dt"), ("dt", "out")]); + assert_eq!( + with_code(&direct, DiagnosticCode::CellCoversDomain).len(), + 1, + "{direct:?}" + ); + + let rewritten = graph_diagnostics( + vec![ + expression_node("calc", "applicant.age", "applicant.age + 1000"), + age_table(), + ], + &[("in", "calc"), ("calc", "dt"), ("dt", "out")], + ); + assert!( + with_code(&rewritten, DiagnosticCode::CellCoversDomain).is_empty(), + "{rewritten:?}" + ); +} + +fn diamonds(count: usize, rewrite: bool) -> Vec { + let mut nodes = Vec::new(); + let mut edges: Vec<(String, String)> = Vec::new(); + let mut previous = "in".to_string(); + for i in 0..count { + let (a, b, join) = (format!("a{i}"), format!("b{i}"), format!("j{i}")); + let key = if rewrite && i == 0 { + "applicant.age" + } else { + "applicant.seen" + }; + nodes.push(expression_node(&a, key, "applicant.age + 1")); + nodes.push(expression_node(&b, "applicant.seen", "applicant.age")); + nodes.push(expression_node(&join, "applicant.seen", "applicant.age")); + edges.push((previous.clone(), a.clone())); + edges.push((previous.clone(), b.clone())); + edges.push((a, join.clone())); + edges.push((b, join.clone())); + previous = join; + } + nodes.push(age_table()); + edges.push((previous, "dt".to_string())); + edges.push(("dt".to_string(), "out".to_string())); + let edges: Vec<(&str, &str)> = edges + .iter() + .map(|(a, b)| (a.as_str(), b.as_str())) + .collect(); + graph_diagnostics(nodes, &edges) +} + +fn summary(diagnostics: &[Diagnostic]) -> Vec { + let mut out: Vec = diagnostics + .iter() + .map(|d| format!("{:?} {:?} {}", d.code, d.severity, d.message)) + .collect(); + out.sort(); + out +} + +#[test] +fn chained_diamonds_stay_linear() { + let started = Instant::now(); + let found = diamonds(30, false); + assert!( + started.elapsed() < Duration::from_secs(1), + "{:?}", + started.elapsed() + ); + let covers = vec![ + "CellCoversDomain Hint this condition accepts every possible Age value, so the cell can be empty" + .to_string(), + ]; + assert_eq!(summary(&found), covers); + assert_eq!(summary(&diamonds(2, false)), covers); + assert_eq!( + summary(&diamonds(2, true)), + vec!["MissingCases Hint no row matches 1 input case: Age > 120".to_string()] + ); +} + +fn policy_content(blocks: Vec) -> DecisionContent { + let mut all = vec![applicant_model()]; + all.extend(blocks); + let policy: zen_engine::policy::PolicyDocument = + serde_json::from_value(json!({ "blocks": all })).expect("policy"); + DecisionContent::Policy(PolicyContent(Arc::new(policy))) +} + +fn gapped_table() -> Value { + json!({ "id": "dt", "type": "decisionTable", "props": { "data": { + "hitPolicy": "first", + "inputs": [ { "id": "i0", "name": "Age", "field": "applicant.age" } ], + "outputs": [ { "id": "o0", "name": "Rate", "field": "applicant.rate" } ], + "rules": [ + { "_id": "r1", "i0": "< 18", "o0": "1" }, + { "_id": "r2", "i0": "< 10", "o0": "2" } + ] + } } }) +} + +#[tokio::test] +async fn evaluate_skips_table_checks_but_keeps_errors() { + let table_codes = |diagnostics: &[Diagnostic]| { + diagnostics + .iter() + .filter(|d| { + matches!( + d.code, + DiagnosticCode::MissingCases | DiagnosticCode::UnreachableRule + ) + }) + .count() + }; + assert_eq!(table_codes(&policy_diagnostics(vec![gapped_table()])), 2); + + let loader = Arc::new(MemoryLoader::default()); + loader.add("ok", policy_content(vec![gapped_table()])); + loader.add( + "broken", + policy_content(vec![ + gapped_table(), + json!({ "id": "calc", "type": "expression", "props": { "data": { "key": "applicant.total", "value": "applicant.missing > 50" } } }), + ]), + ); + let engine = DecisionEngine::default().with_loader(loader); + + let result = engine + .evaluate("ok", json!({ "applicant": { "age": 5 } }).into()) + .await + .expect("evaluate"); + let output: Value = result.result.into(); + assert_eq!(output.pointer("/applicant/rate"), Some(&json!(1))); + + let result = engine + .evaluate("broken", json!({ "applicant": { "age": 5 } }).into()) + .await; + assert!( + format!("{result:?}").contains("CompilationErrors"), + "{result:?}" + ); +} diff --git a/core/engine/tests/table_verification.rs b/core/engine/tests/table_verification.rs index 35e1c2b1..5b07d4ee 100644 --- a/core/engine/tests/table_verification.rs +++ b/core/engine/tests/table_verification.rs @@ -1595,3 +1595,270 @@ fn covered_rows_with_the_same_result_are_redundant_hints() { assert!(d.message.contains("redundant"), "{}", d.message); } } + +#[test] +fn long_value_lists_and_or_chains_fit_a_small_stack() { + std::thread::Builder::new() + .stack_size(1 << 20) + .spawn(|| { + let list: Vec = (0..1500).map(|i| i.to_string()).collect(); + let chain: Vec = (0..1500).map(|i| format!("$ == {i}")).collect(); + for cell in [list.join(", "), chain.join(" or ")] { + let table = Table { + hit: "first", + inputs: &["applicant.age"], + outputs: &["applicant.discount"], + rows: leak_rows(vec![ + ("r1".to_string(), vec![cell.clone()], vec!["1".to_string()]), + ( + "r2".to_string(), + vec!["1499".to_string()], + vec!["2".to_string()], + ), + ]), + }; + table.assert_both(&[ + "UnreachableRule r2 coveredByIds=r1 example={\"applicant\":{\"age\":1499}}", + ]); + } + }) + .expect("thread") + .join() + .expect("no stack overflow"); +} + +#[test] +fn extreme_decimal_bounds_do_not_overflow() { + for (cells, gap) in [ + (["<= 79228162514264337593543950335", "> 0"], None), + ( + ["< -79228162514264337593543950335", "> 0"], + Some("[-79228162514264337593543950335..0]"), + ), + ( + [ + "(79228162514264337593543950333..79228162514264337593543950335]", + "< 0", + ], + Some("[0..79228162514264337593543950333], > 79228162514264337593543950335"), + ), + ( + [ + "> -79228162514264337593543950335 and < -79228162514264337593543950334", + "> 0", + ], + Some("<= -79228162514264337593543950335, [-79228162514264337593543950334..0]"), + ), + ] { + let table = Table { + hit: "first", + inputs: &["applicant.age"], + outputs: &["applicant.discount"], + rows: leak_rows(vec![ + ( + "r1".to_string(), + vec![cells[0].to_string()], + vec!["1".to_string()], + ), + ( + "r2".to_string(), + vec![cells[1].to_string()], + vec!["2".to_string()], + ), + ]), + }; + for gaps in [ + Table::gaps(table.policy_diagnostics()), + Table::gaps(table.graph_diagnostics()), + ] { + let found = gaps.map(|gaps| gaps.cases[0]["cells"]["i0"].clone()); + assert_eq!(found, gap.map(|g| json!(g)), "{cells:?}"); + } + } +} + +#[test] +fn spaced_random_calls_are_never_equal() { + let table = Table { + hit: "first", + inputs: &["", "applicant.age"], + outputs: &["applicant.discount"], + rows: &[ + ("r1", &["rand (10) > 5", ""], &["1"]), + ("r2", &["rand (10) > 5", ""], &["1"]), + ], + }; + table.assert_both(&[]); + table.assert_compressed(None); +} + +#[tokio::test] +async fn compressed_strings_keep_their_quotes_and_backslashes() { + let table = Table { + hit: "first", + inputs: &["applicant.code", "applicant.age"], + outputs: &["applicant.discount"], + rows: &[ + ("r1", &["'a\"b'", "< 18"], &["0.1"]), + ("r2", &["\"c\\d\"", "< 18"], &["0.1"]), + ("r3", &["\"e\"", "< 18"], &["0.1"]), + ("r4", &["\"e\"", ">= 18"], &["0.2"]), + ], + }; + let original = table.content(); + let inputs: Vec = ["a\"b", "c\\d", "e", "f"] + .iter() + .flat_map(|code| { + [10, 30].map(|age| { + json!({ "applicant": { "tier": "gold", "code": code, "age": age, "scores": [], "vip": false } }) + }) + }) + .collect(); + for diagnostics in [table.policy_diagnostics(), table.graph_diagnostics()] { + let (before, rules) = compressed(diagnostics).expect("compressible"); + assert_eq!(before, 4); + let codes = row_summary(&rules); + assert!( + codes[0].starts_with("i0='a\"b', \"c\\d\", \"e\" "), + "{codes:?}" + ); + let mut compact = original.clone(); + compact["rules"] = rules; + let before = outputs_for(original.clone(), &inputs).await; + let after = outputs_for(compact.clone(), &inputs).await; + assert_eq!(before.0, after.0, "policy {compact}"); + assert_eq!(before.1, after.1, "graph {compact}"); + } +} + +#[tokio::test] +async fn compressed_exclusions_still_accept_lists_and_objects() { + let table = Table { + hit: "first", + inputs: &["applicant.code"], + outputs: &["applicant.discount"], + rows: &[ + ("r1", &["!= \"a\" and != \"b\" and != 1"], &["1"]), + ("r2", &["1"], &["1"]), + ("r3", &[""], &["2"]), + ], + }; + let original = table.content(); + let (_, rules) = compressed(table.graph_diagnostics()).expect("compressible"); + assert_eq!( + row_summary(&rules), + vec!["i0=!= \"a\" and != \"b\" o0=1", "i0= o0=2"] + ); + let mut compact = original.clone(); + compact["rules"] = rules; + let decision = |content: &Value| { + let mut graph = table.graph_json(); + for node in graph["nodes"].as_array_mut().expect("nodes") { + if node["id"] == "dt" { + node["content"] = content.clone(); + } + } + let DecisionContent::Graph(graph) = serde_json::from_value(graph).expect("graph") else { + panic!("graph"); + }; + Decision::from(graph) + }; + let (before, after) = (decision(&original), decision(&compact)); + for code in [ + json!([1]), + json!({ "k": 1 }), + Value::Null, + json!(1), + json!(2), + json!("a"), + json!("b"), + json!("c"), + json!(true), + ] { + let input = json!({ "applicant": { "tier": "gold", "code": code, "age": 1, "scores": [], "vip": false } }); + let outcome = |result: Result| -> Value { + match result { + Ok(response) => { + let output: Value = response.result.into(); + output + .pointer("/applicant/discount") + .cloned() + .unwrap_or(Value::Null) + } + Err(_) => json!("error"), + } + }; + assert_eq!( + outcome(after.evaluate(input.clone().into()).await), + outcome(before.evaluate(input.clone().into()).await), + "{code}" + ); + } +} + +#[tokio::test] +async fn collect_compression_keeps_the_order_of_results() { + let crossing = Table { + hit: "collect", + inputs: &["applicant.age"], + outputs: &["applicant.discount"], + rows: &[ + ("r1", &["< 18"], &["1"]), + ("r2", &["[20..40]"], &["2"]), + ("r3", &["[18..30]"], &["1"]), + ], + }; + crossing.assert_compressed(None); + + let apart = Table { + hit: "collect", + inputs: &["applicant.age"], + outputs: &["applicant.discount"], + rows: &[ + ("r1", &["< 18"], &["1"]), + ("r2", &["> 50"], &["2"]), + ("r3", &["[18..30]"], &["1"]), + ], + }; + apart.assert_compressed(Some((3, &["i0=<= 30 o0=1", "i0=> 50 o0=2"]))); + let original = apart.content(); + let (_, rules) = compressed(apart.policy_diagnostics()).expect("compressible"); + let mut compact = original.clone(); + compact["rules"] = rules; + let inputs: Vec = [10, 18, 25, 30, 35, 60] + .iter() + .map(|age| json!({ "applicant": { "tier": "gold", "age": age, "scores": [], "vip": false } })) + .collect(); + assert_eq!( + outputs_for(original.clone(), &inputs).await.0, + outputs_for(compact.clone(), &inputs).await.0 + ); + let graph = |content: &Value| { + let mut graph = apart.graph_json(); + for node in graph["nodes"].as_array_mut().expect("nodes") { + if node["id"] == "dt" { + node["content"] = content.clone(); + } + } + let DecisionContent::Graph(graph) = serde_json::from_value(graph).expect("graph") else { + panic!("graph"); + }; + Decision::from(graph) + }; + let (before, after) = (graph(&original), graph(&compact)); + for input in inputs { + let before: Value = before + .evaluate(input.clone().into()) + .await + .expect("graph") + .result + .into(); + let after: Value = after + .evaluate(input.clone().into()) + .await + .expect("graph") + .result + .into(); + assert_eq!(before, after, "{input}"); + } +}