1use std::future::Future;
9use std::pin::Pin;
10use std::sync::Arc;
11use std::time::Duration;
12
13use rand::Rng;
14use reqwest::StatusCode;
15use tokio::sync::RwLock;
16use tracing::{info, warn};
17
18use super::CloudHomeError;
19use crate::oauth::{self, OAuthConfig, OAuthTokens};
20use coven_foundation::clock::ClockRef;
21#[cfg(test)]
22use coven_keys::keys::StoreKeys;
23use coven_keys::keys::{CloudHomeCredentialCustody, CloudHomeCredentials};
24
25type Sleeper = Arc<dyn Fn(Duration) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
28
29const MAX_TRANSIENT_RETRIES: u32 = 4;
35const RETRY_BASE_DELAY: Duration = Duration::from_millis(500);
37const RETRY_MAX_DELAY: Duration = Duration::from_secs(32);
40
41fn is_transient(status: StatusCode) -> bool {
45 status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
46}
47
48fn parse_retry_after(resp: &reqwest::Response) -> Option<Duration> {
52 let secs: u64 = resp
53 .headers()
54 .get(reqwest::header::RETRY_AFTER)?
55 .to_str()
56 .ok()?
57 .trim()
58 .parse()
59 .ok()?;
60 Some(Duration::from_secs(secs))
61}
62
63pub struct OAuthSession {
67 client: reqwest::Client,
68 tokens: RwLock<OAuthTokens>,
69 credential_custody: Arc<dyn CloudHomeCredentialCustody>,
70 clock: ClockRef,
71 config: OAuthConfig,
72 provider_label: &'static str,
74 sleeper: Sleeper,
75}
76
77pub(crate) struct OAuthRequest<'session> {
81 client: &'session reqwest::Client,
82 token: &'session str,
83}
84
85impl OAuthRequest<'_> {
86 pub(crate) fn get(&self, url: impl reqwest::IntoUrl) -> reqwest::RequestBuilder {
87 self.client.get(url).bearer_auth(self.token)
88 }
89
90 pub(crate) fn post(&self, url: impl reqwest::IntoUrl) -> reqwest::RequestBuilder {
91 self.client.post(url).bearer_auth(self.token)
92 }
93
94 pub(crate) fn put(&self, url: impl reqwest::IntoUrl) -> reqwest::RequestBuilder {
95 self.client.put(url).bearer_auth(self.token)
96 }
97
98 pub(crate) fn patch(&self, url: impl reqwest::IntoUrl) -> reqwest::RequestBuilder {
99 self.client.patch(url).bearer_auth(self.token)
100 }
101
102 pub(crate) fn delete(&self, url: impl reqwest::IntoUrl) -> reqwest::RequestBuilder {
103 self.client.delete(url).bearer_auth(self.token)
104 }
105}
106
107impl OAuthSession {
108 pub fn new(
109 tokens: OAuthTokens,
110 credential_custody: Arc<dyn CloudHomeCredentialCustody>,
111 clock: ClockRef,
112 config: OAuthConfig,
113 provider_label: &'static str,
114 ) -> Self {
115 Self {
116 client: reqwest::Client::new(),
117 tokens: RwLock::new(tokens),
118 credential_custody,
119 clock,
120 config,
121 provider_label,
122 sleeper: Arc::new(|delay| Box::pin(tokio::time::sleep(delay))),
123 }
124 }
125
126 async fn access_token(&self) -> Result<String, CloudHomeError> {
128 let tokens = self.tokens.read().await;
129 match tokens.expires_at {
130 Some(expires_at) if self.clock.now().timestamp() < expires_at - 60 => {
131 return Ok(tokens.access_token.clone());
132 }
133 None => return Ok(tokens.access_token.clone()),
135 _ => {}
136 }
137 drop(tokens);
138 self.refresh().await
139 }
140
141 async fn refresh(&self) -> Result<String, CloudHomeError> {
143 let mut tokens = self.tokens.write().await;
144
145 if let Some(expires_at) = tokens.expires_at {
147 if self.clock.now().timestamp() < expires_at - 60 {
148 return Ok(tokens.access_token.clone());
149 }
150 }
151
152 let refresh_token = tokens.refresh_token.as_deref().ok_or_else(|| {
153 CloudHomeError::Configuration(format!(
154 "Your {} sign-in is missing a refresh token. Reconnect to keep syncing.",
155 self.provider_label,
156 ))
157 })?;
158
159 let new_tokens = oauth::refresh(
160 &self.client,
161 &self.config,
162 refresh_token,
163 self.clock.as_ref(),
164 )
165 .await
166 .map_err(|e| match e {
167 oauth::OAuthError::Reauthorize(detail) => CloudHomeError::Configuration(format!(
168 "Your {} access was revoked or expired. Reconnect to keep syncing. ({detail})",
169 self.provider_label,
170 )),
171 other => CloudHomeError::Transport(format!("OAuth refresh failed: {other}")),
172 })?;
173
174 self.credential_custody
175 .persist(&CloudHomeCredentials::OAuth {
176 tokens: new_tokens.clone(),
177 })
178 .map_err(|e| {
179 CloudHomeError::transport("persist refreshed OAuth tokens".to_string(), e)
180 })?;
181
182 let access_token = new_tokens.access_token.clone();
183 *tokens = new_tokens;
184
185 info!("Refreshed {} OAuth tokens", self.provider_label);
186 Ok(access_token)
187 }
188
189 async fn send_with_refresh<F>(
194 &self,
195 build_request: &F,
196 ) -> Result<reqwest::Response, CloudHomeError>
197 where
198 F: for<'request> Fn(OAuthRequest<'request>) -> reqwest::RequestBuilder,
199 {
200 let token = self.access_token().await?;
201 let resp = build_request(OAuthRequest {
202 client: &self.client,
203 token: &token,
204 })
205 .send()
206 .await
207 .map_err(|e| CloudHomeError::transport("request failed".to_string(), e))?;
208
209 if resp.status() == StatusCode::UNAUTHORIZED {
210 let new_token = self.refresh().await?;
211 build_request(OAuthRequest {
212 client: &self.client,
213 token: &new_token,
214 })
215 .send()
216 .await
217 .map_err(|e| CloudHomeError::transport("retry request failed".to_string(), e))
218 } else {
219 Ok(resp)
220 }
221 }
222
223 pub(crate) async fn api_call<F>(
234 &self,
235 build_request: F,
236 ) -> Result<reqwest::Response, CloudHomeError>
237 where
238 F: for<'request> Fn(OAuthRequest<'request>) -> reqwest::RequestBuilder,
239 {
240 let mut attempt = 0u32;
241 loop {
242 let resp = self.send_with_refresh(&build_request).await?;
243 let status = resp.status();
244 if attempt < MAX_TRANSIENT_RETRIES && is_transient(status) {
245 let delay = self.retry_delay(&resp, attempt);
246 warn!(
247 "{} request returned {status}, retrying in {delay:?} (attempt {}/{})",
248 self.provider_label,
249 attempt + 1,
250 MAX_TRANSIENT_RETRIES,
251 );
252 (self.sleeper)(delay).await;
253 attempt += 1;
254 continue;
255 }
256 return Ok(resp);
257 }
258 }
259
260 pub(crate) async fn api_call_no_transient_retry<F>(
267 &self,
268 build_request: F,
269 ) -> Result<reqwest::Response, CloudHomeError>
270 where
271 F: for<'request> Fn(OAuthRequest<'request>) -> reqwest::RequestBuilder,
272 {
273 self.send_with_refresh(&build_request).await
274 }
275
276 pub(crate) fn range_put_uploader(
277 &self,
278 session_url: String,
279 intermediate_status: u16,
280 total: u64,
281 part_size: usize,
282 key: String,
283 classify: super::resumable::ClassifyWrite,
284 cancellation_succeeded: super::resumable::CancellationSucceeded,
285 ) -> super::resumable::RangePutUploader {
286 super::resumable::RangePutUploader::new(
287 self.client.clone(),
288 session_url,
289 intermediate_status,
290 total,
291 part_size,
292 key,
293 classify,
294 cancellation_succeeded,
295 )
296 }
297
298 pub(crate) fn range_put_sink(
299 &self,
300 session_url: String,
301 intermediate_status: u16,
302 total: u64,
303 part_size: usize,
304 key: String,
305 classify: super::resumable::ClassifyWrite,
306 cancellation_succeeded: super::resumable::CancellationSucceeded,
307 ) -> super::resumable::RangePutSink {
308 super::resumable::RangePutSink::new(
309 self.client.clone(),
310 session_url,
311 intermediate_status,
312 total,
313 part_size,
314 key,
315 classify,
316 cancellation_succeeded,
317 )
318 }
319
320 fn retry_delay(&self, resp: &reqwest::Response, attempt: u32) -> Duration {
326 if let Some(after) = parse_retry_after(resp) {
327 return after.min(RETRY_MAX_DELAY);
328 }
329 let ceiling = RETRY_BASE_DELAY
330 .saturating_mul(1u32 << attempt.min(16))
331 .min(RETRY_MAX_DELAY);
332 let jittered = rand::rng().random_range(0..=ceiling.as_millis() as u64);
333 Duration::from_millis(jittered)
334 }
335}
336
337#[cfg(all(test, feature = "oauth-providers"))]
338mod tests {
339 use super::*;
340 use crate::oauth::test_support::{oauth_config, serve_token_response};
341 use axum::body::Body;
342 use axum::http::Response;
343 use axum::Router;
344 use chrono::{TimeZone, Utc};
345 use coven_foundation::clock::{FixedClock, SystemClock};
346 use std::sync::atomic::{AtomicUsize, Ordering};
347 use std::sync::Arc;
348 use tokio::sync::oneshot;
349
350 fn recording_sleeper() -> (Sleeper, Arc<std::sync::Mutex<Vec<Duration>>>) {
354 let recorded = Arc::new(std::sync::Mutex::new(Vec::new()));
355 let sink = recorded.clone();
356 let sleeper: Sleeper = Arc::new(move |delay| {
357 sink.lock().expect("record sleep delay").push(delay);
358 Box::pin(async {})
359 });
360 (sleeper, recorded)
361 }
362
363 fn retry_test_session(sleeper: Sleeper) -> OAuthSession {
366 let mut session = OAuthSession::new(
367 OAuthTokens {
368 access_token: "access".to_string(),
369 refresh_token: None,
370 expires_at: None,
371 },
372 coven_keys::keys::CloudHomeCredentialsOwner::new(StoreKeys::bind(
373 "oauth-retry".to_string(),
374 ))
375 .current(),
376 Arc::new(SystemClock),
377 oauth_config("http://token.invalid/token".to_string()),
378 "Provider",
379 );
380 session.sleeper = sleeper;
381 session
382 }
383
384 async fn spawn_status_server(
388 responses: Vec<(u16, Option<u64>)>,
389 ) -> (String, Arc<AtomicUsize>, oneshot::Sender<()>) {
390 let hits = Arc::new(AtomicUsize::new(0));
391 let plan = Arc::new(responses);
392
393 let handler_hits = hits.clone();
394 let app = Router::new().fallback(move || {
395 let plan = plan.clone();
396 let hits = handler_hits.clone();
397 async move {
398 let i = hits.fetch_add(1, Ordering::SeqCst);
399 let (status, retry_after) = plan
400 .get(i)
401 .or_else(|| plan.last())
402 .copied()
403 .expect("at least one programmed response");
404 let mut builder = Response::builder().status(status);
405 if let Some(secs) = retry_after {
406 builder = builder.header("Retry-After", secs.to_string());
407 }
408 builder.body(Body::from("body")).expect("build response")
409 }
410 });
411 let (url, shutdown_tx) = crate::cloud::test_server::spawn_test_server(app).await;
412 (url, hits, shutdown_tx)
413 }
414
415 #[tokio::test]
416 async fn retries_after_429_and_honors_retry_after() {
417 let (url, hits, shutdown) = spawn_status_server(vec![(429, Some(2)), (200, None)]).await;
418 let (sleeper, recorded) = recording_sleeper();
419 let session = retry_test_session(sleeper);
420
421 let resp = session
422 .api_call(|oauth| oauth.get(&url))
423 .await
424 .expect("call succeeds after the 429 clears");
425
426 assert_eq!(resp.status(), StatusCode::OK);
427 assert_eq!(hits.load(Ordering::SeqCst), 2, "one retry after the 429");
428 assert_eq!(
429 recorded.lock().expect("read sleeps").as_slice(),
430 &[Duration::from_secs(2)],
431 "the single backoff honored Retry-After",
432 );
433 let _ = shutdown.send(());
434 }
435
436 #[tokio::test]
437 async fn persistent_5xx_exhausts_attempts_and_surfaces_the_failure() {
438 let (url, hits, shutdown) = spawn_status_server(vec![(500, None)]).await;
439 let (sleeper, recorded) = recording_sleeper();
440 let session = retry_test_session(sleeper);
441
442 let resp = session
443 .api_call(|oauth| oauth.get(&url))
444 .await
445 .expect("the exhausted response returns for the caller to fail on");
446
447 assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
448 assert_eq!(
449 hits.load(Ordering::SeqCst),
450 (MAX_TRANSIENT_RETRIES + 1) as usize,
451 "initial attempt plus every retry",
452 );
453 let waits = recorded.lock().expect("read sleeps");
454 assert_eq!(waits.len(), MAX_TRANSIENT_RETRIES as usize);
455 assert!(
456 waits.iter().all(|w| *w <= RETRY_MAX_DELAY),
457 "no computed backoff exceeds the ceiling: {waits:?}",
458 );
459 let _ = shutdown.send(());
460 }
461
462 #[tokio::test]
463 async fn non_transient_4xx_is_not_retried() {
464 let (url, hits, shutdown) = spawn_status_server(vec![(400, None)]).await;
465 let (sleeper, recorded) = recording_sleeper();
466 let session = retry_test_session(sleeper);
467
468 let resp = session
469 .api_call(|oauth| oauth.get(&url))
470 .await
471 .expect("a 400 returns unchanged");
472
473 assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
474 assert_eq!(hits.load(Ordering::SeqCst), 1, "the 400 is sent once");
475 assert!(
476 recorded.lock().expect("read sleeps").is_empty(),
477 "a 400 triggers no backoff",
478 );
479 let _ = shutdown.send(());
480 }
481
482 fn fail_next_cloud_credentials_write(key_service: &StoreKeys) {
483 key_service
484 .fail_next_cloud_home_credentials_operation_for_test(keyring_core::Error::Invalid(
485 "keyring unavailable".to_string(),
486 "test failure".to_string(),
487 ))
488 .expect("configure mock keyring failure");
489 }
490
491 #[tokio::test]
492 async fn refresh_returns_error_when_token_persist_fails() {
493 let (token_url, request_body, server) = serve_token_response(
494 r#"{"access_token":"new-access","refresh_token":"new-refresh","expires_in":3600}"#,
495 )
496 .await;
497 coven_keys::keys::test_keyring::install();
498 let key_service = StoreKeys::bind("oauth-persist-failure".to_string());
499 fail_next_cloud_credentials_write(&key_service);
500 let session = OAuthSession::new(
501 OAuthTokens {
502 access_token: "old-access".to_string(),
503 refresh_token: Some("old-refresh".to_string()),
504 expires_at: Some(1_700_000_000),
505 },
506 coven_keys::keys::CloudHomeCredentialsOwner::new(key_service).current(),
507 Arc::new(FixedClock(Utc.timestamp_opt(1_700_000_120, 0).unwrap())),
508 oauth_config(token_url),
509 "Provider",
510 );
511
512 let error = session
513 .refresh()
514 .await
515 .expect_err("persist failure returns an error");
516 assert!(error.to_string().contains("keyring unavailable"));
517
518 let tokens = session.tokens.read().await;
519 assert_eq!(tokens.access_token, "old-access");
520 assert_eq!(tokens.refresh_token.as_deref(), Some("old-refresh"));
521
522 let _request = request_body.await.expect("receive refresh request");
523 server.await.expect("token server exits");
524 }
525}