Skip to main content

sz_orm_auth/
oauth2.rs

1//! OAuth2 授权码流程(Authorization Code Flow)
2//!
3//! 实现 RFC 6749 Section 4.1 的授权码流程:
4//! 1. 客户端重定向用户到授权服务器
5//! 2. 用户授权后,授权服务器返回授权码
6//! 3. 客户端用授权码交换访问令牌
7//!
8//! 本模块提供流程状态管理,不包含 HTTP 传输层实现。
9
10use parking_lot::Mutex;
11use std::collections::HashMap;
12use std::time::{SystemTime, UNIX_EPOCH};
13
14use crate::error::AuthError;
15
16/// OAuth2 授权请求参数
17#[derive(Debug, Clone)]
18pub struct AuthorizationRequest {
19    /// 客户端 ID
20    pub client_id: String,
21    /// 重定向 URI
22    pub redirect_uri: String,
23    /// 请求的权限范围(空格分隔)
24    pub scope: String,
25    /// CSRF 防护状态值
26    pub state: String,
27    /// 响应类型(固定为 "code")
28    pub response_type: String,
29}
30
31impl AuthorizationRequest {
32    pub fn new(
33        client_id: impl Into<String>,
34        redirect_uri: impl Into<String>,
35        scope: impl Into<String>,
36        state: impl Into<String>,
37    ) -> Self {
38        Self {
39            client_id: client_id.into(),
40            redirect_uri: redirect_uri.into(),
41            scope: scope.into(),
42            state: state.into(),
43            response_type: "code".to_string(),
44        }
45    }
46}
47
48/// OAuth2 授权码
49#[derive(Debug, Clone)]
50pub struct AuthorizationCode {
51    /// 授权码值
52    pub code: String,
53    /// 关联的客户端 ID
54    pub client_id: String,
55    /// 关联的用户 ID
56    pub user_id: i64,
57    /// 重定向 URI(必须与授权请求一致)
58    pub redirect_uri: String,
59    /// 请求的权限范围
60    pub scope: String,
61    /// 创建时间(Unix 秒)
62    pub created_at: i64,
63    /// 过期时间(Unix 秒),默认 600 秒(10 分钟)
64    pub expires_at: i64,
65    /// 是否已使用(一次性消费)
66    pub used: bool,
67}
68
69impl AuthorizationCode {
70    /// 授权码默认有效期:10 分钟(RFC 6749 建议)
71    const DEFAULT_LIFETIME_SECS: i64 = 600;
72
73    pub fn new(
74        code: impl Into<String>,
75        client_id: impl Into<String>,
76        user_id: i64,
77        redirect_uri: impl Into<String>,
78        scope: impl Into<String>,
79    ) -> Self {
80        let now = current_secs();
81        Self {
82            code: code.into(),
83            client_id: client_id.into(),
84            user_id,
85            redirect_uri: redirect_uri.into(),
86            scope: scope.into(),
87            created_at: now,
88            expires_at: now + Self::DEFAULT_LIFETIME_SECS,
89            used: false,
90        }
91    }
92
93    /// 是否已过期
94    pub fn is_expired(&self) -> bool {
95        current_secs() > self.expires_at
96    }
97}
98
99/// OAuth2 令牌交换请求
100#[derive(Debug, Clone)]
101pub struct TokenRequest {
102    pub grant_type: String,
103    pub code: String,
104    pub redirect_uri: String,
105    pub client_id: String,
106}
107
108impl TokenRequest {
109    pub fn new(
110        code: impl Into<String>,
111        redirect_uri: impl Into<String>,
112        client_id: impl Into<String>,
113    ) -> Self {
114        Self {
115            grant_type: "authorization_code".to_string(),
116            code: code.into(),
117            redirect_uri: redirect_uri.into(),
118            client_id: client_id.into(),
119        }
120    }
121}
122
123/// OAuth2 授权服务器:管理授权码的签发、验证与交换。
124///
125/// 内部使用 `Mutex<HashMap>` 存储授权码,支持:
126/// - 创建授权码(`create_authorization_code`)
127/// - 用授权码交换令牌(`exchange_code`)
128/// - 验证客户端凭据
129/// - 授权码一次性消费
130pub struct OAuth2Server {
131    /// 已签发的授权码:code -> AuthorizationCode
132    codes: Mutex<HashMap<String, AuthorizationCode>>,
133    /// 已注册的客户端:client_id -> client_secret
134    clients: HashMap<String, String>,
135}
136
137impl OAuth2Server {
138    /// 创建授权服务器,注册一组客户端
139    pub fn new(clients: HashMap<String, String>) -> Self {
140        Self {
141            codes: Mutex::new(HashMap::new()),
142            clients,
143        }
144    }
145
146    /// 创建空授权服务器,后续通过 `register_client` 注册
147    pub fn empty() -> Self {
148        Self {
149            codes: Mutex::new(HashMap::new()),
150            clients: HashMap::new(),
151        }
152    }
153
154    /// 注册客户端
155    pub fn register_client(
156        &mut self,
157        client_id: impl Into<String>,
158        client_secret: impl Into<String>,
159    ) {
160        self.clients.insert(client_id.into(), client_secret.into());
161    }
162
163    /// 验证客户端凭据
164    pub fn validate_client(&self, client_id: &str, client_secret: &str) -> bool {
165        self.clients
166            .get(client_id)
167            .map(|secret| secret == client_secret)
168            .unwrap_or(false)
169    }
170
171    /// 是否已注册客户端
172    pub fn has_client(&self, client_id: &str) -> bool {
173        self.clients.contains_key(client_id)
174    }
175
176    /// 签发授权码
177    ///
178    /// 验证授权请求中的 client_id 已注册后,创建一次性授权码。
179    pub fn create_authorization_code(
180        &self,
181        req: &AuthorizationRequest,
182        user_id: i64,
183    ) -> Result<AuthorizationCode, AuthError> {
184        if !self.has_client(&req.client_id) {
185            return Err(AuthError::Config(format!(
186                "Unregistered client: {}",
187                req.client_id
188            )));
189        }
190        if req.response_type != "code" {
191            return Err(AuthError::Config(format!(
192                "Unsupported response_type: {}",
193                req.response_type
194            )));
195        }
196        let code_value = generate_code();
197        let auth_code = AuthorizationCode::new(
198            code_value,
199            req.client_id.clone(),
200            user_id,
201            req.redirect_uri.clone(),
202            req.scope.clone(),
203        );
204        self.codes
205            .lock()
206            .insert(auth_code.code.clone(), auth_code.clone());
207        Ok(auth_code)
208    }
209
210    /// 用授权码交换访问令牌
211    ///
212    /// 验证流程:
213    /// 1. 授权码存在
214    /// 2. 授权码未过期
215    /// 3. 授权码未使用(一次性消费)
216    /// 4. redirect_uri 与签发时一致
217    /// 5. client_id 与签发时一致
218    pub fn exchange_code(&self, req: &TokenRequest) -> Result<AuthorizationCode, AuthError> {
219        let mut codes = self.codes.lock();
220        let auth_code = codes
221            .get(&req.code)
222            .ok_or_else(|| AuthError::TokenInvalid("Invalid authorization code".to_string()))?;
223
224        if auth_code.is_expired() {
225            return Err(AuthError::TokenExpired(
226                "Authorization code expired".to_string(),
227            ));
228        }
229
230        if auth_code.used {
231            return Err(AuthError::TokenInvalid(
232                "Authorization code already used".to_string(),
233            ));
234        }
235
236        if auth_code.redirect_uri != req.redirect_uri {
237            return Err(AuthError::TokenInvalid("Redirect URI mismatch".to_string()));
238        }
239
240        if auth_code.client_id != req.client_id {
241            return Err(AuthError::TokenInvalid("Client ID mismatch".to_string()));
242        }
243
244        // 标记为已使用
245        let result = auth_code.clone();
246        codes.get_mut(&req.code).expect("auth_code verified above guarantees existence").used = true;
247        Ok(result)
248    }
249
250    /// 返回当前存储的授权码数量
251    pub fn code_count(&self) -> usize {
252        self.codes.lock().len()
253    }
254
255    /// 清理已过期或已使用的授权码
256    pub fn cleanup(&self) -> usize {
257        let mut codes = self.codes.lock();
258        let before = codes.len();
259        codes.retain(|_, c| !c.is_expired() && !c.used);
260        before - codes.len()
261    }
262}
263
264/// 生成随机授权码(32 字节十六进制)
265fn generate_code() -> String {
266    use std::collections::hash_map::DefaultHasher;
267    use std::hash::{Hash, Hasher};
268    let mut hasher = DefaultHasher::new();
269    current_nanos().hash(&mut hasher);
270    let seed = hasher.finish();
271    format!("{:064x}", seed)
272}
273
274fn current_secs() -> i64 {
275    SystemTime::now()
276        .duration_since(UNIX_EPOCH)
277        .unwrap_or_default()
278        .as_secs() as i64
279}
280
281fn current_nanos() -> u128 {
282    SystemTime::now()
283        .duration_since(UNIX_EPOCH)
284        .unwrap_or_default()
285        .as_nanos()
286}
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291
292    fn make_server() -> OAuth2Server {
293        let mut clients = HashMap::new();
294        clients.insert("client1".to_string(), "secret1".to_string());
295        OAuth2Server::new(clients)
296    }
297
298    fn make_request() -> AuthorizationRequest {
299        AuthorizationRequest::new("client1", "https://app.com/cb", "read write", "xyz123")
300    }
301
302    #[test]
303    fn test_authorization_request_new() {
304        let req = AuthorizationRequest::new("cid", "https://cb", "read", "state");
305        assert_eq!(req.client_id, "cid");
306        assert_eq!(req.redirect_uri, "https://cb");
307        assert_eq!(req.scope, "read");
308        assert_eq!(req.state, "state");
309        assert_eq!(req.response_type, "code");
310    }
311
312    #[test]
313    fn test_oauth2_server_validate_client() {
314        let server = make_server();
315        assert!(server.validate_client("client1", "secret1"));
316        assert!(!server.validate_client("client1", "wrong"));
317        assert!(!server.validate_client("unknown", "secret1"));
318    }
319
320    #[test]
321    fn test_oauth2_server_has_client() {
322        let server = make_server();
323        assert!(server.has_client("client1"));
324        assert!(!server.has_client("unknown"));
325    }
326
327    #[test]
328    fn test_oauth2_server_register_client() {
329        let mut server = OAuth2Server::empty();
330        assert!(!server.has_client("new_client"));
331        server.register_client("new_client", "new_secret");
332        assert!(server.has_client("new_client"));
333        assert!(server.validate_client("new_client", "new_secret"));
334    }
335
336    #[test]
337    fn test_create_authorization_code_success() {
338        let server = make_server();
339        let req = make_request();
340        let code = server.create_authorization_code(&req, 42).unwrap();
341        assert_eq!(code.client_id, "client1");
342        assert_eq!(code.user_id, 42);
343        assert_eq!(code.redirect_uri, "https://app.com/cb");
344        assert_eq!(code.scope, "read write");
345        assert!(!code.used);
346        assert!(!code.is_expired());
347        assert_eq!(server.code_count(), 1);
348    }
349
350    #[test]
351    fn test_create_authorization_code_unregistered_client() {
352        let server = make_server();
353        let req = AuthorizationRequest::new("unknown", "https://cb", "read", "state");
354        let result = server.create_authorization_code(&req, 1);
355        assert!(matches!(result, Err(AuthError::Config(_))));
356    }
357
358    #[test]
359    fn test_create_authorization_code_wrong_response_type() {
360        let server = make_server();
361        let mut req = make_request();
362        req.response_type = "token".to_string();
363        let result = server.create_authorization_code(&req, 1);
364        assert!(matches!(result, Err(AuthError::Config(_))));
365    }
366
367    #[test]
368    fn test_exchange_code_success() {
369        let server = make_server();
370        let req = make_request();
371        let code = server.create_authorization_code(&req, 99).unwrap();
372        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
373        let result = server.exchange_code(&token_req).unwrap();
374        assert_eq!(result.user_id, 99);
375        assert_eq!(result.client_id, "client1");
376    }
377
378    #[test]
379    fn test_exchange_code_invalid_code() {
380        let server = make_server();
381        let token_req = TokenRequest::new("nonexistent", "https://app.com/cb", "client1");
382        let result = server.exchange_code(&token_req);
383        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
384    }
385
386    #[test]
387    fn test_exchange_code_already_used() {
388        let server = make_server();
389        let req = make_request();
390        let code = server.create_authorization_code(&req, 1).unwrap();
391        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
392        // 第一次交换成功
393        server.exchange_code(&token_req).unwrap();
394        // 第二次应失败(一次性消费)
395        let result = server.exchange_code(&token_req);
396        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
397    }
398
399    #[test]
400    fn test_exchange_code_redirect_uri_mismatch() {
401        let server = make_server();
402        let req = make_request();
403        let code = server.create_authorization_code(&req, 1).unwrap();
404        let token_req = TokenRequest::new(&code.code, "https://wrong.com/cb", "client1");
405        let result = server.exchange_code(&token_req);
406        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
407    }
408
409    #[test]
410    fn test_exchange_code_client_id_mismatch() {
411        let server = make_server();
412        let req = make_request();
413        let code = server.create_authorization_code(&req, 1).unwrap();
414        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "wrong_client");
415        let result = server.exchange_code(&token_req);
416        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
417    }
418
419    #[test]
420    fn test_authorization_code_is_expired() {
421        let mut code = AuthorizationCode::new("c", "cid", 1, "uri", "scope");
422        assert!(!code.is_expired());
423        code.expires_at = current_secs() - 100;
424        assert!(code.is_expired());
425    }
426
427    #[test]
428    fn test_oauth2_cleanup_removes_used() {
429        let server = make_server();
430        let req = make_request();
431        let code = server.create_authorization_code(&req, 1).unwrap();
432        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
433        server.exchange_code(&token_req).unwrap();
434        assert_eq!(server.code_count(), 1);
435        let removed = server.cleanup();
436        assert_eq!(removed, 1);
437        assert_eq!(server.code_count(), 0);
438    }
439
440    #[test]
441    fn test_oauth2_cleanup_keeps_valid() {
442        let server = make_server();
443        let req = make_request();
444        server.create_authorization_code(&req, 1).unwrap();
445        assert_eq!(server.code_count(), 1);
446        let removed = server.cleanup();
447        assert_eq!(removed, 0);
448        assert_eq!(server.code_count(), 1);
449    }
450
451    #[test]
452    fn test_generate_code_non_empty() {
453        let code = generate_code();
454        assert!(!code.is_empty());
455        assert_eq!(code.len(), 64);
456    }
457
458    #[test]
459    fn test_generate_code_different_each_call() {
460        let c1 = generate_code();
461        std::thread::sleep(std::time::Duration::from_millis(1));
462        let c2 = generate_code();
463        // 极大概率不同
464        assert_ne!(c1, c2);
465    }
466
467    #[test]
468    fn test_oauth2_empty_server() {
469        let server = OAuth2Server::empty();
470        assert_eq!(server.code_count(), 0);
471        assert!(!server.has_client("any"));
472    }
473
474    #[test]
475    fn test_multiple_clients() {
476        let mut server = OAuth2Server::empty();
477        server.register_client("app1", "secret1");
478        server.register_client("app2", "secret2");
479        let req1 = AuthorizationRequest::new("app1", "https://a1/cb", "read", "s1");
480        let req2 = AuthorizationRequest::new("app2", "https://a2/cb", "write", "s2");
481        let c1 = server.create_authorization_code(&req1, 1).unwrap();
482        let c2 = server.create_authorization_code(&req2, 2).unwrap();
483        assert_eq!(server.code_count(), 2);
484        // app1 不能用 app2 的 code
485        let wrong_req = TokenRequest::new(&c2.code, "https://a2/cb", "app1");
486        assert!(server.exchange_code(&wrong_req).is_err());
487        // 正确交换
488        let right_req = TokenRequest::new(&c1.code, "https://a1/cb", "app1");
489        assert!(server.exchange_code(&right_req).is_ok());
490    }
491}