fraiseql_auth/oauth/
refresh.rs1use std::{sync::Arc, time::Duration as StdDuration};
4
5use chrono::{DateTime, Duration, Utc};
6
7use super::super::error::AuthError;
8
9#[derive(Debug, Clone)]
11pub struct TokenRefreshScheduler {
12 refresh_queue: Arc<std::sync::Mutex<Vec<(String, DateTime<Utc>)>>>,
16}
17
18impl TokenRefreshScheduler {
19 #[must_use]
21 pub fn new() -> Self {
22 Self {
23 refresh_queue: Arc::new(std::sync::Mutex::new(Vec::new())),
24 }
25 }
26
27 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 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 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#[async_trait::async_trait]
87pub trait TokenRefresher: Send + Sync {
88 async fn refresh_session(
94 &self,
95 session_id: &str,
96 ) -> std::result::Result<Option<DateTime<Utc>>, AuthError>;
97}
98
99pub 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 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 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 let remaining = new_expiry - Utc::now();
157 #[allow(clippy::cast_precision_loss, clippy::cast_possible_truncation)]
158 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}