use super::error_policy::*;
use crate::cook::execution::dlq::DeadLetterQueue;
use crate::cook::execution::errors::MapReduceError;
use serde_json::json;
use std::path::PathBuf;
use std::time::Duration;
#[cfg(test)]
mod tests {
use super::*;
fn create_test_policy() -> WorkflowErrorPolicy {
WorkflowErrorPolicy {
on_item_failure: ItemFailureAction::Dlq,
continue_on_failure: true,
max_failures: Some(5),
failure_threshold: Some(0.3),
error_collection: ErrorCollectionStrategy::Aggregate,
circuit_breaker: None,
retry_config: None,
}
}
fn create_test_error(message: &str) -> MapReduceError {
MapReduceError::ProcessingError(message.to_string())
}
#[tokio::test]
async fn test_error_policy_dlq_handling() {
let policy = create_test_policy();
let executor = ErrorPolicyExecutor::new(policy);
let dlq = DeadLetterQueue::new(
"test-job-id".to_string(),
PathBuf::from("/tmp/test-dlq"),
100, 7, None, )
.await
.unwrap();
let item = json!({"id": "test-item"});
let error = create_test_error("Test processing error");
let action = executor
.handle_item_failure("test-item", &item, &error, Some(&dlq))
.await
.unwrap();
assert!(matches!(action, FailureAction::Continue));
let metrics = executor.get_metrics();
assert_eq!(metrics.failed, 1);
assert_eq!(metrics.total_items, 1);
}
#[tokio::test]
async fn test_error_policy_skip_action() {
let mut policy = create_test_policy();
policy.on_item_failure = ItemFailureAction::Skip;
let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
let action = executor
.handle_item_failure("test-item", &item, &error, None)
.await
.unwrap();
assert!(matches!(action, FailureAction::Skip));
}
#[tokio::test]
async fn test_error_policy_stop_action() {
let mut policy = create_test_policy();
policy.on_item_failure = ItemFailureAction::Stop;
let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
let action = executor
.handle_item_failure("test-item", &item, &error, None)
.await
.unwrap();
if let FailureAction::Stop(msg) = action {
assert!(msg.contains("test-item"));
} else {
panic!("Expected Stop action");
}
}
#[tokio::test]
async fn test_error_policy_retry_action() {
let mut policy = create_test_policy();
policy.on_item_failure = ItemFailureAction::Retry;
policy.retry_config = Some(RetryConfig {
max_attempts: 3,
backoff: BackoffStrategy::default(),
});
let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
let action = executor
.handle_item_failure("test-item", &item, &error, None)
.await
.unwrap();
if let FailureAction::Retry(config) = action {
assert_eq!(config.max_attempts, 3);
} else {
panic!("Expected Retry action");
}
}
#[tokio::test]
async fn test_max_failures_threshold() {
let mut policy = create_test_policy();
policy.max_failures = Some(2);
let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
let action1 = executor
.handle_item_failure("item1", &item, &error, None)
.await
.unwrap();
assert!(matches!(action1, FailureAction::Continue));
let action2 = executor
.handle_item_failure("item2", &item, &error, None)
.await
.unwrap();
assert!(matches!(action2, FailureAction::Stop(_)));
}
#[tokio::test]
async fn test_failure_rate_threshold() {
let mut policy = create_test_policy();
policy.failure_threshold = Some(0.25); let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
for _ in 0..10 {
executor.record_success();
}
for i in 1..=3 {
let action = executor
.handle_item_failure(&format!("item{}", i), &item, &error, None)
.await
.unwrap();
assert!(matches!(action, FailureAction::Continue));
}
let action = executor
.handle_item_failure("item4", &item, &error, None)
.await
.unwrap();
assert!(matches!(action, FailureAction::Stop(_)));
}
#[tokio::test]
async fn test_continue_on_failure_false() {
let mut policy = create_test_policy();
policy.continue_on_failure = false;
let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
let action = executor
.handle_item_failure("item1", &item, &error, None)
.await
.unwrap();
assert!(matches!(action, FailureAction::Stop(_)));
}
#[test]
fn test_error_collection_aggregate() {
let mut policy = create_test_policy();
policy.error_collection = ErrorCollectionStrategy::Aggregate;
let executor = ErrorPolicyExecutor::new(policy);
executor.collect_error("Error 1".to_string());
executor.collect_error("Error 2".to_string());
executor.collect_error("Error 3".to_string());
let errors = executor.get_collected_errors();
assert_eq!(errors.len(), 3);
assert!(errors.contains(&"Error 1".to_string()));
assert!(errors.contains(&"Error 2".to_string()));
assert!(errors.contains(&"Error 3".to_string()));
}
#[test]
fn test_error_collection_batched() {
let mut policy = create_test_policy();
policy.error_collection = ErrorCollectionStrategy::Batched { size: 2 };
let executor = ErrorPolicyExecutor::new(policy);
executor.collect_error("Error 1".to_string());
executor.collect_error("Error 2".to_string());
let errors = executor.get_collected_errors();
assert_eq!(errors.len(), 0);
executor.collect_error("Error 3".to_string());
let errors = executor.get_collected_errors();
assert_eq!(errors.len(), 1);
}
#[test]
fn test_circuit_breaker_basic() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
success_threshold: 2,
timeout: Duration::from_millis(100),
half_open_requests: 1,
};
let breaker = CircuitBreaker::new(config);
assert!(!breaker.is_open());
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()); }
#[tokio::test]
async fn test_error_policy_with_circuit_breaker() {
let mut policy = create_test_policy();
policy.circuit_breaker = Some(CircuitBreakerConfig {
failure_threshold: 2,
success_threshold: 2,
timeout: Duration::from_millis(100),
half_open_requests: 1,
});
let executor = ErrorPolicyExecutor::new(policy);
let item = json!({"id": "test-item"});
let error = create_test_error("Test error");
let action1 = executor
.handle_item_failure("item1", &item, &error, None)
.await
.unwrap();
assert!(matches!(action1, FailureAction::Continue));
let action2 = executor
.handle_item_failure("item2", &item, &error, None)
.await
.unwrap();
assert!(matches!(action2, FailureAction::Continue));
let action3 = executor
.handle_item_failure("item3", &item, &error, None)
.await
.unwrap();
if let FailureAction::Stop(msg) = action3 {
assert!(msg.contains("Circuit breaker"));
} else {
panic!("Expected Stop due to circuit breaker");
}
}
#[test]
fn test_error_metrics_tracking() {
let policy = create_test_policy();
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.failed, 0);
assert_eq!(metrics.total_items, 3);
assert_eq!(metrics.failure_rate, 0.0);
}
#[test]
fn test_failure_pattern_detection() {
let policy = create_test_policy();
let executor = ErrorPolicyExecutor::new(policy);
for _ in 0..3 {
executor.update_metrics(&MapReduceError::Timeout);
}
let metrics = executor.get_metrics();
assert!(!metrics.failure_patterns.is_empty());
let timeout_pattern = metrics
.failure_patterns
.iter()
.find(|p| p.pattern_type.contains("Timeout"))
.expect("Should detect timeout pattern");
assert_eq!(timeout_pattern.frequency, 3);
assert!(timeout_pattern.suggested_action.contains("timeout"));
}
}