返回 CodeWhale
mod.rs
根目录 / crates / tui / src / llm_client / mod.rs
1 //! LLM Client Trait and Retry Logic
2 //!
3 //! This module provides a unified interface for LLM providers with robust retry logic,
4 //! exponential backoff, and proper error classification.
5 //!
6 //! # Architecture
7 //!
8 //! - `LlmClient` trait: Async interface for LLM providers (DeepSeek, `OpenAI`, etc.)
9 //! - `RetryConfig`: Configurable retry behavior with exponential backoff and jitter
10 //! - `LlmError`: Classified errors with retryability information
11
12 //! - `with_retry`: Generic retry wrapper for any async operation
13 //!
14 //! # Example
15 //!
16 //! ```ignore
17 //! use crate::llm_client::{LlmClient, RetryConfig, with_retry};
18 //!
19 //! let config = RetryConfig::default();
20 //! let result = with_retry(&config, || async {
21 //! client.create_message(request).await
22 //! }, None).await;
23 //! ```
24
25 use crate::config::RetryPolicy;
26 use anyhow::Result;
27 use codewhale_models::{MessageRequest, MessageResponse, StreamEvent};
28 use serde_json::Value;
29 use std::future::Future;
30 use std::pin::Pin;
31 use std::time::{Duration, Instant};
32 use uuid::Uuid;
33
34 #[cfg(test)]
35 pub mod mock;
36
37 // === LlmClient Trait ===
38
39 /// Type alias for boxed stream of SSE events
40 pub type StreamEventBox =
41 Pin<Box<dyn futures_util::Stream<Item = Result<StreamEvent>> + Send + 'static>>;
42
43 /// Unified interface for LLM providers.
44 ///
45 /// This trait abstracts over different LLM APIs (DeepSeek, `OpenAI`, etc.)
46 /// allowing the agent to work with any provider that implements this interface.
47 ///
48 /// # Implementation Notes
49 ///
50 /// - All methods are async and require `Send + Sync` for thread safety
51 /// - The `create_message_stream` method returns a pinned boxed stream for SSE
52 /// - Implementations should handle their own authentication and base URL configuration
53 #[allow(async_fn_in_trait, dead_code)] // Trait methods are part of the LLM provider interface
54 pub trait LlmClient: Send + Sync {
55 /// Returns the provider name (e.g., "openai", "deepseek")
56 fn provider_name(&self) -> &'static str;
57
58 /// Returns the model identifier being used
59 fn model(&self) -> &str;
60
61 /// Creates a non-streaming message completion
62 fn create_message(
63 &self,
64 request: MessageRequest,
65 ) -> impl Future<Output = Result<MessageResponse>> + Send;
66
67 /// Dispatch a fresh request. Clients with a local response cache must
68 /// override this; authorization decisions cannot reuse earlier answers.
69 fn create_message_uncached(
70 &self,
71 request: MessageRequest,
72 ) -> impl Future<Output = Result<MessageResponse>> + Send {
73 self.create_message(request)
74 }
75
76 /// Creates a streaming message completion
77 ///
78 /// Returns a stream of SSE events that should be consumed until completion.
79 fn create_message_stream(
80 &self,
81 request: MessageRequest,
82 ) -> impl Future<Output = Result<StreamEventBox>> + Send;
83
84 /// Optional health check to verify API connectivity
85 fn health_check(&self) -> impl Future<Output = Result<bool>> + Send {
86 async { Ok(true) }
87 }
88
89 /// The concrete base URL requests go to, when the implementation knows it.
90 ///
91 /// Background cost accrual uses this for billing provenance only: it is
92 /// reduced to a non-secret surface classification and a SHA-256 fingerprint
93 /// before being recorded, and the URL itself is never persisted or logged
94 /// (#4318). The default is `None` so an implementation that cannot report a
95 /// stable endpoint yields "unknown endpoint" — which fails closed — rather
96 /// than being assumed to be the provider's public API.
97 fn billing_base_url(&self) -> Option<&str> {
98 None
99 }
100
101 /// Non-secret limits frozen with the resolved route, when available.
102 fn route_limits(&self) -> Option<codewhale_config::route::RouteLimits> {
103 None
104 }
105
106 /// Output cap for a request sent through this exact client route.
107 fn effective_max_output_tokens(&self, requested_model: &str) -> u32 {
108 let route = self.effective_route_envelope(requested_model, chrono::Utc::now());
109 crate::route_budget::effective_max_output_tokens_for_route(
110 route.provider,
111 &route.model,
112 self.route_limits(),
113 )
114 }
115
116 /// Freeze the non-secret effective route immediately before a request is
117 /// dispatched. Implementations with richer configured identity/billing
118 /// facts should override this fail-closed default.
119 fn effective_route_envelope(
120 &self,
121 requested_model: &str,
122 dispatched_at: chrono::DateTime<chrono::Utc>,
123 ) -> crate::cost_status::EffectiveRouteEnvelope {
124 let provider = crate::config::ProviderKind::parse(self.provider_name())
125 .unwrap_or(crate::config::ProviderKind::Custom);
126 crate::cost_status::EffectiveRouteEnvelope::capture_observed(
127 provider,
128 self.provider_name(),
129 requested_model,
130 self.billing_base_url(),
131 dispatched_at,
132 )
133 }
134 }
135
136 // === Authentication diagnostics ===
137
138 #[derive(Debug, Clone, PartialEq, Eq, Default)]
139 pub struct AuthenticationErrorContext {
140 pub provider: Option<String>,
141 pub base_url_authority: Option<String>,
142 pub model: Option<String>,
143 pub key_source: Option<String>,
144 pub key_fingerprint: Option<String>,
145 pub key_kind: Option<String>,
146 /// The one command or action that replaces the rejected credential.
147 pub fix: Option<String>,
148 }
149
150 impl AuthenticationErrorContext {
151 #[must_use]
152 pub fn from_parts(
153 provider: Option<&str>,
154 base_url: Option<&str>,
155 model: Option<&str>,
156 key_source: Option<&str>,
157 api_key: Option<&str>,
158 ) -> Self {
159 let api_key = api_key.and_then(non_empty_trimmed);
160 Self {
161 provider: provider.and_then(non_empty_trimmed).map(str::to_string),
162 base_url_authority: base_url.and_then(base_url_authority),
163 model: model.and_then(non_empty_trimmed).map(str::to_string),
164 key_source: key_source.and_then(non_empty_trimmed).map(str::to_string),
165 key_fingerprint: api_key.map(redacted_key_fingerprint),
166 key_kind: api_key.map(classify_api_key_prefix).map(str::to_string),
167 fix: None,
168 }
169 }
170
171 #[must_use]
172 pub fn with_fix(mut self, fix: impl Into<String>) -> Self {
173 self.fix = Some(fix.into()).filter(|fix: &String| !fix.trim().is_empty());
174 self
175 }
176
177 fn is_empty(&self) -> bool {
178 self.provider.is_none()
179 && self.base_url_authority.is_none()
180 && self.model.is_none()
181 && self.key_source.is_none()
182 && self.key_fingerprint.is_none()
183 && self.key_kind.is_none()
184 && self.fix.is_none()
185 }
186
187 fn detail_segments(&self) -> Vec<String> {
188 let mut segments = Vec::new();
189 if let Some(provider) = self.provider.as_deref() {
190 segments.push(format!("provider: {provider}"));
191 }
192 if let Some(authority) = self.base_url_authority.as_deref() {
193 segments.push(format!("base URL authority: {authority}"));
194 }
195 if let Some(model) = self.model.as_deref() {
196 segments.push(format!("model: {model}"));
197 }
198 if let Some(source) = self.key_source.as_deref() {
199 segments.push(format!("key source: {source}"));
200 }
201 if let Some(fingerprint) = self.key_fingerprint.as_deref() {
202 segments.push(format!("key fingerprint: {fingerprint}"));
203 }
204 if let Some(kind) = self.key_kind.as_deref() {
205 segments.push(format!("key type: {kind}"));
206 }
207 if let Some(fix) = self.fix.as_deref() {
208 segments.push(format!("fix: {fix}"));
209 }
210 segments
211 }
212 }
213
214 #[derive(Debug, Clone, PartialEq, Eq)]
215 pub struct AuthenticationErrorDetail {
216 message: String,
217 context: Option<AuthenticationErrorContext>,
218 }
219
220 impl AuthenticationErrorDetail {
221 #[must_use]
222 pub fn new(message: impl Into<String>) -> Self {
223 Self {
224 message: message.into(),
225 context: None,
226 }
227 }
228
229 #[must_use]
230 pub fn with_context(
231 message: impl Into<String>,
232 context: Option<AuthenticationErrorContext>,
233 ) -> Self {
234 let context = context.filter(|context| !context.is_empty());
235 Self {
236 message: message.into(),
237 context,
238 }
239 }
240
241 #[must_use]
242 pub fn to_user_message(&self) -> String {
243 let Some(context) = self.context.as_ref() else {
244 return self.message.clone();
245 };
246 let segments = context.detail_segments();
247 if segments.is_empty() {
248 self.message.clone()
249 } else {
250 format!("{} ({})", self.message, segments.join(", "))
251 }
252 }
253 }
254
255 impl From<String> for AuthenticationErrorDetail {
256 fn from(message: String) -> Self {
257 Self::new(message)
258 }
259 }
260
261 impl From<&str> for AuthenticationErrorDetail {
262 fn from(message: &str) -> Self {
263 Self::new(message)
264 }
265 }
266
267 #[must_use]
268 pub fn classify_api_key_prefix(api_key: &str) -> &'static str {
269 if api_key.starts_with("tp-") {
270 "Xiaomi MiMo Token Plan key"
271 } else {
272 "API key"
273 }
274 }
275
276 fn non_empty_trimmed(value: &str) -> Option<&str> {
277 let value = value.trim();
278 if value.is_empty() { None } else { Some(value) }
279 }
280
281 pub(crate) fn base_url_authority(base_url: &str) -> Option<String> {
282 let base_url = non_empty_trimmed(base_url)?;
283 let without_scheme = base_url
284 .split_once("://")
285 .map_or(base_url, |(_, rest)| rest);
286 let authority = without_scheme.split('/').next().unwrap_or(without_scheme);
287 let authority = authority
288 .rsplit_once('@')
289 .map_or(authority, |(_, authority)| authority);
290 non_empty_trimmed(authority).map(str::to_string)
291 }
292
293 fn redacted_key_fingerprint(api_key: &str) -> String {
294 let api_key = api_key.trim();
295 let len = api_key.chars().count();
296 match public_key_prefix(api_key) {
297 Some(prefix) => format!("{prefix}... (len={len})"),
298 None => format!("unprefixed (len={len})"),
299 }
300 }
301
302 fn public_key_prefix(api_key: &str) -> Option<&str> {
303 ["tp-", "sk-", "hf_", "hf-", "ak-", "rk-"]
304 .into_iter()
305 .find(|prefix| api_key.starts_with(prefix))
306 }
307
308 #[cfg(test)]
309 fn redact_api_key_from_message(message: &str, api_key: Option<&str>) -> String {
310 let Some(api_key) = api_key.and_then(non_empty_trimmed) else {
311 return message.to_string();
312 };
313 message.replace(api_key, "[redacted API key]")
314 }
315
316 // === LlmError - Classified Error Types ===
317
318 /// Evidence captured when a provider response explicitly identifies plan quota
319 /// exhaustion. The private field prevents callers outside this parser module
320 /// from manufacturing the durable classification from arbitrary text.
321 #[derive(Debug)]
322 pub struct QuotaExhaustionError {
323 message: String,
324 }
325
326 impl QuotaExhaustionError {
327 fn from_http_message(message: String) -> Self {
328 Self { message }
329 }
330
331 pub(crate) fn into_message(self) -> String {
332 self.message
333 }
334
335 /// Append route guidance (which account hit the limit, how to switch)
336 /// to already-classified evidence. Cannot manufacture the class.
337 #[must_use]
338 pub(crate) fn with_guidance(mut self, guidance: &str) -> Self {
339 self.message = format!("{}\n{guidance}", self.message);
340 self
341 }
342 }
343
344 /// Classified LLM errors with retryability information.
345 ///
346 /// This enum categorizes API errors to enable smart retry decisions.
347 /// Some errors (rate limits, transient server errors) are retryable,
348 /// while others (auth failures, invalid requests) should fail immediately.
349 #[derive(Debug)]
350 pub enum LlmError {
351 /// Rate limit exceeded (HTTP 429)
352 /// Contains optional Retry-After duration from server
353 RateLimited {
354 message: String,
355 retry_after: Option<Duration>,
356 },
357
358 /// The provider explicitly reported that the account's plan quota is exhausted.
359 ///
360 /// Unlike an ordinary 429 rate limit, retrying the same request after a short
361 /// backoff cannot resolve this condition. This variant is constructed only at
362 /// a provider HTTP or structured stream-event boundary from explicit quota evidence.
363 QuotaExhausted(QuotaExhaustionError),
364
365 /// Server error (HTTP 5xx)
366 ServerError { status: u16, message: String },
367
368 /// Network connectivity error
369 NetworkError(String),
370
371 /// Request timed out
372 Timeout(Duration),
373
374 /// Authentication failed (HTTP 401, selected HTTP 403)
375 AuthenticationError(AuthenticationErrorDetail),
376
377 /// Authorization or provider-side blocking failed (HTTP 403)
378 AuthorizationError(String),
379
380 /// Invalid request parameters (HTTP 400)
381 InvalidRequest { status: u16, message: String },
382
383 /// Model-specific error (model not found, etc.)
384 ModelError(String),
385
386 /// Content policy violation (safety filters)
387 ContentPolicyError(String),
388
389 /// Failed to parse API response
390 ParseError(String),
391
392 /// Context length exceeded
393 ContextLengthError(String),
394
395 /// Catch-all for other errors
396 Other(String),
397 }
398
399 impl std::fmt::Display for LlmError {
400 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
401 match self {
402 LlmError::RateLimited { message, .. } => write!(f, "Rate limit exceeded: {message}"),
403 LlmError::QuotaExhausted(error) => {
404 write!(f, "Provider plan quota exhausted: {}", error.message)
405 }
406 LlmError::ServerError { status, message } => {
407 write!(f, "Server error ({status}): {message}")
408 }
409 LlmError::NetworkError(msg) => write!(f, "Network error: {msg}"),
410 LlmError::Timeout(d) => write!(f, "Request timed out after {d:?}"),
411 LlmError::AuthenticationError(auth) => {
412 write!(f, "Authentication failed: {}", auth.to_user_message())
413 }
414 LlmError::AuthorizationError(msg) => write!(f, "Authorization failed: {msg}"),
415 LlmError::InvalidRequest { status, message } => {
416 write!(f, "Invalid request ({status}): {message}")
417 }
418 LlmError::ModelError(msg) => write!(f, "Model error: {msg}"),
419 LlmError::ContentPolicyError(msg) => write!(f, "Content policy violation: {msg}"),
420 LlmError::ParseError(msg) => write!(f, "Response parsing error: {msg}"),
421 LlmError::ContextLengthError(msg) => write!(f, "Context length exceeded: {msg}"),
422 LlmError::Other(msg) => write!(f, "LLM error: {msg}"),
423 }
424 }
425 }
426
427 impl std::error::Error for LlmError {}
428
429 impl LlmError {
430 /// Determines if this error is potentially transient and worth retrying.
431 ///
432 /// Retryable errors:
433 /// - Rate limits (with backoff)
434 /// - Server errors (5xx)
435 /// - Network errors (connection issues)
436 /// - Timeouts
437 ///
438 /// Non-retryable errors:
439 /// - Provider plan quota exhaustion
440 /// - Authentication failures
441 /// - Invalid requests
442 /// - Content policy violations
443 /// - Context length errors
444 pub fn is_retryable(&self) -> bool {
445 matches!(
446 self,
447 LlmError::RateLimited { .. }
448 | LlmError::ServerError { .. }
449 | LlmError::NetworkError(_)
450 | LlmError::Timeout(_)
451 )
452 }
453
454 /// Returns the server-suggested retry delay if available.
455 ///
456 /// This is typically present for rate limit errors when the server
457 /// provides a Retry-After header.
458 pub fn suggested_retry_delay(&self) -> Option<Duration> {
459 match self {
460 LlmError::RateLimited { retry_after, .. } => *retry_after,
461 _ => None,
462 }
463 }
464
465 /// Constructs an `LlmError` from HTTP status code and response body.
466 ///
467 /// Performs heuristic classification based on:
468 /// - Status code (429 = rate limit, 401/403 = auth, 499/5xx = transient upstream error)
469 /// - Response body keywords (`context_length`, `content_policy`, safety, etc.)
470 pub fn from_http_response(status: u16, body: &str) -> Self {
471 if let Some(error) = explicit_quota_code(body)
472 .or_else(|| explicit_quota_code_marker(body))
473 .as_deref()
474 .and_then(Self::from_subscription_sharing_error_code)
475 {
476 return error;
477 }
478 if matches!(status, 400 | 402 | 429) && has_explicit_quota_evidence(body) {
479 return LlmError::QuotaExhausted(QuotaExhaustionError::from_http_message(
480 body.to_string(),
481 ));
482 }
483
484 match status {
485 429 => LlmError::RateLimited {
486 message: body.to_string(),
487 retry_after: None,
488 },
489 401 => Self::authentication_error(body),
490 403 => {
491 if looks_like_authentication_failure(body) {
492 Self::authentication_error(body)
493 } else {
494 LlmError::AuthorizationError(body.to_string())
495 }
496 }
497 400 => {
498 // Classify 400 errors by examining the response body
499 let body_lower = body.to_lowercase();
500 // An "unsupported parameter" 400 names the offending field
501 // (often `max_output_tokens` or another *token* field), which
502 // the generic keyword rules below would misread as a context
503 // window overflow. Parameter shape errors are invalid
504 // requests, not prompt-size errors, so they get their own
505 // branch ahead of the heuristic.
506 if body_lower.contains("unsupported parameter")
507 || body_lower.contains("invalid_request_error")
508 && body_lower.contains("parameter")
509 {
510 LlmError::InvalidRequest {
511 status,
512 message: body.to_string(),
513 }
514 } else if is_context_length_message(&body_lower) {
515 LlmError::ContextLengthError(body.to_string())
516 } else if body_lower.contains("content_policy")
517 || body_lower.contains("safety")
518 || body_lower.contains("harmful")
519 || body_lower.contains("inappropriate")
520 {
521 LlmError::ContentPolicyError(body.to_string())
522 } else if body_lower.contains("model") && body_lower.contains("not found") {
523 LlmError::ModelError(body.to_string())
524 } else {
525 LlmError::InvalidRequest {
526 status,
527 message: body.to_string(),
528 }
529 }
530 }
531 404 => {
532 if body.to_lowercase().contains("model") {
533 LlmError::ModelError(body.to_string())
534 } else {
535 LlmError::InvalidRequest {
536 status,
537 message: body.to_string(),
538 }
539 }
540 }
541 // Several OpenAI-compatible gateways use nginx's non-standard
542 // 499 for an upstream request that was cancelled before response
543 // streaming began. At this boundary no response body stream has
544 // been exposed, so it is eligible for the same bounded retry
545 // policy as a 5xx gateway failure.
546 499..=599 => LlmError::ServerError {
547 status,
548 message: body.to_string(),
549 },
550 _ => LlmError::Other(format!("HTTP {status}: {body}")),
551 }
552 }
553
554 /// Official structured ChatGPT plan errors are terminal account states.
555 /// Plain text containing quota words is never enough to mint this type.
556 #[must_use]
557 pub(crate) fn from_subscription_sharing_error_code(code: &str) -> Option<Self> {
558 let message = match code {
559 "subscription_sharing_usage_limit_exceeded" => {
560 "ChatGPT plan usage limit reached. Check ChatGPT Settings > Usage for your remaining allowance and reset time."
561 }
562 "subscription_sharing_usage_unavailable" => {
563 "ChatGPT plan usage is unavailable. Check ChatGPT Settings > Usage and reconnect if needed."
564 }
565 _ => return None,
566 };
567 Some(Self::QuotaExhausted(
568 QuotaExhaustionError::from_http_message(message.to_string()),
569 ))
570 }
571
572 #[must_use]
573 pub fn authentication_error(message: impl Into<String>) -> Self {
574 LlmError::AuthenticationError(AuthenticationErrorDetail::new(message))
575 }
576
577 #[must_use]
578 pub fn authentication_error_with_context(
579 message: impl Into<String>,
580 context: Option<AuthenticationErrorContext>,
581 ) -> Self {
582 LlmError::AuthenticationError(AuthenticationErrorDetail::with_context(message, context))
583 }
584
585 /// Constructs an `LlmError` from HTTP response data plus request context
586 /// that is safe to display when authentication fails.
587 #[must_use]
588 #[cfg(test)]
589 pub fn from_http_response_with_request_context(
590 status: u16,
591 body: &str,
592 provider: Option<&str>,
593 base_url: Option<&str>,
594 model: Option<&str>,
595 key_source: Option<&str>,
596 api_key: Option<&str>,
597 ) -> Self {
598 let body = redact_api_key_from_message(body, api_key);
599 let context =
600 AuthenticationErrorContext::from_parts(provider, base_url, model, key_source, api_key);
601 Self::from_http_response_with_auth_context(status, &body, Some(context))
602 }
603
604 /// Constructs an `LlmError` from HTTP status code and response body, with
605 /// optional structured details for authentication failures.
606 ///
607 /// The `body` passed here must already be safe for user display. Prefer
608 /// [`Self::from_http_response_with_request_context`] when the raw API key is
609 /// available so the response body can be redacted before rendering.
610 #[must_use]
611 pub fn from_http_response_with_auth_context(
612 status: u16,
613 body: &str,
614 auth_context: Option<AuthenticationErrorContext>,
615 ) -> Self {
616 let classified = Self::from_http_response(status, body);
617 if matches!(classified, Self::QuotaExhausted(_)) {
618 return classified;
619 }
620 match status {
621 401 => Self::authentication_error_with_context(body, auth_context),
622 403 => {
623 if looks_like_authentication_failure(body) {
624 Self::authentication_error_with_context(body, auth_context)
625 } else {
626 LlmError::AuthorizationError(body.to_string())
627 }
628 }
629 _ => classified,
630 }
631 }
632
633 /// Constructs an `LlmError` from HTTP status code, body, and optional Retry-After header.
634 pub fn from_http_response_with_retry_after(
635 status: u16,
636 body: &str,
637 retry_after: Option<Duration>,
638 ) -> Self {
639 let mut error = Self::from_http_response(status, body);
640 if let LlmError::RateLimited {
641 retry_after: ref mut ra,
642 ..
643 } = error
644 {
645 *ra = retry_after;
646 }
647 error
648 }
649
650 /// Constructs an `LlmError` from a reqwest error.
651 pub fn from_reqwest(err: &reqwest::Error) -> Self {
652 if err.is_timeout() {
653 LlmError::Timeout(Duration::from_secs(0))
654 } else if err.is_connect() {
655 LlmError::NetworkError(format!("Connection failed: {err}"))
656 } else if err.is_request() {
657 LlmError::NetworkError(format!("Request failed: {err}"))
658 } else {
659 LlmError::Other(err.to_string())
660 }
661 }
662 }
663
664 /// Format provider HTTP error bodies before they are surfaced in the TUI.
665 ///
666 /// Providers sometimes return whole HTML error pages for gateway/WAF blocks.
667 /// Passing those pages through raw floods the transcript and can also make a
668 /// provider-side 403 look like a broken API key. Keep the useful details and
669 /// cap everything else.
670 #[must_use]
671 pub(crate) fn sanitize_http_error_body(
672 provider_label: Option<&str>,
673 status: u16,
674 body: &str,
675 ) -> String {
676 let json_message = extract_json_error_message(body);
677 let message = json_message.as_deref().unwrap_or(body);
678 // Gate on Google's actual rejection, not the selected provider or model:
679 // compatible gateways may manage signatures themselves (#6048). This
680 // shared boundary covers both streaming and non-streaming HTTP failures.
681 const SIGNATURE_HINT: &str = "Gemini rejected tool-call replay because a thought signature is missing. \
682 Use the built-in `google` provider with its default endpoint, or a gateway that preserves \
683 Google thought signatures, then start a new session before using tools. \
684 Changing reasoning settings will not restore missing signatures.";
685 if status == 400
686 && !is_probably_html(message)
687 && explicit_quota_code(body).is_none()
688 && !message.contains(SIGNATURE_HINT)
689 {
690 let lower = collapse_whitespace(message).to_ascii_lowercase();
691 if lower.contains("missing a thought_signature")
692 || lower.contains("missing thought_signature")
693 || lower.contains("thought_signature is missing")
694 {
695 let detail = truncate_for_error(&collapse_whitespace(message), 900);
696 return format!("{SIGNATURE_HINT} Provider error: {detail}");
697 }
698 }
699
700 if let Some(message) = json_message {
701 let message = truncate_for_error(&collapse_whitespace(&message), 2_000);
702 if let Some(code) = explicit_quota_code(body) {
703 return format!("{message} (provider error code: {code})");
704 }
705 return message;
706 }
707
708 if is_probably_html(body) {
709 let text = html_to_text(body);
710 let lower = text.to_ascii_lowercase();
711 let provider = provider_label.unwrap_or("Provider");
712
713 // Cloudflare's "Access Denied" interstitial strips the literal word
714 // "cloudflare" once tags are removed (it only survives in `<meta>`
715 // attributes and the `<style>`/`<script>` blocks we discard). Arcee's
716 // 403 page is exactly this shape, so also key off the WAF's stock copy
717 // ("security alert", "contact support") and a Cloudflare error/ray ID.
718 let error_id = extract_cloudflare_error_id(&text);
719 let is_cloudflare = lower.contains("cloudflare");
720 let looks_like_access_denied = lower.contains("access denied")
721 && (is_cloudflare
722 || lower.contains("security alert")
723 || lower.contains("contact support")
724 || lower.contains("contact us")
725 || error_id.is_some());
726 if looks_like_access_denied {
727 let label = if is_cloudflare {
728 "Cloudflare Access Denied"
729 } else {
730 "Access Denied"
731 };
732 let mut message = format!(
733 "{provider} API returned {label} (HTTP {status}). \
734 The request was blocked before it reached the model; retry with a \
735 smaller request or fewer tools, or contact provider support"
736 );
737 if let Some(id) = error_id {
738 message.push_str(&format!(" with ID {id}"));
739 }
740 message.push('.');
741 return message;
742 }
743
744 let text = truncate_for_error(&collapse_whitespace(&text), 900);
745 return format!("{provider} API returned an HTML error page (HTTP {status}): {text}");
746 }
747
748 truncate_for_error(&collapse_whitespace(body), 2_000)
749 }
750
751 fn looks_like_authentication_failure(body: &str) -> bool {
752 let lower = body.to_ascii_lowercase();
753 lower.contains("authentication")
754 || lower.contains("unauthorized")
755 || lower.contains("api key")
756 || lower.contains("invalid key")
757 || lower.contains("invalid token")
758 || lower.contains("bearer token")
759 || lower.contains("missing token")
760 }
761
762 /// A provider error is a context overflow only when it says so. Bare
763 /// "token", "too long" or "maximum" also appear in ordinary invalid-request
764 /// errors (`max_tokens must be ...`, a field value too long), which compaction
765 /// or a bigger window cannot fix. This is the one phrase list: the typed 400
766 /// classification here and the engine's string classifier both read it.
767 /// `lower` must already be lowercase.
768 pub(crate) fn is_context_length_message(lower: &str) -> bool {
769 [
770 "context_length",
771 "context length",
772 "context window",
773 "context limit",
774 "maximum context",
775 "prompt is too long",
776 "input is too long",
777 "maximum prompt length",
778 "exceeded model token limit",
779 "tokens exceed",
780 "exceeds the maximum number of tokens",
781 // llama.cpp: "the request exceeds the available context size".
782 "available context size",
783 ]
784 .iter()
785 .any(|phrase| lower.contains(phrase))
786 }
787
788 /// Quota exhaustion is a durable account state, not a generic rate-limit
789 /// synonym. Accept only explicit provider evidence at the HTTP/parser boundary;
790 /// callers holding a stringified error must never promote it to this type.
791 fn has_explicit_quota_evidence(body: &str) -> bool {
792 explicit_quota_code(body).is_some()
793 || explicit_quota_code_marker(body).is_some()
794 || has_explicit_quota_phrase(body)
795 }
796
797 fn explicit_quota_code(body: &str) -> Option<String> {
798 let value: Value = serde_json::from_str(body).ok()?;
799 [
800 "/error/code",
801 "/error/type",
802 "/error/error_code",
803 "/code",
804 "/type",
805 "/error_code",
806 ]
807 .into_iter()
808 .filter_map(|pointer| value.pointer(pointer).and_then(Value::as_str))
809 .find(|code| is_explicit_quota_code(code))
810 .map(ToOwned::to_owned)
811 }
812
813 fn is_explicit_quota_code(code: &str) -> bool {
814 let normalized: String = code
815 .chars()
816 .filter(|ch| ch.is_ascii_alphanumeric())
817 .map(|ch| ch.to_ascii_lowercase())
818 .collect();
819 matches!(
820 normalized.as_str(),
821 "insufficientquota"
822 | "quotaexceeded"
823 | "quotaexhausted"
824 | "billinghardlimitreached"
825 | "billinglimitreached"
826 | "creditbalanceexhausted"
827 // ChatGPT/Codex subscription window (HTTP 429, `error.type`),
828 // as openai/codex `api_bridge.rs` maps it. It resets on the
829 // plan's schedule, not after a short backoff.
830 | "usagelimitreached"
831 // Same backend, same branch: the signed-in plan does not include
832 // Codex. Retrying cannot help; switching accounts can.
833 | "usagenotincluded"
834 | "subscriptionsharingusagelimitexceeded"
835 | "subscriptionsharingusageunavailable"
836 )
837 }
838
839 fn explicit_quota_code_marker(body: &str) -> Option<String> {
840 let lower = body.to_ascii_lowercase();
841 let (_, suffix) = lower.split_once("provider error code:")?;
842 let code = suffix
843 .trim_start()
844 .split(|ch: char| !(ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-')))
845 .next()
846 .unwrap_or_default();
847 is_explicit_quota_code(code).then(|| code.to_string())
848 }
849
850 fn has_explicit_quota_phrase(body: &str) -> bool {
851 let lower = body.to_ascii_lowercase();
852 let current_quota_exhausted = lower.contains("exceeded your current quota")
853 || lower.contains("current quota has been exceeded");
854 let plan_and_billing_guidance = lower.contains("plan") && lower.contains("billing");
855 let durable_scope_exhausted = [
856 "billing quota exceeded",
857 "billing quota exhausted",
858 "billing quota is exhausted",
859 "billing quota has been exceeded",
860 "billing quota has been exhausted",
861 "account quota exceeded",
862 "account quota exhausted",
863 "account quota is exhausted",
864 "account quota has been exceeded",
865 "account quota has been exhausted",
866 "plan quota exceeded",
867 "plan quota exhausted",
868 "plan quota is exhausted",
869 "plan quota has been exceeded",
870 "plan quota has been exhausted",
871 ]
872 .into_iter()
873 .any(|phrase| lower.contains(phrase));
874
875 lower.contains("billing hard limit has been reached")
876 || lower.contains("credit balance exhausted")
877 || lower.contains("credit balance is exhausted")
878 || durable_scope_exhausted
879 || (current_quota_exhausted && plan_and_billing_guidance)
880 }
881
882 fn extract_json_error_message(body: &str) -> Option<String> {
883 let value: Value = serde_json::from_str(body).ok()?;
884 // Flat gateway bodies (`{"error":"Bad Request","message":"Invalid model
885 // name: 'x'"}` — Concentrate, among others) keep the class in `error` and
886 // the detail in `message`. Surfacing only the class hid the one line the
887 // person needed, so carry both when both are present and differ.
888 if let (Some(class), Some(detail)) = (
889 value.pointer("/error").and_then(Value::as_str),
890 value.pointer("/message").and_then(Value::as_str),
891 ) && !class.trim().is_empty()
892 && !detail.trim().is_empty()
893 && !class.trim().eq_ignore_ascii_case(detail.trim())
894 {
895 return Some(format!("{}: {}", class.trim(), detail.trim()));
896 }
897 for pointer in [
898 "/error/message",
899 "/error",
900 "/message",
901 "/detail",
902 "/error_description",
903 ] {
904 let Some(value) = value.pointer(pointer) else {
905 continue;
906 };
907 if let Some(message) = value.as_str() {
908 if !message.trim().is_empty() {
909 return Some(message.to_string());
910 }
911 } else if value.is_object() || value.is_array() {
912 return Some(value.to_string());
913 }
914 }
915 None
916 }
917
918 fn is_probably_html(body: &str) -> bool {
919 let prefix = body
920 .chars()
921 .take(512)
922 .collect::<String>()
923 .to_ascii_lowercase();
924 prefix.contains("<!doctype html") || prefix.contains("<html") || prefix.contains("<head")
925 }
926
927 fn html_to_text(html: &str) -> String {
928 let without_scripts = strip_html_block(html, "script");
929 let without_styles = strip_html_block(&without_scripts, "style");
930 let mut text = String::with_capacity(without_styles.len().min(4096));
931 let mut in_tag = false;
932 for ch in without_styles.chars() {
933 match ch {
934 '<' => {
935 in_tag = true;
936 text.push(' ');
937 }
938 '>' => {
939 in_tag = false;
940 text.push(' ');
941 }
942 _ if !in_tag => text.push(ch),
943 _ => {}
944 }
945 }
946 decode_basic_html_entities(&collapse_whitespace(&text))
947 }
948
949 fn strip_html_block(input: &str, tag: &str) -> String {
950 let mut out = String::with_capacity(input.len());
951 let mut cursor = 0usize;
952 let lower = input.to_ascii_lowercase();
953 let start_marker = format!("<{tag}");
954 let end_marker = format!("</{tag}>");
955
956 while let Some(relative_start) = lower[cursor..].find(&start_marker) {
957 let start = cursor + relative_start;
958 out.push_str(&input[cursor..start]);
959 let after_start = start + start_marker.len();
960 let Some(relative_end) = lower[after_start..].find(&end_marker) else {
961 cursor = input.len();
962 break;
963 };
964 cursor = after_start + relative_end + end_marker.len();
965 out.push(' ');
966 }
967 out.push_str(&input[cursor..]);
968 out
969 }
970
971 fn decode_basic_html_entities(input: &str) -> String {
972 input
973 .replace("&nbsp;", " ")
974 .replace("&amp;", "&")
975 .replace("&lt;", "<")
976 .replace("&gt;", ">")
977 .replace("&quot;", "\"")
978 .replace("&#39;", "'")
979 .replace("&apos;", "'")
980 }
981
982 fn collapse_whitespace(input: &str) -> String {
983 input.split_whitespace().collect::<Vec<_>>().join(" ")
984 }
985
986 fn truncate_for_error(input: &str, max_chars: usize) -> String {
987 let mut out = String::with_capacity(input.len().min(max_chars + 32));
988 for (count, ch) in input.chars().enumerate() {
989 if count >= max_chars {
990 out.push_str("...");
991 return out;
992 }
993 out.push(ch);
994 }
995 out
996 }
997
998 fn extract_cloudflare_error_id(text: &str) -> Option<String> {
999 let mut last = None;
1000 for token in text.split(|ch: char| !ch.is_ascii_hexdigit()) {
1001 if (16..=64).contains(&token.len()) && token.bytes().any(|b| b.is_ascii_alphabetic()) {
1002 last = Some(token.to_string());
1003 }
1004 }
1005 last
1006 }
1007
1008 impl From<reqwest::Error> for LlmError {
1009 fn from(err: reqwest::Error) -> Self {
1010 LlmError::from_reqwest(&err)
1011 }
1012 }
1013
1014 impl From<serde_json::Error> for LlmError {
1015 fn from(err: serde_json::Error) -> Self {
1016 LlmError::ParseError(err.to_string())
1017 }
1018 }
1019
1020 // === RetryConfig - Exponential Backoff Configuration ===
1021
1022 /// Configuration for retry behavior with exponential backoff.
1023 ///
1024 /// This struct controls how retries are performed:
1025 /// - Number of retry attempts
1026 /// - Delay calculation (exponential backoff with optional jitter)
1027 /// - Which HTTP status codes are retryable
1028 /// - Timeout handling
1029 ///
1030 /// # Default Values
1031 ///
1032 /// - `enabled`: true
1033 /// - `max_retries`: 3
1034 /// - `initial_delay`: 1.0 seconds
1035 /// - `max_delay`: 60.0 seconds
1036 /// - `exponential_base`: 2.0
1037 /// - `jitter`: true (adds randomness to prevent thundering herd)
1038 /// - `jitter_factor`: 0.1 (10% variation)
1039 /// - `retryable_status_codes`: [429, 499, 500, 502, 503, 504]
1040 #[derive(Debug, Clone)]
1041 pub struct RetryConfig {
1042 /// Whether retry logic is enabled
1043 pub enabled: bool,
1044
1045 /// Maximum number of retry attempts (0 = no retries, 3 = up to 4 total attempts)
1046 pub max_retries: u32,
1047
1048 /// Initial delay before first retry (seconds)
1049 pub initial_delay: f64,
1050
1051 /// Maximum delay between retries (seconds)
1052 pub max_delay: f64,
1053
1054 /// Base for exponential backoff (delay = initial * base^attempt)
1055 pub exponential_base: f64,
1056
1057 /// Whether to add random jitter to delays
1058 pub jitter: bool,
1059
1060 /// Jitter factor (0.1 = +/- 10% variation)
1061 pub jitter_factor: f64,
1062
1063 /// Whether to respect server's Retry-After header
1064 pub respect_retry_after: bool,
1065
1066 /// HTTP status codes that should trigger a retry
1067 #[allow(dead_code)] // Used in tests via is_retryable_status()
1068 pub retryable_status_codes: Vec<u16>,
1069
1070 /// Timeout for individual requests (seconds, 0 = no timeout)
1071 #[allow(dead_code)] // Configuration field for retry consumers
1072 pub request_timeout: f64,
1073
1074 /// Total timeout for all retry attempts (seconds, 0 = no total timeout)
1075 pub total_timeout: f64,
1076 }
1077
1078 impl Default for RetryConfig {
1079 fn default() -> Self {
1080 Self {
1081 enabled: true,
1082 max_retries: 3,
1083 initial_delay: 1.0,
1084 max_delay: 60.0,
1085 exponential_base: 2.0,
1086 jitter: true,
1087 jitter_factor: 0.1,
1088 respect_retry_after: true,
1089 retryable_status_codes: vec![429, 499, 500, 502, 503, 504],
1090 request_timeout: 120.0,
1091 total_timeout: 0.0, // No total timeout by default
1092 }
1093 }
1094 }
1095
1096 #[allow(dead_code)] // Public builder API, used in tests
1097 impl RetryConfig {
1098 /// Creates a new `RetryConfig` with default values
1099 pub fn new() -> Self {
1100 Self::default()
1101 }
1102
1103 /// Creates a config with retry disabled
1104 pub fn disabled() -> Self {
1105 Self {
1106 enabled: false,
1107 ..Default::default()
1108 }
1109 }
1110
1111 /// Builder method to set max retries
1112 pub fn with_max_retries(mut self, max_retries: u32) -> Self {
1113 self.max_retries = max_retries;
1114 self
1115 }
1116
1117 /// Builder method to set initial delay
1118 pub fn with_initial_delay(mut self, delay: f64) -> Self {
1119 self.initial_delay = delay;
1120 self
1121 }
1122
1123 /// Builder method to set max delay
1124 pub fn with_max_delay(mut self, delay: f64) -> Self {
1125 self.max_delay = delay;
1126 self
1127 }
1128
1129 /// Builder method to enable/disable jitter
1130 pub fn with_jitter(mut self, enabled: bool) -> Self {
1131 self.jitter = enabled;
1132 self
1133 }
1134
1135 /// Calculates the delay for a given retry attempt.
1136 ///
1137 /// Uses exponential backoff: delay = `initial_delay` * `exponential_base^attempt`
1138 /// The result is capped at `max_delay` and optionally has jitter applied.
1139 ///
1140 /// # Arguments
1141 ///
1142 /// * `attempt` - Zero-based attempt number (0 = first retry)
1143 ///
1144 /// # Returns
1145 ///
1146 /// Duration to wait before the next retry attempt
1147 pub fn delay_for_attempt(&self, attempt: u32) -> Duration {
1148 let exponent = i32::try_from(attempt).unwrap_or(i32::MAX);
1149 let base_delay = self.initial_delay * self.exponential_base.powi(exponent);
1150 let capped_delay = base_delay.min(self.max_delay);
1151
1152 let final_delay = if self.jitter {
1153 // Add random jitter to prevent thundering herd problem
1154 let jitter_range = capped_delay * self.jitter_factor;
1155 // Use UUID v4 entropy for jitter randomness.
1156 let bytes = *Uuid::new_v4().as_bytes();
1157 let sample = u16::from_le_bytes([bytes[0], bytes[1]]);
1158 let random_factor = f64::from(sample) / f64::from(u16::MAX); // 0.0 to 1.0
1159 let jitter = jitter_range * (2.0 * random_factor - 1.0); // -range to +range
1160
1161 (capped_delay + jitter).max(0.0)
1162 } else {
1163 capped_delay
1164 };
1165
1166 Duration::from_secs_f64(final_delay)
1167 }
1168
1169 /// Checks if a given HTTP status code should trigger a retry
1170 pub fn is_retryable_status(&self, status: u16) -> bool {
1171 self.retryable_status_codes.contains(&status)
1172 }
1173 }
1174
1175 /// Converts from the existing `RetryPolicy` in config
1176 impl From<RetryPolicy> for RetryConfig {
1177 fn from(policy: RetryPolicy) -> Self {
1178 Self {
1179 enabled: policy.enabled,
1180 max_retries: policy.max_retries,
1181 initial_delay: policy.initial_delay,
1182 max_delay: policy.max_delay,
1183 exponential_base: policy.exponential_base,
1184 jitter: policy.jitter,
1185 jitter_factor: policy.jitter_factor,
1186 respect_retry_after: policy.respect_retry_after,
1187 ..Default::default()
1188 }
1189 }
1190 }
1191
1192 /// Converts back to `RetryPolicy` for compatibility
1193 impl From<RetryConfig> for RetryPolicy {
1194 fn from(config: RetryConfig) -> Self {
1195 Self {
1196 enabled: config.enabled,
1197 max_retries: config.max_retries,
1198 initial_delay: config.initial_delay,
1199 max_delay: config.max_delay,
1200 exponential_base: config.exponential_base,
1201 jitter: config.jitter,
1202 jitter_factor: config.jitter_factor,
1203 respect_retry_after: config.respect_retry_after,
1204 }
1205 }
1206 }
1207
1208 // === Retry Error and Result Types ===
1209
1210 /// Error returned when all retry attempts have been exhausted.
1211 #[derive(Debug)]
1212 pub struct RetryError {
1213 /// The last error encountered
1214 pub last_error: LlmError,
1215
1216 /// Total number of attempts made
1217 pub attempts: u32,
1218
1219 /// Total time spent across all attempts
1220 pub total_time: Duration,
1221 }
1222
1223 impl std::fmt::Display for RetryError {
1224 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1225 write!(
1226 f,
1227 "Retry exhausted after {} attempts ({:?}): {}",
1228 self.attempts, self.total_time, self.last_error
1229 )
1230 }
1231 }
1232
1233 impl std::error::Error for RetryError {
1234 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
1235 Some(&self.last_error)
1236 }
1237 }
1238
1239 /// Result type for retry operations
1240 pub type RetryResult<T> = Result<T, RetryError>;
1241
1242 /// Callback type for retry notifications
1243 ///
1244 /// Called before each retry with:
1245 /// - The error that triggered the retry
1246 /// - The attempt number (0-based)
1247 /// - The delay before the next attempt
1248 pub type RetryCallback = Box<dyn Fn(&LlmError, u32, Duration) + Send + Sync>;
1249
1250 /// An observation of the existing request, never another retry driver. Its
1251 /// lexical scope captures the producing Engine queue, not an ambient session.
1252 pub(crate) type RetryStatusEmitter =
1253 std::sync::Arc<dyn Fn(String) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
1254
1255 #[derive(Clone)]
1256 pub(crate) struct RequestRetryObservation {
1257 pub(crate) retries: std::sync::Arc<std::sync::atomic::AtomicU32>,
1258 pub(crate) emit: RetryStatusEmitter,
1259 }
1260
1261 tokio::task_local! {
1262 static REQUEST_RETRY_OBSERVATION: Option<RequestRetryObservation>;
1263 }
1264
1265 pub(crate) async fn observe_request_retries<F: Future>(
1266 observation: Option<RequestRetryObservation>,
1267 future: F,
1268 ) -> F::Output {
1269 REQUEST_RETRY_OBSERVATION.scope(observation, future).await
1270 }
1271
1272 async fn observe_transport_status(message: String) {
1273 let observation = REQUEST_RETRY_OBSERVATION
1274 .try_with(Clone::clone)
1275 .ok()
1276 .flatten();
1277 if let Some(observation) = observation {
1278 (observation.emit)(message).await;
1279 }
1280 }
1281
1282 fn observe_transport_attempt() {
1283 let _ = REQUEST_RETRY_OBSERVATION.try_with(|observation| {
1284 if let Some(observation) = observation {
1285 let mut count = observation
1286 .retries
1287 .load(std::sync::atomic::Ordering::Relaxed);
1288 loop {
1289 match observation.retries.compare_exchange_weak(
1290 count,
1291 count.saturating_add(1),
1292 std::sync::atomic::Ordering::Relaxed,
1293 std::sync::atomic::Ordering::Relaxed,
1294 ) {
1295 Ok(_) => break,
1296 Err(current) => count = current,
1297 }
1298 }
1299 }
1300 });
1301 }
1302
1303 /// A public receipt must never include the provider-controlled error payload.
1304 /// Raw errors remain intact in the original result and diagnostic/log paths.
1305 pub(crate) fn retry_reason_summary(error: &LlmError) -> String {
1306 match error {
1307 LlmError::RateLimited { .. } => "rate limited".into(),
1308 LlmError::ServerError { status, .. } => format!("upstream {status}"),
1309 LlmError::NetworkError(_) => "network error".into(),
1310 LlmError::Timeout(_) => "timeout".into(),
1311 _ => "non-retryable provider failure".into(),
1312 }
1313 }
1314
1315 async fn observe_transport_stopped(retries: u32, error: &LlmError, exhausted: bool) {
1316 if retries > 0 {
1317 let disposition = if exhausted {
1318 "Retry exhaustion"
1319 } else {
1320 "Retry stopped"
1321 };
1322 observe_transport_status(format!(
1323 "{disposition}: transport request stopped after {retries} retries; {}",
1324 retry_reason_summary(error),
1325 ))
1326 .await;
1327 }
1328 }
1329
1330 // === with_retry - Generic Retry Wrapper ===
1331
1332 /// Executes an async operation with configurable retry logic.
1333 ///
1334 /// This function wraps any async operation that returns `Result<T, LlmError>`
1335 /// and automatically retries on transient failures using exponential backoff.
1336 ///
1337 /// # Arguments
1338 ///
1339 /// * `config` - Retry configuration (delays, max attempts, etc.)
1340 /// * `operation` - Async closure to execute (will be called multiple times on retry)
1341 /// * `callback` - Optional callback for retry notifications (logging, metrics, etc.)
1342 ///
1343 /// # Returns
1344 ///
1345 /// * `Ok(T)` - The successful result from the operation
1346 /// * `Err(RetryError)` - All retries exhausted or non-retryable error encountered
1347 ///
1348 /// # Known limitation: ambiguous failures are replayed
1349 ///
1350 /// A timeout or connection loss after the request was written is retried
1351 /// like a connect failure, so a provider that already accepted the first
1352 /// attempt may bill a second completion. Every caller sends model inference
1353 /// (messages, FIM, translation, speech, provider web search): a replay costs
1354 /// compute but has no external side effect, and the provider APIs used here
1355 /// expose no idempotency key for these requests that could dedupe it. An
1356 /// operation with an external side effect must not be retried through this
1357 /// helper without an idempotency key the server honors.
1358 ///
1359 /// # Example
1360 ///
1361 /// ```ignore
1362 /// let result = with_retry(
1363 /// &config,
1364 /// || async { client.send_request(&req).await },
1365 /// Some(Box::new(|err, attempt, delay| {
1366 /// eprintln!("Retry {} after {:?}: {}", attempt, delay, err);
1367 /// })),
1368 /// ).await;
1369 /// ```
1370 // Keep the structured error inline: this is a public compatibility surface and
1371 // boxing it in a patch release would force every caller to change ownership
1372 // handling for `last_error`.
1373 #[allow(clippy::result_large_err)]
1374 pub async fn with_retry<F, Fut, T>(
1375 config: &RetryConfig,
1376 mut operation: F,
1377 callback: Option<RetryCallback>,
1378 ) -> RetryResult<T>
1379 where
1380 F: FnMut() -> Fut,
1381 Fut: Future<Output = Result<T, LlmError>>,
1382 {
1383 // If retries are disabled, just run once
1384 if !config.enabled {
1385 return operation().await.map_err(|e| RetryError {
1386 last_error: e,
1387 attempts: 1,
1388 total_time: Duration::ZERO,
1389 });
1390 }
1391
1392 let start_time = Instant::now();
1393 let total_timeout = if config.total_timeout > 0.0 {
1394 Some(Duration::from_secs_f64(config.total_timeout))
1395 } else {
1396 None
1397 };
1398
1399 let mut last_error: Option<LlmError> = None;
1400
1401 // Attempt 0 is the first try, then up to max_retries additional attempts
1402 for attempt in 0..=config.max_retries {
1403 // Check total timeout
1404 if let Some(timeout) = total_timeout
1405 && start_time.elapsed() >= timeout
1406 {
1407 let error = last_error.unwrap_or(LlmError::Timeout(timeout));
1408 observe_transport_stopped(attempt.saturating_sub(1), &error, true).await;
1409 return Err(RetryError {
1410 last_error: error,
1411 attempts: attempt,
1412 total_time: start_time.elapsed(),
1413 });
1414 }
1415
1416 if attempt > 0 {
1417 observe_transport_attempt();
1418 }
1419 match operation().await {
1420 Ok(result) => {
1421 if attempt > 0 {
1422 observe_transport_status(format!(
1423 "Retry recovery: transport request recovered after {attempt} retries"
1424 ))
1425 .await;
1426 }
1427 return Ok(result);
1428 }
1429 Err(err) => {
1430 // Non-retryable errors fail immediately
1431 if !err.is_retryable() {
1432 observe_transport_stopped(attempt, &err, false).await;
1433 return Err(RetryError {
1434 last_error: err,
1435 attempts: attempt + 1,
1436 total_time: start_time.elapsed(),
1437 });
1438 }
1439
1440 // Last attempt - no more retries
1441 if attempt >= config.max_retries {
1442 observe_transport_stopped(attempt, &err, true).await;
1443 return Err(RetryError {
1444 last_error: err,
1445 attempts: attempt + 1,
1446 total_time: start_time.elapsed(),
1447 });
1448 }
1449
1450 // Calculate delay
1451 // Use server's Retry-After if available and configured
1452 let base_delay = config.delay_for_attempt(attempt);
1453 let delay = if config.respect_retry_after {
1454 err.suggested_retry_delay().unwrap_or(base_delay)
1455 } else {
1456 base_delay
1457 };
1458
1459 // Notify callback if provided
1460 if let Some(ref cb) = callback {
1461 cb(&err, attempt, delay);
1462 }
1463
1464 observe_transport_status(format!(
1465 "Retry attempt: transport {}/{}; {}; waiting {:.2}s",
1466 attempt + 1,
1467 config.max_retries,
1468 retry_reason_summary(&err),
1469 delay.as_secs_f64(),
1470 ))
1471 .await;
1472 last_error = Some(err);
1473
1474 // Wait before retrying
1475 tokio::time::sleep(delay).await;
1476 }
1477 }
1478 }
1479
1480 // Should not reach here, but handle gracefully
1481 Err(RetryError {
1482 last_error: last_error.unwrap_or(LlmError::Other("Unknown retry error".to_string())),
1483 attempts: config.max_retries + 1,
1484 total_time: start_time.elapsed(),
1485 })
1486 }
1487
1488 // === Utility Functions ===
1489
1490 /// The longest a `Retry-After` value is ever believed. A server (or a proxy
1491 /// in front of it) can send an arbitrarily large delay; without a ceiling a
1492 /// single `Retry-After: 86400` would wedge the turn for a day. One hour is
1493 /// well past any legitimate rate-limit window.
1494 const RETRY_AFTER_MAX: Duration = Duration::from_secs(3600);
1495
1496 /// Parses the Retry-After header value into a Duration.
1497 ///
1498 /// Supports both:
1499 /// - Seconds as integer: "120" -> 120 seconds
1500 /// - HTTP-date format: "Wed, 21 Oct 2015 07:28:00 GMT" (not implemented, returns None)
1501 ///
1502 /// The value is server-controlled, so this never panics and never returns an
1503 /// unbounded delay: negative / NaN / infinite / absurd floats are rejected
1504 /// (`Duration::from_secs_f64` panics on a negative — a remote-triggerable
1505 /// crash before this guard), and any result is clamped to [`RETRY_AFTER_MAX`].
1506 pub fn parse_retry_after(value: &str) -> Option<Duration> {
1507 // Try parsing as seconds
1508 if let Ok(seconds) = value.parse::<u64>() {
1509 return Some(Duration::from_secs(seconds).min(RETRY_AFTER_MAX));
1510 }
1511
1512 // Try parsing as float seconds. Only a finite, non-negative value is a
1513 // meaningful delay; everything else (`-5`, `nan`, `inf`) is "no usable
1514 // hint". Clamp to the ceiling BEFORE `from_secs_f64` so an out-of-range
1515 // float can never reach its overflow-panic path, while keeping the
1516 // sub-second precision a legitimate `1.5` carries.
1517 if let Ok(seconds) = value.parse::<f64>()
1518 && seconds.is_finite()
1519 && seconds >= 0.0
1520 {
1521 let clamped = seconds.min(RETRY_AFTER_MAX.as_secs_f64());
1522 return Some(Duration::from_secs_f64(clamped));
1523 }
1524
1525 // HTTP-date format not supported yet
1526 // Could use chrono or httpdate crate if needed
1527 None
1528 }
1529
1530 /// Extracts Retry-After duration from response headers
1531 pub fn extract_retry_after(headers: &reqwest::header::HeaderMap) -> Option<Duration> {
1532 headers
1533 .get(reqwest::header::RETRY_AFTER)
1534 .and_then(|v| v.to_str().ok())
1535 .and_then(parse_retry_after)
1536 }
1537
1538 #[cfg(test)]
1539 #[path = "tests.rs"]
1540 mod quota_tests;
1541
1542 // === Tests ===
1543
1544 #[cfg(test)]
1545 mod tests {
1546 use super::*;
1547
1548 #[test]
1549 fn official_chatgpt_usage_codes_are_terminal_across_http_boundaries() {
1550 for code in [
1551 "subscription_sharing_usage_limit_exceeded",
1552 "subscription_sharing_usage_unavailable",
1553 ] {
1554 let body = serde_json::json!({"error":{"code":code,"message":"opaque"}}).to_string();
1555 for status in [400, 403, 429] {
1556 for safe_body in [body.clone(), sanitize_http_error_body(None, status, &body)] {
1557 let error =
1558 LlmError::from_http_response_with_auth_context(status, &safe_body, None);
1559 assert!(matches!(error, LlmError::QuotaExhausted(_)), "{error:?}");
1560 assert!(!error.is_retryable());
1561 assert!(error.to_string().contains("ChatGPT Settings > Usage"));
1562 let envelope = crate::error_taxonomy::envelope_for_llm_error(
1563 error.into(),
1564 "allowance unavailable".into(),
1565 );
1566 assert!(!envelope.recoverable);
1567 assert_eq!(envelope.code, "llm_quota_exhausted");
1568 }
1569 }
1570 }
1571 for body in [
1572 "usage limit reached",
1573 r#"{"error":{"code":"subscription_sharing_usage_limit_exceeded_later"}}"#,
1574 ] {
1575 assert!(matches!(
1576 LlmError::from_http_response(429, body),
1577 LlmError::RateLimited { .. }
1578 ));
1579 assert!(LlmError::from_http_response(429, body).is_retryable());
1580 }
1581 }
1582
1583 fn assert_f64_eq(actual: f64, expected: f64) {
1584 assert!(
1585 (actual - expected).abs() < f64::EPSILON,
1586 "expected {expected}, got {actual}"
1587 );
1588 }
1589
1590 fn auth_user_message(error: LlmError) -> String {
1591 match error {
1592 LlmError::AuthenticationError(auth) => auth.to_user_message(),
1593 other => panic!("expected authentication error, got {other}"),
1594 }
1595 }
1596
1597 #[test]
1598 fn test_retry_config_defaults() {
1599 let config = RetryConfig::default();
1600 assert!(config.enabled);
1601 assert_eq!(config.max_retries, 3);
1602 assert_f64_eq(config.initial_delay, 1.0);
1603 assert_f64_eq(config.max_delay, 60.0);
1604 assert_f64_eq(config.exponential_base, 2.0);
1605 assert!(config.jitter);
1606 }
1607
1608 #[test]
1609 fn test_retry_config_disabled() {
1610 let config = RetryConfig::disabled();
1611 assert!(!config.enabled);
1612 }
1613
1614 #[test]
1615 fn test_retry_config_builder() {
1616 let config = RetryConfig::new()
1617 .with_max_retries(5)
1618 .with_initial_delay(2.0)
1619 .with_max_delay(120.0)
1620 .with_jitter(false);
1621
1622 assert_eq!(config.max_retries, 5);
1623 assert_f64_eq(config.initial_delay, 2.0);
1624 assert_f64_eq(config.max_delay, 120.0);
1625 assert!(!config.jitter);
1626 }
1627
1628 #[test]
1629 fn test_delay_for_attempt_exponential() {
1630 let config = RetryConfig::new().with_jitter(false);
1631
1632 // delay = initial * base^attempt
1633 // 1.0 * 2^0 = 1.0
1634 let d0 = config.delay_for_attempt(0);
1635 assert_eq!(d0, Duration::from_secs_f64(1.0));
1636
1637 // 1.0 * 2^1 = 2.0
1638 let d1 = config.delay_for_attempt(1);
1639 assert_eq!(d1, Duration::from_secs_f64(2.0));
1640
1641 // 1.0 * 2^2 = 4.0
1642 let d2 = config.delay_for_attempt(2);
1643 assert_eq!(d2, Duration::from_secs_f64(4.0));
1644
1645 // 1.0 * 2^3 = 8.0
1646 let d3 = config.delay_for_attempt(3);
1647 assert_eq!(d3, Duration::from_secs_f64(8.0));
1648 }
1649
1650 #[test]
1651 fn test_delay_for_attempt_capped() {
1652 let config = RetryConfig::new().with_jitter(false).with_max_delay(5.0);
1653
1654 // 1.0 * 2^3 = 8.0, but capped at 5.0
1655 let d3 = config.delay_for_attempt(3);
1656 assert_eq!(d3, Duration::from_secs_f64(5.0));
1657 }
1658
1659 #[test]
1660 fn test_delay_for_attempt_with_jitter() {
1661 let config = RetryConfig::new().with_jitter(true);
1662
1663 // With jitter, delays should vary slightly
1664 let d1 = config.delay_for_attempt(1);
1665 let d2 = config.delay_for_attempt(1);
1666
1667 // Both should be close to 2.0 seconds (within 10% jitter)
1668 let base = 2.0;
1669 let range = base * 0.1;
1670 assert!(d1.as_secs_f64() >= base - range);
1671 assert!(d1.as_secs_f64() <= base + range);
1672 assert!(d2.as_secs_f64() >= base - range);
1673 assert!(d2.as_secs_f64() <= base + range);
1674 }
1675
1676 #[test]
1677 fn test_is_retryable_status() {
1678 let config = RetryConfig::default();
1679
1680 assert!(config.is_retryable_status(429)); // Rate limit
1681 assert!(config.is_retryable_status(499)); // Upstream request cancelled
1682 assert!(config.is_retryable_status(500)); // Internal server error
1683 assert!(config.is_retryable_status(502)); // Bad gateway
1684 assert!(config.is_retryable_status(503)); // Service unavailable
1685 assert!(config.is_retryable_status(504)); // Gateway timeout
1686
1687 assert!(!config.is_retryable_status(400)); // Bad request
1688 assert!(!config.is_retryable_status(401)); // Unauthorized
1689 assert!(!config.is_retryable_status(403)); // Forbidden
1690 assert!(!config.is_retryable_status(404)); // Not found
1691 }
1692
1693 #[test]
1694 fn auth_error_with_context_includes_provider_authority_model_and_key_source() {
1695 let err = LlmError::from_http_response_with_request_context(
1696 401,
1697 "Invalid API Key",
1698 Some("Xiaomi MiMo"),
1699 Some("https://token-plan-sgp.xiaomimimo.com/v1"),
1700 Some("mimo-v2.5"),
1701 Some("env"),
1702 Some("tp-secret-token-plan-value"),
1703 );
1704 let message = auth_user_message(err);
1705
1706 assert!(message.contains("Invalid API Key"));
1707 assert!(message.contains("provider: Xiaomi MiMo"));
1708 assert!(message.contains("base URL authority: token-plan-sgp.xiaomimimo.com"));
1709 assert!(message.contains("model: mimo-v2.5"));
1710 assert!(message.contains("key source: env"));
1711 assert!(message.contains("key fingerprint: tp-... (len=26)"));
1712 }
1713
1714 #[test]
1715 fn auth_error_redacts_full_api_key_from_body_and_context() {
1716 let api_key = "tp-secret-token-plan-value";
1717 let err = LlmError::from_http_response_with_request_context(
1718 401,
1719 &format!("Invalid API Key: {api_key}"),
1720 Some("Xiaomi MiMo"),
1721 Some("https://token-plan-sgp.xiaomimimo.com/v1"),
1722 Some("mimo-v2.5"),
1723 Some("config-file"),
1724 Some(api_key),
1725 );
1726 let message = auth_user_message(err);
1727
1728 assert!(!message.contains(api_key));
1729 assert!(!message.contains("secret-token-plan-value"));
1730 assert!(message.contains("[redacted API key]"));
1731 assert!(message.contains("key fingerprint: tp-... (len=26)"));
1732 }
1733
1734 #[test]
1735 fn auth_error_classifies_xiaomi_token_plan_key_prefix() {
1736 let token_plan = AuthenticationErrorContext::from_parts(
1737 None,
1738 None,
1739 None,
1740 Some("session"),
1741 Some("tp-secret-token-plan-value"),
1742 );
1743 let generic = AuthenticationErrorContext::from_parts(
1744 None,
1745 None,
1746 None,
1747 Some("session"),
1748 Some("sk-other"),
1749 );
1750 let unprefixed = AuthenticationErrorContext::from_parts(
1751 None,
1752 None,
1753 None,
1754 Some("session"),
1755 Some("plainsecretvalue"),
1756 );
1757
1758 assert_eq!(
1759 token_plan.key_kind.as_deref(),
1760 Some("Xiaomi MiMo Token Plan key")
1761 );
1762 assert_eq!(generic.key_kind.as_deref(), Some("API key"));
1763 assert_eq!(unprefixed.key_kind.as_deref(), Some("API key"));
1764 assert_eq!(
1765 unprefixed.key_fingerprint.as_deref(),
1766 Some("unprefixed (len=16)")
1767 );
1768 }
1769
1770 #[test]
1771 fn authorization_403_is_not_reclassified_by_auth_context() {
1772 let err = LlmError::from_http_response_with_request_context(
1773 403,
1774 "forbidden",
1775 Some("Arcee AI"),
1776 Some("https://api.arcee.ai/v1"),
1777 Some("auto"),
1778 Some("env"),
1779 Some("sk-arcee-secret"),
1780 );
1781
1782 assert!(matches!(err, LlmError::AuthorizationError(_)));
1783 }
1784
1785 #[test]
1786 fn auth_error_without_context_preserves_bare_message() {
1787 let err = LlmError::from_http_response_with_auth_context(
1788 401,
1789 "Invalid API Key",
1790 Some(AuthenticationErrorContext::default()),
1791 );
1792
1793 assert_eq!(auth_user_message(err), "Invalid API Key");
1794 }
1795
1796 /// A flat `{"error","message"}` body keeps its detail: the class alone
1797 /// ("Bad Request") told the person nothing about the rejected model.
1798 #[test]
1799 fn flat_error_and_message_body_surfaces_both_halves() {
1800 let body = r#"{"error":"Bad Request","message":"Invalid model name: 'invalid-model-xyz'"}"#;
1801 assert_eq!(
1802 sanitize_http_error_body(Some("Concentrate"), 400, body),
1803 "Bad Request: Invalid model name: 'invalid-model-xyz'"
1804 );
1805 // Identical halves are not doubled, and the nested OpenAI shape is
1806 // untouched.
1807 assert_eq!(
1808 sanitize_http_error_body(
1809 Some("Concentrate"),
1810 401,
1811 r#"{"error":"Unauthorized","message":"Unauthorized"}"#
1812 ),
1813 "Unauthorized"
1814 );
1815 assert_eq!(
1816 sanitize_http_error_body(
1817 Some("fixture"),
1818 400,
1819 r#"{"error":{"message":"nested detail"}}"#
1820 ),
1821 "nested detail"
1822 );
1823 }
1824
1825 #[test]
1826 fn cloudflare_html_error_is_summarized_without_raw_markup() {
1827 let body = r#"<!DOCTYPE html><html><head><title>Access Denied</title><style>
1828 .hidden { display: none; }
1829 </style></head><body>
1830 <h1>Access Denied</h1>
1831 <p>The action you just performed triggered a security alert.</p>
1832 <script>window.noisy = true;</script>
1833 <span>2600:1700:467:d410:f137:b94f:1dd0:d1e4</span>
1834 <span>a059a2873f3fdf82</span>
1835 <div>Cloudflare Error Pages</div>
1836 </body></html>"#;
1837
1838 let message = sanitize_http_error_body(Some("Arcee AI"), 403, body);
1839
1840 assert!(message.contains("Arcee AI API returned Cloudflare Access Denied"));
1841 assert!(message.contains("ID a059a2873f3fdf82"));
1842 assert!(!message.contains("<!DOCTYPE"));
1843 assert!(!message.contains("tailwindcss"));
1844 assert!(message.len() < 300);
1845 }
1846
1847 #[test]
1848 fn cloudflare_access_denied_403_is_authorization_not_authentication() {
1849 let message = sanitize_http_error_body(
1850 Some("Arcee AI"),
1851 403,
1852 r#"<!doctype html><html><body><h1>Access Denied</h1><p>Cloudflare Error Pages</p></body></html>"#,
1853 );
1854 let err = LlmError::from_http_response(403, &message);
1855
1856 assert!(matches!(err, LlmError::AuthorizationError(_)));
1857 }
1858
1859 #[test]
1860 fn arcee_access_denied_without_literal_cloudflare_is_still_summarized() {
1861 // Mirrors api.arcee.ai's real 403 page: "Cloudflare" appears only in a
1862 // `<meta>` attribute and the `<style>` block, both stripped, so the
1863 // visible text never contains it. The summary must still fire from the
1864 // WAF's stock "security alert" / "Contact Support" copy + error ID.
1865 let body = r#"<!DOCTYPE html><html lang="en"><head>
1866 <meta name="description" content="Cloudflare Error Pages">
1867 <title>Access Denied</title>
1868 <style>:root{--accent:cloudflare}</style></head><body>
1869 <h1>Access Denied</h1>
1870 <p>The action you just performed triggered a security alert.</p>
1871 <p>Please contact us if this was a mistake.</p>
1872 <a>Contact Support</a>
1873 <span>2600:1700:467:d410:f137:b94f:1dd0:d1e4</span>
1874 <span>a059c0d4caf1f9cc</span>
1875 </body></html>"#;
1876
1877 let message = sanitize_http_error_body(Some("Arcee AI"), 403, body);
1878
1879 assert!(
1880 message.contains("Arcee AI API returned Access Denied"),
1881 "got: {message}"
1882 );
1883 assert!(message.contains("ID a059c0d4caf1f9cc"), "got: {message}");
1884 assert!(
1885 !message.to_ascii_lowercase().contains("cloudflare"),
1886 "stripped Arcee page has no literal Cloudflare: {message}"
1887 );
1888 assert!(!message.contains('<'), "no raw markup: {message}");
1889 assert!(message.len() < 300, "stays concise: {message}");
1890
1891 // A WAF block is authorization, not a bad API key.
1892 let err = LlmError::from_http_response(403, &message);
1893 assert!(matches!(err, LlmError::AuthorizationError(_)));
1894 }
1895
1896 #[test]
1897 fn test_llm_error_suggested_retry_delay() {
1898 let err = LlmError::RateLimited {
1899 message: "slow down".to_string(),
1900 retry_after: Some(Duration::from_secs(60)),
1901 };
1902 assert_eq!(err.suggested_retry_delay(), Some(Duration::from_secs(60)));
1903
1904 let err = LlmError::ServerError {
1905 status: 500,
1906 message: "error".to_string(),
1907 };
1908 assert_eq!(err.suggested_retry_delay(), None);
1909 }
1910
1911 #[test]
1912 fn test_parse_retry_after() {
1913 // Integer seconds
1914 assert_eq!(parse_retry_after("120"), Some(Duration::from_secs(120)));
1915 assert_eq!(parse_retry_after("0"), Some(Duration::from_secs(0)));
1916
1917 // Float seconds keep sub-second precision
1918 assert_eq!(parse_retry_after("1.5"), Some(Duration::from_secs_f64(1.5)));
1919
1920 // Invalid
1921 assert_eq!(parse_retry_after("invalid"), None);
1922 assert_eq!(parse_retry_after(""), None);
1923 }
1924
1925 /// A `Retry-After` value is server-controlled. Malformed floats used to
1926 /// reach `Duration::from_secs_f64`, which panics on a negative — a
1927 /// remote-triggerable crash in the request path (2026-08-04 review).
1928 #[test]
1929 fn parse_retry_after_never_panics_and_is_bounded_on_hostile_input() {
1930 // None of these may panic.
1931 assert_eq!(parse_retry_after("-5"), None, "negative is not a delay");
1932 assert_eq!(parse_retry_after("nan"), None);
1933 assert_eq!(parse_retry_after("inf"), None);
1934 assert_eq!(parse_retry_after("-inf"), None);
1935 // Absurdly large values clamp to the ceiling rather than overflowing
1936 // or wedging the turn for a day.
1937 assert_eq!(parse_retry_after("1e300"), Some(RETRY_AFTER_MAX));
1938 assert_eq!(parse_retry_after("86400"), Some(RETRY_AFTER_MAX));
1939 assert_eq!(
1940 parse_retry_after("999999999999"),
1941 Some(RETRY_AFTER_MAX),
1942 "integer path is clamped too"
1943 );
1944 // A normal value still passes through untouched.
1945 assert_eq!(parse_retry_after("30"), Some(Duration::from_secs(30)));
1946 }
1947
1948 #[test]
1949 fn test_retry_policy_conversion() {
1950 let policy = RetryPolicy {
1951 enabled: true,
1952 max_retries: 5,
1953 initial_delay: 2.0,
1954 max_delay: 30.0,
1955 exponential_base: 3.0,
1956 jitter: false,
1957 jitter_factor: 0.25,
1958 respect_retry_after: false,
1959 };
1960
1961 let config: RetryConfig = policy.clone().into();
1962 assert_eq!(config.enabled, policy.enabled);
1963 assert_eq!(config.max_retries, policy.max_retries);
1964 assert_f64_eq(config.initial_delay, policy.initial_delay);
1965 assert_f64_eq(config.max_delay, policy.max_delay);
1966 assert_f64_eq(config.exponential_base, policy.exponential_base);
1967
1968 // #6700: the jitter and Retry-After knobs survive the conversion
1969 // instead of silently resetting to `RetryConfig::default()`.
1970 assert!(!config.jitter);
1971 assert_f64_eq(config.jitter_factor, 0.25);
1972 assert!(!config.respect_retry_after);
1973
1974 // Convert back
1975 let policy2: RetryPolicy = config.into();
1976 assert_eq!(policy2.enabled, policy.enabled);
1977 assert_eq!(policy2.max_retries, policy.max_retries);
1978 assert_eq!(policy2.jitter, policy.jitter);
1979 assert_f64_eq(policy2.jitter_factor, policy.jitter_factor);
1980 assert_eq!(policy2.respect_retry_after, policy.respect_retry_after);
1981 }
1982
1983 #[tokio::test]
1984 async fn test_with_retry_success_first_attempt() {
1985 let config = RetryConfig::default();
1986 let mut call_count = 0;
1987
1988 let result = with_retry(
1989 &config,
1990 || {
1991 call_count += 1;
1992 async { Ok::<_, LlmError>(42) }
1993 },
1994 None,
1995 )
1996 .await;
1997
1998 assert!(result.is_ok());
1999 assert_eq!(result.unwrap(), 42);
2000 assert_eq!(call_count, 1);
2001 }
2002
2003 #[tokio::test]
2004 async fn test_with_retry_disabled() {
2005 let config = RetryConfig::disabled();
2006 let mut call_count = 0;
2007
2008 let result: RetryResult<i32> = with_retry(
2009 &config,
2010 || {
2011 call_count += 1;
2012 async {
2013 Err(LlmError::ServerError {
2014 status: 500,
2015 message: "error".to_string(),
2016 })
2017 }
2018 },
2019 None,
2020 )
2021 .await;
2022
2023 assert!(result.is_err());
2024 assert_eq!(call_count, 1); // No retries when disabled
2025 }
2026
2027 #[tokio::test]
2028 async fn test_with_retry_eventual_success() {
2029 let config = RetryConfig::new()
2030 .with_max_retries(3)
2031 .with_initial_delay(0.01); // Fast for testing
2032
2033 let call_count = std::sync::Arc::new(std::sync::atomic::AtomicU32::new(0));
2034 let cc = call_count.clone();
2035
2036 let result = with_retry(
2037 &config,
2038 || {
2039 let count = cc.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
2040 async move {
2041 if count < 2 {
2042 Err(LlmError::ServerError {
2043 status: 500,
2044 message: "temporary error".to_string(),
2045 })
2046 } else {
2047 Ok::<_, LlmError>(42)
2048 }
2049 }
2050 },
2051 None,
2052 )
2053 .await;
2054
2055 assert!(result.is_ok());
2056 assert_eq!(result.unwrap(), 42);
2057 assert_eq!(call_count.load(std::sync::atomic::Ordering::SeqCst), 3); // 2 failures + 1 success
2058 }
2059
2060 #[tokio::test]
2061 async fn test_with_retry_exhausted() {
2062 let config = RetryConfig::new()
2063 .with_max_retries(2)
2064 .with_initial_delay(0.01);
2065
2066 let call_count = std::sync::Arc::new(std::sync::atomic::AtomicU32::new(0));
2067 let cc = call_count.clone();
2068
2069 let result: RetryResult<i32> = with_retry(
2070 &config,
2071 || {
2072 cc.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
2073 async {
2074 Err(LlmError::ServerError {
2075 status: 500,
2076 message: "persistent error".to_string(),
2077 })
2078 }
2079 },
2080 None,
2081 )
2082 .await;
2083
2084 assert!(result.is_err());
2085 let err = result.unwrap_err();
2086 assert_eq!(err.attempts, 3); // 1 initial + 2 retries
2087 assert_eq!(call_count.load(std::sync::atomic::Ordering::SeqCst), 3);
2088 }
2089
2090 #[tokio::test]
2091 async fn test_with_retry_callback() {
2092 let config = RetryConfig::new()
2093 .with_max_retries(2)
2094 .with_initial_delay(0.01);
2095
2096 let callback_count = std::sync::Arc::new(std::sync::atomic::AtomicU32::new(0));
2097 let cc = callback_count.clone();
2098
2099 let _: RetryResult<i32> = with_retry(
2100 &config,
2101 || async {
2102 Err(LlmError::ServerError {
2103 status: 500,
2104 message: "error".to_string(),
2105 })
2106 },
2107 Some(Box::new(move |_err, _attempt, _delay| {
2108 cc.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
2109 })),
2110 )
2111 .await;
2112
2113 // Callback called once per retry (not for the final failure)
2114 assert_eq!(callback_count.load(std::sync::atomic::Ordering::SeqCst), 2);
2115 }
2116
2117 #[test]
2118 fn test_retry_error_display() {
2119 let err = RetryError {
2120 last_error: LlmError::ServerError {
2121 status: 500,
2122 message: "internal error".to_string(),
2123 },
2124 attempts: 4,
2125 total_time: Duration::from_secs(10),
2126 };
2127
2128 let display = format!("{err}");
2129 assert!(display.contains("4 attempts"));
2130 assert!(display.contains("10"));
2131 assert!(display.contains("Server error"));
2132 }
2133 }
2134
2134 lines RUST