1use std::collections::HashMap;
11use std::sync::Mutex;
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}
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#[derive(Debug, Clone)]
50pub struct AuthorizationCode {
51 pub code: String,
53 pub client_id: String,
55 pub user_id: i64,
57 pub redirect_uri: String,
59 pub scope: String,
61 pub created_at: i64,
63 pub expires_at: i64,
65 pub used: bool,
67}
68
69impl AuthorizationCode {
70 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 pub fn is_expired(&self) -> bool {
95 current_secs() > self.expires_at
96 }
97}
98
99#[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
123pub struct OAuth2Server {
131 codes: Mutex<HashMap<String, AuthorizationCode>>,
133 clients: HashMap<String, String>,
135}
136
137impl OAuth2Server {
138 pub fn new(clients: HashMap<String, String>) -> Self {
140 Self {
141 codes: Mutex::new(HashMap::new()),
142 clients,
143 }
144 }
145
146 pub fn empty() -> Self {
148 Self {
149 codes: Mutex::new(HashMap::new()),
150 clients: HashMap::new(),
151 }
152 }
153
154 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 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 pub fn has_client(&self, client_id: &str) -> bool {
173 self.clients.contains_key(client_id)
174 }
175
176 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 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 let result = auth_code.clone();
247 codes.get_mut(&req.code).unwrap().used = true;
248 Ok(result)
249 }
250
251 pub fn code_count(&self) -> usize {
253 self.codes.lock().unwrap().len()
254 }
255
256 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
265fn 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 server.exchange_code(&token_req).unwrap();
395 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 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 let wrong_req = TokenRequest::new(&c2.code, "https://a2/cb", "app1");
487 assert!(server.exchange_code(&wrong_req).is_err());
488 let right_req = TokenRequest::new(&c1.code, "https://a1/cb", "app1");
490 assert!(server.exchange_code(&right_req).is_ok());
491 }
492}