1use crate::error::AppError;
2use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation};
3use serde::{Deserialize, Serialize};
4use std::time::{SystemTime, UNIX_EPOCH};
5use tracing::debug;
6
7#[derive(Debug, Default, Serialize, Deserialize)]
29pub struct Claims {
30 pub aud: String,
31 pub sub: String,
32 pub session_id: String,
33 pub role: String,
34 #[serde(default)]
35 pub contexts: Vec<String>,
36 pub exp: u64,
37 #[serde(default)]
48 pub iat: u64,
49 #[serde(default, skip_serializing_if = "is_false")]
53 pub tee_attested: bool,
54 #[serde(default, skip_serializing_if = "Vec::is_empty")]
58 pub amr: Vec<String>,
59 #[serde(default, skip_serializing_if = "String::is_empty")]
62 pub acr: String,
63 #[serde(default, skip_serializing_if = "String::is_empty")]
72 pub jti: String,
73}
74
75fn is_false(v: &bool) -> bool {
76 !*v
77}
78
79pub struct JwtKeys {
81 encoding: EncodingKey,
82 decoding: DecodingKey,
83 audience: String,
85}
86
87impl JwtKeys {
88 pub fn from_ed25519_bytes(private_bytes: &[u8; 32], audience: &str) -> Result<Self, AppError> {
95 let signing_key = ed25519_dalek::SigningKey::from_bytes(private_bytes);
97 let public_bytes = signing_key.verifying_key().to_bytes();
98
99 let mut pkcs8 = Vec::with_capacity(48);
107 pkcs8.extend_from_slice(&[
108 0x30, 0x2e, 0x02, 0x01, 0x00, 0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70, 0x04, 0x22, 0x04, 0x20, ]);
113 pkcs8.extend_from_slice(private_bytes);
114
115 let encoding = EncodingKey::from_ed_der(&pkcs8);
116 let decoding = DecodingKey::from_ed_der(&public_bytes);
118
119 Ok(Self {
120 encoding,
121 decoding,
122 audience: audience.to_string(),
123 })
124 }
125
126 pub fn encode(&self, claims: &Claims) -> Result<String, AppError> {
128 let header = Header::new(Algorithm::EdDSA);
129 jsonwebtoken::encode(&header, claims, &self.encoding)
130 .map_err(|e| AppError::Internal(format!("JWT encode failed: {e}")))
131 }
132
133 pub fn decode(&self, token: &str) -> Result<Claims, AppError> {
135 let mut validation = Validation::new(Algorithm::EdDSA);
136 validation.set_audience(&[&self.audience]);
137 validation.set_required_spec_claims(&["exp", "sub", "aud", "session_id", "role"]);
138
139 jsonwebtoken::decode::<Claims>(token, &self.decoding, &validation)
140 .map(|data| data.claims)
141 .map_err(|e| {
142 debug!(error = %e, "JWT decode failed");
143 AppError::Unauthorized(format!("invalid token: {e}"))
144 })
145 }
146
147 pub fn new_claims(
149 &self,
150 sub: String,
151 session_id: String,
152 role: String,
153 contexts: Vec<String>,
154 expiry_secs: u64,
155 tee_attested: bool,
156 ) -> Claims {
157 let now_secs = SystemTime::now()
162 .duration_since(UNIX_EPOCH)
163 .map(|d| d.as_secs())
164 .unwrap_or(0);
165 let exp = now_secs + expiry_secs;
166
167 Claims {
168 aud: self.audience.clone(),
169 sub,
170 session_id,
171 role,
172 contexts,
173 exp,
174 iat: now_secs,
175 tee_attested,
176 amr: Vec::new(),
177 acr: String::new(),
178 jti: String::new(),
179 }
180 }
181}
182
183impl Claims {
184 pub fn with_jti(mut self, jti: impl Into<String>) -> Self {
203 self.jti = jti.into();
204 self
205 }
206
207 pub fn with_aal(mut self, amr: Vec<String>, acr: impl Into<String>) -> Self {
208 self.amr = amr;
209 self.acr = acr.into();
210 self
211 }
212}
213
214#[cfg(test)]
215mod tests {
216 use super::*;
217 use base64::Engine;
218
219 fn init_jwt_provider() {
228 use std::sync::Once;
229 static INIT: Once = Once::new();
230 INIT.call_once(|| {
231 let _ = jsonwebtoken::crypto::aws_lc::DEFAULT_PROVIDER.install_default();
232 });
233 }
234
235 fn test_keys() -> JwtKeys {
236 init_jwt_provider();
237 JwtKeys::from_ed25519_bytes(&[0x42u8; 32], "VTA").unwrap()
238 }
239
240 #[test]
241 fn test_jwt_roundtrip() {
242 let keys = test_keys();
243 let claims = keys.new_claims(
244 "did:key:z6Mk".into(),
245 "sess-1".into(),
246 "admin".into(),
247 vec!["vta".into()],
248 900,
249 false,
250 );
251 let token = keys.encode(&claims).unwrap();
252 let decoded = keys.decode(&token).unwrap();
253 assert_eq!(decoded.sub, "did:key:z6Mk");
254 assert_eq!(decoded.role, "admin");
255 assert!(!decoded.tee_attested);
256 }
257
258 #[test]
259 fn jti_defaults_empty_and_round_trips_when_set() {
260 let keys = test_keys();
261 let plain = keys.new_claims(
263 "did:key:z6Mk".into(),
264 "s".into(),
265 "admin".into(),
266 vec![],
267 900,
268 false,
269 );
270 assert_eq!(plain.jti, "", "jti defaults empty (unpinned)");
271 assert!(!keys.encode(&plain).unwrap().is_empty());
272
273 let pinned = keys
275 .new_claims(
276 "did:key:z6Mk".into(),
277 "s".into(),
278 "admin".into(),
279 vec![],
280 900,
281 false,
282 )
283 .with_jti("tok-abc123");
284 assert_eq!(pinned.jti, "tok-abc123");
285 let decoded = keys.decode(&keys.encode(&pinned).unwrap()).unwrap();
286 assert_eq!(
287 decoded.jti, "tok-abc123",
288 "jti must round-trip so the extractor pin can match"
289 );
290 }
291
292 #[test]
293 fn test_jwt_tee_attested_true() {
294 let keys = test_keys();
295 let claims = keys.new_claims(
296 "did:key:z6Mk".into(),
297 "sess-2".into(),
298 "admin".into(),
299 vec![],
300 900,
301 true,
302 );
303 let token = keys.encode(&claims).unwrap();
304
305 let parts: Vec<&str> = token.split('.').collect();
307 let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD
308 .decode(parts[1])
309 .unwrap();
310 let json: serde_json::Value = serde_json::from_slice(&payload).unwrap();
311 assert_eq!(json["tee_attested"], true);
312
313 let decoded = keys.decode(&token).unwrap();
314 assert!(decoded.tee_attested);
315 }
316
317 #[test]
318 fn test_jwt_tee_attested_false_omitted() {
319 let keys = test_keys();
320 let claims = keys.new_claims(
321 "did:key:z6Mk".into(),
322 "sess-3".into(),
323 "admin".into(),
324 vec![],
325 900,
326 false,
327 );
328 let token = keys.encode(&claims).unwrap();
329
330 let parts: Vec<&str> = token.split('.').collect();
332 let payload = base64::engine::general_purpose::URL_SAFE_NO_PAD
333 .decode(parts[1])
334 .unwrap();
335 let json: serde_json::Value = serde_json::from_slice(&payload).unwrap();
336 assert!(json.get("tee_attested").is_none());
337 }
338
339 #[test]
340 fn test_jwt_audience_parameterized() {
341 let vta_keys = JwtKeys::from_ed25519_bytes(&[0x42u8; 32], "VTA").unwrap();
342 let vtc_keys = JwtKeys::from_ed25519_bytes(&[0x42u8; 32], "VTC").unwrap();
343
344 let claims = vta_keys.new_claims(
346 "did:key:z6Mk".into(),
347 "sess-1".into(),
348 "admin".into(),
349 vec![],
350 900,
351 false,
352 );
353 let token = vta_keys.encode(&claims).unwrap();
354 assert!(vta_keys.decode(&token).is_ok());
355 assert!(vtc_keys.decode(&token).is_err());
357
358 let claims = vtc_keys.new_claims(
360 "did:key:z6Mk".into(),
361 "sess-2".into(),
362 "admin".into(),
363 vec![],
364 900,
365 false,
366 );
367 let token = vtc_keys.encode(&claims).unwrap();
368 assert!(vtc_keys.decode(&token).is_ok());
369 assert!(vta_keys.decode(&token).is_err());
370 }
371
372 use base64::engine::general_purpose::URL_SAFE_NO_PAD as B64URL;
381
382 fn reencode_with<F: FnOnce(&mut serde_json::Value)>(
387 keys: &JwtKeys,
388 claims: &Claims,
389 mutate: F,
390 ) -> String {
391 let mut payload = serde_json::to_value(claims).unwrap();
392 mutate(&mut payload);
393 let header = Header::new(Algorithm::EdDSA);
394 let header_json = serde_json::to_vec(&header).unwrap();
395 let payload_json = serde_json::to_vec(&payload).unwrap();
396 let signing_input = format!(
397 "{}.{}",
398 B64URL.encode(&header_json),
399 B64URL.encode(&payload_json)
400 );
401 let mutated: Claims = serde_json::from_value(payload).unwrap();
405 let _ = signing_input;
406 keys.encode(&mutated).unwrap()
407 }
408
409 #[test]
410 fn decode_rejects_expired_token() {
411 let keys = test_keys();
412 let expired = keys.new_claims(
413 "did:key:z6Mk".into(),
414 "sess-expired".into(),
415 "admin".into(),
416 vec![],
417 900,
418 false,
419 );
420 let past_token = reencode_with(&keys, &expired, |payload| {
422 payload["exp"] = serde_json::json!(1);
423 });
424
425 let err = keys
426 .decode(&past_token)
427 .expect_err("expired token must be rejected");
428 assert!(matches!(err, AppError::Unauthorized(_)), "got {err:?}");
429 }
430
431 #[test]
432 fn decode_rejects_tampered_signature() {
433 let keys = test_keys();
434 let claims = keys.new_claims(
435 "did:key:z6Mk".into(),
436 "sess-tamper".into(),
437 "admin".into(),
438 vec![],
439 900,
440 false,
441 );
442 let token = keys.encode(&claims).unwrap();
443
444 let mut parts: Vec<&str> = token.split('.').collect();
446 assert_eq!(parts.len(), 3);
447 let mut sig_bytes = B64URL.decode(parts[2]).unwrap();
448 sig_bytes[0] ^= 0x01;
449 let tampered_sig = B64URL.encode(&sig_bytes);
450 parts[2] = &tampered_sig;
451 let tampered = parts.join(".");
452
453 let err = keys
454 .decode(&tampered)
455 .expect_err("tampered signature must be rejected");
456 assert!(matches!(err, AppError::Unauthorized(_)), "got {err:?}");
457 }
458
459 #[test]
460 fn decode_rejects_alg_none_header() {
461 let keys = test_keys();
467 let claims = keys.new_claims(
468 "did:key:z6Mk".into(),
469 "sess-none".into(),
470 "admin".into(),
471 vec![],
472 900,
473 false,
474 );
475 let payload = serde_json::to_vec(&claims).unwrap();
476 let none_header = r#"{"typ":"JWT","alg":"none"}"#;
477 let header_b64 = B64URL.encode(none_header.as_bytes());
478 let payload_b64 = B64URL.encode(&payload);
479 for forged in [
482 format!("{header_b64}.{payload_b64}."),
483 format!("{header_b64}.{payload_b64}"),
484 ] {
485 let err = keys.decode(&forged).expect_err("alg=none must be rejected");
486 assert!(
487 matches!(err, AppError::Unauthorized(_)),
488 "got {err:?} for shape {forged:?}"
489 );
490 }
491 }
492
493 #[test]
494 fn decode_rejects_foreign_signer() {
495 let genuine = test_keys();
498 let attacker = JwtKeys::from_ed25519_bytes(&[0xAAu8; 32], "VTA").unwrap();
499
500 let claims = attacker.new_claims(
501 "did:key:zForged".into(),
502 "sess-forged".into(),
503 "admin".into(),
504 vec![],
505 900,
506 false,
507 );
508 let forged = attacker.encode(&claims).unwrap();
509
510 let err = genuine
511 .decode(&forged)
512 .expect_err("token signed by foreign key must be rejected");
513 assert!(matches!(err, AppError::Unauthorized(_)), "got {err:?}");
514 }
515
516 #[test]
517 fn decode_rejects_missing_required_claims() {
518 let keys = test_keys();
523 let claims = keys.new_claims(
524 "did:key:z6Mk".into(),
525 "sess-missing".into(),
526 "admin".into(),
527 vec![],
528 900,
529 false,
530 );
531 let mut payload = serde_json::to_value(&claims).unwrap();
533 payload.as_object_mut().unwrap().remove("exp");
534 let payload_bytes = serde_json::to_vec(&payload).unwrap();
535 let header = Header::new(Algorithm::EdDSA);
536 let header_bytes = serde_json::to_vec(&header).unwrap();
537
538 let signing_input = format!(
541 "{}.{}",
542 B64URL.encode(&header_bytes),
543 B64URL.encode(&payload_bytes)
544 );
545 let signing_key = ed25519_dalek::SigningKey::from_bytes(&[0x42u8; 32]);
547 use ed25519_dalek::Signer;
548 let sig = signing_key.sign(signing_input.as_bytes());
549 let forged = format!("{signing_input}.{}", B64URL.encode(sig.to_bytes()));
550
551 let err = keys
552 .decode(&forged)
553 .expect_err("token missing `exp` must be rejected");
554 assert!(matches!(err, AppError::Unauthorized(_)), "got {err:?}");
555 }
556
557 #[test]
558 fn decode_rejects_empty_token() {
559 let keys = test_keys();
560 assert!(matches!(keys.decode(""), Err(AppError::Unauthorized(_))));
561 }
562
563 #[test]
564 fn decode_rejects_malformed_structure() {
565 let keys = test_keys();
566 for bad in ["not-a-jwt", "only.two", "four.dot.separated.parts"] {
567 let err = keys
568 .decode(bad)
569 .expect_err(&format!("{bad:?} must be rejected"));
570 assert!(matches!(err, AppError::Unauthorized(_)), "got {err:?}");
571 }
572 }
573}