Skip to main content

prism_mcp_rs/auth/
types.rs

1//! OAuth 2.1 Types and Data Structures
2//!
3//! Module contains the core types used in the OAuth 2.1 authorization flow
4//! for MCP, including metadata structures, token types, and discovery responses.
5
6use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9// ============================================================================
10// OAuth 2.0 Protected Resource Metadata (RFC 9728)
11// ============================================================================
12
13/// Protected Resource Metadata as defined in RFC 9728
14#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct ProtectedResourceMetadata {
16    /// The resource indicator (standard URI of the MCP server)
17    pub resource: String,
18
19    /// Array of authorization server issuer identifiers
20    pub authorization_servers: Vec<String>,
21
22    /// Array of OAuth 2.0 bearer token resource access authentication methods
23    #[serde(skip_serializing_if = "Option::is_none")]
24    pub bearer_methods_supported: Option<Vec<String>>,
25
26    /// Array of resource-specific scopes
27    #[serde(skip_serializing_if = "Option::is_none")]
28    pub scopes_supported: Option<Vec<String>>,
29
30    /// Additional metadata
31    #[serde(flatten)]
32    pub additional: HashMap<String, serde_json::Value>,
33}
34
35// ============================================================================
36// OAuth 2.0 Authorization Server Metadata (RFC 8414)
37// ============================================================================
38
39/// Authorization Server Metadata as defined in RFC 8414
40#[derive(Debug, Clone, Serialize, Deserialize)]
41pub struct AuthorizationServerMetadata {
42    /// The authorization server's issuer identifier
43    pub issuer: String,
44
45    /// URL of the authorization endpoint
46    pub authorization_endpoint: String,
47
48    /// URL of the token endpoint
49    pub token_endpoint: String,
50
51    /// URL of the registration endpoint for dynamic client registration
52    #[serde(skip_serializing_if = "Option::is_none")]
53    pub registration_endpoint: Option<String>,
54
55    /// JSON array containing a list of scopes
56    #[serde(skip_serializing_if = "Option::is_none")]
57    pub scopes_supported: Option<Vec<String>>,
58
59    /// JSON array containing a list of response types
60    pub response_types_supported: Vec<String>,
61
62    /// JSON array containing a list of response modes
63    #[serde(skip_serializing_if = "Option::is_none")]
64    pub response_modes_supported: Option<Vec<String>>,
65
66    /// JSON array containing a list of grant types
67    #[serde(skip_serializing_if = "Option::is_none")]
68    pub grant_types_supported: Option<Vec<String>>,
69
70    /// JSON array containing a list of token endpoint authentication methods
71    #[serde(skip_serializing_if = "Option::is_none")]
72    pub token_endpoint_auth_methods_supported: Option<Vec<String>>,
73
74    /// JSON array containing a list of PKCE code challenge methods
75    #[serde(skip_serializing_if = "Option::is_none")]
76    pub code_challenge_methods_supported: Option<Vec<String>>,
77
78    /// URL of the revocation endpoint
79    #[serde(skip_serializing_if = "Option::is_none")]
80    pub revocation_endpoint: Option<String>,
81
82    /// URL of the introspection endpoint
83    #[serde(skip_serializing_if = "Option::is_none")]
84    pub introspection_endpoint: Option<String>,
85
86    /// Additional metadata fields
87    #[serde(flatten)]
88    pub additional: HashMap<String, serde_json::Value>,
89}
90
91// ============================================================================
92// OpenID Connect Discovery Metadata
93// ============================================================================
94
95/// OpenID Connect Provider Metadata (subset relevant for MCP)
96#[derive(Debug, Clone, Serialize, Deserialize)]
97pub struct OpenIDProviderMetadata {
98    /// The issuer identifier
99    pub issuer: String,
100
101    /// URL of the authorization endpoint
102    pub authorization_endpoint: String,
103
104    /// URL of the token endpoint
105    pub token_endpoint: String,
106
107    /// URL of the userinfo endpoint
108    #[serde(skip_serializing_if = "Option::is_none")]
109    pub userinfo_endpoint: Option<String>,
110
111    /// URL of the JWKS endpoint
112    #[serde(skip_serializing_if = "Option::is_none")]
113    pub jwks_uri: Option<String>,
114
115    /// URL of the registration endpoint
116    #[serde(skip_serializing_if = "Option::is_none")]
117    pub registration_endpoint: Option<String>,
118
119    /// Supported scopes
120    #[serde(skip_serializing_if = "Option::is_none")]
121    pub scopes_supported: Option<Vec<String>>,
122
123    /// Supported response types
124    pub response_types_supported: Vec<String>,
125
126    /// PKCE code challenge methods (extension commonly supported)
127    #[serde(skip_serializing_if = "Option::is_none")]
128    pub code_challenge_methods_supported: Option<Vec<String>>,
129
130    /// Additional fields
131    #[serde(flatten)]
132    pub additional: HashMap<String, serde_json::Value>,
133}
134
135// ============================================================================
136// Dynamic Client Registration (RFC 7591)
137// ============================================================================
138
139/// Client Registration Request
140#[derive(Debug, Clone, Serialize, Deserialize)]
141pub struct ClientRegistrationRequest {
142    /// Array of redirect URIs
143    pub redirect_uris: Vec<String>,
144
145    /// Human-readable name of the client
146    #[serde(skip_serializing_if = "Option::is_none")]
147    pub client_name: Option<String>,
148
149    /// URL of the client's homepage
150    #[serde(skip_serializing_if = "Option::is_none")]
151    pub client_uri: Option<String>,
152
153    /// URL of the client's logo
154    #[serde(skip_serializing_if = "Option::is_none")]
155    pub logo_uri: Option<String>,
156
157    /// Array of OAuth 2.0 grant types
158    #[serde(skip_serializing_if = "Option::is_none")]
159    pub grant_types: Option<Vec<String>>,
160
161    /// Array of OAuth 2.0 response types
162    #[serde(skip_serializing_if = "Option::is_none")]
163    pub response_types: Option<Vec<String>>,
164
165    /// Requested authentication method for the token endpoint
166    #[serde(skip_serializing_if = "Option::is_none")]
167    pub token_endpoint_auth_method: Option<String>,
168
169    /// Space-separated list of scope values
170    #[serde(skip_serializing_if = "Option::is_none")]
171    pub scope: Option<String>,
172
173    /// Software identifier
174    #[serde(skip_serializing_if = "Option::is_none")]
175    pub software_id: Option<String>,
176
177    /// Software version
178    #[serde(skip_serializing_if = "Option::is_none")]
179    pub software_version: Option<String>,
180}
181
182/// Client Registration Response
183#[derive(Debug, Clone, Serialize, Deserialize)]
184pub struct ClientRegistrationResponse {
185    /// Unique client identifier
186    pub client_id: String,
187
188    /// Client secret (for confidential clients)
189    #[serde(skip_serializing_if = "Option::is_none")]
190    pub client_secret: Option<String>,
191
192    /// Time at which the client secret expires (0 = no expiration)
193    #[serde(skip_serializing_if = "Option::is_none")]
194    pub client_secret_expires_at: Option<u64>,
195
196    /// Client registration access token
197    #[serde(skip_serializing_if = "Option::is_none")]
198    pub registration_access_token: Option<String>,
199
200    /// Client configuration endpoint
201    #[serde(skip_serializing_if = "Option::is_none")]
202    pub registration_client_uri: Option<String>,
203
204    /// All registered redirect URIs
205    pub redirect_uris: Vec<String>,
206
207    /// All other registration parameters
208    #[serde(flatten)]
209    pub additional: HashMap<String, serde_json::Value>,
210}
211
212// ============================================================================
213// Token Types
214// ============================================================================
215
216/// OAuth 2.0 Token Request
217#[derive(Debug, Clone, Serialize, Deserialize)]
218pub struct TokenRequest {
219    /// Grant type (e.g., "authorization_code", "refresh_token")
220    pub grant_type: String,
221
222    /// Authorization code (for authorization_code grant)
223    #[serde(skip_serializing_if = "Option::is_none")]
224    pub code: Option<String>,
225
226    /// Redirect URI (must match the one used in authorization request)
227    #[serde(skip_serializing_if = "Option::is_none")]
228    pub redirect_uri: Option<String>,
229
230    /// PKCE code verifier
231    #[serde(skip_serializing_if = "Option::is_none")]
232    pub code_verifier: Option<String>,
233
234    /// Refresh token (for refresh_token grant)
235    #[serde(skip_serializing_if = "Option::is_none")]
236    pub refresh_token: Option<String>,
237
238    /// Resource indicator (RFC 8707)
239    #[serde(skip_serializing_if = "Option::is_none")]
240    pub resource: Option<String>,
241
242    /// Client ID (for public clients)
243    #[serde(skip_serializing_if = "Option::is_none")]
244    pub client_id: Option<String>,
245
246    /// Client secret (for confidential clients)
247    #[serde(skip_serializing_if = "Option::is_none")]
248    pub client_secret: Option<String>,
249
250    /// Scope (for refresh_token grant)
251    #[serde(skip_serializing_if = "Option::is_none")]
252    pub scope: Option<String>,
253}
254
255/// OAuth 2.0 Token Response
256#[derive(Debug, Clone, Serialize, Deserialize)]
257pub struct TokenResponse {
258    /// The access token
259    pub access_token: String,
260
261    /// The type of token (typically "Bearer")
262    pub token_type: String,
263
264    /// The lifetime in seconds of the access token
265    #[serde(skip_serializing_if = "Option::is_none")]
266    pub expires_in: Option<u64>,
267
268    /// The refresh token
269    #[serde(skip_serializing_if = "Option::is_none")]
270    pub refresh_token: Option<String>,
271
272    /// The scope of the access token
273    #[serde(skip_serializing_if = "Option::is_none")]
274    pub scope: Option<String>,
275
276    /// Additional parameters
277    #[serde(flatten)]
278    pub additional: HashMap<String, serde_json::Value>,
279}
280
281/// OAuth 2.0 Error Response
282#[derive(Debug, Clone, Serialize, Deserialize)]
283pub struct OAuth2Error {
284    /// Error code
285    pub error: String,
286
287    /// Human-readable error description
288    #[serde(skip_serializing_if = "Option::is_none")]
289    pub error_description: Option<String>,
290
291    /// URI for more information about the error
292    #[serde(skip_serializing_if = "Option::is_none")]
293    pub error_uri: Option<String>,
294}
295
296// ============================================================================
297// WWW-Authenticate Header Components
298// ============================================================================
299
300/// WWW-Authenticate challenge parameters
301#[derive(Debug, Clone)]
302pub struct AuthChallenge {
303    /// Authentication scheme (e.g., "Bearer")
304    pub scheme: String,
305
306    /// Realm parameter
307    pub realm: Option<String>,
308
309    /// Error code
310    pub error: Option<String>,
311
312    /// Error description
313    pub error_description: Option<String>,
314
315    /// Resource metadata URL
316    pub resource_metadata: Option<String>,
317
318    /// Additional parameters
319    pub additional: HashMap<String, String>,
320}
321
322impl AuthChallenge {
323    /// Parse WWW-Authenticate header
324    pub fn parse(header_value: &str) -> Option<Self> {
325        let parts: Vec<&str> = header_value.splitn(2, ' ').collect();
326        if parts.is_empty() {
327            return None;
328        }
329
330        let scheme = parts[0].to_string();
331        let mut challenge = AuthChallenge {
332            scheme,
333            realm: None,
334            error: None,
335            error_description: None,
336            resource_metadata: None,
337            additional: HashMap::new(),
338        };
339
340        if parts.len() > 1 {
341            // Parse parameters
342            let params = parts[1];
343            for param in params.split(',') {
344                let param = param.trim();
345                if let Some(eq_pos) = param.find('=') {
346                    let key = param[..eq_pos].trim();
347                    let value = param[eq_pos + 1..].trim().trim_matches('"');
348
349                    match key {
350                        "realm" => challenge.realm = Some(value.to_string()),
351                        "error" => challenge.error = Some(value.to_string()),
352                        "error_description" => {
353                            challenge.error_description = Some(value.to_string())
354                        }
355                        "resource_metadata" => {
356                            challenge.resource_metadata = Some(value.to_string())
357                        }
358                        _ => {
359                            challenge
360                                .additional
361                                .insert(key.to_string(), value.to_string());
362                        }
363                    }
364                }
365            }
366        }
367
368        Some(challenge)
369    }
370
371    /// Format as WWW-Authenticate header value
372    pub fn format(&self) -> String {
373        let mut result = self.scheme.clone();
374        let mut params = Vec::new();
375
376        if let Some(realm) = &self.realm {
377            params.push(format!(r#"realm="{realm}""#));
378        }
379        if let Some(error) = &self.error {
380            params.push(format!(r#"error="{error}""#));
381        }
382        if let Some(desc) = &self.error_description {
383            params.push(format!(r#"error_description="{desc}""#));
384        }
385        if let Some(metadata) = &self.resource_metadata {
386            params.push(format!(r#"resource_metadata="{metadata}""#));
387        }
388
389        for (key, value) in &self.additional {
390            params.push(format!(r#"{key}="{value}""#));
391        }
392
393        if !params.is_empty() {
394            result.push(' ');
395            result.push_str(&params.join(", "));
396        }
397
398        result
399    }
400}
401
402// ============================================================================
403// Authorization Context
404// ============================================================================
405
406/// Authorization context for a client session
407#[derive(Debug, Clone)]
408pub struct AuthorizationContext {
409    /// Current access token
410    pub access_token: Option<String>,
411
412    /// Current refresh token
413    pub refresh_token: Option<String>,
414
415    /// Token expiration time (Unix timestamp)
416    pub expires_at: Option<u64>,
417
418    /// Authorization server metadata
419    pub auth_server_metadata: Option<AuthorizationServerMetadata>,
420
421    /// Resource metadata
422    pub resource_metadata: Option<ProtectedResourceMetadata>,
423
424    /// Client registration details
425    pub client_registration: Option<ClientRegistrationResponse>,
426
427    /// PKCE verifier for current flow
428    pub pkce_verifier: Option<String>,
429
430    /// State parameter for current flow
431    pub state: Option<String>,
432
433    /// Resource indicator (standard URI of MCP server)
434    pub resource: String,
435}
436
437impl AuthorizationContext {
438    /// Create a new authorization context
439    pub fn new(resource: String) -> Self {
440        Self {
441            access_token: None,
442            refresh_token: None,
443            expires_at: None,
444            auth_server_metadata: None,
445            resource_metadata: None,
446            client_registration: None,
447            pkce_verifier: None,
448            state: None,
449            resource,
450        }
451    }
452
453    /// Check if the access token is expired
454    pub fn is_token_expired(&self) -> bool {
455        if let Some(expires_at) = self.expires_at {
456            let now = std::time::SystemTime::now()
457                .duration_since(std::time::UNIX_EPOCH)
458                .unwrap()
459                .as_secs();
460            now >= expires_at
461        } else {
462            false
463        }
464    }
465
466    /// Check if we have a valid access token
467    pub fn has_valid_token(&self) -> bool {
468        self.access_token.is_some() && !self.is_token_expired()
469    }
470}