use crate::cook::execution::dlq::{DeadLetterQueue, DeadLetteredItem};
use crate::cook::execution::errors::MapReduceError;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use tracing::{debug, info, warn};
#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum ItemFailureAction {
#[default]
Dlq,
Retry,
Skip,
Stop,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq)]
#[serde(rename_all = "snake_case")]
pub enum ErrorCollectionStrategy {
#[default]
Aggregate,
Immediate,
Batched { size: usize },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CircuitBreakerConfig {
#[serde(default = "default_failure_threshold")]
pub failure_threshold: usize,
#[serde(default = "default_success_threshold")]
pub success_threshold: usize,
#[serde(default = "default_circuit_timeout", with = "humantime_serde")]
pub timeout: Duration,
#[serde(default = "default_half_open_requests")]
pub half_open_requests: usize,
}
fn default_failure_threshold() -> usize {
5
}
fn default_success_threshold() -> usize {
3
}
fn default_circuit_timeout() -> Duration {
Duration::from_secs(30)
}
fn default_half_open_requests() -> usize {
3
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: default_failure_threshold(),
success_threshold: default_success_threshold(),
timeout: default_circuit_timeout(),
half_open_requests: default_half_open_requests(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RetryConfig {
#[serde(default = "default_max_attempts")]
pub max_attempts: u32,
#[serde(default)]
pub backoff: BackoffStrategy,
}
fn default_max_attempts() -> u32 {
3
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum BackoffStrategy {
Fixed { delay: Duration },
Linear {
initial: Duration,
increment: Duration,
},
Exponential { initial: Duration, multiplier: f64 },
Fibonacci { initial: Duration },
}
impl Default for BackoffStrategy {
fn default() -> Self {
BackoffStrategy::Exponential {
initial: Duration::from_secs(1),
multiplier: 2.0,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowErrorPolicy {
#[serde(default)]
pub on_item_failure: ItemFailureAction,
#[serde(default = "default_continue_on_failure")]
pub continue_on_failure: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_failures: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub failure_threshold: Option<f64>,
#[serde(default)]
pub error_collection: ErrorCollectionStrategy,
#[serde(skip_serializing_if = "Option::is_none")]
pub circuit_breaker: Option<CircuitBreakerConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retry_config: Option<RetryConfig>,
}
fn default_continue_on_failure() -> bool {
true
}
impl Default for WorkflowErrorPolicy {
fn default() -> Self {
Self {
on_item_failure: ItemFailureAction::default(),
continue_on_failure: default_continue_on_failure(),
max_failures: None,
failure_threshold: None,
error_collection: ErrorCollectionStrategy::default(),
circuit_breaker: None,
retry_config: None,
}
}
}
#[derive(Debug, Clone)]
pub enum FailureAction {
Continue,
Retry(RetryConfig),
Skip,
Stop(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorMetrics {
pub total_items: usize,
pub successful: usize,
pub failed: usize,
pub skipped: usize,
pub failure_rate: f64,
pub error_types: HashMap<String, usize>,
pub failure_patterns: Vec<FailurePattern>,
}
impl Default for ErrorMetrics {
fn default() -> Self {
Self {
total_items: 0,
successful: 0,
failed: 0,
skipped: 0,
failure_rate: 0.0,
error_types: HashMap::new(),
failure_patterns: Vec::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FailurePattern {
pub pattern_type: String,
pub frequency: usize,
pub items: Vec<String>,
pub suggested_action: String,
}
#[derive(Debug, Clone, PartialEq)]
pub enum CircuitState {
Closed,
Open { since: Instant },
HalfOpen { remaining_tests: usize },
}
#[derive(Debug)]
pub struct CircuitBreaker {
config: CircuitBreakerConfig,
state: Arc<Mutex<CircuitState>>,
consecutive_failures: Arc<Mutex<usize>>,
consecutive_successes: Arc<Mutex<usize>>,
}
impl CircuitBreaker {
pub fn new(config: CircuitBreakerConfig) -> Self {
Self {
config,
state: Arc::new(Mutex::new(CircuitState::Closed)),
consecutive_failures: Arc::new(Mutex::new(0)),
consecutive_successes: Arc::new(Mutex::new(0)),
}
}
pub fn is_open(&self) -> bool {
let mut state = self.state.lock().unwrap();
match *state {
CircuitState::Open { since } => {
if since.elapsed() >= self.config.timeout {
*state = CircuitState::HalfOpen {
remaining_tests: self.config.half_open_requests,
};
false
} else {
true
}
}
CircuitState::HalfOpen { .. } => false,
CircuitState::Closed => false,
}
}
pub fn record_success(&self) {
let mut state = self.state.lock().unwrap();
let mut successes = self.consecutive_successes.lock().unwrap();
let mut failures = self.consecutive_failures.lock().unwrap();
*failures = 0;
*successes += 1;
if let CircuitState::HalfOpen { .. } = *state {
if *successes >= self.config.success_threshold {
*state = CircuitState::Closed;
info!("Circuit breaker closed after {} successes", successes);
}
}
}
pub fn record_failure(&self) {
let mut state = self.state.lock().unwrap();
let mut failures = self.consecutive_failures.lock().unwrap();
let mut successes = self.consecutive_successes.lock().unwrap();
*successes = 0;
*failures += 1;
match *state {
CircuitState::Closed => {
if *failures >= self.config.failure_threshold {
*state = CircuitState::Open {
since: Instant::now(),
};
warn!("Circuit breaker opened after {} failures", failures);
}
}
CircuitState::HalfOpen {
mut remaining_tests,
} => {
remaining_tests -= 1;
if remaining_tests == 0 {
*state = CircuitState::Open {
since: Instant::now(),
};
warn!("Circuit breaker re-opened after test failures");
} else {
*state = CircuitState::HalfOpen { remaining_tests };
}
}
_ => {}
}
}
}
pub struct ErrorPolicyExecutor {
policy: WorkflowErrorPolicy,
metrics: Arc<Mutex<ErrorMetrics>>,
circuit_breaker: Option<CircuitBreaker>,
collected_errors: Arc<Mutex<Vec<String>>>,
}
impl ErrorPolicyExecutor {
pub fn new(policy: WorkflowErrorPolicy) -> Self {
let circuit_breaker = policy
.circuit_breaker
.as_ref()
.map(|config| CircuitBreaker::new(config.clone()));
Self {
policy,
metrics: Arc::new(Mutex::new(ErrorMetrics::default())),
circuit_breaker,
collected_errors: Arc::new(Mutex::new(Vec::new())),
}
}
pub async fn handle_item_failure(
&self,
item_id: &str,
item: &serde_json::Value,
error: &MapReduceError,
dlq: Option<&DeadLetterQueue>,
) -> Result<FailureAction, MapReduceError> {
self.update_metrics(error);
if let Some(ref breaker) = self.circuit_breaker {
if breaker.is_open() {
return Ok(FailureAction::Stop("Circuit breaker open".to_string()));
}
breaker.record_failure();
}
if self.should_stop_on_threshold() {
return Ok(FailureAction::Stop(
"Failure threshold exceeded".to_string(),
));
}
match &self.policy.on_item_failure {
ItemFailureAction::Dlq => {
if let Some(dlq) = dlq {
self.send_to_dlq(item_id, item, error, dlq).await?;
}
Ok(FailureAction::Continue)
}
ItemFailureAction::Retry => {
if let Some(ref retry_config) = self.policy.retry_config {
Ok(FailureAction::Retry(retry_config.clone()))
} else {
Ok(FailureAction::Skip)
}
}
ItemFailureAction::Skip => {
debug!("Skipping failed item: {}", item_id);
Ok(FailureAction::Skip)
}
ItemFailureAction::Stop => Ok(FailureAction::Stop(format!("Item {} failed", item_id))),
ItemFailureAction::Custom(handler_name) => {
warn!("Custom handler {} not implemented, skipping", handler_name);
Ok(FailureAction::Skip)
}
}
}
pub fn record_success(&self) {
let mut metrics = self.metrics.lock().unwrap();
metrics.successful += 1;
metrics.total_items += 1;
self.update_failure_rate(&mut metrics);
if let Some(ref breaker) = self.circuit_breaker {
breaker.record_success();
}
}
pub fn update_metrics(&self, error: &MapReduceError) {
let mut metrics = self.metrics.lock().unwrap();
metrics.failed += 1;
metrics.total_items += 1;
let error_type = format!("{:?}", error);
*metrics.error_types.entry(error_type).or_insert(0) += 1;
self.update_failure_rate(&mut metrics);
self.detect_patterns(&mut metrics);
}
fn update_failure_rate(&self, metrics: &mut ErrorMetrics) {
if metrics.total_items > 0 {
metrics.failure_rate = metrics.failed as f64 / metrics.total_items as f64;
}
}
fn detect_patterns(&self, metrics: &mut ErrorMetrics) {
for (error_type, count) in &metrics.error_types {
if *count >= 3 {
let exists = metrics
.failure_patterns
.iter()
.any(|p| p.pattern_type == *error_type);
if !exists {
metrics.failure_patterns.push(FailurePattern {
pattern_type: error_type.clone(),
frequency: *count,
items: Vec::new(), suggested_action: self.suggest_action(error_type),
});
}
}
}
}
fn suggest_action(&self, error_type: &str) -> String {
if error_type.contains("Timeout") {
"Consider increasing timeout_per_agent".to_string()
} else if error_type.contains("Network") {
"Check network connectivity and retry settings".to_string()
} else if error_type.contains("Permission") {
"Verify file permissions and access rights".to_string()
} else {
"Review error logs for more details".to_string()
}
}
fn should_stop_on_threshold(&self) -> bool {
let metrics = self.metrics.lock().unwrap();
if let Some(max_failures) = self.policy.max_failures {
if metrics.failed >= max_failures {
warn!(
"Max failures reached: {} >= {}",
metrics.failed, max_failures
);
return true;
}
}
if let Some(threshold) = self.policy.failure_threshold {
if metrics.total_items >= 10 && metrics.failure_rate > threshold {
warn!(
"Failure rate exceeded: {:.2}% > {:.2}%",
metrics.failure_rate * 100.0,
threshold * 100.0
);
return true;
}
}
if !self.policy.continue_on_failure && metrics.failed > 0 {
warn!("Stopping due to failure (continue_on_failure=false)");
return true;
}
false
}
async fn send_to_dlq(
&self,
item_id: &str,
item: &serde_json::Value,
error: &MapReduceError,
dlq: &DeadLetterQueue,
) -> Result<(), MapReduceError> {
let now = chrono::Utc::now();
let dlq_item = DeadLetteredItem {
item_id: item_id.to_string(),
item_data: item.clone(),
first_attempt: now,
last_attempt: now,
failure_count: 1,
failure_history: vec![crate::cook::execution::dlq::FailureDetail {
attempt_number: 1,
timestamp: now,
error_type: crate::cook::execution::dlq::ErrorType::Unknown,
error_message: error.to_string(),
error_context: None,
stack_trace: None,
agent_id: "map-agent".to_string(),
step_failed: "map_phase".to_string(),
duration_ms: 0,
json_log_location: None,
}],
error_signature: format!("{:?}", error),
worktree_artifacts: None,
reprocess_eligible: true,
manual_review_required: false,
};
dlq.add(dlq_item)
.await
.map_err(|e| MapReduceError::DlqError(e.to_string()))?;
info!("Sent failed item {} to DLQ", item_id);
Ok(())
}
pub fn get_metrics(&self) -> ErrorMetrics {
self.metrics.lock().unwrap().clone()
}
pub fn get_collected_errors(&self) -> Vec<String> {
self.collected_errors.lock().unwrap().clone()
}
pub fn collect_error(&self, error: String) {
match self.policy.error_collection {
ErrorCollectionStrategy::Aggregate => {
self.collected_errors.lock().unwrap().push(error);
}
ErrorCollectionStrategy::Immediate => {
warn!("Item failure: {}", error);
}
ErrorCollectionStrategy::Batched { size } => {
let mut errors = self.collected_errors.lock().unwrap();
errors.push(error);
if errors.len() >= size {
warn!("Batch of {} errors collected", errors.len());
for err in errors.drain(..) {
warn!(" - {}", err);
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_circuit_breaker() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
success_threshold: 2,
timeout: Duration::from_millis(100),
half_open_requests: 1,
};
let breaker = CircuitBreaker::new(config);
assert!(!breaker.is_open());
breaker.record_failure();
breaker.record_failure();
assert!(!breaker.is_open());
breaker.record_failure();
assert!(breaker.is_open());
std::thread::sleep(Duration::from_millis(150));
assert!(!breaker.is_open());
breaker.record_success();
breaker.record_success();
assert!(!breaker.is_open());
}
#[test]
fn test_error_policy_thresholds() {
let policy = WorkflowErrorPolicy {
max_failures: Some(5),
failure_threshold: Some(0.3),
..Default::default()
};
let executor = ErrorPolicyExecutor::new(policy);
executor.record_success();
executor.record_success();
executor.record_success();
let metrics = executor.get_metrics();
assert_eq!(metrics.successful, 3);
assert_eq!(metrics.failure_rate, 0.0);
assert!(!executor.should_stop_on_threshold());
}
}