1use 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#[derive(Debug, Clone)]
17pub struct TokenManager {
18 context: Arc<RwLock<AuthorizationContext>>,
19 http_client: Client,
20}
21
22impl TokenManager {
23 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 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 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 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 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 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 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 let response = self
122 .http_client
123 .post(&token_endpoint)
124 .form(¶ms)
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 self.set_tokens(token_response.clone()).await?;
151
152 Ok(token_response.access_token)
153 }
154
155 pub async fn get_or_refresh_token(&self) -> McpResult<String> {
157 if let Some(token) = self.get_valid_token().await {
159 return Ok(token);
160 }
161
162 self.refresh_token().await
164 }
165
166 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 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 let response = self
223 .http_client
224 .post(&token_endpoint)
225 .form(¶ms)
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 self.set_tokens(token_response.clone()).await?;
252
253 Ok(token_response)
254 }
255
256 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 pub async fn get_context(&self) -> AuthorizationContext {
266 self.context.read().await.clone()
267 }
268
269 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
280pub 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
311pub 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 if let Some(error) = crate::auth::errors::parse_oauth_error(¶ms) {
323 return Err(error.into());
324 }
325
326 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#[derive(Debug, Clone)]
343pub struct CallbackParams {
344 pub code: String,
346 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 assert!(manager.get_valid_token().await.is_none());
403
404 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 assert_eq!(
418 manager.get_valid_token().await,
419 Some("token123".to_string())
420 );
421 }
422}