pub mod backoff;
pub use backoff::ExponentialBackoffRecovery;
use serde::{Deserialize, Serialize};
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, serde::Serialize, serde::Deserialize,
)]
#[serde(rename_all = "snake_case")]
pub enum FailureSeverity {
Low,
Medium,
High,
Critical,
}
impl fmt::Display for FailureSeverity {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Low => write!(f, "low"),
Self::Medium => write!(f, "medium"),
Self::High => write!(f, "high"),
Self::Critical => write!(f, "critical"),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct ReflectionContext {
pub task: String,
pub attempt: u32,
pub max_attempts: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CorrectionType {
InputFix,
ToolChange,
PrerequisiteFix,
ApproachChange,
Escalate,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Correction {
pub correction_type: CorrectionType,
pub description: String,
pub modified_input: Option<serde_json::Value>,
pub alternative_tool: Option<String>,
pub guidance: Option<String>,
}
#[derive(Debug, Clone)]
pub enum CorrectionResult {
Applied,
Failed(String),
Skipped,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct FailureAnalysis {
pub is_recoverable: bool,
pub root_cause: String,
pub severity: FailureSeverity,
pub correction: Option<Correction>,
pub context: String,
}
#[derive(Debug, thiserror::Error)]
pub enum ReflectionError {
#[error("reflection skipped: {0}")]
Skipped(String),
#[error("reflection internal error: {0}")]
Internal(String),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RecoveryAction {
Retry {
delay: Duration,
},
Skip(String),
AskUser(String),
Fail(String),
}
impl RecoveryAction {
#[must_use]
pub fn delay(&self) -> Option<Duration> {
match self {
Self::Retry { delay } => Some(*delay),
_ => None,
}
}
#[must_use]
pub fn is_retry(&self) -> bool {
matches!(self, Self::Retry { .. })
}
#[must_use]
pub fn is_fail(&self) -> bool {
matches!(self, Self::Fail(_))
}
#[must_use]
pub fn is_skip(&self) -> bool {
matches!(self, Self::Skip(_))
}
#[must_use]
pub fn is_ask_user(&self) -> bool {
matches!(self, Self::AskUser(_))
}
}
impl fmt::Display for RecoveryAction {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Retry { delay } => {
write!(f, "retry after {delay:?}")
}
Self::Skip(reason) => write!(f, "skip: {reason}"),
Self::AskUser(prompt) => write!(f, "ask user: {prompt}"),
Self::Fail(reason) => write!(f, "fail: {reason}"),
}
}
}
pub trait Reflector: Send + Sync {
fn analyze(
&self,
error: &str,
tool_name: &str,
tool_input: &serde_json::Value,
context: &ReflectionContext,
) -> Pin<Box<dyn Future<Output = Result<FailureAnalysis, ReflectionError>> + Send + '_>>;
}
pub trait RecoveryStrategy: Send + Sync {
fn decide(
&self,
analysis: &FailureAnalysis,
attempt: u32,
max_attempts: u32,
) -> Pin<Box<dyn Future<Output = RecoveryAction> + Send + '_>>;
}
pub struct NoopReflector;
impl Reflector for NoopReflector {
fn analyze(
&self,
error: &str,
_tool_name: &str,
_tool_input: &serde_json::Value,
_context: &ReflectionContext,
) -> Pin<Box<dyn Future<Output = Result<FailureAnalysis, ReflectionError>> + Send + '_>> {
let root_cause = error.to_string();
Box::pin(async move {
Ok(FailureAnalysis {
is_recoverable: false,
root_cause,
severity: FailureSeverity::Medium,
correction: None,
context: String::new(),
})
})
}
}
impl fmt::Debug for NoopReflector {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NoopReflector").finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn severity_ordering() {
assert!(FailureSeverity::Low < FailureSeverity::Medium);
assert!(FailureSeverity::Medium < FailureSeverity::High);
assert!(FailureSeverity::High < FailureSeverity::Critical);
}
#[test]
fn severity_display() {
assert_eq!(FailureSeverity::Low.to_string(), "low");
assert_eq!(FailureSeverity::Medium.to_string(), "medium");
assert_eq!(FailureSeverity::High.to_string(), "high");
assert_eq!(FailureSeverity::Critical.to_string(), "critical");
}
#[test]
fn severity_serde_round_trip() {
let severity = FailureSeverity::High;
let json = serde_json::to_string(&severity).unwrap();
assert_eq!(json, "\"high\"");
let deserialized: FailureSeverity = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized, severity);
}
#[test]
fn reflection_context_fields() {
let ctx = ReflectionContext {
task: "fix bug".to_string(),
attempt: 2,
max_attempts: 5,
};
assert_eq!(ctx.task, "fix bug");
assert_eq!(ctx.attempt, 2);
assert_eq!(ctx.max_attempts, 5);
}
#[test]
fn failure_analysis_recoverable() {
let analysis = FailureAnalysis {
is_recoverable: true,
root_cause: "timeout".to_string(),
severity: FailureSeverity::Low,
correction: None,
context: "network".to_string(),
};
assert!(analysis.is_recoverable);
assert_eq!(analysis.root_cause, "timeout");
}
#[test]
fn failure_analysis_with_correction() {
let correction = Correction {
correction_type: CorrectionType::InputFix,
description: "fix path".to_string(),
modified_input: Some(serde_json::json!({"path": "/correct/path"})),
alternative_tool: None,
guidance: Some("check paths".to_string()),
};
let analysis = FailureAnalysis {
is_recoverable: true,
root_cause: "file not found".to_string(),
severity: FailureSeverity::Medium,
correction: Some(correction),
context: String::new(),
};
assert!(analysis.correction.is_some());
let c = analysis.correction.unwrap();
assert_eq!(c.description, "fix path");
}
#[test]
fn reflection_error_skipped_display() {
let err = ReflectionError::Skipped("not applicable".to_string());
let s = err.to_string();
assert!(s.contains("skipped"));
assert!(s.contains("not applicable"));
}
#[test]
fn reflection_error_internal_display() {
let err = ReflectionError::Internal("llm timeout".to_string());
let s = err.to_string();
assert!(s.contains("internal"));
assert!(s.contains("llm timeout"));
}
#[test]
fn action_retry_accessors() {
let action = RecoveryAction::Retry {
delay: Duration::from_secs(2),
};
assert!(action.is_retry());
assert!(!action.is_fail());
assert_eq!(action.delay(), Some(Duration::from_secs(2)));
}
#[test]
fn action_skip_accessors() {
let action = RecoveryAction::Skip("not important".to_string());
assert!(!action.is_retry());
assert!(!action.is_fail());
assert_eq!(action.delay(), None);
}
#[test]
fn action_ask_user() {
let action = RecoveryAction::AskUser("choose option".to_string());
assert!(!action.is_retry());
assert_eq!(action.delay(), None);
}
#[test]
fn action_fail_accessors() {
let action = RecoveryAction::Fail("unrecoverable".to_string());
assert!(action.is_fail());
assert!(!action.is_retry());
assert_eq!(action.delay(), None);
}
#[test]
fn action_display() {
let retry = RecoveryAction::Retry {
delay: Duration::from_secs(1),
};
assert!(retry.to_string().contains("retry"));
let skip = RecoveryAction::Skip("reason".to_string());
assert!(skip.to_string().contains("skip: reason"));
let ask = RecoveryAction::AskUser("prompt".to_string());
assert!(ask.to_string().contains("ask user: prompt"));
let fail = RecoveryAction::Fail("bad".to_string());
assert!(fail.to_string().contains("fail: bad"));
}
#[tokio::test]
async fn noop_reflector_marks_non_recoverable() {
let reflector = NoopReflector;
let ctx = ReflectionContext {
task: "test".to_string(),
attempt: 0,
max_attempts: 3,
};
let analysis = reflector
.analyze("some error", "tool", &serde_json::json!({}), &ctx)
.await
.unwrap();
assert!(!analysis.is_recoverable);
assert_eq!(analysis.root_cause, "some error");
assert_eq!(analysis.severity, FailureSeverity::Medium);
assert!(analysis.correction.is_none());
}
#[test]
fn noop_reflector_debug() {
let reflector = NoopReflector;
let debug = format!("{reflector:?}");
assert!(debug.contains("NoopReflector"));
}
}