use crate::dictionary::NodeId;
use crate::error::{Result, TdbError};
use crate::query_hints::IndexType;
use crate::query_optimizer::{QueryOptimizer, QueryPattern, QueryPlan};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
pub struct AdaptiveExecutor {
optimizer: Arc<QueryOptimizer>,
history: Arc<RwLock<ExecutionHistory>>,
config: AdaptiveConfig,
}
#[derive(Debug, Clone)]
pub struct AdaptiveConfig {
pub enabled: bool,
pub reoptimize_threshold: f64,
pub min_samples: usize,
pub max_history_entries: usize,
pub enable_plan_switching: bool,
}
impl Default for AdaptiveConfig {
fn default() -> Self {
Self {
enabled: true,
reoptimize_threshold: 2.0, min_samples: 3,
max_history_entries: 100,
enable_plan_switching: true,
}
}
}
#[derive(Debug, Clone)]
pub struct ExecutionStats {
pub pattern: QueryPattern,
pub index_used: IndexType,
pub execution_time: Duration,
pub actual_results: usize,
pub estimated_results: usize,
pub was_optimal: bool,
pub timestamp: Instant,
}
impl ExecutionStats {
pub fn estimation_error_ratio(&self) -> f64 {
if self.estimated_results == 0 {
return if self.actual_results == 0 {
1.0
} else {
f64::INFINITY
};
}
self.actual_results as f64 / self.estimated_results as f64
}
pub fn has_significant_deviation(&self, threshold: f64) -> bool {
let ratio = self.estimation_error_ratio();
ratio > threshold || ratio < (1.0 / threshold)
}
}
#[derive(Debug, Default)]
struct ExecutionHistory {
entries: HashMap<QueryPattern, Vec<ExecutionStats>>,
}
impl ExecutionHistory {
fn add_entry(&mut self, stats: ExecutionStats, max_entries: usize) {
let entries = self.entries.entry(stats.pattern.clone()).or_default();
entries.push(stats);
if entries.len() > max_entries {
entries.remove(0);
}
}
fn avg_execution_time(&self, pattern: &QueryPattern) -> Option<Duration> {
let entries = self.entries.get(pattern)?;
if entries.is_empty() {
return None;
}
let total_nanos: u128 = entries.iter().map(|e| e.execution_time.as_nanos()).sum();
let avg_nanos = total_nanos / entries.len() as u128;
Some(Duration::from_nanos(avg_nanos as u64))
}
fn avg_estimation_error(&self, pattern: &QueryPattern) -> Option<f64> {
let entries = self.entries.get(pattern)?;
if entries.is_empty() {
return None;
}
let total_error: f64 = entries.iter().map(|e| e.estimation_error_ratio()).sum();
Some(total_error / entries.len() as f64)
}
fn sample_count(&self, pattern: &QueryPattern) -> usize {
self.entries.get(pattern).map_or(0, |e| e.len())
}
fn has_sufficient_history(&self, pattern: &QueryPattern, min_samples: usize) -> bool {
self.sample_count(pattern) >= min_samples
}
}
impl AdaptiveExecutor {
pub fn new(optimizer: Arc<QueryOptimizer>) -> Self {
Self {
optimizer,
history: Arc::new(RwLock::new(ExecutionHistory::default())),
config: AdaptiveConfig::default(),
}
}
pub fn with_config(optimizer: Arc<QueryOptimizer>, config: AdaptiveConfig) -> Self {
Self {
optimizer,
history: Arc::new(RwLock::new(ExecutionHistory::default())),
config,
}
}
pub fn create_plan(&self, pattern: QueryPattern) -> Result<QueryPlan> {
if !self.config.enabled {
let hints = crate::query_hints::QueryHints::new();
return self.optimizer.optimize(pattern, &hints);
}
let history = self.history.read();
if history.has_sufficient_history(&pattern, self.config.min_samples) {
if let Some(avg_error) = history.avg_estimation_error(&pattern) {
let mut hints = crate::query_hints::QueryHints::new();
if avg_error > self.config.reoptimize_threshold {
hints = hints.with_caching(false);
}
return self.optimizer.optimize(pattern, &hints);
}
}
let hints = crate::query_hints::QueryHints::new();
self.optimizer.optimize(pattern, &hints)
}
pub fn record_execution(
&self,
pattern: QueryPattern,
plan: &QueryPlan,
actual_results: usize,
execution_time: Duration,
) {
let stats = ExecutionStats {
pattern: pattern.clone(),
index_used: plan.index,
execution_time,
actual_results,
estimated_results: plan.estimated_results,
was_optimal: !self.should_reoptimize(plan, actual_results),
timestamp: Instant::now(),
};
let mut history = self.history.write();
history.add_entry(stats, self.config.max_history_entries);
}
pub fn should_reoptimize(&self, plan: &QueryPlan, actual_results: usize) -> bool {
if !self.config.enabled {
return false;
}
let ratio = if plan.estimated_results == 0 {
if actual_results == 0 {
1.0
} else {
return true; }
} else {
actual_results as f64 / plan.estimated_results as f64
};
ratio > self.config.reoptimize_threshold || ratio < (1.0 / self.config.reoptimize_threshold)
}
pub fn get_statistics(&self, pattern: &QueryPattern) -> Option<PatternStatistics> {
let history = self.history.read();
if !history.has_sufficient_history(pattern, self.config.min_samples) {
return None;
}
Some(PatternStatistics {
sample_count: history.sample_count(pattern),
avg_execution_time: history.avg_execution_time(pattern)?,
avg_estimation_error: history.avg_estimation_error(pattern)?,
})
}
pub fn clear_history(&self) {
let mut history = self.history.write();
history.entries.clear();
}
pub fn tracked_patterns_count(&self) -> usize {
let history = self.history.read();
history.entries.len()
}
pub fn total_samples(&self) -> usize {
let history = self.history.read();
history.entries.values().map(|v| v.len()).sum()
}
pub fn reoptimize(&self, pattern: QueryPattern, actual_results: usize) -> Result<QueryPlan> {
let mut hints = crate::query_hints::QueryHints::new();
if actual_results > 10000 {
hints = hints.with_caching(false);
}
self.optimizer.optimize(pattern, &hints)
}
}
#[derive(Debug, Clone)]
pub struct PatternStatistics {
pub sample_count: usize,
pub avg_execution_time: Duration,
pub avg_estimation_error: f64,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::statistics::{StatisticsConfig, TripleStatistics};
fn create_test_executor() -> AdaptiveExecutor {
let stats = Arc::new(TripleStatistics::new(StatisticsConfig::default()));
let optimizer = Arc::new(QueryOptimizer::new(stats));
AdaptiveExecutor::new(optimizer)
}
#[test]
fn test_adaptive_executor_creation() {
let executor = create_test_executor();
assert_eq!(executor.tracked_patterns_count(), 0);
assert_eq!(executor.total_samples(), 0);
}
#[test]
fn test_execution_stats_error_ratio() {
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let stats = ExecutionStats {
pattern,
index_used: IndexType::SPO,
execution_time: Duration::from_millis(100),
actual_results: 200,
estimated_results: 100,
was_optimal: false,
timestamp: Instant::now(),
};
assert_eq!(stats.estimation_error_ratio(), 2.0);
assert!(stats.has_significant_deviation(1.5));
}
#[test]
fn test_record_and_retrieve_execution() {
let executor = create_test_executor();
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let plan = QueryPlan::new(pattern.clone(), IndexType::SPO, 100);
executor.record_execution(pattern.clone(), &plan, 100, Duration::from_millis(50));
executor.record_execution(pattern.clone(), &plan, 110, Duration::from_millis(55));
executor.record_execution(pattern.clone(), &plan, 90, Duration::from_millis(45));
let stats = executor.get_statistics(&pattern);
assert!(stats.is_some());
let stats = stats.unwrap();
assert_eq!(stats.sample_count, 3);
assert!(stats.avg_execution_time.as_millis() >= 45);
assert!(stats.avg_execution_time.as_millis() <= 55);
}
#[test]
fn test_reoptimization_trigger() {
let executor = create_test_executor();
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let plan = QueryPlan::new(pattern.clone(), IndexType::SPO, 100);
assert!(!executor.should_reoptimize(&plan, 120));
assert!(executor.should_reoptimize(&plan, 300));
assert!(executor.should_reoptimize(&plan, 30));
}
#[test]
fn test_history_size_limit() {
let config = AdaptiveConfig {
max_history_entries: 5,
..Default::default()
};
let stats = Arc::new(TripleStatistics::new(StatisticsConfig::default()));
let optimizer = Arc::new(QueryOptimizer::new(stats));
let executor = AdaptiveExecutor::with_config(optimizer, config);
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let plan = QueryPlan::new(pattern.clone(), IndexType::SPO, 100);
for i in 0..10 {
executor.record_execution(pattern.clone(), &plan, 100 + i, Duration::from_millis(50));
}
let stats = executor.get_statistics(&pattern).unwrap();
assert_eq!(stats.sample_count, 5);
}
#[test]
fn test_disabled_adaptive_execution() {
let config = AdaptiveConfig {
enabled: false,
..Default::default()
};
let stats = Arc::new(TripleStatistics::new(StatisticsConfig::default()));
let optimizer = Arc::new(QueryOptimizer::new(stats));
let executor = AdaptiveExecutor::with_config(optimizer, config);
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let plan = QueryPlan::new(pattern.clone(), IndexType::SPO, 100);
assert!(!executor.should_reoptimize(&plan, 1000));
}
#[test]
fn test_clear_history() {
let executor = create_test_executor();
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let plan = QueryPlan::new(pattern.clone(), IndexType::SPO, 100);
for _ in 0..5 {
executor.record_execution(pattern.clone(), &plan, 100, Duration::from_millis(50));
}
assert_eq!(executor.total_samples(), 5);
executor.clear_history();
assert_eq!(executor.total_samples(), 0);
assert_eq!(executor.tracked_patterns_count(), 0);
}
#[test]
fn test_insufficient_samples_no_stats() {
let executor = create_test_executor();
let pattern = QueryPattern::new(Some(NodeId::new(1)), None, None);
let plan = QueryPlan::new(pattern.clone(), IndexType::SPO, 100);
executor.record_execution(pattern.clone(), &plan, 100, Duration::from_millis(50));
executor.record_execution(pattern.clone(), &plan, 100, Duration::from_millis(50));
assert!(executor.get_statistics(&pattern).is_none());
}
#[test]
fn test_pattern_specific_tracking() {
let executor = create_test_executor();
let pattern1 = QueryPattern::new(Some(NodeId::new(1)), None, None);
let pattern2 = QueryPattern::new(None, Some(NodeId::new(2)), None);
let plan1 = QueryPlan::new(pattern1.clone(), IndexType::SPO, 100);
let plan2 = QueryPlan::new(pattern2.clone(), IndexType::POS, 200);
for _ in 0..3 {
executor.record_execution(pattern1.clone(), &plan1, 100, Duration::from_millis(50));
executor.record_execution(pattern2.clone(), &plan2, 200, Duration::from_millis(100));
}
assert_eq!(executor.tracked_patterns_count(), 2);
assert_eq!(executor.total_samples(), 6);
let stats1 = executor.get_statistics(&pattern1).unwrap();
let stats2 = executor.get_statistics(&pattern2).unwrap();
assert_eq!(stats1.sample_count, 3);
assert_eq!(stats2.sample_count, 3);
assert!(stats1.avg_execution_time < stats2.avg_execution_time);
}
}