Skip to main content

photon_ai_api/
client.rs

1use std::{
2    convert::Infallible,
3    fmt,
4    future::Future,
5    pin::Pin,
6    sync::Arc,
7    time::{Duration, SystemTime},
8};
9
10use reqwest::{
11    Method, Request, Response, StatusCode,
12    header::{HeaderMap, RETRY_AFTER},
13};
14
15use crate::{Client, Credential, Error, ExecuteFuture, HttpBackend, TransportError};
16
17pub use crate::config_generated::DEFAULT_BASE_URL;
18
19type HeaderFuture = Pin<Box<dyn Future<Output = HeaderMap> + Send>>;
20type DynamicHeaders = dyn Fn() -> HeaderFuture + Send + Sync;
21
22#[derive(Clone, Default)]
23enum HeaderProvider {
24    #[default]
25    None,
26    Static(HeaderMap),
27    Dynamic(Arc<DynamicHeaders>),
28}
29
30impl HeaderProvider {
31    async fn get(&self) -> HeaderMap {
32        match self {
33            Self::None => HeaderMap::new(),
34            Self::Static(headers) => headers.clone(),
35            Self::Dynamic(provider) => provider().await,
36        }
37    }
38}
39
40#[derive(Clone)]
41struct PhotonBackend {
42    client: reqwest::Client,
43    headers: HeaderProvider,
44    timeout: Duration,
45    max_attempts: usize,
46    base_delay: Duration,
47    maximum_delay: Duration,
48    maximum_retry_after: Duration,
49}
50
51/// Longest server-requested `Retry-After` delay honoured by default.
52pub const DEFAULT_MAXIMUM_RETRY_AFTER: Duration = Duration::from_secs(60);
53
54impl fmt::Debug for PhotonBackend {
55    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
56        formatter
57            .debug_struct("PhotonBackend")
58            .field("timeout", &self.timeout)
59            .field("max_attempts", &self.max_attempts)
60            .field("base_delay", &self.base_delay)
61            .field("maximum_delay", &self.maximum_delay)
62            .field("maximum_retry_after", &self.maximum_retry_after)
63            .finish_non_exhaustive()
64    }
65}
66
67impl PhotonBackend {
68    async fn execute_request(self, request: Request) -> Result<Response, TransportError> {
69        let template = request.try_clone();
70        let mut attempts = 1;
71        let mut first = Some(request);
72        let mut attempt = 0;
73
74        while attempt < attempts {
75            attempt += 1;
76            let mut request = if attempt == 1 {
77                first.take().expect("first request is available")
78            } else if let Some(request) = template.as_ref().and_then(Request::try_clone) {
79                request
80            } else {
81                break;
82            };
83            *request.timeout_mut() = Some(self.timeout);
84            let configured_headers = self.headers.get().await;
85            merge_configured_headers(&mut request, &configured_headers);
86            if attempt == 1 && template.is_some() && can_retry(&request) {
87                // Decided once configured headers are merged, so an
88                // Idempotency-Key from the client-wide headers enables retries too.
89                attempts = self.max_attempts;
90            }
91
92            let result = self.client.execute(request).await;
93            let retry =
94                attempt < attempts && retryable_outcome(result.as_ref().ok().map(Response::status));
95            match result {
96                Ok(response) if retry => {
97                    let Some(delay) =
98                        retry_delay(retry_after(&response), self.maximum_retry_after, || {
99                            jitter(&self, attempt)
100                        })
101                    else {
102                        return Ok(response);
103                    };
104                    let _ = response.bytes().await;
105                    tokio::time::sleep(delay).await;
106                }
107                Ok(response) => return Ok(response),
108                Err(source) => return Err(TransportError::new(source)),
109            }
110        }
111
112        unreachable!("the first attempt always returns or advances to a retry")
113    }
114}
115
116impl HttpBackend for PhotonBackend {
117    fn execute(&self, request: Request) -> ExecuteFuture<'_> {
118        let backend = self.clone();
119        Box::pin(async move { backend.execute_request(request).await })
120    }
121}
122
123#[derive(Clone)]
124pub struct PhotonClientBuilder {
125    base_url: String,
126    headers: HeaderProvider,
127    credentials: Vec<(String, Credential)>,
128    timeout: Duration,
129    max_attempts: usize,
130    base_delay: Duration,
131    maximum_delay: Duration,
132    maximum_retry_after: Duration,
133    client: Option<reqwest::Client>,
134}
135
136impl fmt::Debug for PhotonClientBuilder {
137    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
138        formatter
139            .debug_struct("PhotonClientBuilder")
140            .field("base_url", &self.base_url)
141            .field("credential_count", &self.credentials.len())
142            .field("timeout", &self.timeout)
143            .field("max_attempts", &self.max_attempts)
144            .field("base_delay", &self.base_delay)
145            .field("maximum_delay", &self.maximum_delay)
146            .field("maximum_retry_after", &self.maximum_retry_after)
147            .field("has_custom_client", &self.client.is_some())
148            .finish_non_exhaustive()
149    }
150}
151
152impl Default for PhotonClientBuilder {
153    fn default() -> Self {
154        Self {
155            base_url: DEFAULT_BASE_URL.to_owned(),
156            headers: HeaderProvider::None,
157            credentials: Vec::new(),
158            timeout: Duration::from_secs(30),
159            max_attempts: 3,
160            base_delay: Duration::from_millis(250),
161            maximum_delay: Duration::from_secs(2),
162            maximum_retry_after: DEFAULT_MAXIMUM_RETRY_AFTER,
163            client: None,
164        }
165    }
166}
167
168impl PhotonClientBuilder {
169    pub fn new() -> Self {
170        Self::default()
171    }
172
173    pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
174        self.base_url = base_url.into().trim_end_matches('/').to_owned();
175        self
176    }
177
178    pub fn static_headers(mut self, headers: HeaderMap) -> Self {
179        self.headers = HeaderProvider::Static(headers);
180        self
181    }
182
183    pub fn headers<F, Fut>(mut self, provider: F) -> Self
184    where
185        F: Fn() -> Fut + Send + Sync + 'static,
186        Fut: Future<Output = HeaderMap> + Send + 'static,
187    {
188        self.headers = HeaderProvider::Dynamic(Arc::new(move || Box::pin(provider())));
189        self
190    }
191
192    /// Register a credential under an OpenAPI security-scheme name.
193    ///
194    /// The scheme name must match a key in `components.securitySchemes`. The
195    /// generated client validates operation security requirements before the
196    /// transport backend runs, so credentials must be registered here rather
197    /// than supplied only through [`Self::static_headers`] or [`Self::headers`].
198    pub fn credential(mut self, scheme: impl Into<String>, credential: Credential) -> Self {
199        self.credentials.push((scheme.into(), credential));
200        self
201    }
202
203    pub fn timeout(mut self, timeout: Duration) -> Self {
204        self.timeout = timeout;
205        self
206    }
207
208    pub fn max_attempts(mut self, max_attempts: usize) -> Self {
209        self.max_attempts = max_attempts.clamp(1, 3);
210        self
211    }
212
213    /// Longest server-requested `Retry-After` delay to wait for before a retry
214    /// (default 60 seconds). A longer delay is not shortened: the client stops
215    /// retrying and returns that response.
216    pub fn maximum_retry_after(mut self, maximum_retry_after: Duration) -> Self {
217        self.maximum_retry_after = maximum_retry_after;
218        self
219    }
220
221    pub fn reqwest_client(mut self, client: reqwest::Client) -> Self {
222        self.client = Some(client);
223        self
224    }
225
226    // Keep Spargen's native error type in the public builder contract.
227    #[allow(clippy::result_large_err)]
228    pub fn build(self) -> Result<Client, Error<Infallible>> {
229        let client = match self.client {
230            Some(client) => client,
231            None => reqwest::Client::builder()
232                .connect_timeout(self.timeout)
233                .timeout(self.timeout)
234                .build()
235                .map_err(Error::<Infallible>::request_construction)?,
236        };
237        let backend = PhotonBackend {
238            client,
239            headers: self.headers,
240            timeout: self.timeout,
241            max_attempts: self.max_attempts,
242            base_delay: self.base_delay,
243            maximum_delay: self.maximum_delay,
244            maximum_retry_after: self.maximum_retry_after,
245        };
246        let mut client = Client::with_backend(Arc::new(backend), &self.base_url)?;
247        for (scheme, credential) in self.credentials {
248            client = client.with_credential(&scheme, credential);
249        }
250        Ok(client)
251    }
252}
253
254fn can_retry(request: &Request) -> bool {
255    matches!(
256        *request.method(),
257        Method::GET | Method::HEAD | Method::OPTIONS | Method::TRACE
258    ) || request.headers().contains_key("idempotency-key")
259}
260
261fn retryable_status(status: StatusCode) -> bool {
262    matches!(status.as_u16(), 408 | 429 | 502 | 503 | 504)
263}
264
265fn retryable_outcome(status: Option<StatusCode>) -> bool {
266    status.is_some_and(retryable_status)
267}
268
269fn merge_configured_headers(request: &mut Request, configured: &HeaderMap) {
270    for (name, value) in configured {
271        if !request.headers().contains_key(name) {
272            request.headers_mut().insert(name, value.clone());
273        }
274    }
275}
276
277fn retry_after(response: &Response) -> Option<Duration> {
278    let value = response.headers().get(RETRY_AFTER)?.to_str().ok()?;
279    parse_retry_after(value, SystemTime::now())
280}
281
282fn parse_retry_after(value: &str, now: SystemTime) -> Option<Duration> {
283    if let Ok(seconds) = value.parse::<u64>() {
284        return Some(Duration::from_secs(seconds));
285    }
286    let retry_at = httpdate::parse_http_date(value).ok()?;
287    Some(retry_at.duration_since(now).unwrap_or(Duration::ZERO))
288}
289
290/// The delay before the next attempt, or `None` when the server asked for a
291/// longer wait than `maximum_retry_after` and the response should be returned.
292fn retry_delay(
293    retry_after: Option<Duration>,
294    maximum_retry_after: Duration,
295    jitter: impl FnOnce() -> Duration,
296) -> Option<Duration> {
297    match retry_after {
298        Some(delay) if delay > maximum_retry_after => None,
299        Some(delay) => Some(delay),
300        None => Some(jitter()),
301    }
302}
303
304fn backoff_ceiling(backend: &PhotonBackend, attempt: usize) -> Duration {
305    backend
306        .base_delay
307        .saturating_mul(1_u32 << (attempt - 1).min(8))
308        .min(backend.maximum_delay)
309}
310
311fn jitter(backend: &PhotonBackend, attempt: usize) -> Duration {
312    let ceiling = backoff_ceiling(backend, attempt);
313    Duration::from_secs_f64(rand::random_range(0.0..=ceiling.as_secs_f64()))
314}
315
316#[cfg(test)]
317mod tests {
318    use crate::SecretString;
319    use reqwest::header::{AUTHORIZATION, HeaderValue};
320
321    use super::*;
322
323    #[test]
324    fn retries_safe_requests_and_idempotent_mutations_only() {
325        let get = Request::new(Method::GET, "https://example.test".parse().unwrap());
326        assert!(can_retry(&get));
327
328        let post = Request::new(Method::POST, "https://example.test".parse().unwrap());
329        assert!(!can_retry(&post));
330
331        let mut idempotent_post =
332            Request::new(Method::POST, "https://example.test".parse().unwrap());
333        idempotent_post
334            .headers_mut()
335            .insert("idempotency-key", "stable-key".parse().unwrap());
336        assert!(can_retry(&idempotent_post));
337    }
338
339    #[test]
340    fn retry_statuses_match_the_transport_contract() {
341        for status in [408, 429, 502, 503, 504] {
342            assert!(retryable_status(StatusCode::from_u16(status).unwrap()));
343        }
344        for status in [400, 401, 409, 500, 501, 505] {
345            assert!(!retryable_status(StatusCode::from_u16(status).unwrap()));
346        }
347        assert!(
348            !retryable_outcome(None),
349            "transport failures are not retried"
350        );
351    }
352
353    #[test]
354    fn retry_after_and_full_jitter_match_the_transport_contract() {
355        let now = SystemTime::UNIX_EPOCH + Duration::from_secs(1_000_000);
356        assert_eq!(parse_retry_after("3", now), Some(Duration::from_secs(3)));
357        let retry_at = httpdate::fmt_http_date(now + Duration::from_secs(5));
358        assert_eq!(
359            parse_retry_after(&retry_at, now),
360            Some(Duration::from_secs(5))
361        );
362        assert_eq!(parse_retry_after("not-a-date", now), None);
363
364        let backend = PhotonBackend {
365            client: reqwest::Client::new(),
366            headers: HeaderProvider::None,
367            timeout: Duration::from_secs(30),
368            max_attempts: 3,
369            base_delay: Duration::from_millis(250),
370            maximum_delay: Duration::from_secs(2),
371            maximum_retry_after: DEFAULT_MAXIMUM_RETRY_AFTER,
372        };
373        for (attempt, expected) in [250, 500, 1_000, 2_000, 2_000].into_iter().enumerate() {
374            let attempt = attempt + 1;
375            let ceiling = Duration::from_millis(expected);
376            assert_eq!(backoff_ceiling(&backend, attempt), ceiling);
377            assert!(jitter(&backend, attempt) <= ceiling);
378        }
379    }
380
381    #[test]
382    fn retry_after_is_capped_without_shortening_the_server_delay() {
383        let cap = Duration::from_secs(60);
384        let fallback = || Duration::from_millis(7);
385        assert_eq!(
386            retry_delay(Some(Duration::from_secs(60)), cap, fallback),
387            Some(Duration::from_secs(60))
388        );
389        assert_eq!(
390            retry_delay(Some(Duration::from_secs(61)), cap, fallback),
391            None
392        );
393        assert_eq!(
394            retry_delay(Some(Duration::from_secs(u64::MAX)), cap, fallback),
395            None
396        );
397        assert_eq!(
398            retry_delay(None, cap, fallback),
399            Some(Duration::from_millis(7))
400        );
401        assert_eq!(
402            PhotonClientBuilder::new().maximum_retry_after,
403            DEFAULT_MAXIMUM_RETRY_AFTER
404        );
405        assert_eq!(
406            PhotonClientBuilder::new()
407                .maximum_retry_after(Duration::from_secs(5))
408                .maximum_retry_after,
409            Duration::from_secs(5)
410        );
411    }
412
413    #[test]
414    fn operation_headers_take_precedence_over_configured_headers() {
415        let mut request = Request::new(Method::GET, "https://example.test".parse().unwrap());
416        request
417            .headers_mut()
418            .insert(AUTHORIZATION, HeaderValue::from_static("operation"));
419        let mut configured = HeaderMap::new();
420        configured.insert(AUTHORIZATION, HeaderValue::from_static("configured"));
421        configured.insert("x-extra", HeaderValue::from_static("value"));
422
423        merge_configured_headers(&mut request, &configured);
424
425        assert_eq!(request.headers()[AUTHORIZATION], "operation");
426        assert_eq!(request.headers()["x-extra"], "value");
427    }
428
429    #[test]
430    fn registers_credentials_for_arbitrary_security_schemes() {
431        let client = PhotonClientBuilder::new()
432            .credential(
433                "futureSecurityScheme",
434                Credential::ApiKey(SecretString::from("test-secret".to_owned())),
435            )
436            .build()
437            .unwrap();
438
439        assert!(client.core().credential("futureSecurityScheme").is_some());
440    }
441}