Skip to main content

egglog/
scheduler.rs

1use std::sync::Arc;
2use std::sync::Mutex;
3
4use core_relations::{ExecutionState, ExternalFunction, Value};
5use egglog_bridge::{
6    ColumnTy, DefaultVal, FunctionConfig, FunctionId, MergeFn, RuleId, TableAction,
7};
8use egglog_reports::RunReport;
9use numeric_id::define_id;
10
11use crate::{ast::ResolvedVar, core::GenericAtomTerm, core::ResolvedCoreRule, util::IndexMap, *};
12
13/// A scheduler decides which matches to be applied for a rule.
14///
15/// The matches that are not chosen in this iteration will be delayed
16/// to the next iteration.
17pub trait Scheduler: dyn_clone::DynClone + Send + Sync {
18    /// Whether or not the rules can be considered as saturated once no database
19    /// changes were made in the current iteration.
20    ///
21    /// This is only called when the runner is otherwise saturated.
22    /// Default implementation just returns `true`.
23    fn can_stop(&mut self, rules: &[&str], ruleset: &str) -> bool {
24        let _ = (rules, ruleset);
25        true
26    }
27
28    /// Filter the matches for a rule.
29    ///
30    /// Return `true` if the scheduler's next run of the rule should feed
31    /// `filter_matches` with a new iteration of matches.
32    fn filter_matches(&mut self, rule: &str, ruleset: &str, matches: &mut Matches) -> bool;
33}
34
35dyn_clone::clone_trait_object!(Scheduler);
36
37/// A collection of matches produced by a rule.
38/// The user can choose which matches to be fired.
39pub struct Matches {
40    matches: Vec<Value>,
41    chosen: Vec<usize>,
42    vars: Vec<ResolvedVar>,
43    /// Width of each stored tuple in `matches`. This is `vars.len()` for an
44    /// ordinary rule. A rule whose head references no variables would otherwise
45    /// have zero-width tuples, making its match count unrecoverable; for those we
46    /// collect a single unit marker per match, so the width is 1 while `vars` is
47    /// empty.
48    tuple_width: usize,
49    all_chosen: bool,
50}
51
52/// A match is a tuple of values corresponding to the variables in a rule.
53/// It allows you to retrieve the value corresponding to a variable in the match.
54pub struct Match<'a> {
55    values: &'a [Value],
56    vars: &'a [ResolvedVar],
57}
58
59impl Match<'_> {
60    /// Get the value corresponding a variable in this match.
61    pub fn get_value(&self, var: &str) -> Value {
62        let idx = self.vars.iter().position(|v| v.name == var).unwrap();
63        self.values[idx]
64    }
65}
66
67impl Matches {
68    fn new(matches: Vec<Value>, vars: Vec<ResolvedVar>) -> Self {
69        // Variable-free rules collect one unit marker per match (see
70        // `SchedulerRuleInfo::new`), so each stored tuple is one value wide even
71        // though there are no variables.
72        let tuple_width = vars.len().max(1);
73        assert!(matches.len().is_multiple_of(tuple_width));
74        Self {
75            matches,
76            vars,
77            tuple_width,
78            chosen: Vec::new(),
79            all_chosen: false,
80        }
81    }
82
83    /// The number of matches in total.
84    pub fn match_size(&self) -> usize {
85        self.matches.len() / self.tuple_width
86    }
87
88    /// The length of a tuple.
89    pub fn tuple_len(&self) -> usize {
90        self.vars.len()
91    }
92
93    /// Get `idx`-th match.
94    pub fn get_match(&self, idx: usize) -> Match<'_> {
95        Match {
96            values: &self.matches[idx * self.tuple_len()..(idx + 1) * self.tuple_len()],
97            vars: &self.vars,
98        }
99    }
100
101    /// Pick the match at `idx` to be fired.
102    pub fn choose(&mut self, idx: usize) {
103        self.chosen.push(idx);
104    }
105
106    /// Pick all matches to be fired.
107    ///
108    /// This is more efficient than calling `choose` for each match.
109    pub fn choose_all(&mut self) {
110        self.all_chosen = true;
111    }
112
113    /// Apply the chosen matches and return the residual matches.
114    fn instantiate(
115        mut self,
116        state: &mut ExecutionState<'_>,
117        table_action: &TableAction,
118    ) -> Vec<Value> {
119        // Width of the stored tuples (1 for variable-free rules, see `new`) versus
120        // the number of variable columns actually written into the `decided` table.
121        // For a variable-free rule the stored unit marker is dropped and only the
122        // trailing unit is inserted, producing the single `(unit)` row that the
123        // action rule fires on.
124        let tuple_width = self.tuple_width;
125        let var_len = self.vars.len();
126        let unit = state.base_values().get(());
127
128        if self.all_chosen {
129            for row in self.matches.chunks(tuple_width) {
130                table_action.insert(
131                    state,
132                    row[..var_len].iter().cloned().chain(std::iter::once(unit)),
133                );
134            }
135            vec![]
136        } else {
137            for idx in self.chosen.iter() {
138                let row = &self.matches[idx * tuple_width..(idx + 1) * tuple_width];
139                table_action.insert(
140                    state,
141                    row[..var_len].iter().cloned().chain(std::iter::once(unit)),
142                );
143            }
144
145            // swap remove the chosen matches
146            self.chosen.sort_unstable();
147            self.chosen.dedup();
148            let mut p = self.match_size();
149            for c in self.chosen.into_iter().rev() {
150                // It's important to decrement `p` first, because otherwise it might underflow when
151                // matches are exhausted.
152                p -= 1;
153                if c != p {
154                    let idx_c = c * tuple_width;
155                    let idx_p = p * tuple_width;
156                    for i in 0..tuple_width {
157                        self.matches.swap(idx_c + i, idx_p + i);
158                    }
159                }
160            }
161            self.matches.truncate(p * tuple_width);
162
163            self.matches
164        }
165    }
166}
167
168define_id!(
169    pub SchedulerId, u32,
170    "A unique identifier for a scheduler in the EGraph."
171);
172
173impl EGraph {
174    /// Register a new scheduler and return its id.
175    pub fn add_scheduler(&mut self, scheduler: Box<dyn Scheduler>) -> SchedulerId {
176        self.schedulers.push(SchedulerRecord {
177            scheduler,
178            rule_info: Default::default(),
179        })
180    }
181
182    /// Removes a scheduler
183    pub fn remove_scheduler(&mut self, scheduler_id: SchedulerId) -> Option<Box<dyn Scheduler>> {
184        self.schedulers.take(scheduler_id).map(|r| r.scheduler)
185    }
186
187    /// Runs a ruleset for one iteration using the given ruleset
188    pub fn step_rules_with_scheduler(
189        &mut self,
190        scheduler_id: SchedulerId,
191        ruleset: &str,
192    ) -> Result<RunReport, Error> {
193        fn collect_rules<'a>(
194            ruleset: &str,
195            rulesets: &'a IndexMap<String, Ruleset>,
196            ids: &mut Vec<(String, &'a ResolvedCoreRule)>,
197        ) -> Result<(), Error> {
198            let Some(r) = rulesets.get(ruleset) else {
199                return Err(Error::BackendError(format!("no such ruleset: {ruleset}")));
200            };
201            match r {
202                Ruleset::Rules(rules) => {
203                    for (rule_name, (core_rule, _)) in rules.iter() {
204                        ids.push((rule_name.clone(), core_rule));
205                    }
206                }
207                Ruleset::Combined(sub_rulesets) => {
208                    for sub_ruleset in sub_rulesets {
209                        collect_rules(sub_ruleset, rulesets, ids)?;
210                    }
211                }
212            }
213            Ok(())
214        }
215
216        let mut rules = Vec::new();
217        let rulesets = std::mem::take(&mut self.rulesets);
218        let collected = collect_rules(ruleset, &rulesets, &mut rules);
219        // Restore `rulesets` before propagating any error so the EGraph is not
220        // left with its rulesets taken out.
221        if let Err(e) = collected {
222            self.rulesets = rulesets;
223            return Err(e);
224        }
225        let mut schedulers = std::mem::take(&mut self.schedulers);
226
227        // `rulesets` and `schedulers` are now taken out of `self`. The body below
228        // has several fallible steps (rule compilation, `run_rules`), so run it in
229        // a closure and restore both fields afterward no matter how it exits.
230        // Otherwise an early error would leave the EGraph with empty rulesets and
231        // schedulers.
232        let result = (|| -> Result<RunReport, Error> {
233            // Step 1: build all the query/action rules and worklist if have not already
234            let record = &mut schedulers[scheduler_id];
235            for (id, rule) in rules.iter() {
236                if !record.rule_info.contains_key(id) {
237                    let info = SchedulerRuleInfo::new(self, rule, id)?;
238                    record.rule_info.insert((*id).to_owned(), info);
239                }
240            }
241
242            // Step 2: run all the queries for one iteration
243            let query_rules = rules
244                .iter()
245                .filter_map(|(rule_id, _rule)| {
246                    let rule_info = record.rule_info.get(rule_id).unwrap();
247
248                    if rule_info.should_seek {
249                        Some(rule_info.query_rule)
250                    } else {
251                        None
252                    }
253                })
254                .collect::<Vec<_>>();
255
256            let query_iter_report = self
257                .backend
258                .run_rules(&query_rules, Some(&self.type_info))
259                .map_err(|e| Error::BackendError(e.to_string()))?;
260
261            // Step 3: let the scheduler decide which matches need to be kept
262            self.backend
263                .with_execution_state(Some(&self.type_info), |state| {
264                    for (rule_id, _rule) in rules.iter() {
265                        let rule_info = record.rule_info.get_mut(rule_id).unwrap();
266
267                        let matches: Vec<Value> =
268                            std::mem::take(rule_info.matches.lock().unwrap().as_mut());
269                        let mut matches = Matches::new(matches, rule_info.free_vars.clone());
270                        rule_info.should_seek =
271                            record
272                                .scheduler
273                                .filter_matches(rule_id, ruleset, &mut matches);
274                        let table_action = TableAction::new(&self.backend, rule_info.decided);
275                        *rule_info.matches.lock().unwrap() =
276                            matches.instantiate(state, &table_action);
277                    }
278                });
279            self.backend.flush_updates();
280
281            // Step 4: run the action rules
282            let action_rules = rules
283                .iter()
284                .map(|(rule_id, _rule)| {
285                    let rule_info = record.rule_info.get(rule_id).unwrap();
286                    rule_info.action_rule
287                })
288                .collect::<Vec<_>>();
289            let action_iter_report = self
290                .backend
291                .run_rules(&action_rules, Some(&self.type_info))
292                .map_err(|e| Error::BackendError(e.to_string()))?;
293
294            // Step 5: combine the reports
295            let mut query_report = RunReport::singleton(ruleset, query_iter_report);
296            let mut action_report = RunReport::singleton(ruleset, action_iter_report);
297
298            // query matches don't count
299            query_report.updated = false;
300            query_report.num_matches_per_rule.clear();
301            // Scheduler state should not count as database progress. Instead it
302            // determines whether a no-op iteration can be treated as fully stopped.
303            action_report.can_stop = !action_report.updated && {
304                let rule_ids = rules.iter().map(|(id, _)| id.as_str()).collect::<Vec<_>>();
305                record.scheduler.can_stop(&rule_ids, ruleset)
306            };
307
308            query_report.union(action_report);
309
310            Ok(query_report)
311        })();
312
313        self.rulesets = rulesets;
314        self.schedulers = schedulers;
315
316        result
317    }
318}
319
320#[derive(Clone)]
321pub(crate) struct SchedulerRecord {
322    scheduler: Box<dyn Scheduler>,
323    rule_info: HashMap<String, SchedulerRuleInfo>,
324}
325
326/// To enable scheduling without modifying the backend,
327/// we split a rule (rule query action) into a worklist relation
328/// two rules (rule query (worklist vars false)) and
329/// (rule (worklist vars false) (action ... (delete (worklist vars false))))
330#[derive(Clone)]
331struct SchedulerRuleInfo {
332    matches: Arc<Mutex<Vec<Value>>>,
333    should_seek: bool,
334    decided: FunctionId,
335    query_rule: RuleId,
336    action_rule: RuleId,
337    free_vars: Vec<ResolvedVar>,
338}
339
340struct CollectMatches {
341    matches: Arc<Mutex<Vec<Value>>>,
342}
343
344impl Clone for CollectMatches {
345    fn clone(&self) -> Self {
346        Self {
347            matches: Arc::new(Mutex::new(self.matches.lock().unwrap().clone())),
348        }
349    }
350}
351
352impl CollectMatches {
353    fn new(matches: Arc<Mutex<Vec<Value>>>) -> Self {
354        Self { matches }
355    }
356}
357
358impl ExternalFunction for CollectMatches {
359    fn invoke(&self, state: &mut core_relations::ExecutionState, args: &[Value]) -> Option<Value> {
360        self.matches.lock().unwrap().extend(args.iter().copied());
361        Some(state.base_values().get(()))
362    }
363}
364
365impl SchedulerRuleInfo {
366    fn new(
367        egraph: &mut EGraph,
368        rule: &ResolvedCoreRule,
369        name: &str,
370    ) -> Result<SchedulerRuleInfo, Error> {
371        let free_vars = rule.head.get_free_vars().into_iter().collect::<Vec<_>>();
372        let unit_type = egraph.backend.base_values().get_ty::<()>();
373        let unit = egraph.backend.base_values().get(());
374        let unit_entry = egraph.backend.base_value_constant(());
375
376        let matches = Arc::new(Mutex::new(Vec::new()));
377        let collect_matches = egraph
378            .backend
379            .register_external_func(Box::new(CollectMatches::new(matches.clone())));
380        let schema = free_vars
381            .iter()
382            .map(|v| v.sort.column_ty(&egraph.backend))
383            .chain(std::iter::once(ColumnTy::Base(unit_type)))
384            .collect();
385        let decided = egraph.backend.add_table(FunctionConfig {
386            schema,
387            default: DefaultVal::Const(unit),
388            merge: MergeFn::AssertEq,
389            name: "backend".to_string(),
390            can_subsume: false,
391        });
392
393        // Step 1: build the query rule
394        let mut qrule_builder = BackendRule::new(
395            egraph.backend.new_rule(name, true),
396            &egraph.functions,
397            &egraph.type_info,
398            false, // seminaive query: Pure/Write contexts
399        );
400        qrule_builder.query(&rule.body, false)?;
401        let mut entries = free_vars
402            .iter()
403            .map(|fv| qrule_builder.entry(&GenericAtomTerm::Var(span!(), fv.clone())))
404            .collect::<Vec<_>>();
405        // A rule whose head references no variables would otherwise collect empty
406        // tuples, leaving the scheduler unable to tell whether the query matched and
407        // so never applying its actions. Collect a single unit marker per match so
408        // the match count is recoverable.
409        if entries.is_empty() {
410            entries.push(unit_entry.clone());
411        }
412        let _var = qrule_builder.rb.call_external_func(
413            collect_matches,
414            &entries,
415            ColumnTy::Base(unit_type),
416            || "collect_matches".to_string(),
417        );
418        let qrule_id = qrule_builder.build();
419
420        // Step 2: build the action rule
421        let mut arule_builder = BackendRule::new(
422            egraph.backend.new_rule(name, false),
423            &egraph.functions,
424            &egraph.type_info,
425            true, // action rule reads the DB: Read/Full contexts
426        );
427        let mut entries = free_vars
428            .iter()
429            .map(|fv| arule_builder.entry(&GenericAtomTerm::Var(span!(), fv.clone())))
430            .collect::<Vec<_>>();
431        entries.push(unit_entry);
432        arule_builder
433            .rb
434            .query_table(decided, &entries, None)
435            .unwrap();
436        arule_builder.actions(&rule.head)?;
437        // Remove the entry as it's now done
438        entries.pop();
439        arule_builder.rb.remove(decided, &entries);
440        let arule_id = arule_builder.build();
441
442        Ok(SchedulerRuleInfo {
443            free_vars,
444            query_rule: qrule_id,
445            action_rule: arule_id,
446            matches,
447            decided,
448            should_seek: true,
449        })
450    }
451}
452
453#[cfg(test)]
454mod test {
455    use super::*;
456
457    #[derive(Clone)]
458    struct FirstNScheduler {
459        n: usize,
460    }
461
462    impl Scheduler for FirstNScheduler {
463        fn filter_matches(&mut self, _rule: &str, _ruleset: &str, matches: &mut Matches) -> bool {
464            if matches.match_size() <= self.n {
465                matches.choose_all();
466            } else {
467                for i in 0..self.n {
468                    matches.choose(i);
469                }
470            }
471            matches.match_size() < self.n * 2
472        }
473    }
474
475    #[test]
476    fn test_first_n_scheduler() {
477        let mut egraph = EGraph::default();
478        let scheduler = FirstNScheduler { n: 10 };
479        let scheduler_id = egraph.add_scheduler(Box::new(scheduler));
480        let input = r#"
481        (relation R (i64))
482        (R 0)
483        (rule ((R x) (< x 100)) ((R (+ x 1))))
484        (run-schedule (saturate (run)))
485
486        (ruleset test)
487        (relation S (i64))
488        (rule ((R x)) ((S x)) :ruleset test :name "test-rule")
489        "#;
490        egraph.parse_and_run_program(None, input).unwrap();
491        assert_eq!(egraph.get_size("R"), 101);
492        let mut iter = 0;
493        loop {
494            let report = egraph
495                .step_rules_with_scheduler(scheduler_id, "test")
496                .unwrap();
497            let table_size = egraph.get_size("S");
498            iter += 1;
499            assert_eq!(table_size, std::cmp::min(iter * 10, 101));
500
501            let expected_matches = if iter <= 10 { 10 } else { 12 - iter };
502            assert_eq!(
503                report.num_matches_per_rule.iter().collect::<Vec<_>>(),
504                [(&"test-rule".into(), &expected_matches)]
505            );
506
507            // Because of semi-naive, the exact rules that are run are more than just `test-rule`
508            assert!(
509                report
510                    .search_and_apply_time_per_rule
511                    .keys()
512                    .all(|k| k.starts_with("test-rule"))
513            );
514            assert_eq!(
515                report.merge_time_per_ruleset.keys().collect::<Vec<_>>(),
516                [&"test".into()]
517            );
518            assert_eq!(
519                report
520                    .search_and_apply_time_per_ruleset
521                    .keys()
522                    .collect::<Vec<_>>(),
523                [&"test".into()]
524            );
525
526            if report.can_stop {
527                break;
528            }
529        }
530
531        assert_eq!(iter, 12);
532    }
533
534    #[test]
535    fn test_scheduler_does_not_apply_fresh_subsumed_matches() {
536        let mut egraph = EGraph::default();
537        let scheduler_id = egraph.add_scheduler(Box::new(FirstNScheduler { n: 10 }));
538        let input = r#"
539        (ruleset analysis)
540        (ruleset test)
541        (datatype Math
542          (Add Math Math)
543          (Mul Math Math)
544          (Num i64))
545        (relation Hit (i64))
546        (let expr (Add (Mul (Num 0) (Num 1)) (Num 2)))
547        (rewrite (Mul (Num 0) x) (Num 0) :subsume :ruleset analysis)
548        (rewrite (Add (Num 0) x) x :subsume :ruleset analysis)
549        (rule ((= e (Add (Mul (Num a) x) (Num b)))) ((Hit a)) :ruleset test :name "hit-subsumed-affine")
550        (run-schedule (saturate (run analysis)))
551        "#;
552        egraph.parse_and_run_program(None, input).unwrap();
553
554        let report = egraph
555            .step_rules_with_scheduler(scheduler_id, "test")
556            .unwrap();
557
558        assert_eq!(egraph.get_size("Hit"), 0);
559        assert!(
560            !report.updated,
561            "subsumed rows should not be collected as fresh scheduler matches"
562        );
563    }
564
565    #[derive(Clone, Default)]
566    struct DelayStopScheduler {
567        can_stop_calls: usize,
568    }
569
570    impl Scheduler for DelayStopScheduler {
571        fn can_stop(&mut self, _rules: &[&str], _ruleset: &str) -> bool {
572            self.can_stop_calls += 1;
573            self.can_stop_calls > 1
574        }
575
576        fn filter_matches(&mut self, _rule: &str, _ruleset: &str, _matches: &mut Matches) -> bool {
577            false
578        }
579    }
580
581    #[test]
582    fn test_scheduler_progress_is_separate_from_database_progress() {
583        let mut egraph = EGraph::default();
584        let scheduler_id = egraph.add_scheduler(Box::new(DelayStopScheduler::default()));
585        let input = r#"
586        (ruleset test)
587        (relation R (i64))
588        (rule ((R x)) ((R x)) :ruleset test :name "noop")
589        (R 1)
590        (R 2)
591        (R 3)
592        (R 4)
593        "#;
594        egraph.parse_and_run_program(None, input).unwrap();
595
596        let before = egraph.get_size("R");
597        let report = egraph
598            .step_rules_with_scheduler(scheduler_id, "test")
599            .unwrap();
600        let after = egraph.get_size("R");
601
602        assert_eq!(before, after);
603        assert!(!report.updated);
604        assert!(!report.can_stop);
605    }
606
607    #[test]
608    fn test_step_rules_with_scheduler_unknown_ruleset() {
609        let mut egraph = EGraph::default();
610        let scheduler_id = egraph.add_scheduler(Box::new(DelayStopScheduler::default()));
611        let err = egraph
612            .step_rules_with_scheduler(scheduler_id, "does-not-exist")
613            .unwrap_err();
614        assert!(matches!(err, Error::BackendError(_)));
615    }
616
617    /// A scheduler that only inspects `match_size` and never chooses anything.
618    #[derive(Clone)]
619    struct InspectSizeScheduler;
620
621    impl Scheduler for InspectSizeScheduler {
622        fn filter_matches(&mut self, _rule: &str, _ruleset: &str, matches: &mut Matches) -> bool {
623            // Calling `match_size` on a rule with no free variables used to panic
624            // with a divide-by-zero. Just exercise it and stop.
625            let _ = matches.match_size();
626            false
627        }
628    }
629
630    #[test]
631    fn test_match_size_with_no_free_vars() {
632        let mut egraph = EGraph::default();
633        let scheduler_id = egraph.add_scheduler(Box::new(InspectSizeScheduler));
634        // The action `(R 1)` references no variables, so the rule has no free vars.
635        let input = r#"
636        (ruleset test)
637        (relation R (i64))
638        (rule ((R x)) ((R 1)) :ruleset test :name "no-vars")
639        (R 0)
640        "#;
641        egraph.parse_and_run_program(None, input).unwrap();
642        egraph
643            .step_rules_with_scheduler(scheduler_id, "test")
644            .unwrap();
645    }
646
647    /// A scheduler that fires every match.
648    #[derive(Clone)]
649    struct ChooseAllScheduler;
650
651    impl Scheduler for ChooseAllScheduler {
652        fn filter_matches(&mut self, _rule: &str, _ruleset: &str, matches: &mut Matches) -> bool {
653            matches.choose_all();
654            false
655        }
656    }
657
658    #[test]
659    fn test_no_free_vars_rule_applies_actions() {
660        let mut egraph = EGraph::default();
661        let scheduler_id = egraph.add_scheduler(Box::new(ChooseAllScheduler));
662        // The action `(S)` references no variables, so the rule has no free vars.
663        // The scheduler must still apply it when the query matches.
664        let input = r#"
665        (ruleset test)
666        (relation R (i64))
667        (relation S ())
668        (rule ((R x)) ((S)) :ruleset test :name "no-vars")
669        (R 0)
670        "#;
671        egraph.parse_and_run_program(None, input).unwrap();
672        assert_eq!(egraph.get_size("S"), 0);
673        egraph
674            .step_rules_with_scheduler(scheduler_id, "test")
675            .unwrap();
676        assert_eq!(egraph.get_size("S"), 1);
677    }
678}