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
51pub 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 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 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 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 #[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
290fn 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}