1use crate::auth::{error::AuthError, state::Identity};
2
3use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation};
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6
7#[derive(Debug, Serialize, Deserialize, Clone)]
8pub struct Claims {
9 pub iss: Option<String>,
11 pub sub: String,
12 pub aud: Option<String>,
13 pub exp: usize,
14 pub iat: usize,
15 pub nbf: Option<usize>,
16 pub jti: Option<String>,
17
18 pub scope: Option<String>,
20 #[serde(skip_serializing_if = "Option::is_none")]
23 pub identity: Option<Identity>,
24
25 #[serde(flatten)]
27 pub extra: HashMap<String, serde_json::Value>,
28}
29
30#[derive(Clone)]
31pub struct TokenManager {
32 encoding_key: EncodingKey,
33 decoding_key: DecodingKey,
34 issuer: Option<String>,
35 kid: Option<String>,
36 alg: Algorithm,
37 public_jwk: Option<crate::token::jwk::Jwk>,
38}
39
40impl TokenManager {
41 pub fn new(secret: &[u8], issuer: Option<String>) -> Self {
43 Self {
44 encoding_key: EncodingKey::from_secret(secret),
45 decoding_key: DecodingKey::from_secret(secret),
46 issuer,
47 kid: None,
48 alg: Algorithm::HS256,
49 public_jwk: None,
50 }
51 }
52
53 pub fn new_asymmetric(
58 private_key_pem: &[u8],
59 issuer: Option<String>,
60 kid: Option<String>,
61 ) -> Result<Self, AuthError> {
62 let encoding_key = EncodingKey::from_rsa_pem(private_key_pem)
63 .map_err(|e| AuthError::Token(e.to_string()))?;
64 let decoding_key = DecodingKey::from_rsa_pem(private_key_pem)
65 .map_err(|e| AuthError::Token(e.to_string()))?;
66
67 let pem_str = std::str::from_utf8(private_key_pem)
68 .map_err(|_| AuthError::Token("Invalid PEM UTF-8".into()))?;
69
70 use rsa::pkcs1::DecodeRsaPrivateKey;
71 use rsa::pkcs8::DecodePrivateKey;
72 let rsa_key = rsa::RsaPrivateKey::from_pkcs8_pem(pem_str)
73 .or_else(|_| rsa::RsaPrivateKey::from_pkcs1_pem(pem_str))
74 .map_err(|e| AuthError::Token(format!("Failed to parse RSA key: {}", e)))?;
75
76 use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
77 use rsa::traits::PublicKeyParts;
78
79 let n = URL_SAFE_NO_PAD.encode(rsa_key.n().to_bytes_be());
80 let e = URL_SAFE_NO_PAD.encode(rsa_key.e().to_bytes_be());
81
82 let kid_val = kid.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
83
84 let jwk = crate::token::jwk::Jwk {
85 kid: Some(kid_val.clone()),
86 kty: "RSA".to_string(),
87 alg: Some("RS256".to_string()),
88 n: Some(n),
89 e: Some(e),
90 };
91
92 Ok(Self {
93 encoding_key,
94 decoding_key,
95 issuer,
96 kid: Some(kid_val),
97 alg: Algorithm::RS256,
98 public_jwk: Some(jwk),
99 })
100 }
101
102 pub fn public_jwk(&self) -> Option<crate::token::jwk::Jwk> {
103 self.public_jwk.clone()
104 }
105
106 pub fn with_issuer(mut self, issuer: String) -> Self {
107 self.issuer = Some(issuer);
108 self
109 }
110
111 pub fn issue_user_token(
113 &self,
114 identity: Identity,
115 expires_in_secs: u64,
116 scope: Option<String>,
117 aud: Option<String>,
118 ) -> Result<String, AuthError> {
119 let now = chrono::Utc::now().timestamp() as usize;
120 let expiration = now + expires_in_secs as usize;
121
122 let claims = Claims {
123 iss: self.issuer.clone(),
124 sub: identity.external_id.clone(),
125 aud,
126 exp: expiration,
127 iat: now,
128 nbf: Some(now),
129 jti: Some(uuid::Uuid::new_v4().to_string()),
130 scope,
131 identity: Some(identity),
132 extra: HashMap::new(),
133 };
134
135 let mut header = Header::new(self.alg);
136 if let Some(ref kid) = self.kid {
137 header.kid = Some(kid.clone());
138 }
139
140 encode(&header, &claims, &self.encoding_key).map_err(|e| AuthError::Token(e.to_string()))
141 }
142
143 pub fn issue_id_token(
145 &self,
146 identity: Identity,
147 client_id: &str,
148 nonce: Option<String>,
149 expires_in_secs: u64,
150 ) -> Result<String, AuthError> {
151 let now = chrono::Utc::now().timestamp() as usize;
152 let expiration = now + expires_in_secs as usize;
153
154 let mut claims = Claims {
155 iss: self.issuer.clone(),
156 sub: identity.external_id.clone(),
157 aud: Some(client_id.to_string()),
158 exp: expiration,
159 iat: now,
160 nbf: Some(now),
161 jti: Some(uuid::Uuid::new_v4().to_string()),
162 scope: None,
163 identity: Some(identity),
164 extra: HashMap::new(),
165 };
166
167 if let Some(n) = nonce {
168 claims
169 .extra
170 .insert("nonce".to_string(), serde_json::Value::String(n));
171 }
172
173 let mut header = Header::new(self.alg);
174 if let Some(ref kid) = self.kid {
175 header.kid = Some(kid.clone());
176 }
177
178 encode(&header, &claims, &self.encoding_key).map_err(|e| AuthError::Token(e.to_string()))
179 }
180
181 pub fn issue_client_token(
183 &self,
184 client_id: &str,
185 expires_in_secs: u64,
186 scope: Option<String>,
187 aud: Option<String>,
188 ) -> Result<String, AuthError> {
189 let now = chrono::Utc::now().timestamp() as usize;
190 let expiration = now + expires_in_secs as usize;
191
192 let claims = Claims {
193 iss: self.issuer.clone(),
194 sub: client_id.to_string(),
195 aud,
196 exp: expiration,
197 iat: now,
198 nbf: Some(now),
199 jti: Some(uuid::Uuid::new_v4().to_string()),
200 scope,
201 identity: None,
202 extra: HashMap::new(),
203 };
204
205 let mut header = Header::new(self.alg);
206 if let Some(ref kid) = self.kid {
207 header.kid = Some(kid.clone());
208 }
209
210 encode(&header, &claims, &self.encoding_key).map_err(|e| AuthError::Token(e.to_string()))
211 }
212
213 pub fn validate_token(
214 &self,
215 token: &str,
216 expected_aud: Option<&str>,
217 ) -> Result<Claims, AuthError> {
218 let mut validation = Validation::new(self.alg);
219 if let Some(aud) = expected_aud {
220 validation.set_audience(&[aud]);
221 } else {
222 validation.validate_aud = false;
223 }
224 if let Some(ref iss) = self.issuer {
225 validation.set_issuer(&[iss]);
226 }
227
228 let token_data = decode::<Claims>(token, &self.decoding_key, &validation)
229 .map_err(|e| AuthError::Token(e.to_string()))?;
230
231 Ok(token_data.claims)
232 }
233}
234
235#[cfg(test)]
236mod tests {
237
238 use super::*;
239 use crate::auth::state::Identity;
240 use std::collections::HashMap;
241
242 #[test]
243 fn test_claims_serialization() {
244 let mut extra = HashMap::new();
245 extra.insert(
246 "custom".to_string(),
247 serde_json::Value::String("value".to_string()),
248 );
249
250 let claims = Claims {
251 iss: Some("issuer".to_string()),
252 sub: "user123".to_string(),
253 aud: Some("audience".to_string()),
254 exp: 1000,
255 iat: 500,
256 nbf: Some(500),
257 jti: Some("jti".to_string()),
258 scope: Some("openid profile".to_string()),
259 identity: Some(Identity {
260 provider_id: "google".to_string(),
261 external_id: "user123".to_string(),
262 email: Some("user@example.com".to_string()),
263 username: Some("user".to_string()),
264 attributes: HashMap::new(),
265 }),
266 extra,
267 };
268
269 let serialized = serde_json::to_string(&claims).unwrap();
270 let deserialized: Claims = serde_json::from_str(&serialized).unwrap();
271
272 assert_eq!(deserialized.iss, claims.iss);
273 assert_eq!(deserialized.sub, claims.sub);
274 assert_eq!(deserialized.extra.get("custom").unwrap(), "value");
275 }
276
277 #[test]
278 fn test_token_manager_issuance() {
279 let manager = TokenManager::new(b"secret", Some("issuer".to_string()));
280 let identity = Identity {
281 provider_id: "mock".to_string(),
282 external_id: "user123".to_string(),
283 email: None,
284 username: None,
285 attributes: HashMap::new(),
286 };
287
288 let token = manager
289 .issue_user_token(identity, 3600, None, None)
290 .unwrap();
291 let claims = manager.validate_token(&token, None).unwrap();
292
293 assert_eq!(claims.iss, Some("issuer".to_string()));
294 assert_eq!(claims.sub, "user123");
295 assert!(claims.jti.is_some());
296 assert!(claims.nbf.is_some());
297 }
298
299 #[test]
300 fn test_token_manager_asymmetric_issuance() {
301 let pem = b"-----BEGIN PRIVATE KEY-----
302MIIEvAIBADANBgkqhkiG9w0BAQEFAASCBKYwggSiAgEAAoIBAQDA5hJIcQ+2rxMz
303VM8ZH5WAmguCr0xmNDAdy0IzzsUeFLG7BebB7izOkU36J4t8t5tUaQwrBMnx2Fvt
304VqJjbdE242UDpvWF/8m9zJ2HR5298cbwT5cGMKLB0HWzDMahugs+Bbh2lCgwyLZk
305Tr3Diwxp5SwFew/Wb+Ke9cNG9Hu5IFH3BCuJ839d9hfqisIeYrBPfb52xxckM37R
3067zSGu/eDP/HZAeLkQuptZJW4A3u7xni14u4qyqXDqsHsYFNgJaxMSAwWgBRY6HNu
307TnvBArTXCiVfL+F73B2L6mdYr64g+QS9nK9v97MlJu/E3mSduz54pren4mpCHc9m
308/S2+VjCZAgMBAAECggEAASC9qQbGnL7XuExRDOIn/m4bWx92ehjo0lCTibhpY3LW
309umbSbpfbhmmuSj3CjW9VZsaM3hBTgSjoTX72lbY/eIUXD7c0memUK5pV4XcEIrQw
310AZlPIye6ckx4I7ZGnKasO8FoAel9dd7DXw36AuBK3LBzJwtzkEFsBc0e3/wixqmG
311UJBbbt/+5ya7CxyjuePaQhKtkLD5R6DpvN2XnCYq5nHJNJdvSVg1pOzsTHYIf+Ee
3122Rz42fGsfFKqeEQCcBFRZaGb/ELeP4c6UZdktZAvmHb1p1fursVZc6X9JXmiJ2OJ
313Kv2H2tMKuysP8L0fXFOMgkH2SVt6rcdHkO6xhlhWsQKBgQDqR8rAJeEE5BFoXA8T
314VVW6CLMlW51x4ey7PEGOaYh39dTG2Q+GZQBZ9G+SZk3f5Y85UCACSyc//4qaz/c3
3150nWsegZ+JPyymmuc79wzIAFFvXB7pL6wyn0Ed1P620kOZTtA8iBcXrsuxL+KP7iu
316MXfWmU1QiZpbndILtyDnY+70uwKBgQDSyCljWkydQCaPU+fiAXLxP8CvcJTSSNQD
317mVUlwJ+OpHnU+Alsi1rBavMgUtLlYbFqzH7NmYrLC8Yadq3ZOwLt0VEK0r8qstAL
3187QCDUD2WNuQjpZupRnXuMUl3iXB96i2gb+VQKGuUAJvVWjdIbYa4+Gu+sBMfcDcX
319dBihDLuEuwKBgAgX4tEwfc2Fc3R/eaXZVNTQaB/qQk4k1+C//CPHUYeTXn5gEUE7
320S//PiesszZPmgkQgmHp7zidP1KH0fT3Yb2g97ut8q54f54fMYXcCrAiUusYKsuu4
321kwkMdkI8QRHWPW3I74VBYIYFFfjYqrCZ1OH8+cbGeiagFRmCggh8U0zxAoGAVW3u
3226Ge22Z0gg8LcHsu7jG/sZq7Ygool8/d3fT+e669Z+ak2GJo6hF4WgClRdMqtn72W
323PzpV+ImjFyK2v26dd0n48MwN0v56N/ss1Av3iiRhPtlmR6tZLNspDZvUzhPVvkrb
324xCs9vtSoVEamVWKe0eVNthGjDoDqs0TInq2MavUCgYB6REavSJs/CLkSS7iimjxZ
325G7g5YQi9/p1lXLOEUDiwEmvRr0XTwzzxUsIc535IXhh/ZUYpthenW+qBBzn85pEC
326TowIqciHu5redqlQ8rITA8/AOY98vaDIhppDg1rfpnHHaZHFbXD/keYAEbhBtbvf
327a0QMqKUcs8+YTy5R5K6qtw==
328-----END PRIVATE KEY-----";
329
330 let manager = TokenManager::new_asymmetric(
331 pem,
332 Some("issuer".to_string()),
333 Some("my-kid-123".to_string()),
334 )
335 .unwrap();
336
337 let identity = Identity {
338 provider_id: "mock".to_string(),
339 external_id: "user123".to_string(),
340 email: None,
341 username: None,
342 attributes: HashMap::new(),
343 };
344
345 let token = manager
346 .issue_user_token(identity, 3600, None, None)
347 .unwrap();
348
349 let jwk = manager.public_jwk().unwrap();
351 assert_eq!(jwk.kid.as_deref(), Some("my-kid-123"));
352
353 let decoding_key = jwk.to_decoding_key().unwrap();
354 let mut validation = jsonwebtoken::Validation::new(jsonwebtoken::Algorithm::RS256);
355 validation.set_issuer(&["issuer"]);
356
357 let token_data =
358 jsonwebtoken::decode::<Claims>(&token, &decoding_key, &validation).unwrap();
359 assert_eq!(token_data.claims.sub, "user123");
360 assert_eq!(token_data.header.kid.as_deref(), Some("my-kid-123"));
361 }
362
363 #[test]
364 fn test_issue_id_token() {
365 let manager = TokenManager::new(b"secret", Some("issuer".to_string()));
366 let identity = Identity {
367 provider_id: "mock".to_string(),
368 external_id: "user123".to_string(),
369 email: None,
370 username: None,
371 attributes: HashMap::new(),
372 };
373
374 let token = manager
375 .issue_id_token(identity, "client-1", Some("nonce123".to_string()), 3600)
376 .unwrap();
377
378 let claims = manager.validate_token(&token, None).unwrap();
379
380 assert_eq!(claims.iss, Some("issuer".to_string()));
381 assert_eq!(claims.sub, "user123");
382 assert_eq!(claims.aud, Some("client-1".to_string()));
383 assert_eq!(claims.extra.get("nonce").unwrap(), "nonce123");
384 }
385 #[test]
386 fn test_token_manager_audience_validation() {
387 let manager = TokenManager::new(b"secret", Some("issuer".to_string()));
388 let identity = Identity {
389 provider_id: "mock".to_string(),
390 external_id: "user123".to_string(),
391 email: None,
392 username: None,
393 attributes: HashMap::new(),
394 };
395
396 let token = manager
398 .issue_id_token(identity, "client-1", None, 3600)
399 .unwrap();
400
401 let claims = manager.validate_token(&token, Some("client-1")).unwrap();
403 assert_eq!(claims.aud, Some("client-1".to_string()));
404
405 let err = manager
407 .validate_token(&token, Some("client-2"))
408 .unwrap_err();
409 assert!(err.to_string().contains("InvalidAudience"));
410 }
411}
412pub mod jwk;