prism_mcp_rs/auth/
client.rs1use reqwest::Client;
8use std::sync::Arc;
9use tokio::sync::RwLock;
10
11use crate::auth::{
12 discovery::{validate_auth_server_for_mcp, DiscoveryClient},
13 errors::AuthError,
14 pkce::{select_challenge_method, PkceParams},
15 token::{build_authorization_url, parse_callback_url, TokenManager},
16 types::*,
17 AuthConfig,
18};
19use crate::core::error::{McpError, McpResult};
20
21pub struct AuthorizationClient {
23 config: AuthConfig,
24 http_client: Client,
25 token_manager: TokenManager,
26 discovery_client: DiscoveryClient,
27 state: Arc<RwLock<AuthState>>,
28}
29
30#[derive(Debug, Clone)]
32struct AuthState {
33 pkce: Option<PkceParams>,
35 state: Option<String>,
37 resource_metadata: Option<ProtectedResourceMetadata>,
39 auth_server_metadata: Option<AuthorizationServerMetadata>,
41 client_registration: Option<ClientRegistrationResponse>,
43}
44
45impl AuthorizationClient {
46 pub fn new(config: AuthConfig, resource_url: String) -> Self {
48 let http_client = Client::new();
49 Self {
50 config,
51 http_client: http_client.clone(),
52 token_manager: TokenManager::new(resource_url),
53 discovery_client: DiscoveryClient::with_client(http_client),
54 state: Arc::new(RwLock::new(AuthState {
55 pkce: None,
56 state: None,
57 resource_metadata: None,
58 auth_server_metadata: None,
59 client_registration: None,
60 })),
61 }
62 }
63
64 pub async fn handle_unauthorized(&self, www_authenticate: &str) -> McpResult<String> {
69 let metadata_url = self
71 .discovery_client
72 .parse_www_authenticate(www_authenticate)?;
73
74 let resource_url = self.token_manager.get_context().await.resource;
75
76 let resource_metadata = match metadata_url {
78 Some(url) => self.discovery_client.discover_from_resource(&url).await?,
79 None => {
80 self.discovery_client
81 .discover_from_resource(&resource_url)
82 .await?
83 }
84 };
85
86 let auth_server_url = resource_metadata
88 .authorization_servers
89 .first()
90 .ok_or_else(|| McpError::Auth("No authorization servers available".to_string()))?
91 .clone();
92
93 let auth_metadata = self
95 .discovery_client
96 .discover_auth_server(&auth_server_url)
97 .await?;
98
99 validate_auth_server_for_mcp(&auth_metadata)?;
101
102 {
104 let mut state = self.state.write().await;
105 state.resource_metadata = Some(resource_metadata.clone());
106 state.auth_server_metadata = Some(auth_metadata.clone());
107 }
108
109 self.token_manager
111 .update_context(|ctx| {
112 ctx.resource_metadata = Some(resource_metadata);
113 ctx.auth_server_metadata = Some(auth_metadata.clone());
114 })
115 .await?;
116
117 if self.config.client_id.is_none() && self.config.enable_dynamic_registration {
119 self.register_client(&auth_metadata).await?;
120 }
121
122 self.start_authorization_flow(&auth_metadata).await
124 }
125
126 async fn register_client(&self, auth_metadata: &AuthorizationServerMetadata) -> McpResult<()> {
128 let registration_endpoint =
129 auth_metadata
130 .registration_endpoint
131 .as_ref()
132 .ok_or_else(|| {
133 McpError::Auth(
134 "Authorization server does not support dynamic registration".to_string(),
135 )
136 })?;
137
138 let request = ClientRegistrationRequest {
139 redirect_uris: vec![self.config.redirect_uri.clone()],
140 client_name: Some("MCP Client".to_string()),
141 grant_types: Some(vec![
142 "authorization_code".to_string(),
143 "refresh_token".to_string(),
144 ]),
145 response_types: Some(vec!["code".to_string()]),
146 token_endpoint_auth_method: Some("client_secret_basic".to_string()),
147 scope: if self.config.scopes.is_empty() {
148 None
149 } else {
150 Some(self.config.scopes.join(" "))
151 },
152 software_id: Some("mcp-rust-sdk".to_string()),
153 software_version: Some(env!("CARGO_PKG_VERSION").to_string()),
154 client_uri: None,
155 logo_uri: None,
156 };
157
158 let response = self
159 .http_client
160 .post(registration_endpoint)
161 .json(&request)
162 .send()
163 .await
164 .map_err(|e| McpError::Auth(format!("Registration request failed: {e}")))?;
165
166 if !response.status().is_success() {
167 let error_text = response.text().await.unwrap_or_default();
168 return Err(McpError::Auth(format!(
169 "Client registration failed: {error_text}"
170 )));
171 }
172
173 let registration: ClientRegistrationResponse = response
174 .json()
175 .await
176 .map_err(|e| McpError::Auth(format!("Invalid registration response: {e}")))?;
177
178 {
180 let mut state = self.state.write().await;
181 state.client_registration = Some(registration.clone());
182 }
183
184 self.token_manager
186 .update_context(|ctx| {
187 ctx.client_registration = Some(registration);
188 })
189 .await?;
190
191 Ok(())
192 }
193
194 pub async fn start_authorization_flow(
196 &self,
197 auth_metadata: &AuthorizationServerMetadata,
198 ) -> McpResult<String> {
199 let pkce_method = select_challenge_method(auth_metadata)?;
201 let pkce = PkceParams::with_method(pkce_method);
202
203 let state = self.config.generate_state();
205
206 {
208 let mut auth_state = self.state.write().await;
209 auth_state.pkce = Some(pkce.clone());
210 auth_state.state = Some(state.clone());
211 }
212
213 let client_id = if let Some(ref id) = self.config.client_id {
215 id.clone()
216 } else {
217 let auth_state = self.state.read().await;
218 auth_state
219 .client_registration
220 .as_ref()
221 .map(|r| r.client_id.clone())
222 .ok_or_else(|| McpError::Auth("No client ID available".to_string()))?
223 };
224
225 let resource = self.token_manager.get_context().await.resource;
227
228 let auth_url = build_authorization_url(
230 &auth_metadata.authorization_endpoint,
231 &client_id,
232 &self.config.redirect_uri,
233 &state,
234 &pkce.challenge,
235 pkce.method.as_str(),
236 &resource,
237 &self.config.scopes,
238 )?;
239
240 Ok(auth_url)
241 }
242
243 pub async fn handle_callback(&self, callback_url: &str) -> McpResult<String> {
245 let params = parse_callback_url(callback_url)?;
246
247 let stored_state = {
249 let auth_state = self.state.read().await;
250 auth_state.state.clone()
251 };
252
253 if let Some(expected_state) = stored_state {
254 if params.state.as_ref() != Some(&expected_state) {
255 return Err(AuthError::StateMismatch.into());
256 }
257 }
258
259 let pkce_verifier = {
261 let auth_state = self.state.read().await;
262 auth_state.pkce.as_ref().map(|p| p.verifier.clone())
263 };
264
265 let token_response = self
267 .token_manager
268 .exchange_code(params.code, self.config.redirect_uri.clone(), pkce_verifier)
269 .await?;
270
271 {
273 let mut auth_state = self.state.write().await;
274 auth_state.pkce = None;
275 auth_state.state = None;
276 }
277
278 Ok(token_response.access_token)
279 }
280
281 pub async fn get_token(&self) -> McpResult<String> {
283 self.token_manager.get_or_refresh_token().await
284 }
285
286 pub async fn logout(&self) {
288 self.token_manager.clear_tokens().await;
289
290 let mut state = self.state.write().await;
291 *state = AuthState {
292 pkce: None,
293 state: None,
294 resource_metadata: None,
295 auth_server_metadata: None,
296 client_registration: None,
297 };
298 }
299
300 pub fn token_manager(&self) -> &TokenManager {
302 &self.token_manager
303 }
304
305 pub async fn is_authenticated(&self) -> bool {
307 self.token_manager.get_valid_token().await.is_some()
308 }
309}
310
311pub fn add_auth_header(headers: &mut reqwest::header::HeaderMap, token: &str) {
313 use reqwest::header::{HeaderValue, AUTHORIZATION};
314
315 let value = format!("Bearer {token}");
316 if let Ok(header_value) = HeaderValue::from_str(&value) {
317 headers.insert(AUTHORIZATION, header_value);
318 }
319}
320
321pub fn extract_bearer_token(auth_header: &str) -> Option<String> {
323 auth_header.strip_prefix("Bearer ").map(|s| s.to_string())
324}
325
326#[cfg(test)]
327mod tests {
328 use super::*;
329
330 #[test]
331 fn test_add_auth_header() {
332 let mut headers = reqwest::header::HeaderMap::new();
333 add_auth_header(&mut headers, "test_token");
334
335 let auth = headers.get(reqwest::header::AUTHORIZATION).unwrap();
336 assert_eq!(auth.to_str().unwrap(), "Bearer test_token");
337 }
338
339 #[test]
340 fn test_extract_bearer_token() {
341 assert_eq!(
342 extract_bearer_token("Bearer abc123"),
343 Some("abc123".to_string())
344 );
345
346 assert_eq!(extract_bearer_token("Basic abc123"), None);
347
348 assert_eq!(extract_bearer_token("Invalid"), None);
349 }
350
351 #[tokio::test]
352 async fn test_authorization_client_creation() {
353 let config = AuthConfig::new()
354 .with_auth(true)
355 .with_redirect_uri("http://localhost:8080/callback".to_string())
356 .with_scopes(vec!["read".to_string()]);
357
358 let client = AuthorizationClient::new(config, "https://mcp.example.com".to_string());
359
360 assert!(!client.is_authenticated().await);
362 }
363}