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}
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 .insert(auth_code.code.clone(), auth_code.clone());
207 Ok(auth_code)
208 }
209
210 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 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 pub fn code_count(&self) -> usize {
252 self.codes.lock().len()
253 }
254
255 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
264fn 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 server.exchange_code(&token_req).unwrap();
394 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 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 let wrong_req = TokenRequest::new(&c2.code, "https://a2/cb", "app1");
486 assert!(server.exchange_code(&wrong_req).is_err());
487 let right_req = TokenRequest::new(&c1.code, "https://a1/cb", "app1");
489 assert!(server.exchange_code(&right_req).is_ok());
490 }
491}