1use parking_lot::Mutex;
11use std::collections::HashMap;
12use std::time::{SystemTime, UNIX_EPOCH};
13
14use crate::error::AuthError;
15
16#[derive(Debug, Clone)]
18pub struct AuthorizationRequest {
19 pub client_id: String,
21 pub redirect_uri: String,
23 pub scope: String,
25 pub state: String,
27 pub response_type: String,
29 pub code_challenge: Option<String>,
34 pub code_challenge_method: Option<String>,
36}
37
38impl AuthorizationRequest {
39 pub fn new(
40 client_id: impl Into<String>,
41 redirect_uri: impl Into<String>,
42 scope: impl Into<String>,
43 state: impl Into<String>,
44 ) -> Self {
45 Self {
46 client_id: client_id.into(),
47 redirect_uri: redirect_uri.into(),
48 scope: scope.into(),
49 state: state.into(),
50 response_type: "code".to_string(),
51 code_challenge: None,
52 code_challenge_method: None,
53 }
54 }
55
56 pub fn with_pkce(mut self, challenge: impl Into<String>, method: &str) -> Self {
61 self.code_challenge = Some(challenge.into());
62 self.code_challenge_method = Some(method.to_string());
63 self
64 }
65}
66
67#[derive(Debug, Clone)]
69pub struct AuthorizationCode {
70 pub code: String,
72 pub client_id: String,
74 pub user_id: i64,
76 pub redirect_uri: String,
78 pub scope: String,
80 pub created_at: i64,
82 pub expires_at: i64,
84 pub used: bool,
86 pub code_challenge: Option<String>,
88}
89
90impl AuthorizationCode {
91 const DEFAULT_LIFETIME_SECS: i64 = 600;
93
94 pub fn new(
95 code: impl Into<String>,
96 client_id: impl Into<String>,
97 user_id: i64,
98 redirect_uri: impl Into<String>,
99 scope: impl Into<String>,
100 ) -> Self {
101 let now = current_secs();
102 Self {
103 code: code.into(),
104 client_id: client_id.into(),
105 user_id,
106 redirect_uri: redirect_uri.into(),
107 scope: scope.into(),
108 created_at: now,
109 expires_at: now + Self::DEFAULT_LIFETIME_SECS,
110 used: false,
111 code_challenge: None,
112 }
113 }
114
115 pub fn is_expired(&self) -> bool {
117 current_secs() > self.expires_at
118 }
119}
120
121#[derive(Debug, Clone)]
123pub struct TokenRequest {
124 pub grant_type: String,
125 pub code: String,
126 pub redirect_uri: String,
127 pub client_id: String,
128 pub client_secret: Option<String>,
135 pub code_verifier: Option<String>,
140}
141
142impl TokenRequest {
143 pub fn new(
144 code: impl Into<String>,
145 redirect_uri: impl Into<String>,
146 client_id: impl Into<String>,
147 ) -> Self {
148 Self {
149 grant_type: "authorization_code".to_string(),
150 code: code.into(),
151 redirect_uri: redirect_uri.into(),
152 client_id: client_id.into(),
153 client_secret: None,
154 code_verifier: None,
155 }
156 }
157
158 pub fn with_client_secret(mut self, secret: impl Into<String>) -> Self {
160 self.client_secret = Some(secret.into());
161 self
162 }
163
164 pub fn with_code_verifier(mut self, verifier: impl Into<String>) -> Self {
166 self.code_verifier = Some(verifier.into());
167 self
168 }
169}
170
171pub struct OAuth2Server {
179 codes: Mutex<HashMap<String, AuthorizationCode>>,
181 clients: HashMap<String, String>,
183}
184
185impl OAuth2Server {
186 pub fn new(clients: HashMap<String, String>) -> Self {
188 Self {
189 codes: Mutex::new(HashMap::new()),
190 clients,
191 }
192 }
193
194 pub fn empty() -> Self {
196 Self {
197 codes: Mutex::new(HashMap::new()),
198 clients: HashMap::new(),
199 }
200 }
201
202 pub fn register_client(
204 &mut self,
205 client_id: impl Into<String>,
206 client_secret: impl Into<String>,
207 ) {
208 self.clients.insert(client_id.into(), client_secret.into());
209 }
210
211 pub fn validate_client(&self, client_id: &str, client_secret: &str) -> bool {
213 self.clients
214 .get(client_id)
215 .map(|secret| secret == client_secret)
216 .unwrap_or(false)
217 }
218
219 pub fn has_client(&self, client_id: &str) -> bool {
221 self.clients.contains_key(client_id)
222 }
223
224 pub fn create_authorization_code(
228 &self,
229 req: &AuthorizationRequest,
230 user_id: i64,
231 ) -> Result<AuthorizationCode, AuthError> {
232 if !self.has_client(&req.client_id) {
233 return Err(AuthError::Config(format!(
234 "Unregistered client: {}",
235 req.client_id
236 )));
237 }
238 if req.response_type != "code" {
239 return Err(AuthError::Config(format!(
240 "Unsupported response_type: {}",
241 req.response_type
242 )));
243 }
244 let code_value = generate_code();
245 let mut auth_code = AuthorizationCode::new(
246 code_value,
247 req.client_id.clone(),
248 user_id,
249 req.redirect_uri.clone(),
250 req.scope.clone(),
251 );
252 auth_code.code_challenge = req.code_challenge.clone();
254 self.codes
255 .lock()
256 .insert(auth_code.code.clone(), auth_code.clone());
257 Ok(auth_code)
258 }
259
260 pub fn exchange_code(&self, req: &TokenRequest) -> Result<AuthorizationCode, AuthError> {
271 let mut codes = self.codes.lock();
272 let auth_code = codes
273 .get(&req.code)
274 .ok_or_else(|| AuthError::TokenInvalid("Invalid authorization code".to_string()))?;
275
276 if auth_code.is_expired() {
277 return Err(AuthError::TokenExpired(
278 "Authorization code expired".to_string(),
279 ));
280 }
281
282 if auth_code.used {
283 return Err(AuthError::TokenInvalid(
284 "Authorization code already used".to_string(),
285 ));
286 }
287
288 if auth_code.redirect_uri != req.redirect_uri {
289 return Err(AuthError::TokenInvalid("Redirect URI mismatch".to_string()));
290 }
291
292 if auth_code.client_id != req.client_id {
293 return Err(AuthError::TokenInvalid("Client ID mismatch".to_string()));
294 }
295
296 if let Some(secret) = &req.client_secret {
300 if !self.validate_client(&req.client_id, secret) {
301 return Err(AuthError::TokenInvalid(
302 "Invalid client credentials".to_string(),
303 ));
304 }
305 }
306
307 if let Some(challenge) = &auth_code.code_challenge {
311 let verifier = req.code_verifier.as_deref().ok_or_else(|| {
312 AuthError::TokenInvalid("PKCE code_verifier required".to_string())
313 })?;
314 if !verify_pkce_s256(challenge, verifier) {
315 return Err(AuthError::TokenInvalid(
316 "PKCE verification failed".to_string(),
317 ));
318 }
319 }
320
321 let result = auth_code.clone();
323 codes.get_mut(&req.code).unwrap().used = true;
324 Ok(result)
325 }
326
327 pub fn code_count(&self) -> usize {
329 self.codes.lock().len()
330 }
331
332 pub fn cleanup(&self) -> usize {
334 let mut codes = self.codes.lock();
335 let before = codes.len();
336 codes.retain(|_, c| !c.is_expired() && !c.used);
337 before - codes.len()
338 }
339}
340
341fn generate_code() -> String {
349 use rand::rngs::OsRng;
350 use rand::RngCore;
351 let mut bytes = [0u8; 32];
352 OsRng.fill_bytes(&mut bytes);
353 let hex: String = bytes.iter().map(|b| format!("{:02x}", b)).collect();
354 hex
355}
356
357fn verify_pkce_s256(challenge: &str, verifier: &str) -> bool {
362 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
363 use base64::Engine;
364 use sha2::{Digest, Sha256};
365 let digest = Sha256::digest(verifier.as_bytes());
366 let computed = URL_SAFE_NO_PAD.encode(digest);
367 if computed.len() != challenge.len() {
369 return false;
370 }
371 constant_time_eq(computed.as_bytes(), challenge.as_bytes())
372}
373
374fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
376 let mut diff: u8 = (a.len() as u8) ^ (b.len() as u8);
377 for (x, y) in a.iter().zip(b.iter()) {
378 diff |= u8::from(x != y);
379 }
380 diff == 0
381}
382
383fn current_secs() -> i64 {
384 SystemTime::now()
385 .duration_since(UNIX_EPOCH)
386 .unwrap_or_default()
387 .as_secs() as i64
388}
389
390#[cfg(test)]
391mod tests {
392 use super::*;
393
394 fn make_server() -> OAuth2Server {
395 let mut clients = HashMap::new();
396 clients.insert("client1".to_string(), "secret1".to_string());
397 OAuth2Server::new(clients)
398 }
399
400 fn make_request() -> AuthorizationRequest {
401 AuthorizationRequest::new("client1", "https://app.com/cb", "read write", "xyz123")
402 }
403
404 #[test]
405 fn test_authorization_request_new() {
406 let req = AuthorizationRequest::new("cid", "https://cb", "read", "state");
407 assert_eq!(req.client_id, "cid");
408 assert_eq!(req.redirect_uri, "https://cb");
409 assert_eq!(req.scope, "read");
410 assert_eq!(req.state, "state");
411 assert_eq!(req.response_type, "code");
412 }
413
414 #[test]
415 fn test_oauth2_server_validate_client() {
416 let server = make_server();
417 assert!(server.validate_client("client1", "secret1"));
418 assert!(!server.validate_client("client1", "wrong"));
419 assert!(!server.validate_client("unknown", "secret1"));
420 }
421
422 #[test]
423 fn test_oauth2_server_has_client() {
424 let server = make_server();
425 assert!(server.has_client("client1"));
426 assert!(!server.has_client("unknown"));
427 }
428
429 #[test]
430 fn test_oauth2_server_register_client() {
431 let mut server = OAuth2Server::empty();
432 assert!(!server.has_client("new_client"));
433 server.register_client("new_client", "new_secret");
434 assert!(server.has_client("new_client"));
435 assert!(server.validate_client("new_client", "new_secret"));
436 }
437
438 #[test]
439 fn test_create_authorization_code_success() {
440 let server = make_server();
441 let req = make_request();
442 let code = server.create_authorization_code(&req, 42).unwrap();
443 assert_eq!(code.client_id, "client1");
444 assert_eq!(code.user_id, 42);
445 assert_eq!(code.redirect_uri, "https://app.com/cb");
446 assert_eq!(code.scope, "read write");
447 assert!(!code.used);
448 assert!(!code.is_expired());
449 assert_eq!(server.code_count(), 1);
450 }
451
452 #[test]
453 fn test_create_authorization_code_unregistered_client() {
454 let server = make_server();
455 let req = AuthorizationRequest::new("unknown", "https://cb", "read", "state");
456 let result = server.create_authorization_code(&req, 1);
457 assert!(matches!(result, Err(AuthError::Config(_))));
458 }
459
460 #[test]
461 fn test_create_authorization_code_wrong_response_type() {
462 let server = make_server();
463 let mut req = make_request();
464 req.response_type = "token".to_string();
465 let result = server.create_authorization_code(&req, 1);
466 assert!(matches!(result, Err(AuthError::Config(_))));
467 }
468
469 #[test]
470 fn test_exchange_code_success() {
471 let server = make_server();
472 let req = make_request();
473 let code = server.create_authorization_code(&req, 99).unwrap();
474 let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
475 let result = server.exchange_code(&token_req).unwrap();
476 assert_eq!(result.user_id, 99);
477 assert_eq!(result.client_id, "client1");
478 }
479
480 #[test]
481 fn test_exchange_code_invalid_code() {
482 let server = make_server();
483 let token_req = TokenRequest::new("nonexistent", "https://app.com/cb", "client1");
484 let result = server.exchange_code(&token_req);
485 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
486 }
487
488 #[test]
489 fn test_exchange_code_already_used() {
490 let server = make_server();
491 let req = make_request();
492 let code = server.create_authorization_code(&req, 1).unwrap();
493 let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
494 server.exchange_code(&token_req).unwrap();
496 let result = server.exchange_code(&token_req);
498 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
499 }
500
501 #[test]
502 fn test_exchange_code_redirect_uri_mismatch() {
503 let server = make_server();
504 let req = make_request();
505 let code = server.create_authorization_code(&req, 1).unwrap();
506 let token_req = TokenRequest::new(&code.code, "https://wrong.com/cb", "client1");
507 let result = server.exchange_code(&token_req);
508 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
509 }
510
511 #[test]
512 fn test_exchange_code_client_id_mismatch() {
513 let server = make_server();
514 let req = make_request();
515 let code = server.create_authorization_code(&req, 1).unwrap();
516 let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "wrong_client");
517 let result = server.exchange_code(&token_req);
518 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
519 }
520
521 #[test]
522 fn test_authorization_code_is_expired() {
523 let mut code = AuthorizationCode::new("c", "cid", 1, "uri", "scope");
524 assert!(!code.is_expired());
525 code.expires_at = current_secs() - 100;
526 assert!(code.is_expired());
527 }
528
529 #[test]
530 fn test_oauth2_cleanup_removes_used() {
531 let server = make_server();
532 let req = make_request();
533 let code = server.create_authorization_code(&req, 1).unwrap();
534 let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
535 server.exchange_code(&token_req).unwrap();
536 assert_eq!(server.code_count(), 1);
537 let removed = server.cleanup();
538 assert_eq!(removed, 1);
539 assert_eq!(server.code_count(), 0);
540 }
541
542 #[test]
543 fn test_oauth2_cleanup_keeps_valid() {
544 let server = make_server();
545 let req = make_request();
546 server.create_authorization_code(&req, 1).unwrap();
547 assert_eq!(server.code_count(), 1);
548 let removed = server.cleanup();
549 assert_eq!(removed, 0);
550 assert_eq!(server.code_count(), 1);
551 }
552
553 #[test]
554 fn test_generate_code_non_empty() {
555 let code = generate_code();
556 assert!(!code.is_empty());
557 assert_eq!(code.len(), 64);
558 }
559
560 #[test]
561 fn test_generate_code_different_each_call() {
562 let c1 = generate_code();
563 std::thread::sleep(std::time::Duration::from_millis(1));
564 let c2 = generate_code();
565 assert_ne!(c1, c2);
567 }
568
569 #[test]
570 fn test_oauth2_empty_server() {
571 let server = OAuth2Server::empty();
572 assert_eq!(server.code_count(), 0);
573 assert!(!server.has_client("any"));
574 }
575
576 #[test]
577 fn test_multiple_clients() {
578 let mut server = OAuth2Server::empty();
579 server.register_client("app1", "secret1");
580 server.register_client("app2", "secret2");
581 let req1 = AuthorizationRequest::new("app1", "https://a1/cb", "read", "s1");
582 let req2 = AuthorizationRequest::new("app2", "https://a2/cb", "write", "s2");
583 let c1 = server.create_authorization_code(&req1, 1).unwrap();
584 let c2 = server.create_authorization_code(&req2, 2).unwrap();
585 assert_eq!(server.code_count(), 2);
586 let wrong_req = TokenRequest::new(&c2.code, "https://a2/cb", "app1");
588 assert!(server.exchange_code(&wrong_req).is_err());
589 let right_req = TokenRequest::new(&c1.code, "https://a1/cb", "app1");
591 assert!(server.exchange_code(&right_req).is_ok());
592 }
593
594 fn pkce_vector() -> (String, String) {
599 (
600 "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk".to_string(),
601 "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM".to_string(),
602 )
603 }
604
605 #[test]
606 fn test_pkce_s256_rfc7636_vector() {
607 let (verifier, challenge) = pkce_vector();
608 assert!(
609 verify_pkce_s256(&challenge, &verifier),
610 "RFC 7636 §A.1 官方向量必须验证通过"
611 );
612 assert!(
613 !verify_pkce_s256(&challenge, "wrong-verifier"),
614 "错误 verifier 必须失败"
615 );
616 assert!(!verify_pkce_s256(&challenge, ""), "空 verifier 必须失败");
617 }
618
619 #[test]
620 fn test_exchange_code_client_secret_mismatch_rejected() {
621 let server = make_server();
622 let req = make_request();
623 let code = server.create_authorization_code(&req, 1).unwrap();
624
625 let bad = TokenRequest::new(&code.code, "https://app.com/cb", "client1")
627 .with_client_secret("wrong-secret");
628 let result = server.exchange_code(&bad);
629 assert!(
630 matches!(result, Err(AuthError::TokenInvalid(_))),
631 "错误 client_secret 必须被拒绝(M-7 修复失效)"
632 );
633
634 let good = TokenRequest::new(&code.code, "https://app.com/cb", "client1")
636 .with_client_secret("secret1");
637 assert!(server.exchange_code(&good).is_ok());
638 }
639
640 #[test]
641 fn test_pkce_challenge_enforced_on_exchange() {
642 let server = make_server();
643 let (verifier, challenge) = pkce_vector();
644 let req = make_request().with_pkce(challenge.clone(), "S256");
645 let code = server.create_authorization_code(&req, 7).unwrap();
646 assert_eq!(code.code_challenge.as_deref(), Some(challenge.as_str()));
647
648 let no_verifier = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
650 assert!(
651 matches!(
652 server.exchange_code(&no_verifier),
653 Err(AuthError::TokenInvalid(_))
654 ),
655 "缺少 code_verifier 必须被拒绝(M-7 修复失效)"
656 );
657
658 let wrong_verifier = TokenRequest::new(&code.code, "https://app.com/cb", "client1")
660 .with_code_verifier("attacker-guessed-verifier");
661 assert!(
662 matches!(
663 server.exchange_code(&wrong_verifier),
664 Err(AuthError::TokenInvalid(_))
665 ),
666 "错误 code_verifier 必须被拒绝"
667 );
668
669 let correct = TokenRequest::new(&code.code, "https://app.com/cb", "client1")
671 .with_code_verifier(verifier);
672 let exchanged = server.exchange_code(&correct).unwrap();
673 assert_eq!(exchanged.user_id, 7);
674 }
675
676 #[test]
677 fn test_pkce_and_secret_combined() {
678 let server = make_server();
679 let (verifier, challenge) = pkce_vector();
680 let req = make_request().with_pkce(challenge, "S256");
681 let code = server.create_authorization_code(&req, 5).unwrap();
682
683 let ok = TokenRequest::new(&code.code, "https://app.com/cb", "client1")
685 .with_client_secret("secret1")
686 .with_code_verifier(verifier.clone());
687 assert!(server.exchange_code(&ok).is_ok());
688
689 let bad_secret = TokenRequest::new(&code.code, "https://app.com/cb", "client1")
691 .with_client_secret("nope")
692 .with_code_verifier(verifier);
693 assert!(server.exchange_code(&bad_secret).is_err());
694 }
695}