use arc_swap::ArcSwap;
use jsonwebtoken::{decode, Algorithm, DecodingKey, TokenData, Validation};
use moka::future::Cache;
use std::sync::Arc;
#[derive(Clone, serde::Deserialize, serde::Serialize)]
pub struct Claims {
pub sub: String,
pub exp: usize,
#[serde(default)]
pub scope: String,
}
fn now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
pub struct JwtVerifier {
key: ArcSwap<(DecodingKey, u64)>,
validation: Validation,
cache: Cache<String, (Claims, u64)>,
}
impl JwtVerifier {
pub fn new(jwks_pem: &[u8]) -> anyhow::Result<Self> {
Ok(Self {
key: ArcSwap::from_pointee((DecodingKey::from_rsa_pem(jwks_pem)?, 0)),
validation: {
let mut v = Validation::new(Algorithm::RS256);
v.validate_nbf = true;
v
},
cache: Cache::builder()
.max_capacity(10_000)
.time_to_live(std::time::Duration::from_secs(300))
.build(),
})
}
pub fn with_leeway(mut self, secs: u64) -> Self {
self.validation.leeway = secs;
self
}
pub fn with_issuer(mut self, issuer: &str) -> Self {
self.validation.set_issuer(&[issuer]);
self.validation.required_spec_claims.insert("iss".into());
self
}
pub fn with_audience(mut self, audience: &str) -> Self {
self.validation.set_audience(&[audience]);
self.validation.required_spec_claims.insert("aud".into());
self
}
pub async fn verify(&self, token: &str) -> Option<Claims> {
let key = self.key.load_full();
let gen = key.1;
if let Some((claims, g)) = self.cache.get(token).await {
if g == gen {
if (claims.exp as u64) < now_secs().saturating_sub(self.validation.leeway) {
return None;
}
return Some(claims);
}
if g < gen {
self.cache.invalidate(token).await;
}
}
let TokenData { claims, .. } = decode::<Claims>(token, &key.0, &self.validation).ok()?;
self.cache
.insert(token.to_string(), (claims.clone(), gen))
.await;
Some(claims)
}
pub fn reload_key(&self, pem: &[u8]) -> anyhow::Result<()> {
let new = DecodingKey::from_rsa_pem(pem)?;
self.key.rcu(|cur| Arc::new((new.clone(), cur.1 + 1)));
self.cache.invalidate_all();
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use jsonwebtoken::{encode, EncodingKey, Header};
const PRIV_PEM: &[u8] = include_bytes!("../tests/fixtures/jwt-test-priv.pem");
const OTHER_PRIV_PEM: &[u8] = include_bytes!("../tests/fixtures/jwt-test-other-priv.pem");
const PUB_PEM: &[u8] = include_bytes!("../tests/fixtures/jwt-test-pub.pem");
fn sign(key_pem: &[u8], claims: &Claims) -> String {
encode(
&Header::new(Algorithm::RS256),
claims,
&EncodingKey::from_rsa_pem(key_pem).unwrap(),
)
.unwrap()
}
#[tokio::test]
async fn valid_token_verifies() {
let verifier = JwtVerifier::new(PUB_PEM).unwrap();
let token = sign(
PRIV_PEM,
&Claims {
sub: "tenant-a".into(),
exp: now_secs() as usize + 3600,
scope: "read".into(),
},
);
let claims = verifier.verify(&token).await.expect("should verify");
assert_eq!(claims.sub, "tenant-a");
assert_eq!(claims.scope, "read");
}
fn sign_json(value: serde_json::Value) -> String {
encode(
&Header::new(Algorithm::RS256),
&value,
&EncodingKey::from_rsa_pem(PRIV_PEM).unwrap(),
)
.unwrap()
}
#[tokio::test]
async fn issuer_and_audience_enforced_when_configured() {
let exp = now_secs() + 3600;
let verifier = JwtVerifier::new(PUB_PEM)
.unwrap()
.with_issuer("https://issuer.test")
.with_audience("ferryman-edge");
let good = sign_json(serde_json::json!({
"sub": "t", "exp": exp, "iss": "https://issuer.test", "aud": "ferryman-edge"
}));
assert!(verifier.verify(&good).await.is_some());
for bad in [
serde_json::json!({"sub": "t", "exp": exp, "iss": "https://evil.test", "aud": "ferryman-edge"}),
serde_json::json!({"sub": "t", "exp": exp, "iss": "https://issuer.test", "aud": "other-svc"}),
serde_json::json!({"sub": "t", "exp": exp}),
] {
assert!(verifier.verify(&sign_json(bad)).await.is_none());
}
}
#[tokio::test]
async fn token_with_aud_rejected_when_no_audience_configured() {
let verifier = JwtVerifier::new(PUB_PEM).unwrap();
let token =
sign_json(serde_json::json!({"sub": "t", "exp": now_secs() + 3600, "aud": "x"}));
assert!(verifier.verify(&token).await.is_none());
}
#[tokio::test]
async fn not_yet_valid_token_rejected() {
let verifier = JwtVerifier::new(PUB_PEM).unwrap().with_leeway(0);
let now = now_secs();
let token = sign_json(serde_json::json!({"sub": "t", "exp": now + 3600, "nbf": now + 600}));
assert!(verifier.verify(&token).await.is_none());
}
#[tokio::test]
async fn bad_signature_rejected() {
let verifier = JwtVerifier::new(PUB_PEM).unwrap();
let token = sign(
OTHER_PRIV_PEM,
&Claims {
sub: "tenant-a".into(),
exp: now_secs() as usize + 3600,
scope: "read".into(),
},
);
assert!(verifier.verify(&token).await.is_none());
}
#[tokio::test]
async fn expired_token_rejected() {
let verifier = JwtVerifier::new(PUB_PEM).unwrap().with_leeway(0);
let token = sign(
PRIV_PEM,
&Claims {
sub: "tenant-a".into(),
exp: now_secs() as usize - 100,
scope: "read".into(),
},
);
assert!(verifier.verify(&token).await.is_none());
}
#[tokio::test]
async fn token_without_scope_accepts() {
#[derive(serde::Serialize)]
struct NoScope {
sub: String,
exp: usize,
}
let verifier = JwtVerifier::new(PUB_PEM).unwrap();
let token = encode(
&Header::new(Algorithm::RS256),
&NoScope {
sub: "tenant-a".into(),
exp: now_secs() as usize + 3600,
},
&EncodingKey::from_rsa_pem(PRIV_PEM).unwrap(),
)
.unwrap();
let claims = verifier.verify(&token).await.expect("should verify");
assert_eq!(claims.scope, "");
}
#[tokio::test]
async fn cached_token_rejected_after_exp_passes() {
let verifier = JwtVerifier::new(PUB_PEM).unwrap().with_leeway(0);
let token = sign(
PRIV_PEM,
&Claims {
sub: "tenant-a".into(),
exp: now_secs() as usize + 1,
scope: "read".into(),
},
);
assert!(verifier.verify(&token).await.is_some());
tokio::time::sleep(std::time::Duration::from_millis(2100)).await;
assert!(verifier.verify(&token).await.is_none());
}
const OTHER_PUB_PEM: &[u8] = include_bytes!("../tests/fixtures/jwt-test-other-pub.pem");
fn claims() -> Claims {
Claims {
sub: "t".into(),
exp: now_secs() as usize + 3600,
scope: String::new(),
}
}
#[tokio::test]
async fn reload_key_rejects_old_accepts_new_and_keeps_key_on_error() {
let v = JwtVerifier::new(PUB_PEM).unwrap();
let a = sign(PRIV_PEM, &claims());
let b = sign(OTHER_PRIV_PEM, &claims());
assert!(v.verify(&a).await.is_some()); v.reload_key(OTHER_PUB_PEM).unwrap();
assert!(v.verify(&a).await.is_none());
assert!(v.verify(&b).await.is_some());
assert!(v.reload_key(b"garbage").is_err());
assert!(v.verify(&b).await.is_some());
assert!(v.verify(&a).await.is_none());
}
#[test]
fn concurrent_reloads_bump_generation_atomically() {
let v = Arc::new(JwtVerifier::new(PUB_PEM).unwrap());
let hs: Vec<_> = (0..8)
.map(|_| {
let v = v.clone();
std::thread::spawn(move || {
for _ in 0..10 {
v.reload_key(OTHER_PUB_PEM).unwrap();
}
})
})
.collect();
hs.into_iter().for_each(|h| h.join().unwrap());
assert_eq!(v.key.load().1, 80);
}
#[tokio::test]
async fn stale_generation_cache_entry_is_not_served() {
let v = JwtVerifier::new(PUB_PEM).unwrap();
let a = sign(PRIV_PEM, &claims());
v.reload_key(OTHER_PUB_PEM).unwrap();
v.cache.insert(a.clone(), (claims(), 0)).await;
assert!(v.verify(&a).await.is_none());
}
}