Skip to main content

fraiseql_auth/oauth/
refresh.rs

1//! Token refresh scheduler and background worker.
2
3use std::{sync::Arc, time::Duration as StdDuration};
4
5use chrono::{DateTime, Duration, Utc};
6
7use super::super::error::AuthError;
8
9/// Token refresh scheduler
10#[derive(Debug, Clone)]
11pub struct TokenRefreshScheduler {
12    /// Sessions needing refresh
13    // std::sync::Mutex is intentional: this lock is never held across .await.
14    // Switch to tokio::sync::Mutex if that constraint ever changes.
15    refresh_queue: Arc<std::sync::Mutex<Vec<(String, DateTime<Utc>)>>>,
16}
17
18impl TokenRefreshScheduler {
19    /// Create new refresh scheduler
20    #[must_use]
21    pub fn new() -> Self {
22        Self {
23            refresh_queue: Arc::new(std::sync::Mutex::new(Vec::new())),
24        }
25    }
26
27    /// Schedule token refresh for session
28    ///
29    /// # Errors
30    ///
31    /// Returns `AuthError::Internal` if the mutex is poisoned.
32    pub fn schedule_refresh(
33        &self,
34        session_id: String,
35        refresh_time: DateTime<Utc>,
36    ) -> std::result::Result<(), AuthError> {
37        let mut queue = self.refresh_queue.lock().map_err(|_| AuthError::Internal {
38            message: "token refresh scheduler mutex poisoned".to_string(),
39        })?;
40        queue.push((session_id, refresh_time));
41        queue.sort_by_key(|(_, time)| *time);
42        Ok(())
43    }
44
45    /// Get next session to refresh
46    ///
47    /// # Errors
48    ///
49    /// Returns `AuthError::Internal` if the mutex is poisoned.
50    pub fn get_next_refresh(&self) -> std::result::Result<Option<String>, AuthError> {
51        let mut queue = self.refresh_queue.lock().map_err(|_| AuthError::Internal {
52            message: "token refresh scheduler mutex poisoned".to_string(),
53        })?;
54        if let Some((_, refresh_time)) = queue.first() {
55            if *refresh_time <= Utc::now() {
56                let (id, _) = queue.remove(0);
57                return Ok(Some(id));
58            }
59        }
60        Ok(None)
61    }
62
63    /// Cancel scheduled refresh
64    ///
65    /// # Errors
66    ///
67    /// Returns `AuthError::Internal` if the mutex is poisoned.
68    pub fn cancel_refresh(&self, session_id: &str) -> std::result::Result<bool, AuthError> {
69        let mut queue = self.refresh_queue.lock().map_err(|_| AuthError::Internal {
70            message: "token refresh scheduler mutex poisoned".to_string(),
71        })?;
72        let len_before = queue.len();
73        queue.retain(|(id, _)| id != session_id);
74        Ok(queue.len() < len_before)
75    }
76}
77
78impl Default for TokenRefreshScheduler {
79    fn default() -> Self {
80        Self::new()
81    }
82}
83
84/// Callback trait for the token refresh worker to perform provider-specific
85/// token refresh and session updates.
86#[async_trait::async_trait]
87pub trait TokenRefresher: Send + Sync {
88    /// Refresh the token for the given session ID.
89    ///
90    /// Should look up the session, call the appropriate OAuth2 provider's
91    /// `refresh_token()`, update the stored session, and return the new expiry.
92    /// Returns `None` if the session no longer exists or has no refresh token.
93    async fn refresh_session(
94        &self,
95        session_id: &str,
96    ) -> std::result::Result<Option<DateTime<Utc>>, AuthError>;
97}
98
99/// Background worker that polls the `TokenRefreshScheduler` and refreshes
100/// expiring OAuth tokens.
101pub struct TokenRefreshWorker {
102    scheduler:     Arc<TokenRefreshScheduler>,
103    refresher:     Arc<dyn TokenRefresher>,
104    cancel_rx:     tokio::sync::watch::Receiver<bool>,
105    poll_interval: StdDuration,
106}
107
108impl TokenRefreshWorker {
109    /// Create a new token refresh worker.
110    ///
111    /// Returns the worker and a sender to trigger cancellation (send `true` to
112    /// stop).
113    pub fn new(
114        scheduler: Arc<TokenRefreshScheduler>,
115        refresher: Arc<dyn TokenRefresher>,
116        poll_interval: StdDuration,
117    ) -> (Self, tokio::sync::watch::Sender<bool>) {
118        let (cancel_tx, cancel_rx) = tokio::sync::watch::channel(false);
119        (
120            Self {
121                scheduler,
122                refresher,
123                cancel_rx,
124                poll_interval,
125            },
126            cancel_tx,
127        )
128    }
129
130    /// Run the refresh loop until cancelled.
131    pub async fn run(mut self) {
132        tracing::info!(
133            interval_secs = self.poll_interval.as_secs(),
134            "Token refresh worker started"
135        );
136        loop {
137            tokio::select! {
138                result = self.cancel_rx.changed() => {
139                    if result.is_err() || *self.cancel_rx.borrow() {
140                        tracing::info!("Token refresh worker stopped");
141                        break;
142                    }
143                },
144                () = tokio::time::sleep(self.poll_interval) => {
145                    self.process_due_refreshes().await;
146                }
147            }
148        }
149    }
150
151    async fn process_due_refreshes(&self) {
152        while let Ok(Some(session_id)) = self.scheduler.get_next_refresh() {
153            match self.refresher.refresh_session(&session_id).await {
154                Ok(Some(new_expiry)) => {
155                    // Re-schedule at 80% of the remaining time
156                    let remaining = new_expiry - Utc::now();
157                    #[allow(clippy::cast_precision_loss, clippy::cast_possible_truncation)]
158                    // Reason: intentional 80% f64 scaling; sub-second precision loss acceptable for
159                    // scheduling
160                    let next_refresh_secs = (remaining.num_seconds() as f64 * 0.8) as i64;
161                    let next_refresh = Utc::now() + Duration::seconds(next_refresh_secs);
162                    if let Err(e) =
163                        self.scheduler.schedule_refresh(session_id.clone(), next_refresh)
164                    {
165                        tracing::warn!(
166                            session_id = %session_id,
167                            error = %e,
168                            "Failed to re-schedule token refresh"
169                        );
170                    }
171                },
172                Ok(None) => {
173                    tracing::debug!(
174                        session_id = %session_id,
175                        "Session no longer exists, skipping refresh"
176                    );
177                },
178                Err(e) => {
179                    tracing::warn!(
180                        session_id = %session_id,
181                        error = %e,
182                        "Token refresh failed"
183                    );
184                },
185            }
186        }
187    }
188}