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}