1use crate::algebra::{
7 Algebra, BinaryOperator, Expression, Literal, Term, TriplePattern, UnaryOperator, Variable,
8};
9use anyhow::Result;
10use std::collections::{HashMap, HashSet};
11
12pub struct QueryRewriter {
14 config: RewriterConfig,
15 stats: RewriterStats,
16}
17
18#[derive(Debug, Clone)]
20pub struct RewriterConfig {
21 pub constant_folding: bool,
23 pub dead_code_elimination: bool,
25 pub cse_enabled: bool,
27 pub filter_pushdown: bool,
29 pub pattern_simplification: bool,
31 pub empty_pattern_elimination: bool,
33 pub max_iterations: usize,
35}
36
37impl Default for RewriterConfig {
38 fn default() -> Self {
39 Self {
40 constant_folding: true,
41 dead_code_elimination: true,
42 cse_enabled: true,
43 filter_pushdown: true,
44 pattern_simplification: true,
45 empty_pattern_elimination: true,
46 max_iterations: 10,
47 }
48 }
49}
50
51#[derive(Debug, Clone, Default)]
53pub struct RewriterStats {
54 pub constants_folded: usize,
56 pub dead_code_removed: usize,
58 pub cse_count: usize,
60 pub filters_pushed: usize,
62 pub patterns_simplified: usize,
64 pub empty_eliminated: usize,
66 pub iterations: usize,
68}
69
70impl QueryRewriter {
71 pub fn new() -> Self {
73 Self::with_config(RewriterConfig::default())
74 }
75
76 pub fn with_config(config: RewriterConfig) -> Self {
78 Self {
79 config,
80 stats: RewriterStats::default(),
81 }
82 }
83
84 pub fn rewrite(&mut self, algebra: &Algebra) -> Result<Algebra> {
86 let mut current = algebra.clone();
87 let mut iteration = 0;
88
89 while iteration < self.config.max_iterations {
90 let mut changed = false;
91
92 if self.config.constant_folding {
94 let folded = self.fold_constants(¤t)?;
95 if !algebras_equal(¤t, &folded) {
96 current = folded;
97 changed = true;
98 }
99 }
100
101 if self.config.dead_code_elimination {
103 let dce = self.eliminate_dead_code(¤t)?;
104 if !algebras_equal(¤t, &dce) {
105 current = dce;
106 changed = true;
107 }
108 }
109
110 if self.config.empty_pattern_elimination {
112 let empty_elim = self.eliminate_empty_patterns(¤t)?;
113 if !algebras_equal(¤t, &empty_elim) {
114 current = empty_elim;
115 changed = true;
116 }
117 }
118
119 if self.config.pattern_simplification {
121 let simplified = self.simplify_patterns(¤t)?;
122 if !algebras_equal(¤t, &simplified) {
123 current = simplified;
124 changed = true;
125 }
126 }
127
128 if self.config.filter_pushdown {
130 let pushed = self.push_down_filters(¤t)?;
131 if !algebras_equal(¤t, &pushed) {
132 current = pushed;
133 changed = true;
134 }
135 }
136
137 if self.config.cse_enabled {
139 let cse = self.eliminate_common_subexpressions(¤t)?;
140 if !algebras_equal(¤t, &cse) {
141 current = cse;
142 changed = true;
143 }
144 }
145
146 iteration += 1;
147 self.stats.iterations = iteration;
148
149 if !changed {
150 break; }
152 }
153
154 Ok(current)
155 }
156
157 fn fold_constants(&mut self, algebra: &Algebra) -> Result<Algebra> {
159 match algebra {
160 Algebra::Filter { pattern, condition } => {
161 let folded_condition = self.fold_constant_expression(condition)?;
162 let folded_pattern = self.fold_constants(pattern)?;
163
164 if is_constant_true(&folded_condition) {
166 self.stats.constants_folded += 1;
167 return Ok(folded_pattern);
168 }
169
170 if is_constant_false(&folded_condition) {
172 self.stats.constants_folded += 1;
173 return Ok(Algebra::Bgp(vec![])); }
175
176 Ok(Algebra::Filter {
177 pattern: Box::new(folded_pattern),
178 condition: folded_condition,
179 })
180 }
181 Algebra::Join { left, right } => {
182 let folded_left = self.fold_constants(left)?;
183 let folded_right = self.fold_constants(right)?;
184 Ok(Algebra::Join {
185 left: Box::new(folded_left),
186 right: Box::new(folded_right),
187 })
188 }
189 Algebra::LeftJoin {
190 left,
191 right,
192 filter,
193 } => {
194 let folded_left = self.fold_constants(left)?;
195 let folded_right = self.fold_constants(right)?;
196 let folded_filter = filter
197 .as_ref()
198 .map(|f| self.fold_constant_expression(f))
199 .transpose()?;
200 Ok(Algebra::LeftJoin {
201 left: Box::new(folded_left),
202 right: Box::new(folded_right),
203 filter: folded_filter,
204 })
205 }
206 Algebra::Union { left, right } => {
207 let folded_left = self.fold_constants(left)?;
208 let folded_right = self.fold_constants(right)?;
209 Ok(Algebra::Union {
210 left: Box::new(folded_left),
211 right: Box::new(folded_right),
212 })
213 }
214 _ => Ok(algebra.clone()),
215 }
216 }
217
218 fn fold_constant_expression(&self, expr: &Expression) -> Result<Expression> {
220 fold_constant_expression_impl(expr)
221 }
222}
223
224fn fold_constant_expression_impl(expr: &Expression) -> Result<Expression> {
226 match expr {
227 Expression::Binary { op, left, right } => {
229 let folded_left = fold_constant_expression_impl(left)?;
230 let folded_right = fold_constant_expression_impl(right)?;
231
232 match op {
233 BinaryOperator::And => {
234 if is_constant_true(&folded_left) {
236 return Ok(folded_right);
237 }
238 if is_constant_false(&folded_left) {
240 return Ok(folded_left);
241 }
242 if is_constant_true(&folded_right) {
244 return Ok(folded_left);
245 }
246 if is_constant_false(&folded_right) {
248 return Ok(folded_right);
249 }
250 }
251 BinaryOperator::Or => {
252 if is_constant_true(&folded_left) {
254 return Ok(folded_left);
255 }
256 if is_constant_false(&folded_left) {
258 return Ok(folded_right);
259 }
260 if is_constant_true(&folded_right) {
262 return Ok(folded_right);
263 }
264 if is_constant_false(&folded_right) {
266 return Ok(folded_left);
267 }
268 }
269 _ => {}
270 }
271
272 Ok(Expression::Binary {
273 op: op.clone(),
274 left: Box::new(folded_left),
275 right: Box::new(folded_right),
276 })
277 }
278 Expression::Unary { op, operand } => {
280 let folded = fold_constant_expression_impl(operand)?;
281
282 if *op == UnaryOperator::Not {
283 if is_constant_true(&folded) {
285 return Ok(make_boolean_literal(false));
286 }
287 if is_constant_false(&folded) {
289 return Ok(make_boolean_literal(true));
290 }
291 }
292
293 Ok(Expression::Unary {
294 op: op.clone(),
295 operand: Box::new(folded),
296 })
297 }
298 _ => Ok(expr.clone()),
299 }
300}
301
302impl QueryRewriter {
303 fn eliminate_dead_code(&mut self, algebra: &Algebra) -> Result<Algebra> {
305 match algebra {
306 Algebra::Union { left, right } => {
307 let cleaned_left = self.eliminate_dead_code(left)?;
308 let cleaned_right = self.eliminate_dead_code(right)?;
309
310 if is_empty_pattern(&cleaned_left) {
312 self.stats.dead_code_removed += 1;
313 return Ok(cleaned_right);
314 }
315 if is_empty_pattern(&cleaned_right) {
317 self.stats.dead_code_removed += 1;
318 return Ok(cleaned_left);
319 }
320
321 Ok(Algebra::Union {
322 left: Box::new(cleaned_left),
323 right: Box::new(cleaned_right),
324 })
325 }
326 Algebra::Join { left, right } => {
327 let cleaned_left = self.eliminate_dead_code(left)?;
328 let cleaned_right = self.eliminate_dead_code(right)?;
329
330 if is_empty_pattern(&cleaned_left) || is_empty_pattern(&cleaned_right) {
332 self.stats.dead_code_removed += 1;
333 return Ok(Algebra::Bgp(vec![]));
334 }
335
336 Ok(Algebra::Join {
337 left: Box::new(cleaned_left),
338 right: Box::new(cleaned_right),
339 })
340 }
341 _ => Ok(algebra.clone()),
342 }
343 }
344
345 fn eliminate_empty_patterns(&mut self, algebra: &Algebra) -> Result<Algebra> {
347 match algebra {
348 Algebra::Bgp(patterns) if patterns.is_empty() => {
349 self.stats.empty_eliminated += 1;
350 Ok(Algebra::Bgp(vec![]))
351 }
352 Algebra::Join { left, right } => {
353 let left_clean = self.eliminate_empty_patterns(left)?;
354 let right_clean = self.eliminate_empty_patterns(right)?;
355
356 if is_empty_pattern(&left_clean) {
357 self.stats.empty_eliminated += 1;
358 return Ok(right_clean);
359 }
360 if is_empty_pattern(&right_clean) {
361 self.stats.empty_eliminated += 1;
362 return Ok(left_clean);
363 }
364
365 Ok(Algebra::Join {
366 left: Box::new(left_clean),
367 right: Box::new(right_clean),
368 })
369 }
370 _ => Ok(algebra.clone()),
371 }
372 }
373
374 fn simplify_patterns(&mut self, algebra: &Algebra) -> Result<Algebra> {
376 match algebra {
377 Algebra::Bgp(patterns) => {
378 let mut unique_patterns = Vec::new();
380 let mut seen = HashSet::new();
381
382 for pattern in patterns {
383 let key = pattern_key(pattern);
384 if !seen.contains(&key) {
385 seen.insert(key);
386 unique_patterns.push(pattern.clone());
387 } else {
388 self.stats.patterns_simplified += 1;
389 }
390 }
391
392 Ok(Algebra::Bgp(unique_patterns))
393 }
394 Algebra::Join { left, right } => {
395 let simp_left = self.simplify_patterns(left)?;
396 let simp_right = self.simplify_patterns(right)?;
397 Ok(Algebra::Join {
398 left: Box::new(simp_left),
399 right: Box::new(simp_right),
400 })
401 }
402 _ => Ok(algebra.clone()),
403 }
404 }
405
406 fn push_down_filters(&mut self, algebra: &Algebra) -> Result<Algebra> {
408 match algebra {
409 Algebra::Filter { pattern, condition } => {
410 if let Algebra::Join { left, right } = pattern.as_ref() {
412 let left_vars = collect_variables(left);
413 let condition_vars = collect_expression_variables(condition);
414
415 if condition_vars.iter().all(|v| left_vars.contains(v)) {
417 self.stats.filters_pushed += 1;
418 let filtered_left = Algebra::Filter {
419 pattern: left.clone(),
420 condition: condition.clone(),
421 };
422 return Ok(Algebra::Join {
423 left: Box::new(filtered_left),
424 right: right.clone(),
425 });
426 }
427
428 let right_vars = collect_variables(right);
429 if condition_vars.iter().all(|v| right_vars.contains(v)) {
431 self.stats.filters_pushed += 1;
432 let filtered_right = Algebra::Filter {
433 pattern: right.clone(),
434 condition: condition.clone(),
435 };
436 return Ok(Algebra::Join {
437 left: left.clone(),
438 right: Box::new(filtered_right),
439 });
440 }
441 }
442
443 let pushed_pattern = self.push_down_filters(pattern)?;
445 Ok(Algebra::Filter {
446 pattern: Box::new(pushed_pattern),
447 condition: condition.clone(),
448 })
449 }
450 Algebra::Join { left, right } => {
451 let pushed_left = self.push_down_filters(left)?;
452 let pushed_right = self.push_down_filters(right)?;
453 Ok(Algebra::Join {
454 left: Box::new(pushed_left),
455 right: Box::new(pushed_right),
456 })
457 }
458 _ => Ok(algebra.clone()),
459 }
460 }
461
462 fn eliminate_common_subexpressions(&mut self, algebra: &Algebra) -> Result<Algebra> {
464 let subpatterns = find_subpatterns(algebra);
466 let mut pattern_counts: HashMap<String, usize> = HashMap::new();
467
468 for pattern in &subpatterns {
469 *pattern_counts.entry(pattern.clone()).or_insert(0) += 1;
470 }
471
472 for (_pattern, count) in pattern_counts {
475 if count > 1 {
476 self.stats.cse_count += 1;
477 }
478 }
479
480 Ok(algebra.clone())
482 }
483
484 pub fn get_stats(&self) -> &RewriterStats {
486 &self.stats
487 }
488
489 pub fn reset_stats(&mut self) {
491 self.stats = RewriterStats::default();
492 }
493}
494
495impl Default for QueryRewriter {
496 fn default() -> Self {
497 Self::new()
498 }
499}
500
501fn algebras_equal(a: &Algebra, b: &Algebra) -> bool {
505 format!("{:?}", a) == format!("{:?}", b)
507}
508
509fn is_constant_true(expr: &Expression) -> bool {
511 if let Expression::Literal(lit) = expr {
512 lit.value == "true" && lit.datatype.is_none()
513 } else {
514 false
515 }
516}
517
518fn is_constant_false(expr: &Expression) -> bool {
520 if let Expression::Literal(lit) = expr {
521 lit.value == "false" && lit.datatype.is_none()
522 } else {
523 false
524 }
525}
526
527fn make_boolean_literal(value: bool) -> Expression {
529 Expression::Literal(Literal {
530 value: value.to_string(),
531 language: None,
532 datatype: None,
533 })
534}
535
536fn is_empty_pattern(algebra: &Algebra) -> bool {
538 matches!(algebra, Algebra::Bgp(patterns) if patterns.is_empty())
539}
540
541fn pattern_key(pattern: &TriplePattern) -> String {
543 format!("{:?}", pattern)
544}
545
546fn collect_variables(algebra: &Algebra) -> HashSet<Variable> {
548 let mut vars = HashSet::new();
549 collect_variables_recursive(algebra, &mut vars);
550 vars
551}
552
553fn collect_variables_recursive(algebra: &Algebra, vars: &mut HashSet<Variable>) {
555 match algebra {
556 Algebra::Bgp(patterns) => {
557 for pattern in patterns {
558 if let Term::Variable(v) = &pattern.subject {
559 vars.insert(v.clone());
560 }
561 if let Term::Variable(v) = &pattern.predicate {
562 vars.insert(v.clone());
563 }
564 if let Term::Variable(v) = &pattern.object {
565 vars.insert(v.clone());
566 }
567 }
568 }
569 Algebra::Join { left, right }
570 | Algebra::Union { left, right }
571 | Algebra::LeftJoin { left, right, .. } => {
572 collect_variables_recursive(left, vars);
573 collect_variables_recursive(right, vars);
574 }
575 Algebra::Filter { pattern, .. } => {
576 collect_variables_recursive(pattern, vars);
577 }
578 _ => {}
579 }
580}
581
582fn collect_expression_variables(expr: &Expression) -> HashSet<Variable> {
584 let mut vars = HashSet::new();
585 collect_expr_vars_recursive(expr, &mut vars);
586 vars
587}
588
589fn collect_expr_vars_recursive(expr: &Expression, vars: &mut HashSet<Variable>) {
591 match expr {
592 Expression::Variable(v) => {
593 vars.insert(v.clone());
594 }
595 Expression::Unary { operand, .. } => {
596 collect_expr_vars_recursive(operand, vars);
597 }
598 Expression::Binary { left, right, .. } => {
599 collect_expr_vars_recursive(left, vars);
600 collect_expr_vars_recursive(right, vars);
601 }
602 Expression::Function { args, .. } => {
603 for arg in args {
604 collect_expr_vars_recursive(arg, vars);
605 }
606 }
607 Expression::Conditional {
608 condition,
609 then_expr,
610 else_expr,
611 } => {
612 collect_expr_vars_recursive(condition, vars);
613 collect_expr_vars_recursive(then_expr, vars);
614 collect_expr_vars_recursive(else_expr, vars);
615 }
616 Expression::Exists(algebra) | Expression::NotExists(algebra) => {
617 collect_variables_recursive(algebra, vars);
618 }
619 _ => {}
620 }
621}
622
623fn find_subpatterns(algebra: &Algebra) -> Vec<String> {
625 let mut patterns = Vec::new();
626 find_subpatterns_recursive(algebra, &mut patterns);
627 patterns
628}
629
630fn find_subpatterns_recursive(algebra: &Algebra, patterns: &mut Vec<String>) {
632 patterns.push(format!("{:?}", algebra));
633
634 match algebra {
635 Algebra::Join { left, right }
636 | Algebra::Union { left, right }
637 | Algebra::LeftJoin { left, right, .. } => {
638 find_subpatterns_recursive(left, patterns);
639 find_subpatterns_recursive(right, patterns);
640 }
641 Algebra::Filter { pattern, .. } => {
642 find_subpatterns_recursive(pattern, patterns);
643 }
644 _ => {}
645 }
646}
647
648#[cfg(test)]
649mod tests {
650 use super::*;
651
652 #[test]
653 fn test_constant_folding() {
654 let rewriter = QueryRewriter::new();
655
656 let expr = Expression::Binary {
658 op: BinaryOperator::And,
659 left: Box::new(make_boolean_literal(true)),
660 right: Box::new(Expression::Variable(
661 Variable::new("x".to_string()).unwrap(),
662 )),
663 };
664
665 let folded = rewriter.fold_constant_expression(&expr).unwrap();
666 assert!(matches!(folded, Expression::Variable(_)));
667 }
668
669 #[test]
670 fn test_pattern_simplification() {
671 let mut rewriter = QueryRewriter::new();
672
673 let pattern1 = TriplePattern {
675 subject: Term::Variable(Variable::new("s".to_string()).unwrap()),
676 predicate: Term::Variable(Variable::new("p".to_string()).unwrap()),
677 object: Term::Variable(Variable::new("o".to_string()).unwrap()),
678 };
679
680 let pattern2 = pattern1.clone();
681
682 let algebra = Algebra::Bgp(vec![pattern1, pattern2]);
683 let simplified = rewriter.simplify_patterns(&algebra).unwrap();
684
685 if let Algebra::Bgp(patterns) = simplified {
686 assert_eq!(patterns.len(), 1); } else {
688 panic!("Expected Bgp");
689 }
690 }
691
692 #[test]
693 fn test_empty_pattern_elimination() {
694 let mut rewriter = QueryRewriter::new();
695
696 let empty = Algebra::Bgp(vec![]);
697 let non_empty = Algebra::Bgp(vec![TriplePattern {
698 subject: Term::Variable(Variable::new("s".to_string()).unwrap()),
699 predicate: Term::Variable(Variable::new("p".to_string()).unwrap()),
700 object: Term::Variable(Variable::new("o".to_string()).unwrap()),
701 }]);
702
703 let join = Algebra::Join {
704 left: Box::new(empty),
705 right: Box::new(non_empty.clone()),
706 };
707
708 let eliminated = rewriter.eliminate_empty_patterns(&join).unwrap();
709
710 assert!(matches!(eliminated, Algebra::Bgp(_)));
712 }
713
714 #[test]
715 fn test_filter_pushdown() {
716 let mut rewriter = QueryRewriter::new();
717
718 let pattern = TriplePattern {
719 subject: Term::Variable(Variable::new("x".to_string()).unwrap()),
720 predicate: Term::Variable(Variable::new("p".to_string()).unwrap()),
721 object: Term::Variable(Variable::new("o".to_string()).unwrap()),
722 };
723
724 let left = Algebra::Bgp(vec![pattern.clone()]);
725 let right = Algebra::Bgp(vec![pattern]);
726
727 let join = Algebra::Join {
728 left: Box::new(left),
729 right: Box::new(right),
730 };
731
732 let filter_condition = Expression::Variable(Variable::new("x".to_string()).unwrap());
733
734 let filtered_join = Algebra::Filter {
735 pattern: Box::new(join),
736 condition: filter_condition,
737 };
738
739 let pushed = rewriter.push_down_filters(&filtered_join).unwrap();
740
741 assert!(rewriter.stats.filters_pushed > 0 || matches!(pushed, Algebra::Join { .. }));
743 }
744
745 #[test]
746 fn test_rewriter_stats() {
747 let mut rewriter = QueryRewriter::new();
748
749 let pattern = TriplePattern {
750 subject: Term::Variable(Variable::new("s".to_string()).unwrap()),
751 predicate: Term::Variable(Variable::new("p".to_string()).unwrap()),
752 object: Term::Variable(Variable::new("o".to_string()).unwrap()),
753 };
754
755 let algebra = Algebra::Bgp(vec![pattern.clone(), pattern]);
756 let _simplified = rewriter.simplify_patterns(&algebra).unwrap();
757
758 let stats = rewriter.get_stats();
759 assert!(stats.patterns_simplified > 0);
760 }
761}