Skip to main content

roder_core/
reliability.rs

1use roder_api::reliability::{
2    ReliabilityErrorClass, ReliabilityLimitDecision, ReliabilityLimitKind, ReliabilityRequestPolicy,
3};
4use roder_api::transcript::ToolResultRecord;
5
6#[derive(Debug, Clone, PartialEq, Eq)]
7pub struct RuntimeReliabilityConfig {
8    pub max_consecutive_tool_failures: u32,
9    pub max_tool_failures_per_turn: u32,
10    pub max_model_calls_per_turn: u32,
11    pub provider_retry_max_attempts: u32,
12    pub provider_retry_initial_backoff_ms: u64,
13    pub provider_retry_backoff_factor: u32,
14    pub provider_retry_status_codes: Vec<u16>,
15    pub retry_empty_provider_body: bool,
16}
17
18impl Default for RuntimeReliabilityConfig {
19    fn default() -> Self {
20        Self {
21            max_consecutive_tool_failures: 5,
22            max_tool_failures_per_turn: 128,
23            max_model_calls_per_turn: 512,
24            provider_retry_max_attempts: 3,
25            provider_retry_initial_backoff_ms: 1_000,
26            provider_retry_backoff_factor: 2,
27            provider_retry_status_codes: vec![429, 500, 502, 503, 504],
28            retry_empty_provider_body: true,
29        }
30    }
31}
32
33impl From<RuntimeReliabilityConfig> for ReliabilityRequestPolicy {
34    fn from(config: RuntimeReliabilityConfig) -> Self {
35        Self {
36            provider_retry_max_attempts: config.provider_retry_max_attempts,
37            provider_retry_initial_backoff_ms: config.provider_retry_initial_backoff_ms,
38            provider_retry_backoff_factor: config.provider_retry_backoff_factor,
39            retry_empty_provider_body: config.retry_empty_provider_body,
40            provider_retry_status_codes: config.provider_retry_status_codes,
41        }
42    }
43}
44
45#[derive(Debug, Default, Clone, PartialEq, Eq)]
46pub(crate) struct TurnReliabilityState {
47    model_calls: u32,
48    consecutive_tool_failures: u32,
49    tool_failures: u32,
50}
51
52#[derive(Debug, Clone, PartialEq, Eq)]
53pub(crate) struct ReliabilityLimitHit {
54    pub error_class: ReliabilityErrorClass,
55    pub limit_kind: ReliabilityLimitKind,
56    pub decision: ReliabilityLimitDecision,
57    pub current: u32,
58    pub limit: u32,
59    pub message: String,
60}
61
62pub(crate) fn provider_stream_retry_cause(message: &str) -> Option<&'static str> {
63    let lower = message.to_ascii_lowercase();
64    if lower.contains("error decoding response body") {
65        return Some("stream_decode_error");
66    }
67    if lower.contains("stream closed before response.completed") {
68        return Some("stream_closed_before_completed");
69    }
70    if lower.contains("stream closed before message_stop") {
71        return Some("stream_closed_before_message_stop");
72    }
73    None
74}
75
76impl TurnReliabilityState {
77    pub(crate) fn record_model_call(
78        &mut self,
79        cfg: &RuntimeReliabilityConfig,
80        interactive: bool,
81    ) -> Option<ReliabilityLimitHit> {
82        self.model_calls = self.model_calls.saturating_add(1);
83        if self.model_calls > cfg.max_model_calls_per_turn {
84            return Some(limit_hit(
85                ReliabilityErrorClass::ProviderError,
86                ReliabilityLimitKind::ModelCallsPerTurn,
87                self.model_calls,
88                cfg.max_model_calls_per_turn,
89                interactive,
90                "model call limit reached",
91            ));
92        }
93        None
94    }
95
96    pub(crate) fn record_tool_results(
97        &mut self,
98        cfg: &RuntimeReliabilityConfig,
99        results: &[ToolResultRecord],
100        interactive: bool,
101    ) -> Option<ReliabilityLimitHit> {
102        for result in results {
103            if result.is_error {
104                self.tool_failures = self.tool_failures.saturating_add(1);
105                self.consecutive_tool_failures = self.consecutive_tool_failures.saturating_add(1);
106            } else {
107                self.consecutive_tool_failures = 0;
108            }
109        }
110        if self.consecutive_tool_failures >= cfg.max_consecutive_tool_failures {
111            return Some(limit_hit(
112                ReliabilityErrorClass::InvalidArguments,
113                ReliabilityLimitKind::ConsecutiveToolFailures,
114                self.consecutive_tool_failures,
115                cfg.max_consecutive_tool_failures,
116                interactive,
117                "consecutive tool failure limit reached",
118            ));
119        }
120        if self.tool_failures >= cfg.max_tool_failures_per_turn {
121            return Some(limit_hit(
122                ReliabilityErrorClass::InvalidArguments,
123                ReliabilityLimitKind::ToolFailuresPerTurn,
124                self.tool_failures,
125                cfg.max_tool_failures_per_turn,
126                interactive,
127                "tool failure limit reached",
128            ));
129        }
130        None
131    }
132
133    pub(crate) fn tool_failure_count(&self) -> u32 {
134        self.tool_failures
135    }
136}
137
138fn limit_hit(
139    error_class: ReliabilityErrorClass,
140    limit_kind: ReliabilityLimitKind,
141    current: u32,
142    limit: u32,
143    interactive: bool,
144    message: &str,
145) -> ReliabilityLimitHit {
146    ReliabilityLimitHit {
147        error_class,
148        limit_kind,
149        decision: if interactive {
150            ReliabilityLimitDecision::RequestContinuation
151        } else {
152            ReliabilityLimitDecision::StopTurn
153        },
154        current,
155        limit,
156        message: format!("{message}: {current}/{limit}"),
157    }
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163
164    fn result(id: &str, is_error: bool) -> ToolResultRecord {
165        ToolResultRecord {
166            id: id.to_string(),
167            name: Some("test".to_string()),
168            result: if is_error { "error" } else { "ok" }.to_string(),
169            display_payload: None,
170            is_error,
171        }
172    }
173
174    #[test]
175    fn reliability_limits_reset_consecutive_failures_after_success() {
176        let cfg = RuntimeReliabilityConfig {
177            max_consecutive_tool_failures: 2,
178            max_tool_failures_per_turn: 128,
179            ..RuntimeReliabilityConfig::default()
180        };
181        let mut state = TurnReliabilityState::default();
182
183        assert!(
184            state
185                .record_tool_results(&cfg, &[result("first", true)], false)
186                .is_none()
187        );
188        assert!(
189            state
190                .record_tool_results(&cfg, &[result("success", false)], false)
191                .is_none()
192        );
193        assert!(
194            state
195                .record_tool_results(&cfg, &[result("second", true)], false)
196                .is_none()
197        );
198        let limit = state
199            .record_tool_results(&cfg, &[result("third", true)], false)
200            .unwrap();
201        assert_eq!(
202            limit.limit_kind,
203            ReliabilityLimitKind::ConsecutiveToolFailures
204        );
205        assert_eq!(limit.current, 2);
206    }
207
208    #[test]
209    fn default_model_call_limit_allows_long_agentic_turns() {
210        assert_eq!(
211            RuntimeReliabilityConfig::default().max_model_calls_per_turn,
212            512
213        );
214    }
215
216    #[test]
217    fn provider_stream_retry_cause_classifies_transient_stream_failures() {
218        assert_eq!(
219            provider_stream_retry_cause("error decoding response body"),
220            Some("stream_decode_error")
221        );
222        assert_eq!(
223            provider_stream_retry_cause("stream closed before response.completed"),
224            Some("stream_closed_before_completed")
225        );
226        assert_eq!(
227            provider_stream_retry_cause("Anthropic stream closed before message_stop"),
228            Some("stream_closed_before_message_stop")
229        );
230        assert_eq!(provider_stream_retry_cause("invalid request body"), None);
231    }
232}