1use std::collections::{BTreeSet, HashMap, HashSet};
4
5use xlog_core::{Result, XlogError};
6
7use crate::ast::{Atom, BodyLiteral, MagicSetsMode, Program, Rule, Term};
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
11pub enum MagicSetStatus {
12 Disabled,
14 Applied,
16 Declined,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct MagicSetReport {
23 pub status: MagicSetStatus,
25 pub generated_predicates: Vec<String>,
27 pub adorned_predicates: Vec<String>,
29 pub declined_reasons: Vec<String>,
31}
32
33#[derive(Debug, Clone)]
35pub struct MagicSetRewrite {
36 pub program: Program,
38 pub report: MagicSetReport,
40}
41
42#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
43struct Adornment {
44 pred: String,
45 pattern: Vec<bool>,
46}
47
48#[derive(Debug, Clone)]
49struct Seed {
50 pred: String,
51 pattern: Vec<bool>,
52 terms: Vec<Term>,
53}
54
55pub fn rewrite_magic_sets(program: &Program) -> Result<MagicSetRewrite> {
57 let mode = program.directives.magic_sets;
58 if mode.is_none() || mode == Some(MagicSetsMode::Off) {
59 return Ok(with_status(program.clone(), MagicSetStatus::Disabled));
60 }
61 rewrite_enabled(program)
62}
63
64pub fn rewrite_magic_sets_owned(program: Program) -> Result<MagicSetRewrite> {
67 let mode = program.directives.magic_sets;
68 if mode.is_none() || mode == Some(MagicSetsMode::Off) {
69 return Ok(with_status(program, MagicSetStatus::Disabled));
70 }
71 rewrite_enabled(&program)
72}
73
74fn rewrite_enabled(program: &Program) -> Result<MagicSetRewrite> {
75 let mode = program.directives.magic_sets;
76 let mode = mode.expect("checked above");
77
78 let recursive = recursive_predicates(program);
79 let seeds = collect_query_seeds(program, &recursive);
80 if seeds.is_empty() {
81 return decline_or_error(
82 program,
83 mode,
84 vec!["no bound recursive query eligible for magic_sets".to_string()],
85 );
86 }
87
88 if program.is_probabilistic_profile() {
89 return decline_or_error(
90 program,
91 mode,
92 vec!["probabilistic profiles are handled outside deterministic magic_sets".to_string()],
93 );
94 }
95
96 let target_preds: BTreeSet<String> = seeds.iter().map(|seed| seed.pred.clone()).collect();
97 let unsafe_reasons = unsupported_target_reasons(program, &target_preds, &recursive);
98 if !unsafe_reasons.is_empty() {
99 return decline_or_error(program, mode, unsafe_reasons);
100 }
101
102 let mut adornments = initial_adornments(&seeds);
103 expand_adornments(program, &target_preds, &mut adornments)?;
104
105 let mut generated_predicates: BTreeSet<String> = BTreeSet::new();
106 let mut adorned_predicates: BTreeSet<String> = BTreeSet::new();
107 for adornment in &adornments {
108 generated_predicates.insert(magic_predicate(&adornment.pred, &adornment.pattern));
109 adorned_predicates.insert(format!(
110 "{}/{}",
111 adornment.pred,
112 adornment_key(&adornment.pattern)
113 ));
114 }
115
116 let mut rewritten = program.clone();
117 rewritten.rules = rewrite_rules(program, &target_preds, &adornments, &seeds)?;
118
119 Ok(MagicSetRewrite {
120 program: rewritten,
121 report: MagicSetReport {
122 status: MagicSetStatus::Applied,
123 generated_predicates: generated_predicates.into_iter().collect(),
124 adorned_predicates: adorned_predicates.into_iter().collect(),
125 declined_reasons: Vec::new(),
126 },
127 })
128}
129
130fn with_status(program: Program, status: MagicSetStatus) -> MagicSetRewrite {
131 MagicSetRewrite {
132 program,
133 report: MagicSetReport {
134 status,
135 generated_predicates: Vec::new(),
136 adorned_predicates: Vec::new(),
137 declined_reasons: Vec::new(),
138 },
139 }
140}
141
142fn decline_or_error(
143 program: &Program,
144 mode: MagicSetsMode,
145 reasons: Vec<String>,
146) -> Result<MagicSetRewrite> {
147 if mode == MagicSetsMode::On {
148 return Err(magic_error(reasons.join("; ")));
149 }
150 Ok(MagicSetRewrite {
151 program: program.clone(),
152 report: MagicSetReport {
153 status: MagicSetStatus::Declined,
154 generated_predicates: Vec::new(),
155 adorned_predicates: Vec::new(),
156 declined_reasons: reasons,
157 },
158 })
159}
160
161fn magic_error(message: impl Into<String>) -> XlogError {
162 XlogError::Compilation(format!("magic_sets error: {}", message.into()))
163}
164
165fn collect_query_seeds(program: &Program, recursive: &HashSet<String>) -> Vec<Seed> {
166 let mut seen = HashSet::new();
167 let mut out = Vec::new();
168 for query in &program.queries {
169 if !recursive.contains(&query.atom.predicate) {
170 continue;
171 }
172 let pattern: Vec<bool> = query.atom.terms.iter().map(is_seed_term).collect();
173 if !pattern.iter().any(|bound| *bound) {
174 continue;
175 }
176 if !query
177 .atom
178 .terms
179 .iter()
180 .zip(&pattern)
181 .all(|(term, bound)| !*bound || is_supported_magic_term(term))
182 {
183 continue;
184 }
185 let terms = bound_terms(&query.atom, &pattern);
186 let key = format!(
187 "{}:{}:{:?}",
188 query.atom.predicate,
189 adornment_key(&pattern),
190 terms
191 );
192 if seen.insert(key) {
193 out.push(Seed {
194 pred: query.atom.predicate.clone(),
195 pattern,
196 terms,
197 });
198 }
199 }
200 out
201}
202
203fn initial_adornments(seeds: &[Seed]) -> BTreeSet<Adornment> {
204 seeds
205 .iter()
206 .map(|seed| Adornment {
207 pred: seed.pred.clone(),
208 pattern: seed.pattern.clone(),
209 })
210 .collect()
211}
212
213fn expand_adornments(
214 program: &Program,
215 target_preds: &BTreeSet<String>,
216 adornments: &mut BTreeSet<Adornment>,
217) -> Result<()> {
218 let mut changed = true;
219 while changed {
220 changed = false;
221 let snapshot: Vec<Adornment> = adornments.iter().cloned().collect();
222 for adornment in snapshot {
223 for rule in program
224 .rules
225 .iter()
226 .filter(|rule| rule.head.predicate == adornment.pred)
227 {
228 for discovered in discover_body_adornments(rule, &adornment.pattern, target_preds)?
229 {
230 changed |= adornments.insert(discovered);
231 }
232 }
233 }
234 }
235 Ok(())
236}
237
238fn discover_body_adornments(
239 rule: &Rule,
240 head_pattern: &[bool],
241 target_preds: &BTreeSet<String>,
242) -> Result<Vec<Adornment>> {
243 let mut bound = head_bound_variables(&rule.head, head_pattern);
244 let mut out = Vec::new();
245 for lit in &rule.body {
246 match lit {
247 BodyLiteral::Positive(atom) => {
248 if target_preds.contains(&atom.predicate) {
249 let pattern = atom_adornment(atom, &bound);
250 if !pattern.iter().any(|is_bound| *is_bound) {
251 return Err(magic_error(format!(
252 "recursive call {}/{} has no bound argument under supported SIPS",
253 atom.predicate,
254 atom.arity()
255 )));
256 }
257 out.push(Adornment {
258 pred: atom.predicate.clone(),
259 pattern,
260 });
261 }
262 bind_atom_variables(atom, &mut bound);
263 }
264 BodyLiteral::Comparison(_)
265 | BodyLiteral::Epistemic(_)
266 | BodyLiteral::IsExpr(_)
267 | BodyLiteral::Negated(_)
268 | BodyLiteral::Univ(_) => {}
269 }
270 }
271 Ok(out)
272}
273
274fn rewrite_rules(
275 program: &Program,
276 target_preds: &BTreeSet<String>,
277 adornments: &BTreeSet<Adornment>,
278 seeds: &[Seed],
279) -> Result<Vec<Rule>> {
280 let mut out: Vec<Rule> = program
281 .rules
282 .iter()
283 .filter(|rule| !target_preds.contains(&rule.head.predicate))
284 .cloned()
285 .collect();
286
287 let mut emitted = HashSet::new();
288 for seed in seeds {
289 let rule = Rule {
290 head: Atom {
291 predicate: magic_predicate(&seed.pred, &seed.pattern),
292 terms: seed.terms.clone(),
293 },
294 body: Vec::new(),
295 };
296 push_unique_rule(&mut out, &mut emitted, rule);
297 }
298
299 for adornment in adornments {
300 for rule in program
301 .rules
302 .iter()
303 .filter(|rule| rule.head.predicate == adornment.pred)
304 {
305 for magic_rule in propagation_rules(rule, &adornment.pattern, target_preds)? {
306 push_unique_rule(&mut out, &mut emitted, magic_rule);
307 }
308 }
309 }
310
311 for adornment in adornments {
312 for rule in program
313 .rules
314 .iter()
315 .filter(|rule| rule.head.predicate == adornment.pred)
316 {
317 let mut body = vec![BodyLiteral::Positive(magic_atom_for(
318 &rule.head,
319 &adornment.pattern,
320 ))];
321 body.extend(rule.body.clone());
322 out.push(Rule {
323 head: rule.head.clone(),
324 body,
325 });
326 }
327 }
328
329 Ok(out)
330}
331
332fn propagation_rules(
333 rule: &Rule,
334 head_pattern: &[bool],
335 target_preds: &BTreeSet<String>,
336) -> Result<Vec<Rule>> {
337 let caller_magic = magic_atom_for(&rule.head, head_pattern);
338 let mut prefix = vec![BodyLiteral::Positive(caller_magic.clone())];
339 let mut bound = head_bound_variables(&rule.head, head_pattern);
340 let mut out = Vec::new();
341
342 for lit in &rule.body {
343 let BodyLiteral::Positive(atom) = lit else {
344 continue;
345 };
346 if target_preds.contains(&atom.predicate) {
347 let pattern = atom_adornment(atom, &bound);
348 if !pattern.iter().any(|is_bound| *is_bound) {
349 return Err(magic_error(format!(
350 "recursive call {}/{} has no bound argument under supported SIPS",
351 atom.predicate,
352 atom.arity()
353 )));
354 }
355 let head = magic_atom_for(atom, &pattern);
356 let is_trivial = prefix.len() == 1
357 && matches!(&prefix[0], BodyLiteral::Positive(prefix_atom) if *prefix_atom == head);
358 if !is_trivial {
359 out.push(Rule {
360 head,
361 body: prefix.clone(),
362 });
363 }
364 }
365 bind_atom_variables(atom, &mut bound);
366 prefix.push(lit.clone());
367 }
368
369 Ok(out)
370}
371
372fn unsupported_target_reasons(
373 program: &Program,
374 target_preds: &BTreeSet<String>,
375 recursive: &HashSet<String>,
376) -> Vec<String> {
377 let mut reasons = BTreeSet::new();
378 for rule in &program.rules {
379 if !target_preds.contains(&rule.head.predicate) {
380 continue;
381 }
382 if rule.has_negation() {
383 reasons.insert(format!(
384 "negation in recursive rule for {} is outside the supported magic_sets subset",
385 rule.head.predicate
386 ));
387 }
388 if rule.has_aggregation() || rule.body.iter().any(body_literal_has_aggregate) {
389 reasons.insert(format!(
390 "aggregation in recursive rule for {} is outside the supported magic_sets subset",
391 rule.head.predicate
392 ));
393 }
394 for lit in &rule.body {
395 match lit {
396 BodyLiteral::Positive(atom) => {
397 if recursive.contains(&atom.predicate) && atom.predicate != rule.head.predicate
398 {
399 reasons.insert(format!(
400 "mutual recursion through {} is outside the supported magic_sets subset",
401 atom.predicate
402 ));
403 }
404 if atom.predicate.starts_with("__xlog_meta_")
405 || atom.predicate.starts_with("__xlog_list_")
406 {
407 reasons.insert(format!(
408 "meta/list helper {} in recursive rule is outside the supported magic_sets subset",
409 atom.predicate
410 ));
411 }
412 }
413 BodyLiteral::Negated(_) => {}
414 BodyLiteral::Comparison(_)
415 | BodyLiteral::Epistemic(_)
416 | BodyLiteral::IsExpr(_)
417 | BodyLiteral::Univ(_) => {
418 reasons.insert(format!(
419 "non-positive literal in recursive rule for {} is outside the supported magic_sets subset",
420 rule.head.predicate
421 ));
422 }
423 }
424 }
425 }
426 reasons.into_iter().collect()
427}
428
429fn recursive_predicates(program: &Program) -> HashSet<String> {
430 let mut deps: HashMap<String, HashSet<String>> = HashMap::new();
431 for rule in &program.rules {
432 let entry = deps.entry(rule.head.predicate.clone()).or_default();
433 for pred in rule.body_predicates() {
434 entry.insert(pred.to_string());
435 }
436 }
437 deps.keys()
438 .filter(|pred| reaches(pred, pred, &deps, &mut HashSet::new()))
439 .cloned()
440 .collect()
441}
442
443fn reaches(
444 start: &str,
445 target: &str,
446 deps: &HashMap<String, HashSet<String>>,
447 seen: &mut HashSet<String>,
448) -> bool {
449 let Some(next) = deps.get(start) else {
450 return false;
451 };
452 for pred in next {
453 if pred == target {
454 return true;
455 }
456 if seen.insert(pred.clone()) && reaches(pred, target, deps, seen) {
457 return true;
458 }
459 }
460 false
461}
462
463fn head_bound_variables(atom: &Atom, pattern: &[bool]) -> HashSet<String> {
464 atom.terms
465 .iter()
466 .zip(pattern)
467 .filter(|(_, bound)| **bound)
468 .flat_map(|(term, _)| term.variables().into_iter().map(str::to_string))
469 .collect()
470}
471
472fn atom_adornment(atom: &Atom, bound: &HashSet<String>) -> Vec<bool> {
473 atom.terms
474 .iter()
475 .map(|term| term_is_bound(term, bound))
476 .collect()
477}
478
479fn term_is_bound(term: &Term, bound: &HashSet<String>) -> bool {
480 match term {
481 Term::Variable(name) => bound.contains(name),
482 Term::Anonymous => false,
483 Term::List(items) => items.iter().all(|item| term_is_bound(item, bound)),
484 Term::Cons { head, tail } => term_is_bound(head, bound) && term_is_bound(tail, bound),
485 Term::Compound { args, .. } => args.iter().all(|arg| term_is_bound(arg, bound)),
486 Term::Integer(_)
487 | Term::Float(_)
488 | Term::String(_)
489 | Term::Symbol(_)
490 | Term::PredRef(_) => true,
491 Term::Aggregate(_) => false,
492 }
493}
494
495fn bind_atom_variables(atom: &Atom, bound: &mut HashSet<String>) {
496 for name in atom.variables() {
497 bound.insert(name.to_string());
498 }
499}
500
501fn body_literal_has_aggregate(lit: &BodyLiteral) -> bool {
502 match lit {
503 BodyLiteral::Positive(atom) | BodyLiteral::Negated(atom) => atom_has_aggregate(atom),
504 BodyLiteral::Epistemic(lit) => atom_has_aggregate(&lit.atom),
505 BodyLiteral::Comparison(comparison) => {
506 term_has_aggregate(&comparison.left) || term_has_aggregate(&comparison.right)
507 }
508 BodyLiteral::IsExpr(_) => false,
509 BodyLiteral::Univ(univ) => {
510 term_has_aggregate(&univ.term) || term_has_aggregate(&univ.parts)
511 }
512 }
513}
514
515fn atom_has_aggregate(atom: &Atom) -> bool {
516 atom.terms.iter().any(term_has_aggregate)
517}
518
519fn term_has_aggregate(term: &Term) -> bool {
520 match term {
521 Term::Aggregate(_) => true,
522 Term::List(items) => items.iter().any(term_has_aggregate),
523 Term::Cons { head, tail } => term_has_aggregate(head) || term_has_aggregate(tail),
524 Term::Compound { args, .. } => args.iter().any(term_has_aggregate),
525 Term::Variable(_)
526 | Term::Anonymous
527 | Term::Integer(_)
528 | Term::Float(_)
529 | Term::String(_)
530 | Term::Symbol(_)
531 | Term::PredRef(_) => false,
532 }
533}
534
535fn is_seed_term(term: &Term) -> bool {
536 is_supported_magic_term(term) && !term.is_any_variable()
537}
538
539fn is_supported_magic_term(term: &Term) -> bool {
540 matches!(
541 term,
542 Term::Integer(_) | Term::Float(_) | Term::String(_) | Term::Symbol(_)
543 )
544}
545
546fn bound_terms(atom: &Atom, pattern: &[bool]) -> Vec<Term> {
547 atom.terms
548 .iter()
549 .zip(pattern)
550 .filter(|(_, bound)| **bound)
551 .map(|(term, _)| term.clone())
552 .collect()
553}
554
555fn magic_atom_for(atom: &Atom, pattern: &[bool]) -> Atom {
556 Atom {
557 predicate: magic_predicate(&atom.predicate, pattern),
558 terms: bound_terms(atom, pattern),
559 }
560}
561
562fn magic_predicate(pred: &str, pattern: &[bool]) -> String {
563 format!("__xlog_magic_{}_{}", pred, adornment_key(pattern))
564}
565
566fn adornment_key(pattern: &[bool]) -> String {
567 pattern
568 .iter()
569 .map(|bound| if *bound { 'b' } else { 'f' })
570 .collect()
571}
572
573fn push_unique_rule(out: &mut Vec<Rule>, emitted: &mut HashSet<String>, rule: Rule) {
574 let key = format!("{:?}", rule);
575 if emitted.insert(key) {
576 out.push(rule);
577 }
578}