use crate::error_handling::{ErrorContext, SuggestedAction};
use crate::metrics::Metrics;
use crate::types::AiLibError;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Mutex};
use std::time::Instant;
pub struct ErrorRecoveryManager {
error_history: Arc<Mutex<VecDeque<ErrorRecord>>>,
recovery_strategies: HashMap<ErrorType, Box<dyn RecoveryStrategy>>,
metrics: Option<Arc<dyn Metrics>>,
#[allow(dead_code)] start_time: Instant,
error_patterns: Arc<Mutex<HashMap<ErrorType, ErrorPattern>>>,
}
impl Default for ErrorRecoveryManager {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorPattern {
pub error_type: ErrorType,
pub count: u32,
pub first_occurrence: chrono::DateTime<chrono::Utc>,
pub last_occurrence: chrono::DateTime<chrono::Utc>,
pub frequency: f64, pub suggested_action: SuggestedAction,
pub recovery_attempts: u32,
pub successful_recoveries: u32,
}
#[derive(Debug, Clone)]
pub struct ErrorRecord {
pub error_type: ErrorType,
pub context: ErrorContext,
pub timestamp: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum ErrorType {
RateLimit,
Network,
Authentication,
Provider,
Timeout,
Configuration,
Validation,
Serialization,
Deserialization,
FileOperation,
ModelNotFound,
ContextLengthExceeded,
UnsupportedFeature,
Unknown,
}
#[async_trait]
pub trait RecoveryStrategy: Send + Sync {
async fn can_recover(&self, error: &AiLibError) -> bool;
async fn recover(&self, error: &AiLibError, context: &ErrorContext) -> Result<(), AiLibError>;
}
impl ErrorRecoveryManager {
pub fn new() -> Self {
Self {
error_history: Arc::new(Mutex::new(VecDeque::new())),
recovery_strategies: HashMap::new(),
metrics: None,
start_time: Instant::now(),
error_patterns: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn with_metrics(metrics: Arc<dyn Metrics>) -> Self {
Self {
error_history: Arc::new(Mutex::new(VecDeque::new())),
recovery_strategies: HashMap::new(),
metrics: Some(metrics),
start_time: Instant::now(),
error_patterns: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn register_strategy(
&mut self,
error_type: ErrorType,
strategy: Box<dyn RecoveryStrategy>,
) {
self.recovery_strategies.insert(error_type, strategy);
}
pub async fn handle_error(
&self,
error: &AiLibError,
context: &ErrorContext,
) -> Result<(), AiLibError> {
let error_type = self.classify_error(error);
self.record_error(error_type.clone(), context.clone()).await;
if let Some(strategy) = self.recovery_strategies.get(&error_type) {
if strategy.can_recover(error).await {
return strategy.recover(error, context).await;
}
}
Err((*error).clone())
}
fn classify_error(&self, error: &AiLibError) -> ErrorType {
match error {
AiLibError::RateLimitExceeded(_) => ErrorType::RateLimit,
AiLibError::NetworkError(_) => ErrorType::Network,
AiLibError::AuthenticationError(_) => ErrorType::Authentication,
AiLibError::ProviderError(_) => ErrorType::Provider,
AiLibError::TimeoutError(_) => ErrorType::Timeout,
AiLibError::ConfigurationError(_) => ErrorType::Configuration,
AiLibError::InvalidRequest(_) => ErrorType::Validation,
AiLibError::SerializationError(_) => ErrorType::Serialization,
AiLibError::DeserializationError(_) => ErrorType::Deserialization,
AiLibError::FileError(_) => ErrorType::FileOperation,
AiLibError::ModelNotFound(_) => ErrorType::ModelNotFound,
AiLibError::ContextLengthExceeded(_) => ErrorType::ContextLengthExceeded,
AiLibError::UnsupportedFeature(_) => ErrorType::UnsupportedFeature,
_ => ErrorType::Unknown,
}
}
fn generate_suggested_action(
&self,
error_type: &ErrorType,
pattern: &ErrorPattern,
) -> SuggestedAction {
match error_type {
ErrorType::RateLimit => {
if pattern.frequency > 10.0 {
SuggestedAction::SwitchProvider {
alternative_providers: vec!["groq".to_string(), "anthropic".to_string()],
}
} else {
SuggestedAction::Retry {
delay_ms: 60000,
max_attempts: 3,
}
}
}
ErrorType::Network => SuggestedAction::Retry {
delay_ms: 2000,
max_attempts: 5,
},
ErrorType::Authentication => SuggestedAction::CheckCredentials,
ErrorType::Provider => SuggestedAction::SwitchProvider {
alternative_providers: vec!["openai".to_string(), "groq".to_string()],
},
ErrorType::Timeout => SuggestedAction::Retry {
delay_ms: 5000,
max_attempts: 3,
},
ErrorType::ContextLengthExceeded => SuggestedAction::ReduceRequestSize {
max_tokens: Some(1000),
},
ErrorType::ModelNotFound => SuggestedAction::ContactSupport {
reason: "Model not found - please verify model name".to_string(),
},
_ => SuggestedAction::NoAction,
}
}
async fn record_error(&self, error_type: ErrorType, mut context: ErrorContext) {
let now = chrono::Utc::now();
let record = ErrorRecord {
error_type: error_type.clone(),
context: context.clone(),
timestamp: now,
};
self.update_error_pattern(&error_type, now).await;
let suggested_action = self.get_suggested_action_for_error(&error_type).await;
context.suggested_action = suggested_action;
{
let mut history = self.error_history.lock().unwrap();
history.push_back(record);
if history.len() > 1000 {
history.pop_front();
}
}
if let Some(metrics) = &self.metrics {
metrics
.incr_counter(&format!("errors.{}", self.error_type_name(&error_type)), 1)
.await;
}
}
async fn update_error_pattern(
&self,
error_type: &ErrorType,
timestamp: chrono::DateTime<chrono::Utc>,
) {
let mut patterns = self.error_patterns.lock().unwrap();
let entry = patterns.entry(error_type.clone());
use std::collections::hash_map::Entry;
match entry {
Entry::Occupied(mut occ) => {
let pattern = occ.get_mut();
pattern.count += 1;
pattern.last_occurrence = timestamp;
let duration = pattern
.last_occurrence
.signed_duration_since(pattern.first_occurrence);
if duration.num_minutes() > 0 {
pattern.frequency = pattern.count as f64 / duration.num_minutes() as f64;
}
pattern.suggested_action = self.generate_suggested_action(error_type, pattern);
}
Entry::Vacant(vac) => {
vac.insert(ErrorPattern {
error_type: error_type.clone(),
count: 1,
first_occurrence: timestamp,
last_occurrence: timestamp,
frequency: 0.0,
suggested_action: SuggestedAction::NoAction,
recovery_attempts: 0,
successful_recoveries: 0,
});
}
}
}
async fn get_suggested_action_for_error(&self, error_type: &ErrorType) -> SuggestedAction {
let patterns = self.error_patterns.lock().unwrap();
patterns
.get(error_type)
.map(|p| p.suggested_action.clone())
.unwrap_or(SuggestedAction::NoAction)
}
fn error_type_name(&self, error_type: &ErrorType) -> String {
match error_type {
ErrorType::RateLimit => "rate_limit".to_string(),
ErrorType::Network => "network".to_string(),
ErrorType::Authentication => "authentication".to_string(),
ErrorType::Provider => "provider".to_string(),
ErrorType::Timeout => "timeout".to_string(),
ErrorType::Configuration => "configuration".to_string(),
ErrorType::Validation => "validation".to_string(),
ErrorType::Serialization => "serialization".to_string(),
ErrorType::Deserialization => "deserialization".to_string(),
ErrorType::FileOperation => "file_operation".to_string(),
ErrorType::ModelNotFound => "model_not_found".to_string(),
ErrorType::ContextLengthExceeded => "context_length_exceeded".to_string(),
ErrorType::UnsupportedFeature => "unsupported_feature".to_string(),
ErrorType::Unknown => "unknown".to_string(),
}
}
pub fn get_error_patterns(&self) -> HashMap<ErrorType, ErrorPattern> {
self.error_patterns.lock().unwrap().clone()
}
pub fn get_error_statistics(&self) -> ErrorStatistics {
let patterns = self.error_patterns.lock().unwrap();
let total_errors: u32 = patterns.values().map(|p| p.count).sum();
let most_common_error = patterns
.values()
.max_by_key(|p| p.count)
.map(|p| p.error_type.clone());
ErrorStatistics {
total_errors,
unique_error_types: patterns.len(),
most_common_error,
patterns: patterns.clone(),
}
}
pub fn reset(&self) {
let mut history = self.error_history.lock().unwrap();
history.clear();
let mut patterns = self.error_patterns.lock().unwrap();
patterns.clear();
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorStatistics {
pub total_errors: u32,
pub unique_error_types: usize,
pub most_common_error: Option<ErrorType>,
pub patterns: HashMap<ErrorType, ErrorPattern>,
}