use crate::error::{Result, TdbError};
use crate::statistics::StatisticsSnapshot;
use anyhow::Context;
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct TriplePattern {
pub subject: String,
pub predicate: String,
pub object: String,
pub estimated_cardinality: Option<u64>,
}
impl TriplePattern {
pub fn new(subject: &str, predicate: &str, object: &str) -> Self {
Self {
subject: subject.to_string(),
predicate: predicate.to_string(),
object: object.to_string(),
estimated_cardinality: None,
}
}
pub fn is_variable(term: &str) -> bool {
term.starts_with('?')
}
pub fn variables(&self) -> HashSet<String> {
let mut vars = HashSet::new();
if Self::is_variable(&self.subject) {
vars.insert(self.subject.clone());
}
if Self::is_variable(&self.predicate) {
vars.insert(self.predicate.clone());
}
if Self::is_variable(&self.object) {
vars.insert(self.object.clone());
}
vars
}
pub fn join_variables(&self, other: &TriplePattern) -> HashSet<String> {
self.variables()
.intersection(&other.variables())
.cloned()
.collect()
}
pub fn bound_count(&self) -> usize {
let mut count = 0;
if !Self::is_variable(&self.subject) {
count += 1;
}
if !Self::is_variable(&self.predicate) {
count += 1;
}
if !Self::is_variable(&self.object) {
count += 1;
}
count
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct JoinNode {
pub left: Box<JoinPlan>,
pub right: Box<JoinPlan>,
pub join_vars: HashSet<String>,
pub cost: f64,
pub cardinality: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum JoinPlan {
Pattern(TriplePattern),
Join(JoinNode),
}
impl JoinPlan {
pub fn variables(&self) -> HashSet<String> {
match self {
JoinPlan::Pattern(pattern) => pattern.variables(),
JoinPlan::Join(node) => {
let mut vars = node.left.variables();
vars.extend(node.right.variables());
vars
}
}
}
pub fn cardinality(&self) -> u64 {
match self {
JoinPlan::Pattern(pattern) => pattern.estimated_cardinality.unwrap_or(1000),
JoinPlan::Join(node) => node.cardinality,
}
}
pub fn cost(&self) -> f64 {
match self {
JoinPlan::Pattern(_) => 0.0,
JoinPlan::Join(node) => node.cost,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum JoinAlgorithm {
Greedy,
DynamicProgramming,
Genetic,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct JoinOptimizerConfig {
pub algorithm: JoinAlgorithm,
pub dp_max_patterns: usize,
pub default_join_selectivity: f64,
pub cost_based: bool,
}
impl Default for JoinOptimizerConfig {
fn default() -> Self {
Self {
algorithm: JoinAlgorithm::Greedy,
dp_max_patterns: 12,
default_join_selectivity: 0.1,
cost_based: true,
}
}
}
pub struct JoinOptimizer {
config: JoinOptimizerConfig,
stats: Option<Arc<StatisticsSnapshot>>,
opt_stats: Arc<Mutex<OptimizationStats>>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct OptimizationStats {
pub total_optimizations: u64,
pub greedy_optimizations: u64,
pub dp_optimizations: u64,
pub avg_optimization_time_ms: f64,
total_time_ms: f64,
}
impl JoinOptimizer {
pub fn new(config: JoinOptimizerConfig) -> Self {
Self {
config,
stats: None,
opt_stats: Arc::new(Mutex::new(OptimizationStats::default())),
}
}
pub fn set_statistics(&mut self, stats: Arc<StatisticsSnapshot>) {
self.stats = Some(stats);
}
pub fn optimize(&mut self, mut patterns: Vec<TriplePattern>) -> Result<JoinPlan> {
if patterns.is_empty() {
return Err(TdbError::Other("No patterns to optimize".to_string()));
}
if patterns.len() == 1 {
return Ok(JoinPlan::Pattern(
patterns
.pop()
.expect("collection validated to be non-empty"),
));
}
self.estimate_cardinalities(&mut patterns)?;
let algorithm = if patterns.len() > self.config.dp_max_patterns {
JoinAlgorithm::Greedy
} else {
self.config.algorithm
};
let start = std::time::Instant::now();
let plan = match algorithm {
JoinAlgorithm::Greedy => self.greedy_optimize(patterns)?,
JoinAlgorithm::DynamicProgramming => self.dp_optimize(patterns)?,
JoinAlgorithm::Genetic => {
self.greedy_optimize(patterns)?
}
};
let duration = start.elapsed().as_millis() as f64;
let mut opt_stats = self.opt_stats.lock();
opt_stats.total_optimizations += 1;
opt_stats.total_time_ms += duration;
opt_stats.avg_optimization_time_ms =
opt_stats.total_time_ms / opt_stats.total_optimizations as f64;
match algorithm {
JoinAlgorithm::Greedy => opt_stats.greedy_optimizations += 1,
JoinAlgorithm::DynamicProgramming => opt_stats.dp_optimizations += 1,
JoinAlgorithm::Genetic => opt_stats.greedy_optimizations += 1,
}
Ok(plan)
}
fn estimate_cardinalities(&self, patterns: &mut [TriplePattern]) -> Result<()> {
for pattern in patterns.iter_mut() {
let cardinality = self.estimate_pattern_cardinality(pattern);
pattern.estimated_cardinality = Some(cardinality);
}
Ok(())
}
fn estimate_pattern_cardinality(&self, pattern: &TriplePattern) -> u64 {
if let Some(ref stats) = self.stats {
let bound_count = pattern.bound_count();
match bound_count {
3 => 1, 2 => stats.total_triples / 100, 1 => stats.total_triples / 10, 0 => stats.total_triples, _ => stats.total_triples / 10,
}
} else {
match pattern.bound_count() {
3 => 1,
2 => 1000,
1 => 10000,
0 => 100000,
_ => 10000,
}
}
}
fn greedy_optimize(&self, patterns: Vec<TriplePattern>) -> Result<JoinPlan> {
if patterns.len() == 1 {
return Ok(JoinPlan::Pattern(patterns[0].clone()));
}
let mut remaining: Vec<_> = patterns.clone();
let start_idx = remaining
.iter()
.enumerate()
.min_by_key(|(_, p)| p.estimated_cardinality.unwrap_or(u64::MAX))
.map(|(i, _)| i)
.expect("collection validated to be non-empty");
let mut current_plan = JoinPlan::Pattern(remaining.remove(start_idx));
let mut current_vars = current_plan.variables();
while !remaining.is_empty() {
let (best_idx, _best_cost) = remaining
.iter()
.enumerate()
.map(|(i, pattern)| {
let join_vars = current_vars
.intersection(&pattern.variables())
.cloned()
.collect();
let cost = self.estimate_join_cost(¤t_plan, pattern, &join_vars);
(i, cost)
})
.min_by(|(_, cost1), (_, cost2)| {
cost1
.partial_cmp(cost2)
.unwrap_or(std::cmp::Ordering::Equal)
})
.expect("collection validated to be non-empty");
let next_pattern = remaining.remove(best_idx);
let join_vars = current_vars
.intersection(&next_pattern.variables())
.cloned()
.collect();
current_plan = self.create_join(current_plan, next_pattern, join_vars);
current_vars = current_plan.variables();
}
Ok(current_plan)
}
fn dp_optimize(&self, patterns: Vec<TriplePattern>) -> Result<JoinPlan> {
let n = patterns.len();
if n == 1 {
return Ok(JoinPlan::Pattern(patterns[0].clone()));
}
let mut dp: HashMap<u64, JoinPlan> = HashMap::new();
for (i, pattern) in patterns.iter().enumerate() {
let mask = 1u64 << i;
dp.insert(mask, JoinPlan::Pattern(pattern.clone()));
}
for size in 2..=n {
let subsets = self.generate_subsets(n, size);
for subset in subsets {
let mut best_plan: Option<JoinPlan> = None;
let mut best_cost = f64::MAX;
for left_mask in 1..subset {
if (left_mask & subset) != left_mask {
continue;
}
let right_mask = subset & !left_mask;
if right_mask == 0 {
continue;
}
let left_plan = dp.get(&left_mask).cloned();
let right_plan = dp.get(&right_mask).cloned();
if let (Some(left), Some(right)) = (left_plan, right_plan) {
let join_vars = left
.variables()
.intersection(&right.variables())
.cloned()
.collect();
let cost = self.estimate_join_cost_plans(&left, &right, &join_vars);
if cost < best_cost {
best_cost = cost;
best_plan = Some(self.create_join_plans(left, right, join_vars));
}
}
}
if let Some(plan) = best_plan {
dp.insert(subset, plan);
}
}
}
let full_mask = (1u64 << n) - 1;
dp.remove(&full_mask)
.ok_or_else(|| TdbError::Other("DP optimization failed".to_string()))
}
fn generate_subsets(&self, n: usize, k: usize) -> Vec<u64> {
let mut subsets = Vec::new();
self.generate_subsets_recursive(n, k, 0, 0, &mut subsets);
subsets
}
#[allow(clippy::only_used_in_recursion)]
fn generate_subsets_recursive(
&self,
n: usize,
k: usize,
start: usize,
current: u64,
subsets: &mut Vec<u64>,
) {
if k == 0 {
subsets.push(current);
return;
}
for i in start..n {
let next = current | (1u64 << i);
self.generate_subsets_recursive(n, k - 1, i + 1, next, subsets);
}
}
fn estimate_join_cost(
&self,
plan: &JoinPlan,
pattern: &TriplePattern,
join_vars: &HashSet<String>,
) -> f64 {
let left_card = plan.cardinality() as f64;
let right_card = pattern.estimated_cardinality.unwrap_or(1000) as f64;
let selectivity = if join_vars.is_empty() {
1.0
} else {
self.config.default_join_selectivity
};
left_card * right_card * selectivity
}
fn estimate_join_cost_plans(
&self,
left: &JoinPlan,
right: &JoinPlan,
join_vars: &HashSet<String>,
) -> f64 {
let left_card = left.cardinality() as f64;
let right_card = right.cardinality() as f64;
let selectivity = if join_vars.is_empty() {
1.0
} else {
self.config.default_join_selectivity
};
left_card * right_card * selectivity + left.cost() + right.cost()
}
fn create_join(
&self,
left: JoinPlan,
right: TriplePattern,
join_vars: HashSet<String>,
) -> JoinPlan {
let right_plan = JoinPlan::Pattern(right);
self.create_join_plans(left, right_plan, join_vars)
}
fn create_join_plans(
&self,
left: JoinPlan,
right: JoinPlan,
join_vars: HashSet<String>,
) -> JoinPlan {
let cost = self.estimate_join_cost_plans(&left, &right, &join_vars);
let left_card = left.cardinality();
let right_card = right.cardinality();
let cardinality = if join_vars.is_empty() {
left_card.saturating_mul(right_card)
} else {
((left_card as f64 * right_card as f64 * self.config.default_join_selectivity) as u64)
.max(1)
};
JoinPlan::Join(JoinNode {
left: Box::new(left),
right: Box::new(right),
join_vars,
cost,
cardinality,
})
}
pub fn stats(&self) -> OptimizationStats {
self.opt_stats.lock().clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_triple_pattern_variables() {
let pattern = TriplePattern::new("?s", "type", "Person");
let vars = pattern.variables();
assert_eq!(vars.len(), 1);
assert!(vars.contains("?s"));
}
#[test]
fn test_triple_pattern_bound_count() {
let pattern1 = TriplePattern::new("?s", "type", "Person");
assert_eq!(pattern1.bound_count(), 2);
let pattern2 = TriplePattern::new("?s", "?p", "?o");
assert_eq!(pattern2.bound_count(), 0);
let pattern3 = TriplePattern::new("Alice", "name", "?name");
assert_eq!(pattern3.bound_count(), 2);
}
#[test]
fn test_join_variables() {
let pattern1 = TriplePattern::new("?s", "type", "Person");
let pattern2 = TriplePattern::new("?s", "name", "?name");
let join_vars = pattern1.join_variables(&pattern2);
assert_eq!(join_vars.len(), 1);
assert!(join_vars.contains("?s"));
}
#[test]
fn test_optimizer_creation() {
let config = JoinOptimizerConfig::default();
let optimizer = JoinOptimizer::new(config);
assert_eq!(optimizer.config.algorithm, JoinAlgorithm::Greedy);
}
#[test]
fn test_single_pattern_optimization() {
let config = JoinOptimizerConfig::default();
let mut optimizer = JoinOptimizer::new(config);
let patterns = vec![TriplePattern::new("?s", "type", "Person")];
let plan = optimizer.optimize(patterns).unwrap();
match plan {
JoinPlan::Pattern(pattern) => {
assert_eq!(pattern.subject, "?s");
assert_eq!(pattern.predicate, "type");
assert_eq!(pattern.object, "Person");
}
_ => panic!("Expected pattern plan"),
}
}
#[test]
fn test_greedy_optimization() {
let config = JoinOptimizerConfig {
algorithm: JoinAlgorithm::Greedy,
..Default::default()
};
let mut optimizer = JoinOptimizer::new(config);
let patterns = vec![
TriplePattern::new("?s", "type", "Person"),
TriplePattern::new("?s", "name", "?name"),
TriplePattern::new("?s", "age", "?age"),
];
let plan = optimizer.optimize(patterns).unwrap();
match plan {
JoinPlan::Join(_) => {
}
_ => panic!("Expected join plan"),
}
let stats = optimizer.stats();
assert_eq!(stats.total_optimizations, 1);
assert_eq!(stats.greedy_optimizations, 1);
}
#[test]
fn test_dp_optimization() {
let config = JoinOptimizerConfig {
algorithm: JoinAlgorithm::DynamicProgramming,
..Default::default()
};
let mut optimizer = JoinOptimizer::new(config);
let patterns = vec![
TriplePattern::new("?s", "type", "Person"),
TriplePattern::new("?s", "name", "?name"),
TriplePattern::new("?s", "age", "?age"),
];
let plan = optimizer.optimize(patterns).unwrap();
match plan {
JoinPlan::Join(_) => {
}
_ => panic!("Expected join plan"),
}
let stats = optimizer.stats();
assert_eq!(stats.dp_optimizations, 1);
}
#[test]
fn test_cardinality_estimation() {
let config = JoinOptimizerConfig::default();
let optimizer = JoinOptimizer::new(config);
let pattern1 = TriplePattern::new("Alice", "name", "\"Alice\"");
let card1 = optimizer.estimate_pattern_cardinality(&pattern1);
assert_eq!(card1, 1);
let pattern2 = TriplePattern::new("?s", "name", "\"Alice\"");
let card2 = optimizer.estimate_pattern_cardinality(&pattern2);
assert!(card2 > 1 && card2 < 10000);
let pattern3 = TriplePattern::new("?s", "?p", "?o");
let card3 = optimizer.estimate_pattern_cardinality(&pattern3);
assert!(card3 > 10000);
}
#[test]
fn test_join_plan_variables() {
let pattern1 = TriplePattern::new("?s", "type", "Person");
let pattern2 = TriplePattern::new("?s", "name", "?name");
let plan1 = JoinPlan::Pattern(pattern1.clone());
let plan2 = JoinPlan::Pattern(pattern2.clone());
let join_vars: HashSet<String> = vec!["?s".to_string()].into_iter().collect();
let join_plan = JoinPlan::Join(JoinNode {
left: Box::new(plan1),
right: Box::new(plan2),
join_vars,
cost: 100.0,
cardinality: 1000,
});
let vars = join_plan.variables();
assert_eq!(vars.len(), 2);
assert!(vars.contains("?s"));
assert!(vars.contains("?name"));
}
#[test]
fn test_multiple_optimizations() {
let config = JoinOptimizerConfig::default();
let mut optimizer = JoinOptimizer::new(config);
let patterns1 = vec![
TriplePattern::new("?s", "type", "Person"),
TriplePattern::new("?s", "name", "?name"),
];
optimizer.optimize(patterns1).unwrap();
let patterns2 = vec![
TriplePattern::new("?x", "knows", "?y"),
TriplePattern::new("?y", "age", "?age"),
];
optimizer.optimize(patterns2).unwrap();
let stats = optimizer.stats();
assert_eq!(stats.total_optimizations, 2);
assert!(stats.avg_optimization_time_ms >= 0.0);
}
#[test]
fn test_empty_patterns() {
let config = JoinOptimizerConfig::default();
let mut optimizer = JoinOptimizer::new(config);
let patterns = vec![];
let result = optimizer.optimize(patterns);
assert!(result.is_err());
}
#[test]
fn test_large_pattern_set() {
let config = JoinOptimizerConfig {
algorithm: JoinAlgorithm::DynamicProgramming,
dp_max_patterns: 5,
..Default::default()
};
let mut optimizer = JoinOptimizer::new(config);
let patterns: Vec<_> = (0..10)
.map(|i| TriplePattern::new(&format!("?s{}", i), "type", "Person"))
.collect();
let plan = optimizer.optimize(patterns).unwrap();
let stats = optimizer.stats();
assert_eq!(stats.greedy_optimizations, 1);
assert_eq!(stats.dp_optimizations, 0);
}
}