use std::collections::{HashMap, VecDeque};
use std::time::{Duration, Instant};
use crate::error::{OptimError, Result};
use super::types::{
ComputationId, ErrorStatistics, ErrorType, ExecutionUtilization, PerformanceSample,
RecoveryStrategy, TPUBackendConfig, TaskExecutionResult,
};
const MAX_RETAINED_SAMPLES: usize = 1024;
#[derive(Debug)]
pub struct TPUErrorHandler {
recovery_enabled: bool,
error_statistics: ErrorStatistics,
recovery_strategies: HashMap<ErrorType, RecoveryStrategy>,
max_retry_attempts: usize,
operations_observed: usize,
recovery_attempts: usize,
recoveries_succeeded: usize,
}
impl TPUErrorHandler {
pub fn new(config: &TPUBackendConfig) -> Self {
let mut recovery_strategies = HashMap::new();
recovery_strategies.insert(ErrorType::TimeoutError, RecoveryStrategy::Retry);
recovery_strategies.insert(ErrorType::ResourceError, RecoveryStrategy::Retry);
recovery_strategies.insert(ErrorType::CommunicationError, RecoveryStrategy::Retry);
recovery_strategies.insert(ErrorType::MemoryError, RecoveryStrategy::Fallback);
recovery_strategies.insert(ErrorType::DeviceError, RecoveryStrategy::Migrate);
recovery_strategies.insert(ErrorType::ComputationError, RecoveryStrategy::Restart);
Self {
recovery_enabled: config.enable_error_recovery,
error_statistics: ErrorStatistics::default(),
recovery_strategies,
max_retry_attempts: config.max_retry_attempts,
operations_observed: 0,
recovery_attempts: 0,
recoveries_succeeded: 0,
}
}
pub fn get_error_rate(&self) -> f64 {
self.error_statistics.error_rate
}
pub fn error_statistics(&self) -> &ErrorStatistics {
&self.error_statistics
}
pub fn classify(error: &OptimError) -> ErrorType {
match error {
OptimError::DeviceError(_) => ErrorType::DeviceError,
OptimError::MemoryError(_) | OptimError::AllocationError(_) => ErrorType::MemoryError,
OptimError::TimeoutError(_) => ErrorType::TimeoutError,
OptimError::MutexError(_) => ErrorType::ResourceError,
_ => ErrorType::ComputationError,
}
}
pub fn record_operation(&mut self) {
self.operations_observed += 1;
self.refresh_error_rate();
}
pub fn record_error(&mut self, error: &OptimError, attempt: usize) -> bool {
let class = Self::classify(error);
self.error_statistics.total_errors += 1;
*self
.error_statistics
.errors_by_type
.entry(class)
.or_insert(0) += 1;
self.refresh_error_rate();
let retry = self.recovery_enabled
&& attempt + 1 < self.max_retry_attempts.max(1)
&& matches!(
self.recovery_strategies.get(&class),
Some(RecoveryStrategy::Retry)
);
if retry {
self.recovery_attempts += 1;
}
retry
}
pub fn record_recovery_success(&mut self) {
self.recoveries_succeeded += 1;
self.error_statistics.recovery_success_rate = if self.recovery_attempts == 0 {
0.0
} else {
self.recoveries_succeeded as f64 / self.recovery_attempts as f64
};
}
pub fn max_attempts(&self) -> usize {
if self.recovery_enabled {
self.max_retry_attempts.max(1)
} else {
1
}
}
fn refresh_error_rate(&mut self) {
let denominator = self
.operations_observed
.max(self.error_statistics.total_errors);
self.error_statistics.error_rate = if denominator == 0 {
0.0
} else {
self.error_statistics.total_errors as f64 / denominator as f64
};
}
}
#[derive(Debug)]
pub struct PerformanceMonitor {
enabled: bool,
pub total_executions: usize,
pub average_execution_time: Duration,
performance_history: VecDeque<PerformanceSample>,
collection_interval: Duration,
}
impl PerformanceMonitor {
pub fn new(config: &TPUBackendConfig) -> Self {
Self {
enabled: config.enable_performance_monitoring,
total_executions: 0,
average_execution_time: Duration::from_millis(0),
performance_history: VecDeque::new(),
collection_interval: Duration::from_millis(1000),
}
}
pub fn is_enabled(&self) -> bool {
self.enabled
}
pub fn performance_history(&self) -> &VecDeque<PerformanceSample> {
&self.performance_history
}
pub fn record_execution(
&mut self,
computation_id: ComputationId,
time: Duration,
results: &TaskExecutionResult,
utilization: ExecutionUtilization,
) {
self.total_executions += 1;
let count = self.total_executions as u32;
let previous_total = self
.average_execution_time
.saturating_mul(count.saturating_sub(1));
self.average_execution_time = (previous_total + time) / count.max(1);
if !self.enabled {
return;
}
let due = match self.performance_history.back() {
Some(last) => last.timestamp.elapsed() >= self.collection_interval,
None => true,
};
if !due {
return;
}
let seconds = time.as_secs_f64();
let throughput = if seconds > 0.0 {
results.output_data.len() as f64 / seconds
} else {
0.0
};
self.performance_history.push_back(PerformanceSample {
timestamp: Instant::now(),
computation: computation_id,
execution_time: results.execution_time,
throughput,
device_utilization: utilization.device,
memory_utilization: utilization.memory,
});
while self.performance_history.len() > MAX_RETAINED_SAMPLES {
self.performance_history.pop_front();
}
}
pub fn flush_metrics(&mut self) -> Result<()> {
self.performance_history.clear();
Ok(())
}
}