1use crate::Claims;
2use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation};
3use serde::{Deserialize, Serialize};
4use thiserror::Error;
5
6#[derive(Debug, Clone)]
8pub struct JwtConfig {
9 pub private_key_pem: String,
11 pub public_key_pem: String,
13 pub expiration_seconds: u64,
15}
16
17#[derive(Debug, Error)]
19pub enum JwtError {
20 #[error("Failed to encode JWT: {0}")]
22 EncodingFailed(String),
23
24 #[error("Failed to decode JWT: {0}")]
26 DecodingFailed(String),
27
28 #[error("Invalid key format: {0}")]
30 InvalidKey(String),
31
32 #[error("Token expired")]
34 TokenExpired,
35
36 #[error("Invalid token: {0}")]
38 InvalidToken(String),
39}
40
41pub fn encode_jwt<E>(claims: &Claims<E>, config: &JwtConfig) -> Result<String, JwtError>
43where
44 E: Serialize,
45{
46 let encoding_key = EncodingKey::from_ec_pem(config.private_key_pem.as_bytes())
47 .map_err(|e| JwtError::InvalidKey(e.to_string()))?;
48
49 encode(&Header::new(Algorithm::ES256), claims, &encoding_key)
50 .map_err(|e| JwtError::EncodingFailed(e.to_string()))
51}
52
53pub fn decode_jwt<E>(token: &str, config: &JwtConfig) -> Result<Claims<E>, JwtError>
55where
56 E: for<'de> Deserialize<'de>,
57{
58 let decoding_key = DecodingKey::from_ec_pem(config.public_key_pem.as_bytes())
59 .map_err(|e| JwtError::InvalidKey(e.to_string()))?;
60
61 let validation = Validation::new(Algorithm::ES256);
62
63 decode::<Claims<E>>(token, &decoding_key, &validation)
64 .map(|token_data| token_data.claims)
65 .map_err(|e| {
66 if e.kind() == &jsonwebtoken::errors::ErrorKind::ExpiredSignature {
67 JwtError::TokenExpired
68 } else {
69 JwtError::DecodingFailed(e.to_string())
70 }
71 })
72}
73
74#[cfg(test)]
75mod tests {
76 use super::*;
77
78 const PRIVATE_KEY_PEM: &str = "-----BEGIN PRIVATE KEY-----
80MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgtgbDmCbWzH1rPZlb
81qucYzcKQppWx4YxRh0TfnEd0wd6hRANCAATbjOo4G431D+jMHWgoGXaW/vr20Qxn
82QuoeHrU++Hh7LgqOwXbpqEmKfJa5Os5GQfdQ579fyDqZ/MepnZz2ijhz
83-----END PRIVATE KEY-----";
84
85 const PUBLIC_KEY_PEM: &str = "-----BEGIN PUBLIC KEY-----
86MFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAE24zqOBuN9Q/ozB1oKBl2lv769tEM
87Z0LqHh61Pvh4ey4KjsF26ahJinyWuTrORkH3UOe/X8g6mfzHqZ2c9oo4cw==
88-----END PUBLIC KEY-----";
89
90 const OTHER_PUBLIC_KEY_PEM: &str = "-----BEGIN PUBLIC KEY-----
92MFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEQYaZ+hmOmyIcf6OlLbdfrdRDIQVP
93WvgpcJQZdAq9Q3dsB0xGIC4Ea8ps7xzypEj0W6wXZ/zgKyK9NSmDMtgzPg==
94-----END PUBLIC KEY-----";
95
96 fn config() -> JwtConfig {
97 JwtConfig {
98 private_key_pem: PRIVATE_KEY_PEM.to_string(),
99 public_key_pem: PUBLIC_KEY_PEM.to_string(),
100 expiration_seconds: 3600,
101 }
102 }
103
104 fn now() -> usize {
105 std::time::SystemTime::now()
106 .duration_since(std::time::UNIX_EPOCH)
107 .unwrap()
108 .as_secs() as usize
109 }
110
111 fn claims(exp: usize) -> Claims<()> {
112 Claims {
113 sub: "user-1".to_string(),
114 session_id: "session-1".to_string(),
115 exp,
116 ext: (),
117 }
118 }
119
120 #[test]
121 fn test_encode_decode_round_trip_success() {
122 let config = config();
123 let original = claims(now() + config.expiration_seconds as usize);
124 let token = encode_jwt(&original, &config).expect("encode should succeed");
125 let decoded: Claims<()> = decode_jwt(&token, &config).expect("decode should succeed");
126 assert_eq!(decoded.sub, original.sub);
127 assert_eq!(decoded.session_id, original.session_id);
128 assert_eq!(decoded.exp, original.exp);
129 }
130
131 #[test]
132 fn test_decode_expired_token_returns_token_expired() {
133 let config = config();
134 let expired = claims(now().saturating_sub(3600));
135 let token = encode_jwt(&expired, &config).expect("encode should succeed");
136 let result: Result<Claims<()>, JwtError> = decode_jwt(&token, &config);
137 assert!(matches!(result, Err(JwtError::TokenExpired)));
138 }
139
140 #[test]
141 fn test_decode_malformed_token_string_returns_decoding_failed() {
142 let config = config();
143 let result: Result<Claims<()>, JwtError> = decode_jwt("not-a-valid-jwt", &config);
144 assert!(matches!(result, Err(JwtError::DecodingFailed(_))));
145 }
146
147 #[test]
148 fn test_decode_with_mismatched_public_key_returns_decoding_failed() {
149 let signing_config = config();
150 let original = claims(now() + 3600);
151 let token = encode_jwt(&original, &signing_config).expect("encode should succeed");
152
153 let mut verifying_config = signing_config.clone();
154 verifying_config.public_key_pem = OTHER_PUBLIC_KEY_PEM.to_string();
155
156 let result: Result<Claims<()>, JwtError> = decode_jwt(&token, &verifying_config);
157 assert!(matches!(result, Err(JwtError::DecodingFailed(_))));
158 }
159
160 #[test]
161 fn test_encode_with_invalid_pem_returns_invalid_key() {
162 let mut config = config();
163 config.private_key_pem = "not a pem".to_string();
164 let original = claims(now() + 3600);
165 let result = encode_jwt(&original, &config);
166 assert!(matches!(result, Err(JwtError::InvalidKey(_))));
167 }
168
169 #[test]
170 fn test_decode_with_invalid_pem_returns_invalid_key() {
171 let mut config = config();
172 config.public_key_pem = "not a pem".to_string();
173 let result: Result<Claims<()>, JwtError> = decode_jwt("irrelevant.token.value", &config);
174 assert!(matches!(result, Err(JwtError::InvalidKey(_))));
175 }
176
177 #[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
178 struct CustomExt {
179 email: String,
180 role: String,
181 }
182
183 #[test]
184 fn test_encode_decode_round_trip_with_flatten_ext_struct() {
185 let config = config();
186 let original = Claims {
187 sub: "user-2".to_string(),
188 session_id: "session-2".to_string(),
189 exp: now() + 3600,
190 ext: CustomExt {
191 email: "user@example.com".to_string(),
192 role: "admin".to_string(),
193 },
194 };
195 let token = encode_jwt(&original, &config).expect("encode should succeed");
196 let decoded: Claims<CustomExt> =
197 decode_jwt(&token, &config).expect("decode should succeed");
198 assert_eq!(decoded.ext, original.ext);
199 assert_eq!(decoded.sub, original.sub);
200 }
201}