1use std::sync::LazyLock;
29use std::time::Duration;
30
31use pi_ai::{AssistantMessage, StopReason};
32use regex::Regex;
33
34use super::AgentSession;
35use super::events::AgentSessionEvent;
36
37impl AgentSession {
38 #[must_use]
44 pub fn retry_attempt(&self) -> u32 {
45 self.lock_inner().retry_attempt
46 }
47
48 pub(super) fn is_retryable_error(message: &AssistantMessage) -> bool {
57 if is_context_overflow(message) {
58 return false;
59 }
60 is_retryable_assistant_error(message)
61 }
62
63 pub(super) async fn prepare_retry(
73 self: &std::sync::Arc<Self>,
74 message: &AssistantMessage,
75 ) -> bool {
76 let (enabled, max_retries) = {
80 let inner = self.lock_inner();
81 (inner.auto_retry_enabled, inner.max_retries)
82 };
83 if !enabled {
84 return false;
85 }
86
87 let base_delay_ms = self.lock_settings().get_retry_settings().base_delay_ms;
88
89 let attempt = {
91 let mut inner = self.lock_inner();
92 inner.retry_attempt = inner.retry_attempt.saturating_add(1);
93 if inner.retry_attempt > max_retries {
94 inner.retry_attempt = inner.retry_attempt.saturating_sub(1);
95 return false;
96 }
97 inner.retry_attempt
98 };
99
100 let delay_ms = backoff_delay_ms(base_delay_ms, attempt);
101
102 self.emit_public(AgentSessionEvent::AutoRetryStart {
103 attempt,
104 max_attempts: max_retries,
105 delay_ms,
106 error_message: message
107 .error_message
108 .clone()
109 .unwrap_or_else(|| "Unknown error".to_owned()),
110 });
111
112 let _ = self.agent.pop_last_if_assistant();
115
116 let token = self.begin_retry_abort();
118 tokio::select! {
119 () = token.cancelled() => {
120 let attempt = {
121 let mut inner = self.lock_inner();
122 let prev = inner.retry_attempt;
123 inner.retry_attempt = 0;
124 prev
125 };
126 self.clear_retry_abort();
127 self.emit_public(AgentSessionEvent::AutoRetryEnd {
128 success: false,
129 attempt,
130 final_error: Some("Retry cancelled".to_owned()),
131 });
132 false
133 }
134 () = tokio::time::sleep(Duration::from_millis(delay_ms)) => {
135 self.clear_retry_abort();
136 true
137 }
138 }
139 }
140
141 pub(super) fn emit_retry_exhausted(&self, message: &AssistantMessage) {
149 let (attempt, final_error) = {
150 let mut inner = self.lock_inner();
151 if message.stop_reason == StopReason::Error && inner.retry_attempt > 0 {
152 let attempt = inner.retry_attempt;
153 inner.retry_attempt = 0;
154 (attempt, message.error_message.clone())
155 } else {
156 return;
157 }
158 };
159 self.emit_public(AgentSessionEvent::AutoRetryEnd {
160 success: false,
161 attempt,
162 final_error,
163 });
164 }
165}
166
167fn backoff_delay_ms(base_delay_ms: u64, attempt: u32) -> u64 {
173 if attempt == 0 {
174 return 0;
175 }
176 let exp = attempt.saturating_sub(1);
177 base_delay_ms.saturating_mul(2u64.saturating_pow(exp))
178}
179
180fn is_context_overflow(message: &AssistantMessage) -> bool {
187 if message.stop_reason != StopReason::Error {
188 return false;
189 }
190 let Some(err) = message.error_message.as_deref() else {
191 return false;
192 };
193 let lower = err.to_ascii_lowercase();
194
195 let is_non_overflow = lower.contains("rate_limit")
197 || lower.contains("rate limit")
198 || lower.contains("too many requests")
199 || lower.contains("throttling")
200 || lower.contains("service unavailable")
201 || lower.contains("overloaded");
202
203 if is_non_overflow {
204 return false;
205 }
206
207 lower.contains("context overflow")
208 || lower.contains("context length")
209 || lower.contains("maximum context")
210 || lower.contains("too long")
211 || lower.contains("exceeds the limit")
212 || lower.contains("exceeds the context")
213 || lower.contains("exceeds model")
214 || lower.contains("prompt has")
215 || lower.contains("token count")
216 || lower.contains("input is too long")
217 || lower.contains("input length")
218}
219
220fn contains_status_code(text: &str, status: &str) -> bool {
222 text.match_indices(status).any(|(index, _)| {
223 let before_is_digit = text[..index]
224 .chars()
225 .next_back()
226 .is_some_and(|ch| ch.is_ascii_digit());
227 let after = index.saturating_add(status.len());
228 let after_is_digit = text[after..]
229 .chars()
230 .next()
231 .is_some_and(|ch| ch.is_ascii_digit());
232 !before_is_digit && !after_is_digit
233 })
234}
235
236static NON_RETRYABLE_PROVIDER_LIMIT_ERROR_PATTERN: LazyLock<Result<Regex, regex::Error>> =
237 LazyLock::new(|| {
238 Regex::new(
239 r"(?i)invalid_api_key|invalid api key|gousagelimiterror|freeusagelimiterror|monthly usage limit reached|available balance|insufficient_quota|out of budget|quota exceeded|billing|context overflow|context length|maximum context",
240 )
241 });
242
243static RETRYABLE_PROVIDER_ERROR_PATTERN: LazyLock<Result<Regex, regex::Error>> = LazyLock::new(
244 || {
245 Regex::new(
246 r"(?i)overloaded|rate.?limit|too many requests|service.?unavailable|server.?error|internal.?error|provider.?returned.?error|network.?error|connection.?error|connection.?refused|connection.?lost|other side closed|fetch failed|upstream.?connect|reset before headers|socket hang up|socket connection was closed|timed? out|timeout|terminated|websocket.?closed|websocket.?error|ended without|stream ended before message_stop|http2 request did not get a response|retry delay|you can retry your request|try your request again|please retry your request|resourceexhausted|temporarily unavailable",
247 )
248 },
249);
250
251fn is_retryable_assistant_error(message: &AssistantMessage) -> bool {
257 if message.stop_reason != StopReason::Error {
258 return false;
259 }
260 let Some(err) = message.error_message.as_deref() else {
261 return false;
262 };
263
264 match NON_RETRYABLE_PROVIDER_LIMIT_ERROR_PATTERN.as_ref() {
267 Ok(pattern) if pattern.is_match(err) => return false,
268 Ok(_) => {}
269 Err(_) => return false,
272 }
273
274 RETRYABLE_PROVIDER_ERROR_PATTERN
277 .as_ref()
278 .is_ok_and(|pattern| pattern.is_match(err))
279 || ["429", "500", "502", "503", "504", "524"]
280 .iter()
281 .any(|status| contains_status_code(err, status))
282}
283
284#[cfg(test)]
285mod tests {
286 use super::*;
287
288 #[test]
289 fn backoff_progression() {
290 assert_eq!(backoff_delay_ms(1000, 1), 1000);
291 assert_eq!(backoff_delay_ms(1000, 2), 2000);
292 assert_eq!(backoff_delay_ms(1000, 3), 4000);
293 assert_eq!(backoff_delay_ms(1000, 0), 0);
294 assert_eq!(backoff_delay_ms(u64::MAX, 1), u64::MAX);
296 }
297
298 #[test]
299 fn overloaded_is_retryable() {
300 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
301 msg.stop_reason = StopReason::Error;
302 msg.error_message = Some("overloaded_error".to_owned());
303 assert!(is_retryable_assistant_error(&msg));
304 assert!(!is_context_overflow(&msg));
305 }
306
307 #[test]
308 fn context_overflow_not_retryable() {
309 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
310 msg.stop_reason = StopReason::Error;
311 msg.error_message = Some("This model's maximum context length is 8192 tokens.".to_owned());
312 assert!(!is_retryable_assistant_error(&msg));
313 assert!(is_context_overflow(&msg));
314 }
315
316 #[test]
317 fn rate_limit_is_retryable_but_not_overflow() {
318 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
319 msg.stop_reason = StopReason::Error;
320 msg.error_message = Some("rate_limit exceeded".to_owned());
321 assert!(is_retryable_assistant_error(&msg));
322 assert!(!is_context_overflow(&msg));
323 }
324
325 #[test]
326 fn status_first_429_is_retryable() {
327 for error in [
328 "OpenAI API error: 429: {}",
329 "HTTP 429: ",
330 "provider failed with status 429",
331 ] {
332 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
333 msg.stop_reason = StopReason::Error;
334 msg.error_message = Some(error.to_owned());
335 assert!(is_retryable_assistant_error(&msg), "{error}");
336 }
337
338 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
339 msg.stop_reason = StopReason::Error;
340 msg.error_message = Some("internal code 14290".to_owned());
341 assert!(!is_retryable_assistant_error(&msg));
342 }
343
344 #[test]
345 fn network_connection_lost_retryable() {
346 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
347 msg.stop_reason = StopReason::Error;
348 msg.error_message = Some("Network connection lost.".to_owned());
349 assert!(is_retryable_assistant_error(&msg));
350 }
351
352 #[test]
353 fn bedrock_retry_text() {
354 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
355 msg.stop_reason = StopReason::Error;
356 msg.error_message = Some("Please try your request again later.".to_owned());
357 assert!(is_retryable_assistant_error(&msg));
358 }
359
360 #[test]
361 fn openai_retry_text() {
362 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
363 msg.stop_reason = StopReason::Error;
364 msg.error_message = Some("The server had an error while processing your request. Sorry about that! Please retry your request.".to_owned());
365 assert!(is_retryable_assistant_error(&msg));
366 }
367
368 #[test]
369 fn throttling_not_overflow() {
370 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
371 msg.stop_reason = StopReason::Error;
372 msg.error_message = Some("Throttling error: Too many tokens, please wait.".to_owned());
373 assert!(!is_context_overflow(&msg));
374 }
375
376 #[test]
377 fn stop_reason_non_error_not_retryable() {
378 let msg = AssistantMessage::new("api", "provider", "m", 0);
379 assert!(!is_retryable_assistant_error(&msg));
380 }
381
382 #[test]
383 fn permanent_provider_limits_are_not_retryable() {
384 for error in [
385 "invalid_api_key",
386 "GoUsageLimitError: status 429",
387 "FreeUsageLimitError: too many requests",
388 "Monthly usage limit reached: status 429",
389 "Enable available balance usage; too many requests",
390 "insufficient_quota: HTTP 429",
391 "429 quota exceeded",
392 "account is out of budget; please retry your request",
393 "billing limit reached: service unavailable",
394 ] {
395 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
396 msg.stop_reason = StopReason::Error;
397 msg.error_message = Some(error.to_owned());
398 assert!(!is_retryable_assistant_error(&msg), "{error}");
399 }
400 }
401
402 #[test]
403 fn transient_provider_and_network_errors_are_retryable() {
404 for error in [
405 "overloaded_error",
406 "rate-limited by provider",
407 "too many requests",
408 "HTTP 429: throttled",
409 "HTTP 500: status code",
410 "HTTP 502: status code",
411 "HTTP 503: status code",
412 "HTTP 504: status code",
413 "524 status code (no body)",
414 "service unavailable",
415 "server_error",
416 "internal-error",
417 "Provider returned error",
418 "network error",
419 "connection error",
420 "connection refused",
421 "connection lost",
422 "other side closed",
423 "fetch failed",
424 "upstream connect error",
425 "reset before headers",
426 "socket hang up",
427 "socket connection was closed unexpectedly",
428 "request timed out",
429 "request timeout",
430 "stream terminated",
431 "websocket closed",
432 "websocket error",
433 "stream ended without a final message",
434 "stream ended before message_stop",
435 "HTTP2 request did not get a response",
436 "retry delay exceeded",
437 "you can retry your request",
438 "try your request again",
439 "please retry your request",
440 "ResourceExhausted: worker request limit reached",
441 ] {
442 let mut msg = AssistantMessage::new("api", "provider", "m", 0);
443 msg.stop_reason = StopReason::Error;
444 msg.error_message = Some(error.to_owned());
445 assert!(is_retryable_assistant_error(&msg), "{error}");
446 }
447 }
448}