pub mod backoff;
pub use backoff::ExponentialBackoffRecovery;
pub mod llm;
pub use llm::LlmReflector;
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,
}
impl crate::structured::StructuredOutput for FailureAnalysis {
fn name() -> &'static str {
"failure_analysis"
}
fn schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"is_recoverable": {"type": "boolean"},
"root_cause": {"type": "string"},
"severity": {
"type": "string",
"enum": ["low", "medium", "high", "critical"]
},
"correction": {
"anyOf": [
{
"type": "object",
"properties": {
"correction_type": {
"type": "string",
"enum": [
"input_fix",
"tool_change",
"prerequisite_fix",
"approach_change",
"escalate"
]
},
"description": {"type": "string"},
"modified_input": {"type": ["object", "null"]},
"alternative_tool": {"type": ["string", "null"]},
"guidance": {"type": ["string", "null"]}
},
"required": [
"correction_type",
"description",
"modified_input",
"alternative_tool",
"guidance"
],
"additionalProperties": false
},
{"type": "null"}
]
},
"context": {"type": "string"}
},
"required": [
"is_recoverable",
"root_cause",
"severity",
"correction",
"context"
],
"additionalProperties": false
})
}
}
#[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,
tool_schema: Option<&crate::tool::ToolSchema>,
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,
_tool_schema: Option<&crate::tool::ToolSchema>,
_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!({}), None, &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"));
}
#[test]
fn failure_analysis_structured_round_trip() {
use crate::structured::StructuredOutput;
let v = serde_json::json!({
"is_recoverable": true,
"root_cause": "timeout",
"severity": "low",
"correction": {
"correction_type": "input_fix",
"description": "fix the path",
"modified_input": {"path": "/x"},
"alternative_tool": null,
"guidance": null
},
"context": "open call"
});
let analysis = FailureAnalysis::from_value(v).expect("should deserialize");
assert!(analysis.is_recoverable);
assert_eq!(analysis.root_cause, "timeout");
assert_eq!(analysis.severity, FailureSeverity::Low);
let correction = analysis.correction.expect("correction");
assert_eq!(correction.description, "fix the path");
assert_eq!(
correction.modified_input,
Some(serde_json::json!({"path": "/x"}))
);
}
#[test]
fn failure_analysis_schema_is_valid_json() {
let schema = <FailureAnalysis as crate::structured::StructuredOutput>::schema();
let obj = schema.as_object().expect("schema must be a JSON object");
assert_eq!(obj["type"], "object");
let required = obj["required"]
.as_array()
.expect("required must be an array");
assert_eq!(required.len(), 5);
}
#[test]
fn failure_analysis_schema_correction_type_enum() {
let schema = <FailureAnalysis as crate::structured::StructuredOutput>::schema();
let enum_values = schema
.pointer("/properties/correction/anyOf/0/properties/correction_type/enum")
.expect("correction_type enum must be present")
.as_array()
.expect("enum must be an array");
let values: Vec<&str> = enum_values.iter().map(|v| v.as_str().unwrap()).collect();
assert_eq!(
values,
vec![
"input_fix",
"tool_change",
"prerequisite_fix",
"approach_change",
"escalate"
]
);
}
#[test]
fn failure_analysis_name_is_stable() {
use crate::structured::StructuredOutput;
let name = FailureAnalysis::name();
assert_eq!(name, "failure_analysis");
assert!(
name.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-'),
"name must match ^[a-zA-Z0-9_-]+$: {name}"
);
}
#[tokio::test]
async fn noop_reflector_accepts_new_signature() {
let reflector = NoopReflector;
let ctx = ReflectionContext {
task: "t".to_string(),
attempt: 0,
max_attempts: 1,
};
let analysis = reflector
.analyze("err", "tool", &serde_json::json!({}), None, &ctx)
.await
.unwrap();
assert!(!analysis.is_recoverable);
assert_eq!(analysis.severity, FailureSeverity::Medium);
}
}