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(skip_serializing_if = "Option::is_none")]
45 pub aud: Option<String>,
46 #[serde(default)]
47 pub roles: Vec<String>,
48 #[serde(default)]
49 pub permissions: Vec<String>,
50 #[serde(default, skip_serializing_if = "Option::is_none")]
55 pub user_id: Option<i64>,
56}
57
58impl JwtClaims {
59 pub fn new(sub: impl Into<String>, exp: i64) -> Self {
60 Self {
61 sub: sub.into(),
62 exp,
63 iat: current_timestamp(),
64 iss: None,
65 aud: None,
66 roles: Vec::new(),
67 permissions: Vec::new(),
68 user_id: None,
69 }
70 }
71
72 pub fn with_issuer(mut self, iss: impl Into<String>) -> Self {
73 self.iss = Some(iss.into());
74 self
75 }
76
77 pub fn with_audience(mut self, aud: impl Into<String>) -> Self {
79 self.aud = Some(aud.into());
80 self
81 }
82
83 pub fn with_roles(mut self, roles: Vec<String>) -> Self {
84 self.roles = roles;
85 self
86 }
87
88 pub fn with_permissions(mut self, permissions: Vec<String>) -> Self {
89 self.permissions = permissions;
90 self
91 }
92
93 pub fn with_user_id(mut self, user_id: i64) -> Self {
95 self.user_id = Some(user_id);
96 self
97 }
98
99 pub fn is_expired(&self) -> bool {
100 current_timestamp() > self.exp
101 }
102}
103
104pub struct JwtEncoder {
105 secret: String,
106}
107
108impl JwtEncoder {
109 pub fn new(secret: impl Into<String>) -> Self {
110 Self {
111 secret: secret.into(),
112 }
113 }
114
115 pub fn secret(&self) -> &str {
116 &self.secret
117 }
118
119 pub fn encode(&self, claims: &JwtClaims) -> Result<String, AuthError> {
120 let header = JwtHeader::default();
121 let header_json = serde_json::to_string(&header)
122 .map_err(|e| AuthError::TokenInvalid(format!("Header serialization failed: {}", e)))?;
123 let claims_json = serde_json::to_string(claims)
124 .map_err(|e| AuthError::TokenInvalid(format!("Claims serialization failed: {}", e)))?;
125
126 let header_b64 = base64_url_encode(header_json.as_bytes());
127 let claims_b64 = base64_url_encode(claims_json.as_bytes());
128
129 let signing_input = format!("{}.{}", header_b64, claims_b64);
130 let signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
131 let signature_b64 = base64_url_encode(&signature);
132
133 Ok(format!("{}.{}.{}", header_b64, claims_b64, signature_b64))
134 }
135
136 pub fn decode(&self, token: &str) -> Result<JwtClaims, AuthError> {
137 if token.is_empty() {
138 return Err(AuthError::TokenInvalid("Token is empty".to_string()));
139 }
140
141 let parts: Vec<&str> = token.split('.').collect();
142 if parts.len() != 3 {
143 return Err(AuthError::TokenInvalid(
144 "Invalid JWT format: expected 3 parts".to_string(),
145 ));
146 }
147
148 let header_b64 = parts[0];
149 let claims_b64 = parts[1];
150 let signature_b64 = parts[2];
151
152 use subtle::ConstantTimeEq;
161 let signing_input = format!("{}.{}", header_b64, claims_b64);
162 let expected_signature = hmac_sha256(self.secret.as_bytes(), signing_input.as_bytes());
163 let expected_signature_b64 = base64_url_encode(&expected_signature);
164
165 let sig_bytes = signature_b64.as_bytes();
166 let expected_bytes = expected_signature_b64.as_bytes();
167 if sig_bytes.len() != expected_bytes.len() {
169 return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
170 }
171 let sig_match: bool = sig_bytes.ct_eq(expected_bytes).into();
173 if !sig_match {
174 return Err(AuthError::TokenInvalid("Invalid signature".to_string()));
175 }
176
177 let header_bytes = base64_url_decode(header_b64)
179 .map_err(|e| AuthError::TokenInvalid(format!("Header decode failed: {}", e)))?;
180 let header: JwtHeader = serde_json::from_slice(&header_bytes)
181 .map_err(|e| AuthError::TokenInvalid(format!("Header parse failed: {}", e)))?;
182
183 if header.alg != "HS256" {
184 return Err(AuthError::TokenInvalid(format!(
185 "Unsupported algorithm: {}",
186 header.alg
187 )));
188 }
189 if header.typ != "JWT" {
190 return Err(AuthError::TokenInvalid(format!(
191 "Unsupported token type: {}",
192 header.typ
193 )));
194 }
195
196 let claims_bytes = base64_url_decode(claims_b64)
198 .map_err(|e| AuthError::TokenInvalid(format!("Claims decode failed: {}", e)))?;
199 let claims: JwtClaims = serde_json::from_slice(&claims_bytes)
200 .map_err(|e| AuthError::TokenInvalid(format!("Claims parse failed: {}", e)))?;
201
202 if claims.is_expired() {
203 return Err(AuthError::TokenExpired("Token has expired".to_string()));
204 }
205
206 Ok(claims)
207 }
208}
209
210fn current_timestamp() -> i64 {
211 use std::time::{SystemTime, UNIX_EPOCH};
212 SystemTime::now()
213 .duration_since(UNIX_EPOCH)
214 .unwrap_or_default()
215 .as_secs() as i64
216}
217
218fn base64_url_encode(input: &[u8]) -> String {
225 URL_SAFE_NO_PAD.encode(input)
226}
227
228fn base64_url_decode(input: &str) -> Result<Vec<u8>, String> {
229 if input.contains('=') {
231 return Err("base64url must not contain padding '='".to_string());
232 }
233 URL_SAFE_NO_PAD
234 .decode(input)
235 .map_err(|e| format!("Invalid base64url: {}", e))
236}
237
238#[cfg(test)]
245fn sha256(data: &[u8]) -> [u8; 32] {
246 use sha2::Digest;
247 let mut hasher = Sha256::new();
248 hasher.update(data);
249 let result = hasher.finalize();
250 let mut out = [0u8; 32];
251 out.copy_from_slice(&result);
252 out
253}
254
255fn hmac_sha256(key: &[u8], message: &[u8]) -> [u8; 32] {
262 let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts any key length");
263 mac.update(message);
264 let result = mac.finalize().into_bytes();
265 let mut out = [0u8; 32];
266 out.copy_from_slice(&result);
267 out
268}
269
270#[cfg(test)]
271mod tests {
272 use super::*;
273
274 fn now() -> i64 {
275 current_timestamp()
276 }
277
278 #[test]
281 fn test_sha256_empty() {
282 let hash = sha256(b"");
284 let expected_hex = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855";
285 assert_eq!(hex_str(&hash), expected_hex);
286 }
287
288 #[test]
289 fn test_sha256_abc() {
290 let hash = sha256(b"abc");
292 let expected_hex = "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad";
293 assert_eq!(hex_str(&hash), expected_hex);
294 }
295
296 #[test]
297 fn test_sha256_longer_message() {
298 let input = b"abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq";
300 let hash = sha256(input);
301 let expected_hex = "248d6a61d20638b8e5c026930c3e6039a33ce45964ff2167f6ecedd419db06c1";
302 assert_eq!(hex_str(&hash), expected_hex);
303 }
304
305 #[test]
308 fn test_hmac_sha256_rfc4231_case1() {
309 let key = [0x0bu8; 20];
311 let message = b"Hi There";
312 let mac = hmac_sha256(&key, message);
313 let expected_hex = "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7";
314 assert_eq!(hex_str(&mac), expected_hex);
315 }
316
317 #[test]
318 fn test_hmac_sha256_rfc4231_case2() {
319 let mac = hmac_sha256(b"Jefe", b"what do ya want for nothing?");
321 let expected_hex = "5bdcc146bf60754e6a042426089575c75a003f089d2739839dec58b964ec3843";
322 assert_eq!(hex_str(&mac), expected_hex);
323 }
324
325 #[test]
326 fn test_hmac_sha256_long_key() {
327 let key = [0xaau8; 131];
329 let data = b"Test Using Larger Than Block-Size Key - Hash Key First";
330 let mac = hmac_sha256(&key, data);
331 let expected_hex = "60e431591ee0b67f0d8a26aacbf5b77f8e0bc6213728c5140546040f0ee37f54";
332 assert_eq!(hex_str(&mac), expected_hex);
333 }
334
335 #[test]
338 fn test_base64_url_encode_known() {
339 assert_eq!(base64_url_encode(b""), "");
343 assert_eq!(base64_url_encode(b"f"), "Zg");
344 assert_eq!(base64_url_encode(b"fo"), "Zm8");
345 assert_eq!(base64_url_encode(b"foo"), "Zm9v");
346 assert_eq!(base64_url_encode(b"foob"), "Zm9vYg");
347 assert_eq!(base64_url_encode(b"fooba"), "Zm9vYmE");
348 assert_eq!(base64_url_encode(b"foobar"), "Zm9vYmFy");
349 }
350
351 #[test]
352 fn test_base64_url_decode_known() {
353 assert_eq!(base64_url_decode("").unwrap(), b"");
354 assert_eq!(base64_url_decode("Zg").unwrap(), b"f");
355 assert_eq!(base64_url_decode("Zm8").unwrap(), b"fo");
356 assert_eq!(base64_url_decode("Zm9v").unwrap(), b"foo");
357 assert_eq!(base64_url_decode("Zm9vYg").unwrap(), b"foob");
358 assert_eq!(base64_url_decode("Zm9vYmE").unwrap(), b"fooba");
359 assert_eq!(base64_url_decode("Zm9vYmFy").unwrap(), b"foobar");
360 }
361
362 #[test]
363 fn test_base64_url_roundtrip() {
364 let cases: &[&[u8]] = &[
365 b"",
366 b"a",
367 b"ab",
368 b"abc",
369 b"abcd",
370 b"hello world",
371 &[0xffu8; 64],
372 &[0x00u8; 64],
373 &(0u8..=255).collect::<Vec<u8>>(),
374 ];
375 for c in cases {
376 let encoded = base64_url_encode(c);
377 let decoded = base64_url_decode(&encoded).unwrap();
378 assert_eq!(decoded.as_slice(), *c, "roundtrip failed for {:?}", c);
379 }
380 }
381
382 #[test]
383 fn test_base64_url_rejects_padding() {
384 assert!(base64_url_decode("Zg==").is_err());
385 }
386
387 #[test]
388 fn test_base64_url_rejects_invalid_char() {
389 assert!(base64_url_decode("Zm9v*").is_err());
390 }
391
392 #[test]
395 fn test_jwt_encode_decode_roundtrip() {
396 let encoder = JwtEncoder::new("my-secret");
397 let claims = JwtClaims::new("user123", now() + 3600)
398 .with_issuer("test-issuer")
399 .with_roles(vec!["user".to_string(), "editor".to_string()])
400 .with_permissions(vec!["read:posts".to_string(), "write:posts".to_string()]);
401
402 let token = encoder.encode(&claims).expect("encode");
403 assert!(!token.is_empty());
404
405 let parts: Vec<&str> = token.split('.').collect();
406 assert_eq!(parts.len(), 3);
407
408 let decoded = encoder.decode(&token).expect("decode");
409 assert_eq!(decoded.sub, "user123");
410 assert_eq!(decoded.iss, Some("test-issuer".to_string()));
411 assert_eq!(
412 decoded.roles,
413 vec!["user".to_string(), "editor".to_string()]
414 );
415 assert_eq!(
416 decoded.permissions,
417 vec!["read:posts".to_string(), "write:posts".to_string()]
418 );
419 }
420
421 #[test]
422 fn test_jwt_format_is_header_payload_signature() {
423 let encoder = JwtEncoder::new("secret");
424 let claims = JwtClaims::new("alice", now() + 60);
425 let token = encoder.encode(&claims).unwrap();
426 let parts: Vec<&str> = token.split('.').collect();
427 assert_eq!(parts.len(), 3);
428
429 let header_bytes = base64_url_decode(parts[0]).unwrap();
431 let header: JwtHeader = serde_json::from_slice(&header_bytes).unwrap();
432 assert_eq!(header.alg, "HS256");
433 assert_eq!(header.typ, "JWT");
434 }
435
436 #[test]
437 fn test_jwt_signature_changes_with_secret() {
438 let encoder_a = JwtEncoder::new("secret-a");
439 let encoder_b = JwtEncoder::new("secret-b");
440 let claims = JwtClaims::new("user", now() + 3600);
441
442 let token_a = encoder_a.encode(&claims).unwrap();
443 let token_b = encoder_b.encode(&claims).unwrap();
444
445 let parts_a: Vec<&str> = token_a.split('.').collect();
447 let parts_b: Vec<&str> = token_b.split('.').collect();
448 assert_eq!(parts_a[0], parts_b[0]); assert_eq!(parts_a[1], parts_b[1]); assert_ne!(parts_a[2], parts_b[2]); }
452
453 #[test]
454 fn test_jwt_verify_with_wrong_secret_fails() {
455 let encoder_a = JwtEncoder::new("secret-a");
456 let encoder_b = JwtEncoder::new("secret-b");
457 let claims = JwtClaims::new("user", now() + 3600);
458
459 let token = encoder_a.encode(&claims).unwrap();
460 let result = encoder_b.decode(&token);
461 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
462 }
463
464 #[test]
465 fn test_jwt_expired_token_rejected() {
466 let encoder = JwtEncoder::new("secret");
467 let claims = JwtClaims::new("user", now() - 100); let token = encoder.encode(&claims).unwrap();
469 let result = encoder.decode(&token);
470 assert!(matches!(result, Err(AuthError::TokenExpired(_))));
471 }
472
473 #[test]
474 fn test_jwt_decode_invalid_format() {
475 let encoder = JwtEncoder::new("secret");
476 assert!(matches!(
477 encoder.decode(""),
478 Err(AuthError::TokenInvalid(_))
479 ));
480 assert!(matches!(
481 encoder.decode("not.a.jwt.token"),
482 Err(AuthError::TokenInvalid(_))
483 ));
484 assert!(matches!(
485 encoder.decode("only.two"),
486 Err(AuthError::TokenInvalid(_))
487 ));
488 }
489
490 #[test]
491 fn test_jwt_tampered_payload_rejected() {
492 let encoder = JwtEncoder::new("secret");
493 let claims = JwtClaims::new("alice", now() + 3600);
494 let token = encoder.encode(&claims).unwrap();
495
496 let parts: Vec<&str> = token.split('.').collect();
498 let tampered_payload = base64_url_encode(
499 br#"{"sub":"mallory","exp":9999999999,"iat":0,"roles":[],"permissions":[]}"#,
500 );
501 let tampered = format!("{}.{}.{}", parts[0], tampered_payload, parts[2]);
502 let result = encoder.decode(&tampered);
503 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
504 }
505
506 #[test]
507 fn test_jwt_tampered_signature_rejected() {
508 let encoder = JwtEncoder::new("secret");
509 let claims = JwtClaims::new("alice", now() + 3600);
510 let token = encoder.encode(&claims).unwrap();
511
512 let parts: Vec<&str> = token.split('.').collect();
513 let mut sig = parts[2].to_string();
515 let first = sig.chars().next().unwrap();
516 let replacement = if first == 'A' { 'B' } else { 'A' };
517 sig.replace_range(0..first.len_utf8(), &replacement.to_string());
518 let tampered = format!("{}.{}.{}", parts[0], parts[1], sig);
519 let result = encoder.decode(&tampered);
520 assert!(matches!(result, Err(AuthError::TokenInvalid(_))));
521 }
522
523 fn hex_str(bytes: &[u8]) -> String {
524 let mut s = String::with_capacity(bytes.len() * 2);
525 for b in bytes {
526 s.push_str(&format!("{:02x}", b));
527 }
528 s
529 }
530}