Skip to main content

xlog_logic/
magic_sets.rs

1//! Magic-set rewriting for the deterministic language-completeness subset.
2
3use std::collections::{BTreeSet, HashMap, HashSet};
4
5use xlog_core::{Result, XlogError};
6
7use crate::ast::{Atom, BodyLiteral, MagicSetsMode, Program, Rule, Term};
8
9/// Status of a magic-set rewrite attempt.
10#[derive(Debug, Clone, Copy, PartialEq, Eq)]
11pub enum MagicSetStatus {
12    /// Rewriting was disabled by source configuration or there was no directive.
13    Disabled,
14    /// A supported bound recursive query was rewritten.
15    Applied,
16    /// `auto` mode found an unsafe or inapplicable program and left it unchanged.
17    Declined,
18}
19
20/// Human- and test-readable metadata for a magic-set rewrite attempt.
21#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct MagicSetReport {
23    /// Final status.
24    pub status: MagicSetStatus,
25    /// Generated magic predicate names.
26    pub generated_predicates: Vec<String>,
27    /// Adorned recursive predicates, formatted as `predicate/adornment`.
28    pub adorned_predicates: Vec<String>,
29    /// Reasons the rewrite declined.
30    pub declined_reasons: Vec<String>,
31}
32
33/// Rewritten program plus its report.
34#[derive(Debug, Clone)]
35pub struct MagicSetRewrite {
36    /// Program after rewriting, or the original program when disabled/declined.
37    pub program: Program,
38    /// Rewrite metadata.
39    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
55/// Rewrite supported bound recursive queries with magic predicates.
56pub 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
64/// [`rewrite_magic_sets`] taking the program by value: when magic sets are
65/// off (the default) the program is returned as is, without a clone.
66pub 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}