Skip to main content

prism_mcp_rs/auth/
pkce.rs

1//! PKCE (Proof Key for Code Exchange) Implementation
2//!
3//! Module implements PKCE as defined in RFC 7636 for OAuth 2.1.
4//! PKCE is mandatory for MCP authorization to prevent authorization code
5//! interception attacks
6
7use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
8use rand::RngExt;
9use sha2::{Digest, Sha256};
10
11use crate::core::error::{McpError, McpResult};
12
13/// PKCE code challenge methods
14#[derive(Debug, Clone, PartialEq)]
15pub enum CodeChallengeMethod {
16    /// Plain text (not recommended, only for compatibility)
17    Plain,
18    /// SHA-256 hash (recommended)
19    S256,
20}
21
22impl CodeChallengeMethod {
23    /// Get the string representation for OAuth parameters
24    pub fn as_str(&self) -> &str {
25        match self {
26            Self::Plain => "plain",
27            Self::S256 => "S256",
28        }
29    }
30
31    /// Parse from string
32    pub fn parse(s: &str) -> Option<Self> {
33        match s {
34            "plain" => Some(Self::Plain),
35            "S256" => Some(Self::S256),
36            _ => None,
37        }
38    }
39}
40
41/// PKCE parameters for authorization flow
42#[derive(Debug, Clone)]
43pub struct PkceParams {
44    /// The code verifier (random string)
45    pub verifier: String,
46    /// The code challenge (derived from verifier)
47    pub challenge: String,
48    /// The challenge method used
49    pub method: CodeChallengeMethod,
50}
51
52impl PkceParams {
53    /// Generate new PKCE parameters with S256 method (recommended)
54    pub fn new() -> Self {
55        Self::with_method(CodeChallengeMethod::S256)
56    }
57
58    /// Generate new PKCE parameters with specified method
59    pub fn with_method(method: CodeChallengeMethod) -> Self {
60        let verifier = Self::generate_verifier();
61        let challenge = Self::compute_challenge(&verifier, &method);
62
63        Self {
64            verifier,
65            challenge,
66            method,
67        }
68    }
69
70    /// Generate a code verifier
71    ///
72    /// According to RFC 7636, the verifier should be a cryptographically random
73    /// string using unreserved characters [A-Z] / [a-z] / [0-9] / "-" / "." / "_" / "~"
74    /// with a minimum length of 43 characters and maximum of 128 characters.
75    fn generate_verifier() -> String {
76        let mut rng = rand::rng();
77        let _length = rng.random_range(43..=128);
78
79        // Use URL-safe base64 alphabet which matches the unreserved characters
80        let mut bytes = [0u8; 32];
81        for byte in &mut bytes {
82            *byte = rng.random::<u8>();
83        }
84
85        // Convert to URL-safe base64 without padding
86        URL_SAFE_NO_PAD.encode(&bytes[..32]) // Use 32 bytes = 43 chars in base64
87    }
88
89    /// Compute the code challenge from the verifier
90    fn compute_challenge(verifier: &str, method: &CodeChallengeMethod) -> String {
91        match method {
92            CodeChallengeMethod::Plain => verifier.to_string(),
93            CodeChallengeMethod::S256 => {
94                let mut hasher = Sha256::new();
95                hasher.update(verifier.as_bytes());
96                let hash = hasher.finalize();
97                URL_SAFE_NO_PAD.encode(hash)
98            }
99        }
100    }
101
102    /// Verify that a verifier matches a challenge
103    pub fn verify(verifier: &str, challenge: &str, method: &CodeChallengeMethod) -> bool {
104        let computed = Self::compute_challenge(verifier, method);
105        // Use constant-time comparison to prevent timing attacks
106        constant_time_eq(&computed, challenge)
107    }
108}
109
110impl Default for PkceParams {
111    fn default() -> Self {
112        Self::new()
113    }
114}
115
116/// Constant-time string comparison to prevent timing attacks
117fn constant_time_eq(a: &str, b: &str) -> bool {
118    if a.len() != b.len() {
119        return false;
120    }
121
122    let a_bytes = a.as_bytes();
123    let b_bytes = b.as_bytes();
124    let mut result = 0u8;
125
126    for i in 0..a.len() {
127        result |= a_bytes[i] ^ b_bytes[i];
128    }
129
130    result == 0
131}
132
133/// Check if the authorization server supports PKCE
134pub fn check_pkce_support(metadata: &crate::auth::types::AuthorizationServerMetadata) -> bool {
135    metadata
136        .code_challenge_methods_supported
137        .as_ref()
138        .map(|methods| !methods.is_empty())
139        .unwrap_or(false)
140}
141
142/// Check if the authorization server supports S256 method
143pub fn supports_s256(metadata: &crate::auth::types::AuthorizationServerMetadata) -> bool {
144    metadata
145        .code_challenge_methods_supported
146        .as_ref()
147        .map(|methods| methods.contains(&"S256".to_string()))
148        .unwrap_or(false)
149}
150
151/// Select the best available PKCE method from server metadata
152pub fn select_challenge_method(
153    metadata: &crate::auth::types::AuthorizationServerMetadata,
154) -> McpResult<CodeChallengeMethod> {
155    let methods = metadata
156        .code_challenge_methods_supported
157        .as_ref()
158        .ok_or_else(|| {
159            McpError::Auth(
160                "Authorization server does not support PKCE (required for MCP)".to_string(),
161            )
162        })?;
163
164    if methods.is_empty() {
165        return Err(McpError::Auth(
166            "Authorization server does not support any PKCE methods".to_string(),
167        ));
168    }
169
170    // Prefer S256 over plain
171    if methods.contains(&"S256".to_string()) {
172        Ok(CodeChallengeMethod::S256)
173    } else if methods.contains(&"plain".to_string()) {
174        Ok(CodeChallengeMethod::Plain)
175    } else {
176        Err(McpError::Auth(format!(
177            "No supported PKCE methods. Server supports: {methods:?}"
178        )))
179    }
180}
181
182#[cfg(test)]
183mod tests {
184    use super::*;
185
186    #[test]
187    fn test_pkce_generation() {
188        let pkce = PkceParams::new();
189
190        // Verifier should be at least 43 characters
191        assert!(pkce.verifier.len() >= 43);
192        assert!(pkce.verifier.len() <= 128);
193
194        // Challenge should be different from verifier for S256
195        assert_ne!(pkce.verifier, pkce.challenge);
196        assert_eq!(pkce.method, CodeChallengeMethod::S256);
197
198        // Should be URL-safe base64
199        assert!(!pkce.verifier.contains('+'));
200        assert!(!pkce.verifier.contains('/'));
201        assert!(!pkce.verifier.contains('='));
202        assert!(!pkce.challenge.contains('+'));
203        assert!(!pkce.challenge.contains('/'));
204        assert!(!pkce.challenge.contains('='));
205    }
206
207    #[test]
208    fn test_pkce_verification() {
209        let pkce = PkceParams::new();
210
211        // Should verify correctly
212        assert!(PkceParams::verify(
213            &pkce.verifier,
214            &pkce.challenge,
215            &pkce.method
216        ));
217
218        // Should fail with wrong verifier
219        assert!(!PkceParams::verify(
220            "wrong_verifier",
221            &pkce.challenge,
222            &pkce.method
223        ));
224
225        // Should fail with wrong challenge
226        assert!(!PkceParams::verify(
227            &pkce.verifier,
228            "wrong_challenge",
229            &pkce.method
230        ));
231    }
232
233    #[test]
234    fn test_plain_method() {
235        let pkce = PkceParams::with_method(CodeChallengeMethod::Plain);
236
237        // For plain method, challenge should equal verifier
238        assert_eq!(pkce.verifier, pkce.challenge);
239        assert_eq!(pkce.method, CodeChallengeMethod::Plain);
240    }
241
242    #[test]
243    fn test_constant_time_comparison() {
244        assert!(constant_time_eq("hello", "hello"));
245        assert!(!constant_time_eq("hello", "world"));
246        assert!(!constant_time_eq("hello", "hello!"));
247        assert!(!constant_time_eq("", "a"));
248    }
249
250    #[test]
251    fn test_s256_challenge() {
252        // Test with known values from RFC 7636 Appendix B
253        let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
254        let expected_challenge = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM";
255
256        let challenge = PkceParams::compute_challenge(verifier, &CodeChallengeMethod::S256);
257
258        assert_eq!(challenge, expected_challenge);
259    }
260}