Skip to main content

nanocodex_service/middleware/
retry.rs

1use std::{
2    future::Future,
3    pin::Pin,
4    sync::{Arc, atomic::Ordering},
5    time::Duration,
6};
7
8use nanocodex_core::AgentEventKind;
9use tower::retry::{Policy, Retry};
10use web_time::Instant;
11
12use crate::{
13    attempt::{ResponsesAttempt, ResponsesServiceResponse},
14    service::ResponsesService,
15    service_error::{FailurePhase, ResponsesServiceError},
16    telemetry::{AttemptRetrying, duration_ns, elapsed_ns},
17};
18
19#[cfg(not(target_family = "wasm"))]
20type RetryFuture = Pin<Box<dyn Future<Output = ()> + Send>>;
21#[cfg(target_family = "wasm")]
22type RetryFuture = Pin<Box<dyn Future<Output = ()>>>;
23
24#[derive(Clone, Copy, Default)]
25pub struct ResponsesRetryPolicy;
26
27impl Policy<ResponsesAttempt, ResponsesServiceResponse, ResponsesServiceError>
28    for ResponsesRetryPolicy
29{
30    type Future = RetryFuture;
31
32    fn retry(
33        &mut self,
34        request: &mut ResponsesAttempt,
35        result: &mut Result<ResponsesServiceResponse, ResponsesServiceError>,
36    ) -> Option<Self::Future> {
37        let failure = result.as_ref().err()?;
38        let checkpoint_missing =
39            failure.is_checkpoint_missing() && request.previous_response_id().is_some();
40        let advice = failure.retry_advice;
41        if !checkpoint_missing && advice.is_none() {
42            return None;
43        }
44        if request.attempt >= request.max_attempts {
45            return None;
46        }
47        let delay = if checkpoint_missing {
48            Duration::ZERO
49        } else {
50            advice
51                .and_then(|advice| advice.server_delay)
52                .unwrap_or_else(|| retry_delay(request.attempt, request.call_index))
53        };
54        let error_class = if checkpoint_missing {
55            "checkpoint_missing"
56        } else {
57            advice.map_or("unknown", |advice| advice.class)
58        };
59        let message = failure.source.to_string();
60        if let Err(error) = request.observer.emit(
61            AgentEventKind::ModelAttemptRetrying,
62            AttemptRetrying {
63                phase: request.kind,
64                model_call_index: request.call_index,
65                attempt: request.attempt,
66                next_attempt: request.attempt + 1,
67                max_attempts: request.max_attempts,
68                failure_phase: failure.phase,
69                error_class,
70                delay_ns: duration_ns(delay),
71                server_requested_delay: advice.is_some_and(|advice| advice.server_delay.is_some()),
72                opens_new_socket: !checkpoint_missing,
73                replay_mode: "full_history",
74                connection_generation: failure.connection_generation,
75                error: &message,
76            },
77        ) {
78            *result = Err(ResponsesServiceError::event(
79                error,
80                FailurePhase::Output,
81                failure.connection_generation,
82            ));
83            return None;
84        }
85        request
86            .observer
87            .stats
88            .response_retries
89            .fetch_add(1, Ordering::Relaxed);
90        tracing::warn!(
91            target: "nanocodex_service",
92            phase = request.kind.phase(),
93            model.call_index = request.call_index,
94            attempt = request.attempt,
95            next_attempt = request.attempt + 1,
96            error.class = error_class,
97            delay_ms = u64::try_from(delay.as_millis()).unwrap_or(u64::MAX),
98            server_requested_delay = advice.is_some_and(|advice| advice.server_delay.is_some()),
99            "retrying Responses attempt"
100        );
101        if !request.prepare_retry() {
102            return None;
103        }
104        let stats = Arc::clone(&request.observer.stats);
105        Some(Box::pin(async move {
106            let started_at = Instant::now();
107            sleep(delay).await;
108            stats
109                .retry_backoff_duration_ns
110                .fetch_add(elapsed_ns(started_at), Ordering::Relaxed);
111        }))
112    }
113
114    fn clone_request(&mut self, request: &ResponsesAttempt) -> Option<ResponsesAttempt> {
115        Some(request.clone())
116    }
117}
118
119#[cfg(not(target_family = "wasm"))]
120async fn sleep(delay: Duration) {
121    tokio::time::sleep(delay).await;
122}
123
124#[cfg(target_family = "wasm")]
125async fn sleep(delay: Duration) {
126    use wasm_bindgen::prelude::*;
127    use wasm_bindgen_futures::JsFuture;
128
129    #[wasm_bindgen]
130    extern "C" {
131        #[wasm_bindgen(js_namespace = ["globalThis", "nanocodexHost"], js_name = sleep)]
132        fn host_sleep(milliseconds: u32) -> js_sys::Promise;
133    }
134
135    let milliseconds = u32::try_from(delay.as_millis()).unwrap_or(u32::MAX);
136    drop(JsFuture::from(host_sleep(milliseconds)).await);
137}
138
139pub type DefaultResponsesService = Retry<ResponsesRetryPolicy, ResponsesService>;
140
141fn retry_delay(attempt: u32, call_index: Option<u32>) -> Duration {
142    let base_ms = if cfg!(test) { 1 } else { 200 };
143    let exponent = attempt.saturating_sub(1).min(4);
144    let raw_ms = base_ms * 2_u64.pow(exponent);
145    let seed = u64::from(call_index.unwrap_or_default()) * 31 + u64::from(attempt) * 17;
146    let jitter_percent = 90 + seed % 21;
147    Duration::from_millis(raw_ms * jitter_percent / 100)
148}
149
150#[cfg(test)]
151mod tests {
152    use super::retry_delay;
153
154    #[test]
155    fn local_retry_delay_is_bounded_and_exponential() {
156        let first = retry_delay(1, Some(7));
157        let second = retry_delay(2, Some(7));
158        assert!(first.as_millis() <= 2);
159        assert!(second > first);
160    }
161}