Skip to main content

prism_mcp_rs/auth/
token.rs

1//! Token Management and Refresh
2//!
3//! Module handles access token management, including automatic refresh
4//! when tokens expire
5
6use reqwest::Client;
7use std::sync::Arc;
8use tokio::sync::RwLock;
9use url::Url;
10
11use crate::auth::errors::AuthError;
12use crate::auth::types::*;
13use crate::core::error::{McpError, McpResult};
14
15/// Token manager for handling access and refresh tokens
16#[derive(Debug, Clone)]
17pub struct TokenManager {
18    context: Arc<RwLock<AuthorizationContext>>,
19    http_client: Client,
20}
21
22impl TokenManager {
23    /// Create a new token manager
24    pub fn new(resource: String) -> Self {
25        Self {
26            context: Arc::new(RwLock::new(AuthorizationContext::new(resource))),
27            http_client: Client::new(),
28        }
29    }
30
31    /// Create with existing context
32    pub fn with_context(context: AuthorizationContext) -> Self {
33        Self {
34            context: Arc::new(RwLock::new(context)),
35            http_client: Client::new(),
36        }
37    }
38
39    /// Get current access token if valid
40    pub async fn get_valid_token(&self) -> Option<String> {
41        let ctx = self.context.read().await;
42        if ctx.has_valid_token() {
43            ctx.access_token.clone()
44        } else {
45            None
46        }
47    }
48
49    /// Set tokens from a token response
50    pub async fn set_tokens(&self, response: TokenResponse) -> McpResult<()> {
51        let mut ctx = self.context.write().await;
52
53        ctx.access_token = Some(response.access_token);
54        ctx.refresh_token = response.refresh_token;
55
56        // Calculate expiration time
57        if let Some(expires_in) = response.expires_in {
58            let now = std::time::SystemTime::now()
59                .duration_since(std::time::UNIX_EPOCH)
60                .unwrap()
61                .as_secs();
62            ctx.expires_at = Some(now + expires_in);
63        }
64
65        Ok(())
66    }
67
68    /// Refresh the access token using the refresh token
69    pub async fn refresh_token(&self) -> McpResult<String> {
70        let (refresh_token, token_endpoint, client_id, client_secret, resource) = {
71            let ctx = self.context.read().await;
72
73            let refresh_token = ctx
74                .refresh_token
75                .as_ref()
76                .ok_or_else(|| McpError::Auth("No refresh token available".to_string()))?
77                .clone();
78
79            let token_endpoint = ctx
80                .auth_server_metadata
81                .as_ref()
82                .ok_or_else(|| McpError::Auth("No authorization server metadata".to_string()))?
83                .token_endpoint
84                .clone();
85
86            let client_id = ctx
87                .client_registration
88                .as_ref()
89                .map(|r| r.client_id.clone());
90
91            let client_secret = ctx
92                .client_registration
93                .as_ref()
94                .and_then(|r| r.client_secret.clone());
95
96            (
97                refresh_token,
98                token_endpoint,
99                client_id,
100                client_secret,
101                ctx.resource.clone(),
102            )
103        };
104
105        // Build refresh token request
106        let mut params = vec![
107            ("grant_type".to_string(), "refresh_token".to_string()),
108            ("refresh_token".to_string(), refresh_token),
109            ("resource".to_string(), resource),
110        ];
111
112        if let Some(client_id) = client_id {
113            params.push(("client_id".to_string(), client_id));
114        }
115
116        if let Some(client_secret) = client_secret {
117            params.push(("client_secret".to_string(), client_secret));
118        }
119
120        // Send token request
121        let response = self
122            .http_client
123            .post(&token_endpoint)
124            .form(&params)
125            .send()
126            .await
127            .map_err(|e| McpError::Auth(format!("Failed to refresh token: {e}")))?;
128
129        if !response.status().is_success() {
130            let error_text = response.text().await.unwrap_or_default();
131            if let Ok(oauth_error) = serde_json::from_str::<OAuth2Error>(&error_text) {
132                return Err(AuthError::OAuthError {
133                    error: oauth_error.error,
134                    description: oauth_error.error_description,
135                    uri: oauth_error.error_uri,
136                }
137                .into());
138            }
139            return Err(McpError::Auth(format!(
140                "Token refresh failed: {error_text}"
141            )));
142        }
143
144        let token_response: TokenResponse = response
145            .json()
146            .await
147            .map_err(|e| McpError::Auth(format!("Invalid token response: {e}")))?;
148
149        // Update stored tokens
150        self.set_tokens(token_response.clone()).await?;
151
152        Ok(token_response.access_token)
153    }
154
155    /// Get or refresh access token
156    pub async fn get_or_refresh_token(&self) -> McpResult<String> {
157        // First check if we have a valid token
158        if let Some(token) = self.get_valid_token().await {
159            return Ok(token);
160        }
161
162        // Try to refresh
163        self.refresh_token().await
164    }
165
166    /// Exchange authorization code for tokens
167    pub async fn exchange_code(
168        &self,
169        code: String,
170        redirect_uri: String,
171        code_verifier: Option<String>,
172    ) -> McpResult<TokenResponse> {
173        let (token_endpoint, client_id, client_secret, resource) = {
174            let ctx = self.context.read().await;
175
176            let token_endpoint = ctx
177                .auth_server_metadata
178                .as_ref()
179                .ok_or_else(|| McpError::Auth("No authorization server metadata".to_string()))?
180                .token_endpoint
181                .clone();
182
183            let client_id = ctx
184                .client_registration
185                .as_ref()
186                .map(|r| r.client_id.clone());
187
188            let client_secret = ctx
189                .client_registration
190                .as_ref()
191                .and_then(|r| r.client_secret.clone());
192
193            (
194                token_endpoint,
195                client_id,
196                client_secret,
197                ctx.resource.clone(),
198            )
199        };
200
201        // Build token request
202        let mut params = vec![
203            ("grant_type".to_string(), "authorization_code".to_string()),
204            ("code".to_string(), code),
205            ("redirect_uri".to_string(), redirect_uri),
206            ("resource".to_string(), resource),
207        ];
208
209        if let Some(verifier) = code_verifier {
210            params.push(("code_verifier".to_string(), verifier));
211        }
212
213        if let Some(client_id) = client_id {
214            params.push(("client_id".to_string(), client_id));
215        }
216
217        if let Some(client_secret) = client_secret {
218            params.push(("client_secret".to_string(), client_secret));
219        }
220
221        // Send token request
222        let response = self
223            .http_client
224            .post(&token_endpoint)
225            .form(&params)
226            .send()
227            .await
228            .map_err(|e| McpError::Auth(format!("Failed to exchange code: {e}")))?;
229
230        if !response.status().is_success() {
231            let error_text = response.text().await.unwrap_or_default();
232            if let Ok(oauth_error) = serde_json::from_str::<OAuth2Error>(&error_text) {
233                return Err(AuthError::OAuthError {
234                    error: oauth_error.error,
235                    description: oauth_error.error_description,
236                    uri: oauth_error.error_uri,
237                }
238                .into());
239            }
240            return Err(McpError::Auth(format!(
241                "Code exchange failed: {error_text}"
242            )));
243        }
244
245        let token_response: TokenResponse = response
246            .json()
247            .await
248            .map_err(|e| McpError::Auth(format!("Invalid token response: {e}")))?;
249
250        // Store tokens
251        self.set_tokens(token_response.clone()).await?;
252
253        Ok(token_response)
254    }
255
256    /// Clear all tokens
257    pub async fn clear_tokens(&self) {
258        let mut ctx = self.context.write().await;
259        ctx.access_token = None;
260        ctx.refresh_token = None;
261        ctx.expires_at = None;
262    }
263
264    /// Get the authorization context
265    pub async fn get_context(&self) -> AuthorizationContext {
266        self.context.read().await.clone()
267    }
268
269    /// Update the authorization context
270    pub async fn update_context<F>(&self, f: F) -> McpResult<()>
271    where
272        F: FnOnce(&mut AuthorizationContext),
273    {
274        let mut ctx = self.context.write().await;
275        f(&mut ctx);
276        Ok(())
277    }
278}
279
280/// Build authorization URL for OAuth flow
281pub fn build_authorization_url(
282    auth_endpoint: &str,
283    client_id: &str,
284    redirect_uri: &str,
285    state: &str,
286    code_challenge: &str,
287    code_challenge_method: &str,
288    resource: &str,
289    scopes: &[String],
290) -> McpResult<String> {
291    let mut url = Url::parse(auth_endpoint)
292        .map_err(|e| McpError::Auth(format!("Invalid authorization endpoint: {e}")))?;
293
294    url.query_pairs_mut()
295        .append_pair("response_type", "code")
296        .append_pair("client_id", client_id)
297        .append_pair("redirect_uri", redirect_uri)
298        .append_pair("state", state)
299        .append_pair("code_challenge", code_challenge)
300        .append_pair("code_challenge_method", code_challenge_method)
301        .append_pair("resource", resource);
302
303    if !scopes.is_empty() {
304        url.query_pairs_mut()
305            .append_pair("scope", &scopes.join(" "));
306    }
307
308    Ok(url.to_string())
309}
310
311/// Parse authorization callback URL
312pub fn parse_callback_url(callback_url: &str) -> McpResult<CallbackParams> {
313    let url = Url::parse(callback_url)
314        .map_err(|e| McpError::Auth(format!("Invalid callback URL: {e}")))?;
315
316    let params: Vec<(String, String)> = url
317        .query_pairs()
318        .map(|(k, v)| (k.to_string(), v.to_string()))
319        .collect();
320
321    // Check for error response
322    if let Some(error) = crate::auth::errors::parse_oauth_error(&params) {
323        return Err(error.into());
324    }
325
326    // Extract code and state
327    let code = params
328        .iter()
329        .find(|(k, _)| k == "code")
330        .map(|(_, v)| v.clone())
331        .ok_or_else(|| McpError::Auth("No authorization code in callback".to_string()))?;
332
333    let state = params
334        .iter()
335        .find(|(k, _)| k == "state")
336        .map(|(_, v)| v.clone());
337
338    Ok(CallbackParams { code, state })
339}
340
341/// Parameters extracted from OAuth callback
342#[derive(Debug, Clone)]
343pub struct CallbackParams {
344    /// Authorization code
345    pub code: String,
346    /// State parameter (for CSRF protection)
347    pub state: Option<String>,
348}
349
350#[cfg(test)]
351mod tests {
352    use super::*;
353
354    #[test]
355    fn test_build_authorization_url() {
356        let url = build_authorization_url(
357            "https://auth.example.com/authorize",
358            "client123",
359            "http://localhost:8080/callback",
360            "random_state",
361            "challenge123",
362            "S256",
363            "https://mcp.example.com",
364            &["read".to_string(), "write".to_string()],
365        )
366        .unwrap();
367
368        assert!(url.contains("response_type=code"));
369        assert!(url.contains("client_id=client123"));
370        assert!(url.contains("redirect_uri="));
371        assert!(url.contains("state=random_state"));
372        assert!(url.contains("code_challenge=challenge123"));
373        assert!(url.contains("code_challenge_method=S256"));
374        assert!(url.contains("resource="));
375        assert!(url.contains("scope=read+write"));
376    }
377
378    #[test]
379    fn test_parse_callback_success() {
380        let callback = "http://localhost:8080/callback?code=auth123&state=random_state";
381        let params = parse_callback_url(callback).unwrap();
382
383        assert_eq!(params.code, "auth123");
384        assert_eq!(params.state, Some("random_state".to_string()));
385    }
386
387    #[test]
388    fn test_parse_callback_error() {
389        let callback = "http://localhost:8080/callback?error=access_denied&error_description=User+denied+access";
390        let result = parse_callback_url(callback);
391
392        assert!(result.is_err());
393        let error = result.unwrap_err();
394        assert!(error.to_string().contains("access_denied"));
395    }
396
397    #[tokio::test]
398    async fn test_token_manager() {
399        let manager = TokenManager::new("https://mcp.example.com".to_string());
400
401        // Initially no token
402        assert!(manager.get_valid_token().await.is_none());
403
404        // Set a token
405        let response = TokenResponse {
406            access_token: "token123".to_string(),
407            token_type: "Bearer".to_string(),
408            expires_in: Some(3600),
409            refresh_token: Some("refresh123".to_string()),
410            scope: None,
411            additional: Default::default(),
412        };
413
414        manager.set_tokens(response).await.unwrap();
415
416        // Now should have a valid token
417        assert_eq!(
418            manager.get_valid_token().await,
419            Some("token123".to_string())
420        );
421    }
422}