Skip to main content

oxirs_arq/
query_rewriter.rs

1//! Advanced Query Rewriting
2//!
3//! This module implements sophisticated query rewriting optimizations that
4//! transform algebra expressions into more efficient forms while preserving semantics.
5
6use crate::algebra::{
7    Algebra, BinaryOperator, Expression, Literal, Term, TriplePattern, UnaryOperator, Variable,
8};
9use anyhow::Result;
10use std::collections::{HashMap, HashSet};
11
12/// Advanced query rewriter with multiple optimization passes
13pub struct QueryRewriter {
14    config: RewriterConfig,
15    stats: RewriterStats,
16}
17
18/// Configuration for query rewriting
19#[derive(Debug, Clone)]
20pub struct RewriterConfig {
21    /// Enable constant folding
22    pub constant_folding: bool,
23    /// Enable dead code elimination
24    pub dead_code_elimination: bool,
25    /// Enable common subexpression elimination
26    pub cse_enabled: bool,
27    /// Enable filter pushdown
28    pub filter_pushdown: bool,
29    /// Enable pattern simplification
30    pub pattern_simplification: bool,
31    /// Enable empty pattern elimination
32    pub empty_pattern_elimination: bool,
33    /// Maximum rewrite iterations
34    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/// Statistics about rewriting operations
52#[derive(Debug, Clone, Default)]
53pub struct RewriterStats {
54    /// Number of constant folding optimizations
55    pub constants_folded: usize,
56    /// Number of dead code eliminations
57    pub dead_code_removed: usize,
58    /// Number of common subexpressions eliminated
59    pub cse_count: usize,
60    /// Number of filters pushed down
61    pub filters_pushed: usize,
62    /// Number of patterns simplified
63    pub patterns_simplified: usize,
64    /// Number of empty patterns eliminated
65    pub empty_eliminated: usize,
66    /// Total rewrite iterations
67    pub iterations: usize,
68}
69
70impl QueryRewriter {
71    /// Create a new query rewriter with default configuration
72    pub fn new() -> Self {
73        Self::with_config(RewriterConfig::default())
74    }
75
76    /// Create a new query rewriter with custom configuration
77    pub fn with_config(config: RewriterConfig) -> Self {
78        Self {
79            config,
80            stats: RewriterStats::default(),
81        }
82    }
83
84    /// Rewrite an algebra expression with all enabled optimizations
85    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            // Pass 1: Constant folding
93            if self.config.constant_folding {
94                let folded = self.fold_constants(&current)?;
95                if !algebras_equal(&current, &folded) {
96                    current = folded;
97                    changed = true;
98                }
99            }
100
101            // Pass 2: Dead code elimination
102            if self.config.dead_code_elimination {
103                let dce = self.eliminate_dead_code(&current)?;
104                if !algebras_equal(&current, &dce) {
105                    current = dce;
106                    changed = true;
107                }
108            }
109
110            // Pass 3: Empty pattern elimination
111            if self.config.empty_pattern_elimination {
112                let empty_elim = self.eliminate_empty_patterns(&current)?;
113                if !algebras_equal(&current, &empty_elim) {
114                    current = empty_elim;
115                    changed = true;
116                }
117            }
118
119            // Pass 4: Pattern simplification
120            if self.config.pattern_simplification {
121                let simplified = self.simplify_patterns(&current)?;
122                if !algebras_equal(&current, &simplified) {
123                    current = simplified;
124                    changed = true;
125                }
126            }
127
128            // Pass 5: Filter pushdown
129            if self.config.filter_pushdown {
130                let pushed = self.push_down_filters(&current)?;
131                if !algebras_equal(&current, &pushed) {
132                    current = pushed;
133                    changed = true;
134                }
135            }
136
137            // Pass 6: Common subexpression elimination
138            if self.config.cse_enabled {
139                let cse = self.eliminate_common_subexpressions(&current)?;
140                if !algebras_equal(&current, &cse) {
141                    current = cse;
142                    changed = true;
143                }
144            }
145
146            iteration += 1;
147            self.stats.iterations = iteration;
148
149            if !changed {
150                break; // Fixed point reached
151            }
152        }
153
154        Ok(current)
155    }
156
157    /// Fold constant expressions
158    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 expression is constant true, remove filter
165                if is_constant_true(&folded_condition) {
166                    self.stats.constants_folded += 1;
167                    return Ok(folded_pattern);
168                }
169
170                // If expression is constant false, return empty result
171                if is_constant_false(&folded_condition) {
172                    self.stats.constants_folded += 1;
173                    return Ok(Algebra::Bgp(vec![])); // Empty pattern
174                }
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    /// Fold constant expressions
219    fn fold_constant_expression(&self, expr: &Expression) -> Result<Expression> {
220        fold_constant_expression_impl(expr)
221    }
222}
223
224/// Helper function for constant expression folding (separated to avoid clippy warning)
225fn fold_constant_expression_impl(expr: &Expression) -> Result<Expression> {
226    match expr {
227        // Binary operations (including And/Or)
228        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                    // true && x = x
235                    if is_constant_true(&folded_left) {
236                        return Ok(folded_right);
237                    }
238                    // false && x = false
239                    if is_constant_false(&folded_left) {
240                        return Ok(folded_left);
241                    }
242                    // x && true = x
243                    if is_constant_true(&folded_right) {
244                        return Ok(folded_left);
245                    }
246                    // x && false = false
247                    if is_constant_false(&folded_right) {
248                        return Ok(folded_right);
249                    }
250                }
251                BinaryOperator::Or => {
252                    // true || x = true
253                    if is_constant_true(&folded_left) {
254                        return Ok(folded_left);
255                    }
256                    // false || x = x
257                    if is_constant_false(&folded_left) {
258                        return Ok(folded_right);
259                    }
260                    // x || true = true
261                    if is_constant_true(&folded_right) {
262                        return Ok(folded_right);
263                    }
264                    // x || false = x
265                    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        // Unary operations (including Not)
279        Expression::Unary { op, operand } => {
280            let folded = fold_constant_expression_impl(operand)?;
281
282            if *op == UnaryOperator::Not {
283                // !true = false
284                if is_constant_true(&folded) {
285                    return Ok(make_boolean_literal(false));
286                }
287                // !false = true
288                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    /// Eliminate dead code (unreachable algebra nodes)
304    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 left is empty, return right
311                if is_empty_pattern(&cleaned_left) {
312                    self.stats.dead_code_removed += 1;
313                    return Ok(cleaned_right);
314                }
315                // If right is empty, return left
316                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                // Join with empty pattern is empty
331                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    /// Eliminate empty patterns
346    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    /// Simplify patterns
375    fn simplify_patterns(&mut self, algebra: &Algebra) -> Result<Algebra> {
376        match algebra {
377            Algebra::Bgp(patterns) => {
378                // Remove duplicate patterns
379                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    /// Push filters down closer to data sources
407    fn push_down_filters(&mut self, algebra: &Algebra) -> Result<Algebra> {
408        match algebra {
409            Algebra::Filter { pattern, condition } => {
410                // Try to push filter into join operands
411                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 filter only uses variables from left, push to left
416                    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 filter only uses variables from right, push to right
430                    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                // Recursively process inner pattern
444                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    /// Eliminate common subexpressions
463    fn eliminate_common_subexpressions(&mut self, algebra: &Algebra) -> Result<Algebra> {
464        // Find common subpatterns
465        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        // If there are patterns that appear multiple times, we could optimize
473        // For now, just count them
474        for (_pattern, count) in pattern_counts {
475            if count > 1 {
476                self.stats.cse_count += 1;
477            }
478        }
479
480        // CSE implementation would be more complex in practice
481        Ok(algebra.clone())
482    }
483
484    /// Get rewriter statistics
485    pub fn get_stats(&self) -> &RewriterStats {
486        &self.stats
487    }
488
489    /// Reset statistics
490    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
501// Helper functions
502
503/// Check if two algebra expressions are equal
504fn algebras_equal(a: &Algebra, b: &Algebra) -> bool {
505    // Simplified equality check
506    format!("{:?}", a) == format!("{:?}", b)
507}
508
509/// Check if expression is constant true
510fn 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
518/// Check if expression is constant false
519fn 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
527/// Create a boolean literal
528fn make_boolean_literal(value: bool) -> Expression {
529    Expression::Literal(Literal {
530        value: value.to_string(),
531        language: None,
532        datatype: None,
533    })
534}
535
536/// Check if algebra is an empty pattern
537fn is_empty_pattern(algebra: &Algebra) -> bool {
538    matches!(algebra, Algebra::Bgp(patterns) if patterns.is_empty())
539}
540
541/// Generate a key for pattern deduplication
542fn pattern_key(pattern: &TriplePattern) -> String {
543    format!("{:?}", pattern)
544}
545
546/// Collect all variables from an algebra expression
547fn collect_variables(algebra: &Algebra) -> HashSet<Variable> {
548    let mut vars = HashSet::new();
549    collect_variables_recursive(algebra, &mut vars);
550    vars
551}
552
553/// Recursively collect variables
554fn 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
582/// Collect variables from an expression
583fn 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
589/// Recursively collect variables from expression
590fn 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
623/// Find all subpatterns in an algebra expression
624fn find_subpatterns(algebra: &Algebra) -> Vec<String> {
625    let mut patterns = Vec::new();
626    find_subpatterns_recursive(algebra, &mut patterns);
627    patterns
628}
629
630/// Recursively find subpatterns
631fn 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        // true AND x = x
657        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        // Create duplicate patterns
674        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); // Duplicates removed
687        } 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        // Should return the non-empty pattern
711        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        // Check that rewriting occurred
742        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}