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
13pub trait Scheduler: dyn_clone::DynClone + Send + Sync {
18 fn can_stop(&mut self, rules: &[&str], ruleset: &str) -> bool {
24 let _ = (rules, ruleset);
25 true
26 }
27
28 fn filter_matches(&mut self, rule: &str, ruleset: &str, matches: &mut Matches) -> bool;
33}
34
35dyn_clone::clone_trait_object!(Scheduler);
36
37pub struct Matches {
40 matches: Vec<Value>,
41 chosen: Vec<usize>,
42 vars: Vec<ResolvedVar>,
43 tuple_width: usize,
49 all_chosen: bool,
50}
51
52pub struct Match<'a> {
55 values: &'a [Value],
56 vars: &'a [ResolvedVar],
57}
58
59impl Match<'_> {
60 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 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 pub fn match_size(&self) -> usize {
85 self.matches.len() / self.tuple_width
86 }
87
88 pub fn tuple_len(&self) -> usize {
90 self.vars.len()
91 }
92
93 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 pub fn choose(&mut self, idx: usize) {
103 self.chosen.push(idx);
104 }
105
106 pub fn choose_all(&mut self) {
110 self.all_chosen = true;
111 }
112
113 fn instantiate(
115 mut self,
116 state: &mut ExecutionState<'_>,
117 table_action: &TableAction,
118 ) -> Vec<Value> {
119 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 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 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 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 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 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 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 let result = (|| -> Result<RunReport, Error> {
233 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 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 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 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 let mut query_report = RunReport::singleton(ruleset, query_iter_report);
296 let mut action_report = RunReport::singleton(ruleset, action_iter_report);
297
298 query_report.updated = false;
300 query_report.num_matches_per_rule.clear();
301 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#[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 let mut qrule_builder = BackendRule::new(
395 egraph.backend.new_rule(name, true),
396 &egraph.functions,
397 &egraph.type_info,
398 false, );
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 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 let mut arule_builder = BackendRule::new(
422 egraph.backend.new_rule(name, false),
423 &egraph.functions,
424 &egraph.type_info,
425 true, );
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 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 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 #[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 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 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 #[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 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}