1use crate::error::AuthError;
10use base64::engine::general_purpose::URL_SAFE_NO_PAD;
11use base64::Engine;
12use hmac::{Hmac, Mac};
13use serde::{Deserialize, Serialize};
14use sha2::Sha256;
15
16type HmacSha256 = Hmac<Sha256>;
18
19#[derive(Debug, Clone, Serialize, Deserialize)]
20pub struct JwtHeader {
21 pub alg: String,
22 pub typ: String,
23}
24
25impl Default for JwtHeader {
26 fn default() -> Self {
27 Self {
28 alg: "HS256".to_string(),
29 typ: "JWT".to_string(),
30 }
31 }
32}
33
34#[derive(Debug, Clone, Serialize, Deserialize)]
35pub struct JwtClaims {
36 pub sub: String,
37 pub exp: i64,
38 pub iat: i64,
39 #[serde(skip_serializing_if = "Option::is_none")]
40 pub iss: Option<String>,
41 #[serde(default)]
42 pub roles: Vec<String>,
43 #[serde(default)]
44 pub permissions: Vec<String>,
45 #[serde(default, skip_serializing_if = "Option::is_none")]
50 pub user_id: Option<i64>,
51}
52
53impl JwtClaims {
54 pub fn new(sub: impl Into<String>, exp: i64) -> Self {
55 Self {
56 sub: sub.into(),
57 exp,
58 iat: current_timestamp(),
59 iss: None,
60 roles: Vec::new(),
61 permissions: Vec::new(),
62 user_id: None,
63 }
64 }
65
66 pub fn with_issuer(mut self, iss: impl Into<String>) -> Self {
67 self.iss = Some(iss.into());
68 self
69 }
70
71 pub fn with_roles(mut self, roles: Vec<String>) -> Self {
72 self.roles = roles;
73 self
74 }
75
76 pub fn with_permissions(mut self, permissions: Vec<String>) -> Self {
77 self.permissions = permissions;
78 self
79 }
80
81 pub fn with_user_id(mut self, user_id: i64) -> Self {
83 self.user_id = Some(user_id);
84 self
85 }
86
87 pub fn is_expired(&self) -> bool {
88 current_timestamp() > self.exp
89 }
90}
91
92pub struct JwtEncoder {
93 secret: String,
94}
95
96impl JwtEncoder {
97 pub fn new(secret: impl Into<String>) -> Self {
98 Self {
99 secret: secret.into(),
100 }
101 }
102
103 pub fn secret(&self) -> &str {
104 &self.secret
105 }
106
107 pub fn encode(&self, claims: &JwtClaims) -> Result<String, AuthError> {
108 let header = JwtHeader::default();
109 let header_json = serde_json::to_string(&header)
110 .map_err(|e| AuthError::TokenInvalid(format!("Header serialization failed: {}", e)))?;
111 let claims_json = serde_json::to_string(claims)
112 .map_err(|e| AuthError::TokenInvalid(format!("Claims serialization failed: {}", e)))?;
113
114 let header_b64 = base64_url_encode(header_json.as_bytes());
115 let claims_b64 = base64_url_encode(claims_json.as_bytes());
116
117 let signing_input = format!("{}.{}", header_b64, claims_b64);
118 let signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
119 let signature_b64 = base64_url_encode(&signature);
120
121 Ok(format!("{}.{}.{}", header_b64, claims_b64, signature_b64))
122 }
123
124 pub fn decode(&self, token: &str) -> Result<JwtClaims, AuthError> {
125 if token.is_empty() {
126 return Err(AuthError::TokenInvalid("Token is empty".to_string()));
127 }
128
129 let parts: Vec<&str> = token.split('.').collect();
130 if parts.len() != 3 {
131 return Err(AuthError::TokenInvalid(
132 "Invalid JWT format: expected 3 parts".to_string(),
133 ));
134 }
135
136 let header_b64 = parts[0];
137 let claims_b64 = parts[1];
138 let signature_b64 = parts[2];
139
140 use subtle::ConstantTimeEq;
149 let signing_input = format!("{}.{}", header_b64, claims_b64);
150 let expected_signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
151 let expected_signature_b64 = base64_url_encode(&expected_signature);
152
153 let sig_bytes = signature_b64.as_bytes();
154 let expected_bytes = expected_signature_b64.as_bytes();
155 if sig_bytes.len() != expected_bytes.len() {
157 return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
158 }
159 let sig_match: bool = sig_bytes.ct_eq(expected_bytes).into();
161 if !sig_match {
162 return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
163 }
164
165 let header_bytes = base64_url_decode(header_b64)
167 .map_err(|e| AuthError::TokenInvalid(format!("Header decode failed: {}", e)))?;
168 let header: JwtHeader = serde_json::from_slice(&header_bytes)
169 .map_err(|e| AuthError::TokenInvalid(format!("Header parse failed: {}", e)))?;
170
171 if header.alg != "HS256" {
172 return Err(AuthError::TokenInvalid(format!(
173 "Unsupported algorithm: {}",
174 header.alg
175 )));
176 }
177 if header.typ != "JWT" {
178 return Err(AuthError::TokenInvalid(format!(
179 "Unsupported token type: {}",
180 header.typ
181 )));
182 }
183
184 let claims_bytes = base64_url_decode(claims_b64)
186 .map_err(|e| AuthError::TokenInvalid(format!("Claims decode failed: {}", e)))?;
187 let claims: JwtClaims = serde_json::from_slice(&claims_bytes)
188 .map_err(|e| AuthError::TokenInvalid(format!("Claims parse failed: {}", e)))?;
189
190 if claims.is_expired() {
191 return Err(AuthError::TokenExpired("Token has expired".to_string()));
192 }
193
194 Ok(claims)
195 }
196}
197
198fn current_timestamp() -> i64 {
199 use std::time::{SystemTime, UNIX_EPOCH};
200 SystemTime::now()
201 .duration_since(UNIX_EPOCH)
202 .unwrap_or_default()
203 .as_secs() as i64
204}
205
206fn base64_url_encode(input: &[u8]) -> String {
213 URL_SAFE_NO_PAD.encode(input)
214}
215
216fn base64_url_decode(input: &str) -> Result<Vec<u8>, String> {
217 if input.contains('=') {
219 return Err("base64url must not contain padding '='".to_string());
220 }
221 URL_SAFE_NO_PAD
222 .decode(input)
223 .map_err(|e| format!("Invalid base64url: {}", e))
224}
225
226#[cfg(test)]
233fn sha256(data: &[u8]) -> [u8; 32] {
234 use sha2::Digest;
235 let mut hasher = Sha256::new();
236 hasher.update(data);
237 let result = hasher.finalize();
238 let mut out = [0u8; 32];
239 out.copy_from_slice(&result);
240 out
241}
242
243fn hmac_sha256(key: &[u8], message: &[u8]) -> [u8; 32] {
250 let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts any key length");
251 mac.update(message);
252 let result = mac.finalize().into_bytes();
253 let mut out = [0u8; 32];
254 out.copy_from_slice(&result);
255 out
256}
257
258#[cfg(test)]
259mod tests {
260 use super::*;
261
262 fn now() -> i64 {
263 current_timestamp()
264 }
265
266 #[test]
269 fn test_sha256_empty() {
270 let hash = sha256(b"");
272 let expected_hex = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
273 assert_eq!(hex_str(&hash), expected_hex);
274 }
275
276 #[test]
277 fn test_sha256_abc() {
278 let hash = sha256(b"abc");
280 let expected_hex = "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad";
281 assert_eq!(hex_str(&hash), expected_hex);
282 }
283
284 #[test]
285 fn test_sha256_longer_message() {
286 let input = b"abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq";
288 let hash = sha256(input);
289 let expected_hex = "248d6a61d20638b8e5c026930c3e6039a33ce45964ff2167f6ecedd419db06c1";
290 assert_eq!(hex_str(&hash), expected_hex);
291 }
292
293 #[test]
296 fn test_hmac_sha256_rfc4231_case1() {
297 let key = [0x0bu8; 20];
299 let message = b"Hi There";
300 let mac = hmac_sha256(&key, message);
301 let expected_hex = "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7";
302 assert_eq!(hex_str(&mac), expected_hex);
303 }
304
305 #[test]
306 fn test_hmac_sha256_rfc4231_case2() {
307 let mac = hmac_sha256(b"Jefe", b"what do ya want for nothing?");
309 let expected_hex = "5bdcc146bf60754e6a042426089575c75a003f089d2739839dec58b964ec3843";
310 assert_eq!(hex_str(&mac), expected_hex);
311 }
312
313 #[test]
314 fn test_hmac_sha256_long_key() {
315 let key = [0xaau8; 131];
317 let data = b"Test Using Larger Than Block-Size Key - Hash Key First";
318 let mac = hmac_sha256(&key, data);
319 let expected_hex = "60e431591ee0b67f0d8a26aacbf5b77f8e0bc6213728c5140546040f0ee37f54";
320 assert_eq!(hex_str(&mac), expected_hex);
321 }
322
323 #[test]
326 fn test_base64_url_encode_known() {
327 assert_eq!(base64_url_encode(b""), "");
331 assert_eq!(base64_url_encode(b"f"), "Zg");
332 assert_eq!(base64_url_encode(b"fo"), "Zm8");
333 assert_eq!(base64_url_encode(b"foo"), "Zm9v");
334 assert_eq!(base64_url_encode(b"foob"), "Zm9vYg");
335 assert_eq!(base64_url_encode(b"fooba"), "Zm9vYmE");
336 assert_eq!(base64_url_encode(b"foobar"), "Zm9vYmFy");
337 }
338
339 #[test]
340 fn test_base64_url_decode_known() {
341 assert_eq!(base64_url_decode("").unwrap(), b"");
342 assert_eq!(base64_url_decode("Zg").unwrap(), b"f");
343 assert_eq!(base64_url_decode("Zm8").unwrap(), b"fo");
344 assert_eq!(base64_url_decode("Zm9v").unwrap(), b"foo");
345 assert_eq!(base64_url_decode("Zm9vYg").unwrap(), b"foob");
346 assert_eq!(base64_url_decode("Zm9vYmE").unwrap(), b"fooba");
347 assert_eq!(base64_url_decode("Zm9vYmFy").unwrap(), b"foobar");
348 }
349
350 #[test]
351 fn test_base64_url_roundtrip() {
352 let cases: &[&[u8]] = &[
353 b"",
354 b"a",
355 b"ab",
356 b"abc",
357 b"abcd",
358 b"hello world",
359 &[0xffu8; 64],
360 &[0x00u8; 64],
361 &(0u8..=255).collect::<Vec<u8>>(),
362 ];
363 for c in cases {
364 let encoded = base64_url_encode(c);
365 let decoded = base64_url_decode(&encoded).unwrap();
366 assert_eq!(decoded.as_slice(), *c, "roundtrip failed for {:?}", c);
367 }
368 }
369
370 #[test]
371 fn test_base64_url_rejects_padding() {
372 assert!(base64_url_decode("Zg==").is_err());
373 }
374
375 #[test]
376 fn test_base64_url_rejects_invalid_char() {
377 assert!(base64_url_decode("Zm9v*").is_err());
378 }
379
380 #[test]
383 fn test_jwt_encode_decode_roundtrip() {
384 let encoder = JwtEncoder::new("my-secret");
385 let claims = JwtClaims::new("user123", now() + 3600)
386 .with_issuer("test-issuer")
387 .with_roles(vec!["user".to_string(), "editor".to_string()])
388 .with_permissions(vec!["read:posts".to_string(), "write:posts".to_string()]);
389
390 let token = encoder.encode(&claims).expect("encode");
391 assert!(!token.is_empty());
392
393 let parts: Vec<&str> = token.split('.').collect();
394 assert_eq!(parts.len(), 3);
395
396 let decoded = encoder.decode(&token).expect("decode");
397 assert_eq!(decoded.sub, "user123");
398 assert_eq!(decoded.iss, Some("test-issuer".to_string()));
399 assert_eq!(
400 decoded.roles,
401 vec!["user".to_string(), "editor".to_string()]
402 );
403 assert_eq!(
404 decoded.permissions,
405 vec!["read:posts".to_string(), "write:posts".to_string()]
406 );
407 }
408
409 #[test]
410 fn test_jwt_format_is_header_payload_signature() {
411 let encoder = JwtEncoder::new("secret");
412 let claims = JwtClaims::new("alice", now() + 60);
413 let token = encoder.encode(&claims).unwrap();
414 let parts: Vec<&str> = token.split('.').collect();
415 assert_eq!(parts.len(), 3);
416
417 let header_bytes = base64_url_decode(parts[0]).unwrap();
419 let header: JwtHeader = serde_json::from_slice(&header_bytes).unwrap();
420 assert_eq!(header.alg, "HS256");
421 assert_eq!(header.typ, "JWT");
422 }
423
424 #[test]
425 fn test_jwt_signature_changes_with_secret() {
426 let encoder_a = JwtEncoder::new("secret-a");
427 let encoder_b = JwtEncoder::new("secret-b");
428 let claims = JwtClaims::new("user", now() + 3600);
429
430 let token_a = encoder_a.encode(&claims).unwrap();
431 let token_b = encoder_b.encode(&claims).unwrap();
432
433 let parts_a: Vec<&str> = token_a.split('.').collect();
435 let parts_b: Vec<&str> = token_b.split('.').collect();
436 assert_eq!(parts_a[0], parts_b[0]); assert_eq!(parts_a[1], parts_b[1]); assert_ne!(parts_a[2], parts_b[2]); }
440
441 #[test]
442 fn test_jwt_verify_with_wrong_secret_fails() {
443 let encoder_a = JwtEncoder::new("secret-a");
444 let encoder_b = JwtEncoder::new("secret-b");
445 let claims = JwtClaims::new("user", now() + 3600);
446
447 let token = encoder_a.encode(&claims).unwrap();
448 let result = encoder_b.decode(&token);
449 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
450 }
451
452 #[test]
453 fn test_jwt_expired_token_rejected() {
454 let encoder = JwtEncoder::new("secret");
455 let claims = JwtClaims::new("user", now() - 100); let token = encoder.encode(&claims).unwrap();
457 let result = encoder.decode(&token);
458 assert!(matches!(result, Err(AuthError::TokenExpired(_))));
459 }
460
461 #[test]
462 fn test_jwt_decode_invalid_format() {
463 let encoder = JwtEncoder::new("secret");
464 assert!(matches!(
465 encoder.decode(""),
466 Err(AuthError::TokenInvalid(_))
467 ));
468 assert!(matches!(
469 encoder.decode("not.a.jwt.token"),
470 Err(AuthError::TokenInvalid(_))
471 ));
472 assert!(matches!(
473 encoder.decode("only.two"),
474 Err(AuthError::TokenInvalid(_))
475 ));
476 }
477
478 #[test]
479 fn test_jwt_tampered_payload_rejected() {
480 let encoder = JwtEncoder::new("secret");
481 let claims = JwtClaims::new("alice", now() + 3600);
482 let token = encoder.encode(&claims).unwrap();
483
484 let parts: Vec<&str> = token.split('.').collect();
486 let tampered_payload = base64_url_encode(
487 br#"{"sub":"mallory","exp":9999999999,"iat":0,"roles":[],"permissions":[]}"#,
488 );
489 let tampered = format!("{}.{}.{}", parts[0], tampered_payload, parts[2]);
490 let result = encoder.decode(&tampered);
491 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
492 }
493
494 #[test]
495 fn test_jwt_tampered_signature_rejected() {
496 let encoder = JwtEncoder::new("secret");
497 let claims = JwtClaims::new("alice", now() + 3600);
498 let token = encoder.encode(&claims).unwrap();
499
500 let parts: Vec<&str> = token.split('.').collect();
501 let mut sig = parts[2].to_string();
503 let first = sig.chars().next().unwrap();
504 let replacement = if first == 'A' { 'B' } else { 'A' };
505 sig.replace_range(0..first.len_utf8(), &replacement.to_string());
506 let tampered = format!("{}.{}.{}", parts[0], parts[1], sig);
507 let result = encoder.decode(&tampered);
508 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
509 }
510
511 fn hex_str(bytes: &[u8]) -> String {
512 let mut s = String::with_capacity(bytes.len() * 2);
513 for b in bytes {
514 s.push_str(&format!("{:02x}", b));
515 }
516 s
517 }
518}