Skip to main content

prism_mcp_rs/auth/
client.rs

1//! OAuth 2.1 Authorization Client
2//!
3//! Module provides the main authorization client for MCP,
4//! handling the full OAuth 2.1 flow including discovery, registration,
5//! and token management
6
7use 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
21/// OAuth 2.1 Authorization Client for MCP
22pub 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/// Internal state for authorization flow
31#[derive(Debug, Clone)]
32struct AuthState {
33    /// Current PKCE parameters
34    pkce: Option<PkceParams>,
35    /// Current state parameter
36    state: Option<String>,
37    /// Resource metadata
38    resource_metadata: Option<ProtectedResourceMetadata>,
39    /// Authorization server metadata
40    auth_server_metadata: Option<AuthorizationServerMetadata>,
41    /// Client registration
42    client_registration: Option<ClientRegistrationResponse>,
43}
44
45impl AuthorizationClient {
46    /// Create a new authorization client
47    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    /// Handle authorization challenge response and initiate authorization.
65    ///
66    /// Works for both 401 and 403 challenge responses as long as a
67    /// `WWW-Authenticate` header is present.
68    pub async fn handle_unauthorized(&self, www_authenticate: &str) -> McpResult<String> {
69        // Parse WWW-Authenticate header
70        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        // Discover resource metadata
77        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        // Select authorization server (use first for now)
87        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        // Discover authorization server metadata
94        let auth_metadata = self
95            .discovery_client
96            .discover_auth_server(&auth_server_url)
97            .await?;
98
99        // Validate for MCP requirements
100        validate_auth_server_for_mcp(&auth_metadata)?;
101
102        // Store metadata
103        {
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        // Update token manager context
110        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        // Perform dynamic registration if needed
118        if self.config.client_id.is_none() && self.config.enable_dynamic_registration {
119            self.register_client(&auth_metadata).await?;
120        }
121
122        // Start authorization flow
123        self.start_authorization_flow(&auth_metadata).await
124    }
125
126    /// Perform dynamic client registration
127    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        // Store registration
179        {
180            let mut state = self.state.write().await;
181            state.client_registration = Some(registration.clone());
182        }
183
184        // Update token manager
185        self.token_manager
186            .update_context(|ctx| {
187                ctx.client_registration = Some(registration);
188            })
189            .await?;
190
191        Ok(())
192    }
193
194    /// Start the authorization flow
195    pub async fn start_authorization_flow(
196        &self,
197        auth_metadata: &AuthorizationServerMetadata,
198    ) -> McpResult<String> {
199        // Select PKCE method
200        let pkce_method = select_challenge_method(auth_metadata)?;
201        let pkce = PkceParams::with_method(pkce_method);
202
203        // Generate state
204        let state = self.config.generate_state();
205
206        // Store PKCE and state
207        {
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        // Get client ID
214        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        // Get resource URL
226        let resource = self.token_manager.get_context().await.resource;
227
228        // Build authorization URL
229        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    /// Handle authorization callback
244    pub async fn handle_callback(&self, callback_url: &str) -> McpResult<String> {
245        let params = parse_callback_url(callback_url)?;
246
247        // Verify state
248        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        // Get PKCE verifier
260        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        // Exchange code for tokens
266        let token_response = self
267            .token_manager
268            .exchange_code(params.code, self.config.redirect_uri.clone(), pkce_verifier)
269            .await?;
270
271        // Clear temporary state
272        {
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    /// Get current access token (refreshing if needed)
282    pub async fn get_token(&self) -> McpResult<String> {
283        self.token_manager.get_or_refresh_token().await
284    }
285
286    /// Clear all tokens and state
287    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    /// Get the token manager
301    pub fn token_manager(&self) -> &TokenManager {
302        &self.token_manager
303    }
304
305    /// Check if we have a valid token
306    pub async fn is_authenticated(&self) -> bool {
307        self.token_manager.get_valid_token().await.is_some()
308    }
309}
310
311/// Helper to add authorization header to HTTP requests
312pub 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
321/// Extract bearer token from Authorization header
322pub 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        // Initially not authenticated
361        assert!(!client.is_authenticated().await);
362    }
363}