prism_mcp_rs/auth/
pkce.rs1use 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#[derive(Debug, Clone, PartialEq)]
15pub enum CodeChallengeMethod {
16 Plain,
18 S256,
20}
21
22impl CodeChallengeMethod {
23 pub fn as_str(&self) -> &str {
25 match self {
26 Self::Plain => "plain",
27 Self::S256 => "S256",
28 }
29 }
30
31 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#[derive(Debug, Clone)]
43pub struct PkceParams {
44 pub verifier: String,
46 pub challenge: String,
48 pub method: CodeChallengeMethod,
50}
51
52impl PkceParams {
53 pub fn new() -> Self {
55 Self::with_method(CodeChallengeMethod::S256)
56 }
57
58 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 fn generate_verifier() -> String {
76 let mut rng = rand::rng();
77 let _length = rng.random_range(43..=128);
78
79 let mut bytes = [0u8; 32];
81 for byte in &mut bytes {
82 *byte = rng.random::<u8>();
83 }
84
85 URL_SAFE_NO_PAD.encode(&bytes[..32]) }
88
89 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 pub fn verify(verifier: &str, challenge: &str, method: &CodeChallengeMethod) -> bool {
104 let computed = Self::compute_challenge(verifier, method);
105 constant_time_eq(&computed, challenge)
107 }
108}
109
110impl Default for PkceParams {
111 fn default() -> Self {
112 Self::new()
113 }
114}
115
116fn 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
133pub 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
142pub 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
151pub 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 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 assert!(pkce.verifier.len() >= 43);
192 assert!(pkce.verifier.len() <= 128);
193
194 assert_ne!(pkce.verifier, pkce.challenge);
196 assert_eq!(pkce.method, CodeChallengeMethod::S256);
197
198 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 assert!(PkceParams::verify(
213 &pkce.verifier,
214 &pkce.challenge,
215 &pkce.method
216 ));
217
218 assert!(!PkceParams::verify(
220 "wrong_verifier",
221 &pkce.challenge,
222 &pkce.method
223 ));
224
225 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 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 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}