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 std::collections::HashMap;
11use std::sync::Mutex;
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            .unwrap()
207            .insert(auth_code.code.clone(), auth_code.clone());
208        Ok(auth_code)
209    }
210
211    /// 用授权码交换访问令牌
212    ///
213    /// 验证流程:
214    /// 1. 授权码存在
215    /// 2. 授权码未过期
216    /// 3. 授权码未使用(一次性消费)
217    /// 4. redirect_uri 与签发时一致
218    /// 5. client_id 与签发时一致
219    pub fn exchange_code(&self, req: &TokenRequest) -> Result<AuthorizationCode, AuthError> {
220        let mut codes = self.codes.lock().unwrap();
221        let auth_code = codes
222            .get(&req.code)
223            .ok_or_else(|| AuthError::TokenInvalid("Invalid authorization code".to_string()))?;
224
225        if auth_code.is_expired() {
226            return Err(AuthError::TokenExpired(
227                "Authorization code expired".to_string(),
228            ));
229        }
230
231        if auth_code.used {
232            return Err(AuthError::TokenInvalid(
233                "Authorization code already used".to_string(),
234            ));
235        }
236
237        if auth_code.redirect_uri != req.redirect_uri {
238            return Err(AuthError::TokenInvalid("Redirect URI mismatch".to_string()));
239        }
240
241        if auth_code.client_id != req.client_id {
242            return Err(AuthError::TokenInvalid("Client ID mismatch".to_string()));
243        }
244
245        // 标记为已使用
246        let result = auth_code.clone();
247        codes.get_mut(&req.code).unwrap().used = true;
248        Ok(result)
249    }
250
251    /// 返回当前存储的授权码数量
252    pub fn code_count(&self) -> usize {
253        self.codes.lock().unwrap().len()
254    }
255
256    /// 清理已过期或已使用的授权码
257    pub fn cleanup(&self) -> usize {
258        let mut codes = self.codes.lock().unwrap();
259        let before = codes.len();
260        codes.retain(|_, c| !c.is_expired() && !c.used);
261        before - codes.len()
262    }
263}
264
265/// 生成随机授权码(32 字节十六进制)
266fn generate_code() -> String {
267    use std::collections::hash_map::DefaultHasher;
268    use std::hash::{Hash, Hasher};
269    let mut hasher = DefaultHasher::new();
270    current_nanos().hash(&mut hasher);
271    let seed = hasher.finish();
272    format!("{:064x}", seed)
273}
274
275fn current_secs() -> i64 {
276    SystemTime::now()
277        .duration_since(UNIX_EPOCH)
278        .unwrap_or_default()
279        .as_secs() as i64
280}
281
282fn current_nanos() -> u128 {
283    SystemTime::now()
284        .duration_since(UNIX_EPOCH)
285        .unwrap_or_default()
286        .as_nanos()
287}
288
289#[cfg(test)]
290mod tests {
291    use super::*;
292
293    fn make_server() -> OAuth2Server {
294        let mut clients = HashMap::new();
295        clients.insert("client1".to_string(), "secret1".to_string());
296        OAuth2Server::new(clients)
297    }
298
299    fn make_request() -> AuthorizationRequest {
300        AuthorizationRequest::new("client1", "https://app.com/cb", "read write", "xyz123")
301    }
302
303    #[test]
304    fn test_authorization_request_new() {
305        let req = AuthorizationRequest::new("cid", "https://cb", "read", "state");
306        assert_eq!(req.client_id, "cid");
307        assert_eq!(req.redirect_uri, "https://cb");
308        assert_eq!(req.scope, "read");
309        assert_eq!(req.state, "state");
310        assert_eq!(req.response_type, "code");
311    }
312
313    #[test]
314    fn test_oauth2_server_validate_client() {
315        let server = make_server();
316        assert!(server.validate_client("client1", "secret1"));
317        assert!(!server.validate_client("client1", "wrong"));
318        assert!(!server.validate_client("unknown", "secret1"));
319    }
320
321    #[test]
322    fn test_oauth2_server_has_client() {
323        let server = make_server();
324        assert!(server.has_client("client1"));
325        assert!(!server.has_client("unknown"));
326    }
327
328    #[test]
329    fn test_oauth2_server_register_client() {
330        let mut server = OAuth2Server::empty();
331        assert!(!server.has_client("new_client"));
332        server.register_client("new_client", "new_secret");
333        assert!(server.has_client("new_client"));
334        assert!(server.validate_client("new_client", "new_secret"));
335    }
336
337    #[test]
338    fn test_create_authorization_code_success() {
339        let server = make_server();
340        let req = make_request();
341        let code = server.create_authorization_code(&req, 42).unwrap();
342        assert_eq!(code.client_id, "client1");
343        assert_eq!(code.user_id, 42);
344        assert_eq!(code.redirect_uri, "https://app.com/cb");
345        assert_eq!(code.scope, "read write");
346        assert!(!code.used);
347        assert!(!code.is_expired());
348        assert_eq!(server.code_count(), 1);
349    }
350
351    #[test]
352    fn test_create_authorization_code_unregistered_client() {
353        let server = make_server();
354        let req = AuthorizationRequest::new("unknown", "https://cb", "read", "state");
355        let result = server.create_authorization_code(&req, 1);
356        assert!(matches!(result, Err(AuthError::Config(_))));
357    }
358
359    #[test]
360    fn test_create_authorization_code_wrong_response_type() {
361        let server = make_server();
362        let mut req = make_request();
363        req.response_type = "token".to_string();
364        let result = server.create_authorization_code(&req, 1);
365        assert!(matches!(result, Err(AuthError::Config(_))));
366    }
367
368    #[test]
369    fn test_exchange_code_success() {
370        let server = make_server();
371        let req = make_request();
372        let code = server.create_authorization_code(&req, 99).unwrap();
373        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
374        let result = server.exchange_code(&token_req).unwrap();
375        assert_eq!(result.user_id, 99);
376        assert_eq!(result.client_id, "client1");
377    }
378
379    #[test]
380    fn test_exchange_code_invalid_code() {
381        let server = make_server();
382        let token_req = TokenRequest::new("nonexistent", "https://app.com/cb", "client1");
383        let result = server.exchange_code(&token_req);
384        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
385    }
386
387    #[test]
388    fn test_exchange_code_already_used() {
389        let server = make_server();
390        let req = make_request();
391        let code = server.create_authorization_code(&req, 1).unwrap();
392        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
393        // 第一次交换成功
394        server.exchange_code(&token_req).unwrap();
395        // 第二次应失败(一次性消费)
396        let result = server.exchange_code(&token_req);
397        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
398    }
399
400    #[test]
401    fn test_exchange_code_redirect_uri_mismatch() {
402        let server = make_server();
403        let req = make_request();
404        let code = server.create_authorization_code(&req, 1).unwrap();
405        let token_req = TokenRequest::new(&code.code, "https://wrong.com/cb", "client1");
406        let result = server.exchange_code(&token_req);
407        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
408    }
409
410    #[test]
411    fn test_exchange_code_client_id_mismatch() {
412        let server = make_server();
413        let req = make_request();
414        let code = server.create_authorization_code(&req, 1).unwrap();
415        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "wrong_client");
416        let result = server.exchange_code(&token_req);
417        assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
418    }
419
420    #[test]
421    fn test_authorization_code_is_expired() {
422        let mut code = AuthorizationCode::new("c", "cid", 1, "uri", "scope");
423        assert!(!code.is_expired());
424        code.expires_at = current_secs() - 100;
425        assert!(code.is_expired());
426    }
427
428    #[test]
429    fn test_oauth2_cleanup_removes_used() {
430        let server = make_server();
431        let req = make_request();
432        let code = server.create_authorization_code(&req, 1).unwrap();
433        let token_req = TokenRequest::new(&code.code, "https://app.com/cb", "client1");
434        server.exchange_code(&token_req).unwrap();
435        assert_eq!(server.code_count(), 1);
436        let removed = server.cleanup();
437        assert_eq!(removed, 1);
438        assert_eq!(server.code_count(), 0);
439    }
440
441    #[test]
442    fn test_oauth2_cleanup_keeps_valid() {
443        let server = make_server();
444        let req = make_request();
445        server.create_authorization_code(&req, 1).unwrap();
446        assert_eq!(server.code_count(), 1);
447        let removed = server.cleanup();
448        assert_eq!(removed, 0);
449        assert_eq!(server.code_count(), 1);
450    }
451
452    #[test]
453    fn test_generate_code_non_empty() {
454        let code = generate_code();
455        assert!(!code.is_empty());
456        assert_eq!(code.len(), 64);
457    }
458
459    #[test]
460    fn test_generate_code_different_each_call() {
461        let c1 = generate_code();
462        std::thread::sleep(std::time::Duration::from_millis(1));
463        let c2 = generate_code();
464        // 极大概率不同
465        assert_ne!(c1, c2);
466    }
467
468    #[test]
469    fn test_oauth2_empty_server() {
470        let server = OAuth2Server::empty();
471        assert_eq!(server.code_count(), 0);
472        assert!(!server.has_client("any"));
473    }
474
475    #[test]
476    fn test_multiple_clients() {
477        let mut server = OAuth2Server::empty();
478        server.register_client("app1", "secret1");
479        server.register_client("app2", "secret2");
480        let req1 = AuthorizationRequest::new("app1", "https://a1/cb", "read", "s1");
481        let req2 = AuthorizationRequest::new("app2", "https://a2/cb", "write", "s2");
482        let c1 = server.create_authorization_code(&req1, 1).unwrap();
483        let c2 = server.create_authorization_code(&req2, 2).unwrap();
484        assert_eq!(server.code_count(), 2);
485        // app1 不能用 app2 的 code
486        let wrong_req = TokenRequest::new(&c2.code, "https://a2/cb", "app1");
487        assert!(server.exchange_code(&wrong_req).is_err());
488        // 正确交换
489        let right_req = TokenRequest::new(&c1.code, "https://a1/cb", "app1");
490        assert!(server.exchange_code(&right_req).is_ok());
491    }
492}