1use std::collections::{HashMap, HashSet};
4use std::path::{Path, PathBuf};
5
6use xlog_core::{symbol, Result, ScalarType, XlogError};
7use xlog_logic::ast::{BodyLiteral, CompOp, Program, Term};
8use xlog_logic::{
9 compare_arithmetic_values, evaluate_arithmetic_expression, format_atom, format_term,
10 generated_function_variable_sources, source_format_normalized_alternative, ArithmeticValue,
11 Lowerer,
12};
13
14use super::{json_escape, json_string_array};
15
16pub(super) struct GeneratedRuleDiagnostic {
17 rule_head: String,
18 source_relation: String,
19 row_decisions: Vec<GeneratedRuleRowDecision>,
20}
21
22struct GeneratedRuleRowDecision {
23 row_key: String,
24 accepted: bool,
25 failed_predicates: Vec<String>,
26 threshold_comparisons: Vec<ThresholdComparison>,
27 aggregate_inputs: Vec<String>,
28}
29
30#[derive(Clone)]
31struct ThresholdComparison {
32 predicate: String,
33 left: String,
34 op: String,
35 right: String,
36 left_value: String,
37 right_value: String,
38 passed: bool,
39}
40
41struct GeneratedRuleEvaluation {
42 accepted: bool,
43 failed_predicates: Vec<String>,
44 threshold_comparisons: Vec<ThresholdComparison>,
45}
46
47enum GeneratedRuleSearchTask {
48 Continue {
49 literal_index: usize,
50 bindings: DiagnosticBindings,
51 threshold_comparisons: Vec<ThresholdComparison>,
52 },
53 TryPositiveRows {
54 literal_index: usize,
55 row_index: usize,
56 bindings: DiagnosticBindings,
57 threshold_comparisons: Vec<ThresholdComparison>,
58 },
59}
60
61#[derive(Clone)]
62struct DiagnosticScalar {
63 value: ArithmeticValue,
64 scalar_type: ScalarType,
65 label: String,
66}
67
68type DiagnosticBindings = HashMap<String, DiagnosticScalar>;
69type DiagnosticRow = Vec<DiagnosticScalar>;
70type DiagnosticRelationRows = HashMap<(String, usize), Vec<DiagnosticRow>>;
71
72pub(super) fn explain_generated_rule_diagnostics(
73 source_program: &Program,
74 analysis_program: &Program,
75 source_path: Option<&Path>,
76) -> Result<Vec<GeneratedRuleDiagnostic>> {
77 let function_variable_sources =
78 generated_function_variable_sources(source_program, analysis_program);
79 let source_predicates = analysis_program
80 .rules
81 .iter()
82 .filter(|rule| !rule.body.is_empty() && generated_rule_candidate(rule))
83 .filter_map(diagnostic_source_atom)
84 .map(|atom| atom.predicate.clone())
85 .collect::<HashSet<_>>();
86 let external_rows = source_path
87 .map(|path| load_external_relation_rows(analysis_program, path, &source_predicates))
88 .transpose()?
89 .unwrap_or_default();
90 let relation_rows = diagnostic_relation_rows(analysis_program, &external_rows)?;
91 let mut diagnostics = Vec::new();
92 for rule in analysis_program
93 .rules
94 .iter()
95 .filter(|rule| !rule.body.is_empty() && generated_rule_candidate(rule))
96 {
97 let Some(source_atom) = diagnostic_source_atom(rule) else {
98 continue;
99 };
100 ensure_extensional_generated_rule_support(analysis_program, rule)?;
101 let mut row_decisions = Vec::new();
102 for source_row in relation_rows_for_atom(&relation_rows, source_atom) {
103 let Some(bindings) = bindings_for_source_row(source_atom, source_row)? else {
104 continue;
105 };
106 let evaluation = evaluate_generated_rule(
107 &relation_rows,
108 rule,
109 source_atom,
110 bindings,
111 &function_variable_sources,
112 )?;
113 row_decisions.push(GeneratedRuleRowDecision {
114 row_key: source_row
115 .first()
116 .map(|value| value.label.clone())
117 .unwrap_or_else(|| source_atom.predicate.clone()),
118 accepted: evaluation.accepted,
119 failed_predicates: evaluation.failed_predicates,
120 threshold_comparisons: evaluation.threshold_comparisons,
121 aggregate_inputs: vec![format!(
122 "{}({})",
123 source_atom.predicate,
124 source_row
125 .iter()
126 .map(|value| value.label.clone())
127 .collect::<Vec<_>>()
128 .join(", ")
129 )],
130 });
131 }
132
133 if !row_decisions.is_empty() {
134 diagnostics.push(GeneratedRuleDiagnostic {
135 rule_head: rule.head.predicate.clone(),
136 source_relation: source_atom.predicate.clone(),
137 row_decisions,
138 });
139 }
140 }
141 Ok(diagnostics)
142}
143
144fn ensure_extensional_generated_rule_support(
145 program: &Program,
146 rule: &xlog_logic::ast::Rule,
147) -> Result<()> {
148 let probabilistic_predicates = program
149 .prob_facts
150 .iter()
151 .map(|fact| fact.atom.predicate.as_str())
152 .chain(
153 program
154 .annotated_disjunctions
155 .iter()
156 .flat_map(|disjunction| {
157 disjunction
158 .choices
159 .iter()
160 .map(|choice| choice.atom.predicate.as_str())
161 }),
162 )
163 .collect::<HashSet<_>>();
164 let intensional_predicates = program
165 .rules
166 .iter()
167 .filter(|candidate| !candidate.body.is_empty())
168 .map(|candidate| candidate.head.predicate.as_str())
169 .collect::<HashSet<_>>();
170 let body_predicates = || {
171 rule.body.iter().filter_map(|literal| match literal {
172 BodyLiteral::Positive(atom) | BodyLiteral::Negated(atom) => {
173 Some(atom.predicate.as_str())
174 }
175 BodyLiteral::Epistemic(_)
176 | BodyLiteral::Comparison(_)
177 | BodyLiteral::IsExpr(_)
178 | BodyLiteral::Univ(_) => None,
179 })
180 };
181 if let Some(predicate) =
182 body_predicates().find(|predicate| probabilistic_predicates.contains(predicate))
183 {
184 return Err(XlogError::Compilation(format!(
185 "generated-rule diagnostics do not assign deterministic row decisions to probabilistic predicate '{predicate}'"
186 )));
187 }
188 if let Some(predicate) =
189 body_predicates().find(|predicate| intensional_predicates.contains(predicate))
190 {
191 return Err(XlogError::Compilation(format!(
192 "generated-rule diagnostics require materialized rows for derived predicate '{predicate}'"
193 )));
194 }
195 Ok(())
196}
197
198fn generated_rule_candidate(rule: &xlog_logic::ast::Rule) -> bool {
199 rule.head.predicate.starts_with("generated_")
200 || rule.head.predicate.starts_with("xlog_accepted_")
201 || rule.head.predicate.starts_with("xlog_rejected_")
202 || rule.body.iter().any(|literal| match literal {
203 BodyLiteral::Positive(atom) | BodyLiteral::Negated(atom) => {
204 diagnostic_source_predicate(&atom.predicate)
205 }
206 BodyLiteral::Epistemic(_) => false,
207 BodyLiteral::Comparison(_) | BodyLiteral::IsExpr(_) | BodyLiteral::Univ(_) => false,
208 })
209}
210
211fn diagnostic_source_atom(rule: &xlog_logic::ast::Rule) -> Option<&xlog_logic::ast::Atom> {
212 rule.body.iter().find_map(|literal| match literal {
213 BodyLiteral::Positive(atom) if diagnostic_source_predicate(&atom.predicate) => Some(atom),
214 _ => None,
215 })
216}
217
218fn diagnostic_source_predicate(predicate: &str) -> bool {
219 predicate.starts_with("generated_")
220 || predicate.ends_with("_candidate_input")
221 || (predicate.contains("candidate") && predicate.ends_with("_input"))
222}
223
224fn diagnostic_relation_rows(
225 program: &Program,
226 external_rows: &HashMap<String, Vec<DiagnosticRow>>,
227) -> Result<DiagnosticRelationRows> {
228 let mut lowerer = Lowerer::new();
229 lowerer.infer_and_validate_schemas(program)?;
230 let schemas = lowerer.schemas();
231 let mut relation_rows = DiagnosticRelationRows::new();
232 for fact in program.rules.iter().filter(|rule| rule.body.is_empty()) {
233 let row = diagnostic_row(&fact.head.predicate, &fact.head.terms, schemas)?;
234 relation_rows
235 .entry((fact.head.predicate.clone(), fact.head.terms.len()))
236 .or_default()
237 .push(row);
238 }
239 for (predicate, rows) in external_rows {
240 for row in rows {
241 relation_rows
242 .entry((predicate.clone(), row.len()))
243 .or_default()
244 .push(row.clone());
245 }
246 }
247 Ok(relation_rows)
248}
249
250fn diagnostic_row(
251 predicate: &str,
252 terms: &[Term],
253 schemas: &HashMap<String, xlog_core::Schema>,
254) -> Result<DiagnosticRow> {
255 let schema = schemas.get(predicate).ok_or_else(|| {
256 XlogError::Compilation(format!(
257 "generated-rule diagnostics require a schema for predicate '{predicate}'"
258 ))
259 })?;
260 if schema.arity() != terms.len() {
261 return Err(XlogError::Compilation(format!(
262 "generated-rule diagnostic row for '{predicate}' has {} values but its schema has {}",
263 terms.len(),
264 schema.arity()
265 )));
266 }
267 terms
268 .iter()
269 .enumerate()
270 .map(|(index, term)| {
271 let scalar_type = schema.column_type(index).ok_or_else(|| {
272 XlogError::Compilation(format!(
273 "generated-rule diagnostics require a type for '{predicate}' column {}",
274 index + 1
275 ))
276 })?;
277 let value = ArithmeticValue::from_typed_term(term, scalar_type)?;
278 Ok(DiagnosticScalar {
279 label: arithmetic_value_label(&value),
280 value,
281 scalar_type,
282 })
283 })
284 .collect()
285}
286
287fn relation_rows_for_atom<'a>(
288 relation_rows: &'a DiagnosticRelationRows,
289 atom: &xlog_logic::ast::Atom,
290) -> &'a [DiagnosticRow] {
291 relation_rows
292 .get(&(atom.predicate.clone(), atom.terms.len()))
293 .map(Vec::as_slice)
294 .unwrap_or(&[])
295}
296
297fn bindings_for_source_row(
298 atom: &xlog_logic::ast::Atom,
299 row: &[DiagnosticScalar],
300) -> Result<Option<DiagnosticBindings>> {
301 extend_bindings_for_row(atom, row, &HashMap::new())
302}
303
304fn extend_bindings_for_row(
305 atom: &xlog_logic::ast::Atom,
306 row: &[DiagnosticScalar],
307 initial: &DiagnosticBindings,
308) -> Result<Option<DiagnosticBindings>> {
309 if atom.terms.len() != row.len() {
310 return Ok(None);
311 }
312 let mut bindings = initial.clone();
313 for (pattern, value) in atom.terms.iter().zip(row) {
314 match pattern {
315 Term::Variable(name) => {
316 if let Some(existing) = bindings.get(name) {
317 if existing.scalar_type != value.scalar_type
318 || !compare_arithmetic_values(&existing.value, CompOp::Eq, &value.value)?
319 {
320 return Ok(None);
321 }
322 } else {
323 bindings.insert(name.clone(), value.clone());
324 }
325 }
326 Term::Anonymous => {}
327 _ => {
328 let pattern_value = ArithmeticValue::from_typed_term(pattern, value.scalar_type)?;
329 if !compare_arithmetic_values(&pattern_value, CompOp::Eq, &value.value)? {
330 return Ok(None);
331 }
332 }
333 }
334 }
335 Ok(Some(bindings))
336}
337
338fn load_external_relation_rows(
339 program: &Program,
340 source_path: &Path,
341 source_predicates: &HashSet<String>,
342) -> Result<HashMap<String, Vec<DiagnosticRow>>> {
343 let mut lowerer = Lowerer::new();
344 lowerer.infer_and_validate_schemas(program)?;
345 let manifest = external_relation_manifest(source_path);
346 if manifest.is_some() && source_predicates.len() > 1 {
347 let mut predicates = source_predicates.iter().cloned().collect::<Vec<_>>();
348 predicates.sort();
349 return Err(XlogError::Compilation(format!(
350 "candidate relation manifest is ambiguous for generated-rule source predicates: {}",
351 predicates.join(", ")
352 )));
353 }
354 let manifest_predicate = source_predicates.iter().next();
355 let mut loaded = HashMap::new();
356 for decl in &program.predicates {
357 let manifest_source = manifest_predicate
358 .filter(|predicate| *predicate == &decl.name)
359 .and(manifest.as_ref());
360 let Some((relation_path, columns)) =
361 external_relation_source(source_path, decl, manifest_source)
362 else {
363 continue;
364 };
365 if columns.len() != decl.arity() {
366 continue;
367 }
368 let schema = lowerer.schemas().get(&decl.name).ok_or_else(|| {
369 XlogError::Compilation(format!(
370 "generated-rule diagnostics require a schema for external predicate '{}'",
371 decl.name
372 ))
373 })?;
374 let source = std::fs::read_to_string(&relation_path).map_err(|error| {
375 XlogError::Compilation(format!(
376 "cannot read external relation '{}': {error}",
377 relation_path.display()
378 ))
379 })?;
380 let json = serde_json::from_str::<serde_json::Value>(&source).map_err(|error| {
381 XlogError::Compilation(format!(
382 "invalid JSON in external relation '{}': {error}",
383 relation_path.display()
384 ))
385 })?;
386 let rows = json
387 .get("rows")
388 .and_then(serde_json::Value::as_array)
389 .ok_or_else(|| {
390 XlogError::Compilation(format!(
391 "external relation '{}' must contain a rows array",
392 relation_path.display()
393 ))
394 })?;
395 let mut relation_rows = Vec::new();
396 for (row_index, row) in rows.iter().enumerate() {
397 let object = row.as_object().ok_or_else(|| {
398 XlogError::Compilation(format!(
399 "external relation '{}' row {} must be an object",
400 relation_path.display(),
401 row_index + 1
402 ))
403 })?;
404 let mut values = Vec::with_capacity(columns.len());
405 for (column_index, column) in columns.iter().enumerate() {
406 let scalar_type = schema.column_type(column_index).ok_or_else(|| {
407 XlogError::Compilation(format!(
408 "external predicate '{}' has no type for column '{}'",
409 decl.name, column
410 ))
411 })?;
412 let value = object.get(column).ok_or_else(|| {
413 XlogError::Compilation(format!(
414 "external relation '{}' row {} is missing column '{}'",
415 relation_path.display(),
416 row_index + 1,
417 column
418 ))
419 })?;
420 values.push(json_diagnostic_scalar(value, scalar_type).map_err(|error| {
421 XlogError::Compilation(format!(
422 "external relation '{}' row {} column '{}': {error}",
423 relation_path.display(),
424 row_index + 1,
425 column
426 ))
427 })?);
428 }
429 relation_rows.push(values);
430 }
431 if !relation_rows.is_empty() {
432 loaded.insert(decl.name.clone(), relation_rows);
433 }
434 }
435 Ok(loaded)
436}
437
438fn external_relation_source(
439 source_path: &Path,
440 decl: &xlog_logic::ast::PredDecl,
441 manifest: Option<&(PathBuf, Vec<String>)>,
442) -> Option<(PathBuf, Vec<String>)> {
443 if let Some((relation_path, columns)) = manifest {
444 if columns.len() == decl.arity() {
445 return Some((relation_path.clone(), columns.clone()));
446 }
447 }
448 let columns = declared_column_names(decl)?;
449 let source_dir = source_path.parent()?;
450 for candidate in relation_json_candidates(source_dir, &decl.name) {
451 if candidate.exists() {
452 return Some((candidate, columns));
453 }
454 }
455 None
456}
457
458fn external_relation_manifest(source_path: &Path) -> Option<(PathBuf, Vec<String>)> {
459 let source_dir = source_path.parent()?;
460 let mut manifests = vec![source_dir.join("xlog_hypothesis_execution.json")];
461 if let Some(parent) = source_dir.parent() {
462 manifests.push(parent.join("xlog_hypothesis_execution.json"));
463 }
464 for manifest_path in manifests {
465 let Ok(source) = std::fs::read_to_string(&manifest_path) else {
466 continue;
467 };
468 let Ok(json) = serde_json::from_str::<serde_json::Value>(&source) else {
469 continue;
470 };
471 let Some(columns) = json
472 .get("relation_input_columns")
473 .and_then(serde_json::Value::as_array)
474 .map(|items| {
475 items
476 .iter()
477 .filter_map(serde_json::Value::as_str)
478 .map(ToString::to_string)
479 .collect::<Vec<_>>()
480 })
481 else {
482 continue;
483 };
484 let Some(path_value) = json
485 .get("relation_input_path")
486 .and_then(serde_json::Value::as_str)
487 else {
488 continue;
489 };
490 let relation_path = PathBuf::from(path_value);
491 let relation_path = if relation_path.is_absolute() {
492 relation_path
493 } else {
494 manifest_path
495 .parent()
496 .unwrap_or_else(|| Path::new("."))
497 .join(relation_path)
498 };
499 if relation_path.exists() {
500 return Some((relation_path, columns));
501 }
502 }
503 None
504}
505
506fn declared_column_names(decl: &xlog_logic::ast::PredDecl) -> Option<Vec<String>> {
507 decl.schema_columns()
508 .into_iter()
509 .map(|column| column.name)
510 .collect()
511}
512
513fn relation_json_candidates(source_dir: &Path, predicate: &str) -> Vec<PathBuf> {
514 let mut candidates = vec![source_dir.join(format!("{predicate}.json"))];
515 if let Some(stem) = predicate.strip_suffix("_input") {
516 candidates.push(source_dir.join(format!("{stem}_relation.json")));
517 }
518 candidates
519}
520
521fn json_diagnostic_scalar(
522 value: &serde_json::Value,
523 scalar_type: ScalarType,
524) -> Result<DiagnosticScalar> {
525 let arithmetic = match scalar_type {
526 ScalarType::I32 => ArithmeticValue::I32(
527 value
528 .as_i64()
529 .and_then(|value| i32::try_from(value).ok())
530 .ok_or_else(|| XlogError::Compilation("expected an i32 JSON value".to_string()))?,
531 ),
532 ScalarType::I64 => ArithmeticValue::I64(
533 value
534 .as_i64()
535 .ok_or_else(|| XlogError::Compilation("expected an i64 JSON value".to_string()))?,
536 ),
537 ScalarType::U32 => ArithmeticValue::U32(
538 value
539 .as_u64()
540 .and_then(|value| u32::try_from(value).ok())
541 .ok_or_else(|| XlogError::Compilation("expected a u32 JSON value".to_string()))?,
542 ),
543 ScalarType::U64 => ArithmeticValue::U64(
544 value
545 .as_u64()
546 .ok_or_else(|| XlogError::Compilation("expected a u64 JSON value".to_string()))?,
547 ),
548 ScalarType::F32 => ArithmeticValue::F32(
549 value
550 .as_f64()
551 .filter(|value| value.is_finite() && value.abs() <= f64::from(f32::MAX))
552 .ok_or_else(|| {
553 XlogError::Compilation("expected a finite f32 JSON value".to_string())
554 })? as f32,
555 ),
556 ScalarType::F64 => ArithmeticValue::F64(
557 value
558 .as_f64()
559 .ok_or_else(|| XlogError::Compilation("expected an f64 JSON value".to_string()))?,
560 ),
561 ScalarType::Bool => {
562 ArithmeticValue::Bool(value.as_bool().ok_or_else(|| {
563 XlogError::Compilation("expected a boolean JSON value".to_string())
564 })?)
565 }
566 ScalarType::Symbol => {
567 ArithmeticValue::Symbol(symbol::intern(value.as_str().ok_or_else(|| {
568 XlogError::Compilation("expected a string JSON value".to_string())
569 })?))
570 }
571 };
572 Ok(DiagnosticScalar {
573 label: match value {
574 serde_json::Value::String(value) => value.clone(),
575 _ => value.to_string(),
576 },
577 value: arithmetic,
578 scalar_type,
579 })
580}
581
582fn evaluate_generated_rule(
583 relation_rows: &DiagnosticRelationRows,
584 rule: &xlog_logic::ast::Rule,
585 source_atom: &xlog_logic::ast::Atom,
586 bindings: DiagnosticBindings,
587 function_variable_sources: &HashMap<String, String>,
588) -> Result<GeneratedRuleEvaluation> {
589 let mut tasks = vec![GeneratedRuleSearchTask::Continue {
590 literal_index: 0,
591 bindings,
592 threshold_comparisons: Vec::new(),
593 }];
594 let mut first_failure = None;
595
596 while let Some(task) = tasks.pop() {
597 match task {
598 GeneratedRuleSearchTask::Continue {
599 literal_index,
600 mut bindings,
601 mut threshold_comparisons,
602 } => {
603 let Some(literal) = rule.body.get(literal_index) else {
604 return Ok(GeneratedRuleEvaluation {
605 accepted: true,
606 failed_predicates: Vec::new(),
607 threshold_comparisons,
608 });
609 };
610 match literal {
611 BodyLiteral::Positive(atom) if std::ptr::eq(atom, source_atom) => {
612 tasks.push(GeneratedRuleSearchTask::Continue {
613 literal_index: literal_index + 1,
614 bindings,
615 threshold_comparisons,
616 });
617 }
618 BodyLiteral::Positive(_) => {
619 tasks.push(GeneratedRuleSearchTask::TryPositiveRows {
620 literal_index,
621 row_index: 0,
622 bindings,
623 threshold_comparisons,
624 });
625 }
626 BodyLiteral::Negated(atom) => {
627 let mut matched = false;
628 for row in relation_rows_for_atom(relation_rows, atom) {
629 if extend_bindings_for_row(atom, row, &bindings)?.is_some() {
630 matched = true;
631 break;
632 }
633 }
634 if matched {
635 first_failure.get_or_insert_with(|| GeneratedRuleEvaluation {
636 accepted: false,
637 failed_predicates: vec![source_format_normalized_alternative(
638 &format!("not {}", format_atom(atom)),
639 function_variable_sources,
640 )],
641 threshold_comparisons,
642 });
643 } else {
644 tasks.push(GeneratedRuleSearchTask::Continue {
645 literal_index: literal_index + 1,
646 bindings,
647 threshold_comparisons,
648 });
649 }
650 }
651 BodyLiteral::Comparison(comparison) => {
652 let report =
653 threshold_comparison(comparison, &bindings, function_variable_sources)?;
654 let passed = report.passed;
655 let predicate = report.predicate.clone();
656 threshold_comparisons.push(report);
657 if passed {
658 tasks.push(GeneratedRuleSearchTask::Continue {
659 literal_index: literal_index + 1,
660 bindings,
661 threshold_comparisons,
662 });
663 } else {
664 first_failure.get_or_insert(GeneratedRuleEvaluation {
665 accepted: false,
666 failed_predicates: vec![predicate],
667 threshold_comparisons,
668 });
669 }
670 }
671 BodyLiteral::IsExpr(binding) => {
672 let value = arithmetic_expression_value(&binding.expr, &bindings)?;
673 let compatible = match bindings.get(&binding.target) {
674 Some(existing) => {
675 existing.scalar_type == value.scalar_type
676 && compare_arithmetic_values(
677 &existing.value,
678 CompOp::Eq,
679 &value.value,
680 )?
681 }
682 None => {
683 bindings.insert(binding.target.clone(), value);
684 true
685 }
686 };
687 if compatible {
688 tasks.push(GeneratedRuleSearchTask::Continue {
689 literal_index: literal_index + 1,
690 bindings,
691 threshold_comparisons,
692 });
693 } else {
694 first_failure.get_or_insert(GeneratedRuleEvaluation {
695 accepted: false,
696 failed_predicates: vec![source_format_normalized_alternative(
697 &binding.target,
698 function_variable_sources,
699 )],
700 threshold_comparisons,
701 });
702 }
703 }
704 BodyLiteral::Epistemic(_) => {
705 return Err(XlogError::Compilation(
706 "generated-rule diagnostics do not evaluate epistemic literals"
707 .to_string(),
708 ));
709 }
710 BodyLiteral::Univ(_) => {
711 return Err(XlogError::Compilation(
712 "generated-rule diagnostics do not evaluate univ literals".to_string(),
713 ));
714 }
715 }
716 }
717 GeneratedRuleSearchTask::TryPositiveRows {
718 literal_index,
719 mut row_index,
720 bindings,
721 threshold_comparisons,
722 } => {
723 let BodyLiteral::Positive(atom) = &rule.body[literal_index] else {
724 return Err(XlogError::Compilation(
725 "invalid generated-rule diagnostic search state".to_string(),
726 ));
727 };
728 let rows = relation_rows_for_atom(relation_rows, atom);
729 let mut matched = None;
730 while let Some(row) = rows.get(row_index) {
731 row_index += 1;
732 if let Some(next_bindings) = extend_bindings_for_row(atom, row, &bindings)? {
733 matched = Some(next_bindings);
734 break;
735 }
736 }
737 if let Some(next_bindings) = matched {
738 tasks.push(GeneratedRuleSearchTask::TryPositiveRows {
739 literal_index,
740 row_index,
741 bindings,
742 threshold_comparisons: threshold_comparisons.clone(),
743 });
744 tasks.push(GeneratedRuleSearchTask::Continue {
745 literal_index: literal_index + 1,
746 bindings: next_bindings,
747 threshold_comparisons,
748 });
749 } else {
750 first_failure.get_or_insert_with(|| GeneratedRuleEvaluation {
751 accepted: false,
752 failed_predicates: vec![source_format_normalized_alternative(
753 &format_atom(atom),
754 function_variable_sources,
755 )],
756 threshold_comparisons,
757 });
758 }
759 }
760 }
761 }
762
763 Ok(first_failure.unwrap_or(GeneratedRuleEvaluation {
764 accepted: false,
765 failed_predicates: Vec::new(),
766 threshold_comparisons: Vec::new(),
767 }))
768}
769
770fn threshold_comparison(
771 comparison: &xlog_logic::ast::Comparison,
772 bindings: &DiagnosticBindings,
773 function_variable_sources: &HashMap<String, String>,
774) -> Result<ThresholdComparison> {
775 let left = source_format_normalized_alternative(
776 &format_term(&comparison.left),
777 function_variable_sources,
778 );
779 let right = source_format_normalized_alternative(
780 &format_term(&comparison.right),
781 function_variable_sources,
782 );
783 let bound_left = bound_scalar(&comparison.left, bindings);
784 let bound_right = bound_scalar(&comparison.right, bindings);
785 let left_value = comparison_scalar(
786 &comparison.left,
787 bound_left.as_ref(),
788 bound_right.as_ref().map(|value| value.scalar_type),
789 )?;
790 let right_value = comparison_scalar(
791 &comparison.right,
792 bound_right.as_ref(),
793 bound_left.as_ref().map(|value| value.scalar_type),
794 )?;
795 let passed = compare_arithmetic_values(&left_value.value, comparison.op, &right_value.value)?;
796 Ok(ThresholdComparison {
797 predicate: format!("{left} {} {right}", comp_op_label(comparison.op)),
798 left,
799 op: comp_op_label(comparison.op).to_string(),
800 right,
801 left_value: left_value.label,
802 right_value: right_value.label,
803 passed,
804 })
805}
806
807fn arithmetic_expression_value(
808 expression: &xlog_logic::ast::ArithExpr,
809 bindings: &DiagnosticBindings,
810) -> Result<DiagnosticScalar> {
811 let evaluator_bindings = bindings
812 .iter()
813 .map(|(name, value)| (name.clone(), value.value.clone()))
814 .collect::<HashMap<_, _>>();
815 let value = evaluate_arithmetic_expression(expression, &evaluator_bindings)?;
816 let scalar_type = value.scalar_type().ok_or_else(|| {
817 XlogError::Compilation(
818 "generated-rule arithmetic produced a value without a runtime scalar type".to_string(),
819 )
820 })?;
821 Ok(DiagnosticScalar {
822 label: arithmetic_value_label(&value),
823 scalar_type,
824 value,
825 })
826}
827
828fn bound_scalar(term: &Term, bindings: &DiagnosticBindings) -> Option<DiagnosticScalar> {
829 match term {
830 Term::Variable(name) => bindings.get(name).cloned(),
831 _ => None,
832 }
833}
834
835fn comparison_scalar(
836 term: &Term,
837 bound: Option<&DiagnosticScalar>,
838 peer_type: Option<ScalarType>,
839) -> Result<DiagnosticScalar> {
840 if let Some(bound) = bound {
841 return Ok(bound.clone());
842 }
843 if matches!(term, Term::Variable(_)) {
844 return Err(XlogError::Compilation(format!(
845 "Unbound variable {} in generated-rule comparison",
846 format_term(term)
847 )));
848 }
849 let value = if let Some(expected) = peer_type {
850 ArithmeticValue::from_typed_term(term, expected)?
851 } else {
852 ArithmeticValue::from_term(term)?
853 };
854 let scalar_type = value.scalar_type().ok_or_else(|| {
855 XlogError::Compilation(
856 "generated-rule comparison requires a runtime scalar type".to_string(),
857 )
858 })?;
859 Ok(DiagnosticScalar {
860 value,
861 scalar_type,
862 label: format_term(term),
863 })
864}
865
866fn arithmetic_value_label(value: &ArithmeticValue) -> String {
867 match value {
868 ArithmeticValue::I32(value) => value.to_string(),
869 ArithmeticValue::I64(value) => value.to_string(),
870 ArithmeticValue::U32(value) => value.to_string(),
871 ArithmeticValue::U64(value) => value.to_string(),
872 ArithmeticValue::F32(value) => value.to_string(),
873 ArithmeticValue::F64(value) => value.to_string(),
874 ArithmeticValue::Bool(value) => value.to_string(),
875 ArithmeticValue::Symbol(value) => symbol::resolve(*value),
876 ArithmeticValue::String(value) => value.clone(),
877 }
878}
879
880fn comp_op_label(op: CompOp) -> &'static str {
881 match op {
882 CompOp::Eq => "==",
883 CompOp::Ne => "!=",
884 CompOp::Lt => "<",
885 CompOp::Le => "<=",
886 CompOp::Gt => ">",
887 CompOp::Ge => ">=",
888 }
889}
890
891pub(super) fn print_generated_rule_diagnostics_json(entries: &[GeneratedRuleDiagnostic]) {
892 println!(" \"generated_rule_diagnostics\": [");
893 for (idx, entry) in entries.iter().enumerate() {
894 let suffix = if idx + 1 == entries.len() { "" } else { "," };
895 println!(" {{");
896 println!(
897 " \"rule_head\": \"{}\",",
898 json_escape(&entry.rule_head)
899 );
900 println!(
901 " \"source_relation\": \"{}\",",
902 json_escape(&entry.source_relation)
903 );
904 println!(" \"row_decisions\": [");
905 for (row_idx, row) in entry.row_decisions.iter().enumerate() {
906 let row_suffix = if row_idx + 1 == entry.row_decisions.len() {
907 ""
908 } else {
909 ","
910 };
911 println!(" {{");
912 println!(" \"row_key\": \"{}\",", json_escape(&row.row_key));
913 println!(" \"accepted\": {},", row.accepted);
914 println!(
915 " \"failed_predicates\": {},",
916 json_string_array(&row.failed_predicates)
917 );
918 println!(" \"threshold_comparisons\": [");
919 for (comparison_idx, comparison) in row.threshold_comparisons.iter().enumerate() {
920 let comparison_suffix = if comparison_idx + 1 == row.threshold_comparisons.len() {
921 ""
922 } else {
923 ","
924 };
925 println!(" {{");
926 println!(
927 " \"predicate\": \"{}\",",
928 json_escape(&comparison.predicate)
929 );
930 println!(
931 " \"left\": \"{}\",",
932 json_escape(&comparison.left)
933 );
934 println!(" \"op\": \"{}\",", json_escape(&comparison.op));
935 println!(
936 " \"right\": \"{}\",",
937 json_escape(&comparison.right)
938 );
939 println!(
940 " \"left_value\": \"{}\",",
941 json_escape(&comparison.left_value)
942 );
943 println!(
944 " \"right_value\": \"{}\",",
945 json_escape(&comparison.right_value)
946 );
947 println!(" \"passed\": {}", comparison.passed);
948 println!(" }}{}", comparison_suffix);
949 }
950 println!(" ],");
951 println!(
952 " \"aggregate_inputs\": {}",
953 json_string_array(&row.aggregate_inputs)
954 );
955 println!(" }}{}", row_suffix);
956 }
957 println!(" ]");
958 println!(" }}{}", suffix);
959 }
960 println!(" ]");
961}