Skip to main content

pi/core/agent_session/
retry.rs

1//! Auto-retry with exponential backoff for transient provider errors.
2//!
3//! Owns the retry classification, backoff schedule, abortable sleep, and the
4//! `auto_retry_start` / `auto_retry_end` lifecycle events. `agent_end.willRetry`
5//! is computed here so the event pump can annotate the public event.
6//!
7//! # Event ordering
8//!
9//! Single transient error then success:
10//! `auto_retry_start{1}` → (backoff) → successful assistant `message_end` →
11//! `auto_retry_end{success:true, attempt:1}` (emitted from persistence).
12//!
13//! Exhausted retries (maxRetries = 2, three consecutive errors):
14//! `auto_retry_start{1}` → `auto_retry_start{2}` →
15//! `auto_retry_end{success:false, attempt:2, final_error}` (emitted from
16//! [`AgentSession::handle_post_agent_run`]).
17//!
18//! Aborted during sleep:
19//! `auto_retry_start{1}` → abort →
20//! `auto_retry_end{success:false, attempt:1, final_error:"Retry cancelled"}`.
21//!
22//! # Lock discipline
23//!
24//! `retry_attempt` lives on `AgentSessionInner` under the std `Mutex`. The
25//! persistence half (pump task) resets it on success; the prompt half (caller
26//! task) increments and reads it. Never hold the inner mutex across `.await`.
27
28use 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    // -----------------------------------------------------------------
39    // Public accessors
40    // -----------------------------------------------------------------
41
42    /// Current retry attempt (0 when not retrying).
43    #[must_use]
44    pub fn retry_attempt(&self) -> u32 {
45        self.lock_inner().retry_attempt
46    }
47
48    // -----------------------------------------------------------------
49    // Classification
50    // -----------------------------------------------------------------
51
52    /// Whether `message` is a transient error that auto-retry should handle.
53    ///
54    /// Context overflow and auth errors are NOT retryable: overflow is owned by
55    /// compaction, and auth errors require user action.
56    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    // -----------------------------------------------------------------
64    // Backoff + sleep
65    // -----------------------------------------------------------------
66
67    /// Prepare a retry: increment attempt, emit `auto_retry_start`, drop the
68    /// trailing error assistant from agent state, then sleep with backoff.
69    ///
70    /// Returns `Ok(true)` when the caller should `continue_run`.
71    /// Returns `Ok(false)` when retries are disabled, exhausted, or aborted.
72    pub(super) async fn prepare_retry(
73        self: &std::sync::Arc<Self>,
74        message: &AssistantMessage,
75    ) -> bool {
76        // Runtime flags (`auto_retry_enabled` / `max_retries`) are the source of
77        // truth for the live session — same as will_retry / set_auto_retry_enabled.
78        // Settings only supply base_delay_ms (not mutated by the toggle).
79        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        // Increment attempt; bail (preserving count) when over the limit.
90        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        // Remove the trailing error assistant from agent state so the retry
113        // re-generates it. Session history keeps the original entry.
114        let _ = self.agent.pop_last_if_assistant();
115
116        // Abortable exponential backoff sleep.
117        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    // -----------------------------------------------------------------
142    // Terminal failure emit
143    // -----------------------------------------------------------------
144
145    /// Emit `auto_retry_end{success:false}` when an error follows retries and
146    /// retry did not fire again. Resets the counter. No-op when `retry_attempt`
147    /// is already 0.
148    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
167// -----------------------------------------------------------------------
168// Free functions
169// -----------------------------------------------------------------------
170
171/// Compute exponential backoff: `base * 2^(attempt-1)` (saturating).
172fn 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
180/// Whether the assistant message is a context-overflow error.
181///
182/// Context overflow is handled by compaction, not retry. The detection is
183/// conservative: it only fires for `stop_reason == Error` messages whose error
184/// text matches overflow patterns, while excluding rate-limit and throttling
185/// messages that mention tokens but are transient.
186fn 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    // Exclude transient errors that may mention "tokens" or "limit".
196    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
220/// Whether `text` contains `status` as a standalone decimal status token.
221fn 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
251/// Whether the assistant message is a retryable transient error.
252///
253/// Non-retryable: authentication failures, permanent account/quota/billing
254/// limits, and context overflow (overflow is caught by [`is_context_overflow`]
255/// first, but this function is also safe to call directly).
256fn 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    // Permanent account and subscription limits take precedence over transient
265    // markers such as HTTP 429 or provider retry guidance in the same message.
266    match NON_RETRYABLE_PROVIDER_LIMIT_ERROR_PATTERN.as_ref() {
267        Ok(pattern) if pattern.is_match(err) => return false,
268        Ok(_) => {}
269        // Conservatively disable retries if a developer introduces an invalid
270        // permanent-error pattern.
271        Err(_) => return false,
272    }
273
274    // Transient / retryable patterns. Status tokens are checked separately so
275    // unrelated numeric identifiers such as `14290` do not match HTTP 429.
276    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        // Saturates rather than overflowing.
295        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}