Skip to main content

coven_storage/cloud/
oauth_session.rs

1//! Shared OAuth token lifecycle for the consumer-cloud backends.
2//!
3//! Google Drive, Dropbox, and OneDrive all cache an access token, refresh it on
4//! expiry (persisting the new tokens through its credential custody), and retry
5//! a request once on a 401. This holds that logic in one place; each backend owns an
6//! `OAuthSession` and routes its requests through `api_call`.
7
8use 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
25/// Waits out one retry delay. A field so tests exercise the retry schedule
26/// (attempt count, honored `Retry-After`) without sleeping real seconds.
27type Sleeper = Arc<dyn Fn(Duration) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
28
29/// A 429 or 5xx retries this many times before the failure surfaces. Chosen with
30/// [`RETRY_BASE_DELAY`]/[`RETRY_MAX_DELAY`] so quota pressure (routine 429s on a
31/// large Drive sync) degrades a cycle to slow rather than failed, while a hard
32/// outage exhausts the attempts in a few seconds and fails loud to the cycle,
33/// which then applies its own minutes-long backoff.
34const MAX_TRANSIENT_RETRIES: u32 = 4;
35/// The first retry waits this long (before jitter); each further retry doubles it.
36const RETRY_BASE_DELAY: Duration = Duration::from_millis(500);
37/// Ceiling on any single wait, applied to both the exponential term and a
38/// server-supplied `Retry-After`, so one response can't stall a cycle for minutes.
39const RETRY_MAX_DELAY: Duration = Duration::from_secs(32);
40
41/// 429 and 5xx are the transient failures worth retrying; every other 4xx
42/// (including 401, handled separately by token refresh) is a caller error a retry
43/// won't fix.
44fn is_transient(status: StatusCode) -> bool {
45    status == StatusCode::TOO_MANY_REQUESTS || status.is_server_error()
46}
47
48/// The `Retry-After` header as a duration when the server sent one as an integer
49/// number of seconds — the form Drive/Dropbox/OneDrive use. Absent or non-integer
50/// header falls through to computed backoff.
51fn 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
63/// Owns a provider's OAuth tokens (refreshing them as needed) and the
64/// `reqwest::Client` its requests go out on — every OAuth backend shared the same
65/// client field and token lifecycle, so both live here once.
66pub struct OAuthSession {
67    client: reqwest::Client,
68    tokens: RwLock<OAuthTokens>,
69    credential_custody: Arc<dyn CloudHomeCredentialCustody>,
70    clock: ClockRef,
71    config: OAuthConfig,
72    /// Human-readable provider name, used only in log lines.
73    provider_label: &'static str,
74    sleeper: Sleeper,
75}
76
77/// One token-bound request construction scope. Provider code may configure the
78/// request builder it receives, but cannot extract either the session's HTTP
79/// client or its access token.
80pub(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    /// The current access token, refreshing if it's expired or about to expire.
127    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            // No expiry info: assume valid.
134            None => return Ok(tokens.access_token.clone()),
135            _ => {}
136        }
137        drop(tokens);
138        self.refresh().await
139    }
140
141    /// Refresh the tokens and persist them through this provider's custody.
142    async fn refresh(&self) -> Result<String, CloudHomeError> {
143        let mut tokens = self.tokens.write().await;
144
145        // Another task may have refreshed while we waited for the write lock.
146        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    /// Build and send a request with the current token, refreshing and retrying
190    /// once on a 401. This is the single send path; [`api_call`](Self::api_call)
191    /// wraps it with transient-failure retries, and the resumable part sinks call
192    /// it through [`api_call_no_transient_retry`](Self::api_call_no_transient_retry).
193    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    /// Send a request, retrying transient failures (429 and 5xx) with bounded
224    /// exponential backoff, jitter, and honored `Retry-After`, so Drive, Dropbox,
225    /// and OneDrive all inherit quota/outage tolerance. Non-transient responses
226    /// (2xx, 404, other 4xx) return unchanged for the caller to interpret.
227    ///
228    /// The request is rebuilt from `build_request` on every attempt, so its body
229    /// must be replayable — which the `Fn` signature enforces. Requests whose
230    /// success mutates non-idempotent server state (resumable upload parts) must
231    /// not come through here; see
232    /// [`api_call_no_transient_retry`](Self::api_call_no_transient_retry).
233    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    /// Send a request whose transient-failure retry must be owned by a higher
261    /// layer rather than this session. A resumable upload part advances a
262    /// server-side session offset on success, so re-sending it after a lost
263    /// response collides with that advanced offset; the recovery is to re-run the
264    /// whole upload from the source, which the blob engine does when this call's
265    /// error surfaces. Still refreshes and retries once on a 401.
266    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    /// The wait before the next retry: a server-supplied `Retry-After` when
321    /// present, else exponential backoff with full jitter — a random point in
322    /// `[0, base·2^attempt]` so a fleet hitting the same quota window doesn't
323    /// resynchronize onto the same retry instant. Both are clamped to
324    /// [`RETRY_MAX_DELAY`].
325    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    /// A [`Sleeper`] that records the delays it's asked to wait instead of
351    /// sleeping, so a test can assert the retry schedule and honored `Retry-After`
352    /// without spending real seconds.
353    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    /// A session pointed at a mock endpoint with a non-expiring token (so
364    /// `access_token` never refreshes) and the injected `sleeper`.
365    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    /// Serve the programmed `(status, retry_after_secs)` responses in order, then
385    /// repeat the last one for any further request. Returns the URL, a shared hit
386    /// counter, and a shutdown sender.
387    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}