Skip to main content

mcp_gmailcal/
auth.rs

1use crate::config::{get_token_expiry_buffer_seconds, get_token_expiry_seconds, get_token_refresh_threshold_seconds, Config, OAUTH_TOKEN_URL};
2use crate::errors::{GmailApiError, GmailResult};
3use log::{debug, error, info, warn};
4use reqwest::Client;
5use serde::Deserialize;
6use std::time::{Duration, SystemTime};
7use std::cmp::min;
8
9// Alias for backward compatibility within this module
10type Result<T> = GmailResult<T>;
11
12// Token response for OAuth2
13#[derive(Debug, Deserialize)]
14struct TokenResponse {
15    access_token: String,
16    expires_in: u64,
17    #[serde(default)]
18    #[allow(dead_code)]
19    token_type: String,
20}
21
22use crate::token_cache::{TokenCache, TokenCacheConfig};
23
24// OAuth token manager
25#[derive(Debug, Clone)]
26pub struct TokenManager {
27    access_token: String,
28    expiry: SystemTime,
29    refresh_token: String,
30    client_id: String,
31    client_secret: String,
32    retry_count: u8,
33    max_retries: u8,
34    base_retry_delay_ms: u64,
35    cache: Option<TokenCache>,
36}
37
38impl TokenManager {
39    // Test-only accessors - used by the test extension trait
40    #[cfg(test)]
41    pub(crate) fn get_access_token(&self) -> &str {
42        &self.access_token
43    }
44    
45    #[cfg(test)]
46    pub(crate) fn get_expiry(&self) -> SystemTime {
47        self.expiry
48    }
49    
50    #[cfg(test)]
51    pub(crate) fn get_refresh_threshold(&self) -> u64 {
52        get_token_refresh_threshold_seconds()
53    }
54    
55    #[cfg(test)]
56    pub(crate) fn get_expiry_buffer(&self) -> u64 {
57        get_token_expiry_buffer_seconds()
58    }
59    
60    #[cfg(test)]
61    pub(crate) fn get_retry_count(&self) -> u8 {
62        self.retry_count
63    }
64    
65    #[cfg(test)]
66    pub(crate) fn get_max_retries(&self) -> u8 {
67        self.max_retries
68    }
69    
70    #[cfg(test)]
71    pub(crate) fn get_base_retry_delay_ms(&self) -> u64 {
72        self.base_retry_delay_ms
73    }
74    
75    #[cfg(test)]
76    pub(crate) fn set_access_token(&mut self, token: String) {
77        self.access_token = token;
78    }
79    
80    #[cfg(test)]
81    pub(crate) fn set_expiry(&mut self, expiry: SystemTime) {
82        self.expiry = expiry;
83    }
84    
85    #[cfg(test)]
86    pub(crate) fn increment_retry_count(&mut self) {
87        self.retry_count += 1;
88    }
89    
90    #[cfg(test)]
91    pub(crate) fn set_retry_count(&mut self, count: u8) {
92        self.retry_count = count;
93    }
94    
95    #[cfg(test)]
96    pub(crate) fn set_max_retries(&mut self, max: u8) {
97        self.max_retries = max;
98    }
99    
100    #[cfg(test)]
101    pub(crate) fn reset_retry_count(&mut self) {
102        self.retry_count = 0;
103    }
104    pub fn new(config: &Config) -> Self {
105        // Initialize token cache if enabled
106        let cache = match TokenCacheConfig::from_env() {
107            Ok(cache_config) => {
108                if cache_config.enabled {
109                    match TokenCache::new(cache_config) {
110                        Ok(cache) => {
111                            debug!("Token cache initialized successfully");
112                            Some(cache)
113                        }
114                        Err(e) => {
115                            error!("Failed to initialize token cache: {}", e);
116                            None
117                        }
118                    }
119                } else {
120                    debug!("Token caching is disabled");
121                    None
122                }
123            }
124            Err(e) => {
125                error!("Failed to load token cache configuration: {}", e);
126                None
127            }
128        };
129
130        // Try to load token from cache first
131        let (access_token, expiry, loaded_from_cache) = if let Some(cache) = &cache {
132            match cache.load_token() {
133                Ok(Some(cached_token)) => {
134                    if cache.is_token_valid(&cached_token) {
135                        info!("Loaded valid token from cache");
136                        let expiry_system_time = SystemTime::UNIX_EPOCH + Duration::from_secs(cached_token.expiry_timestamp);
137                        (cached_token.access_token, expiry_system_time, true)
138                    } else {
139                        debug!("Cached token exists but is expired or nearly expired");
140                        let default_token = config.access_token.clone().unwrap_or_default();
141                        let default_expiry = if config.access_token.is_some() {
142                            SystemTime::now() + Duration::from_secs(get_token_expiry_seconds())
143                        } else {
144                            SystemTime::now() // Force refresh
145                        };
146                        (default_token, default_expiry, false)
147                    }
148                }
149                Ok(None) => {
150                    debug!("No cached token found");
151                    let default_token = config.access_token.clone().unwrap_or_default();
152                    let default_expiry = if config.access_token.is_some() {
153                        SystemTime::now() + Duration::from_secs(get_token_expiry_seconds())
154                    } else {
155                        SystemTime::now() // Force refresh
156                    };
157                    (default_token, default_expiry, false)
158                }
159                Err(e) => {
160                    warn!("Error loading token from cache: {}", e);
161                    let default_token = config.access_token.clone().unwrap_or_default();
162                    let default_expiry = if config.access_token.is_some() {
163                        SystemTime::now() + Duration::from_secs(get_token_expiry_seconds())
164                    } else {
165                        SystemTime::now() // Force refresh
166                    };
167                    (default_token, default_expiry, false)
168                }
169            }
170        } else {
171            // No cache, use config values
172            let default_token = config.access_token.clone().unwrap_or_default();
173            let default_expiry = if config.access_token.is_some() {
174                SystemTime::now() + Duration::from_secs(get_token_expiry_seconds())
175            } else {
176                SystemTime::now() // Force refresh
177            };
178            (default_token, default_expiry, false)
179        };
180
181        debug!(
182            "Creating TokenManager with refresh threshold: {}s, expiry buffer: {}s", 
183            config.token_refresh_threshold, config.token_expiry_buffer
184        );
185        
186        if loaded_from_cache {
187            debug!("Using cached token, expires at {:?}", expiry);
188        } else if !access_token.is_empty() {
189            debug!("Using token from config, expires at {:?}", expiry);
190        } else {
191            debug!("No valid token available, will refresh on first use");
192        }
193
194        Self {
195            access_token,
196            expiry,
197            refresh_token: config.refresh_token.clone(),
198            client_id: config.client_id.clone(),
199            client_secret: config.client_secret.clone(),
200            retry_count: 0,
201            max_retries: 5,  // Default maximum retries
202            base_retry_delay_ms: 1000, // Start with 1 second delay
203            cache,
204        }
205    }
206
207    // Calculate time until token expiry in seconds (can be negative if expired)
208    fn time_until_expiry(&self) -> i64 {
209        match self.expiry.duration_since(SystemTime::now()) {
210            Ok(duration) => duration.as_secs() as i64,
211            Err(_) => -1, // Token has expired
212        }
213    }
214    
215    // Check if token needs refresh based on refresh threshold
216    fn needs_refresh(&self) -> bool {
217        let refresh_threshold = get_token_refresh_threshold_seconds();
218        let seconds_until_expiry = self.time_until_expiry();
219        
220        if seconds_until_expiry < 0 {
221            debug!("Token has expired");
222            return true;
223        }
224        
225        if seconds_until_expiry < refresh_threshold as i64 {
226            debug!("Token will expire in {} seconds (refresh threshold: {} seconds), needs refresh", 
227                   seconds_until_expiry, refresh_threshold);
228            return true;
229        }
230        
231        debug!("Token valid for {} more seconds (refresh threshold: {} seconds), no refresh needed", 
232               seconds_until_expiry, refresh_threshold);
233        false
234    }
235    
236    // Reset retry counter after successful operation - accessible both from tests and internal code
237    #[cfg(not(test))]
238    fn reset_retry_count(&mut self) {
239        if self.retry_count > 0 {
240            debug!("Resetting retry count from {} to 0", self.retry_count);
241            self.retry_count = 0;
242        }
243    }
244    
245    // Calculate exponential backoff delay based on retry count
246    fn get_backoff_delay(&self) -> Duration {
247        if self.retry_count == 0 {
248            return Duration::from_millis(0);
249        }
250        
251        // Calculate delay with exponential backoff: base_delay * 2^(retry_count-1)
252        // Example: 1000ms base -> 1s, 2s, 4s, 8s, 16s
253        let exponent = min(self.retry_count - 1, 16); // Prevent potential overflow
254        let delay_ms = self.base_retry_delay_ms * (1u64 << exponent);
255        
256        // Cap maximum delay at 64 seconds
257        let capped_delay_ms = min(delay_ms, 64_000);
258        
259        debug!("Backoff delay for retry {}: {}ms", self.retry_count, capped_delay_ms);
260        Duration::from_millis(capped_delay_ms)
261    }
262
263    pub async fn get_token(&mut self, client: &Client) -> Result<String> {
264        // Log token expiration details
265        let seconds_until_expiry = self.time_until_expiry();
266        let has_token = !self.access_token.is_empty();
267        
268        debug!(
269            "Token status check - have token: {}, expires in: {} seconds",
270            has_token,
271            seconds_until_expiry
272        );
273
274        // Check if current token is still valid and not near expiration
275        if has_token && !self.needs_refresh() {
276            debug!("Using existing token, not near expiration");
277            return Ok(self.access_token.clone());
278        }
279
280        // Token is missing, expired, or near expiration
281        if has_token {
282            debug!("OAuth token expiring soon, proactively refreshing");
283        } else {
284            debug!("OAuth token not set, obtaining new token");
285        }
286
287        // Apply exponential backoff if retrying
288        if self.retry_count > 0 {
289            let backoff_delay = self.get_backoff_delay();
290            warn!(
291                "Applying exponential backoff delay of {}ms for retry attempt {}",
292                backoff_delay.as_millis(),
293                self.retry_count
294            );
295            tokio::time::sleep(backoff_delay).await;
296        }
297
298        // Check if we've exceeded max retries
299        if self.retry_count >= self.max_retries && self.retry_count > 0 {
300            error!(
301                "Maximum retry attempts ({}) exceeded for token refresh",
302                self.max_retries
303            );
304            return Err(GmailApiError::AuthError(
305                "Maximum retry attempts exceeded for token refresh".to_string(),
306            ));
307        }
308
309        // Increment retry counter
310        self.retry_count += 1;
311        if self.retry_count > 1 {
312            debug!("Retry attempt {} of {}", self.retry_count, self.max_retries);
313        }
314
315        // Prepare token refresh request
316        let params = [
317            ("client_id", self.client_id.as_str()),
318            ("client_secret", self.client_secret.as_str()),
319            ("refresh_token", self.refresh_token.as_str()),
320            ("grant_type", "refresh_token"),
321        ];
322
323        // Log request details for troubleshooting (but hide credentials)
324        debug!("Requesting token from {}", OAUTH_TOKEN_URL);
325        // Securely log truncated credential information - never log full credentials
326        if log::log_enabled!(log::Level::Debug) {
327            let client_id_trunc = if self.client_id.len() > 8 {
328                format!(
329                    "{}...{}",
330                    &self.client_id[..4],
331                    &self.client_id[self.client_id.len().saturating_sub(4)..]
332                )
333            } else {
334                "<short-id>".to_string()
335            };
336
337            let refresh_token_trunc = if self.refresh_token.len() > 8 {
338                format!("{}...", &self.refresh_token[..4])
339            } else {
340                "<short-token>".to_string()
341            };
342
343            debug!("Using client_id: {} (truncated)", client_id_trunc);
344            debug!(
345                "Using refresh_token starting with: {} (truncated)",
346                refresh_token_trunc
347            );
348        }
349
350        // Send token refresh request
351        let response = match client
352            .post(OAUTH_TOKEN_URL)
353            .form(&params)
354            .send()
355            .await
356        {
357            Ok(resp) => resp,
358            Err(e) => {
359                warn!("Network error during token refresh: {}", e);
360                return Err(GmailApiError::NetworkError(e.to_string()));
361            }
362        };
363
364        let status = response.status();
365        debug!("Token response status: {}", status);
366
367        if !status.is_success() {
368            let error_text = response
369                .text()
370                .await
371                .unwrap_or_else(|_| "<no response body>".to_string());
372
373            error!(
374                "Token refresh failed. Status: {}, Error: {}",
375                status, error_text
376            );
377            
378            // Some errors shouldn't be retried (e.g., invalid_grant)
379            if error_text.contains("invalid_grant") {
380                error!("Invalid grant error detected, not retrying");
381                self.retry_count = self.max_retries; // Prevent further retries
382            }
383            
384            return Err(GmailApiError::AuthError(format!(
385                "Failed to refresh token. Status: {}, Error: {}",
386                status, error_text
387            )));
388        }
389
390        let response_text = match response.text().await {
391            Ok(text) => text,
392            Err(e) => {
393                warn!("Failed to get token response text: {}", e);
394                return Err(GmailApiError::ApiError(format!(
395                    "Failed to get token response: {}",
396                    e
397                )));
398            }
399        };
400
401        debug!("Token response received, parsing JSON");
402
403        let token_data: TokenResponse = match serde_json::from_str(&response_text) {
404            Ok(data) => data,
405            Err(e) => {
406                error!(
407                    "Failed to parse token response: {}. Response: {}",
408                    e, response_text
409                );
410                return Err(GmailApiError::ApiError(format!(
411                    "Failed to parse token response: {}",
412                    e
413                )));
414            }
415        };
416
417        // Update token and expiry
418        self.access_token = token_data.access_token.clone();
419        
420        // Apply buffer to expiry time from config
421        let buffer = get_token_expiry_buffer_seconds();
422        let expires_in = token_data.expires_in.saturating_sub(buffer);
423        self.expiry = SystemTime::now() + Duration::from_secs(expires_in);
424
425        // Save token to cache if enabled
426        if let Some(cache) = &self.cache {
427            match cache.save_token(&self.access_token, &self.refresh_token, self.expiry) {
428                Ok(_) => debug!("Token successfully saved to cache"),
429                Err(e) => warn!("Failed to save token to cache: {}", e),
430            }
431        }
432
433        // Reset retry counter after success
434        self.reset_retry_count();
435
436        // Calculate when we'll need to refresh this token
437        let refresh_threshold = get_token_refresh_threshold_seconds();
438        let effective_lifetime = expires_in.saturating_sub(refresh_threshold);
439
440        info!(
441            "Token refreshed successfully, valid for {} seconds (with {}s buffer)",
442            expires_in,
443            buffer
444        );
445        debug!(
446            "Token will be refreshed after {} seconds (threshold: {}s)",
447            effective_lifetime,
448            refresh_threshold
449        );
450        
451        // Securely log truncated token - never log the full token
452        if log::log_enabled!(log::Level::Debug) {
453            let token_trunc = if self.access_token.len() > 10 {
454                format!(
455                    "{}...{}",
456                    &self.access_token[..4],
457                    &self.access_token[self.access_token.len().saturating_sub(4)..]
458                )
459            } else {
460                "<short-token>".to_string()
461            };
462            debug!("Token (truncated): {}", token_trunc);
463        };
464
465        Ok(self.access_token.clone())
466    }
467}