use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use p256::ecdsa::VerifyingKey;
use parking_lot::RwLock;
use rand::RngCore;
use serde::{Deserialize, Serialize};
use super::AuthenticatedUser;
use crate::crypto::{
base64url_decode, base64url_encode, constant_time_eq, sha256_digest, sign_p256_raw,
verify_p256_raw,
};
use crate::error::{CryptoError, IntegrationError, TokenError};
use crate::session::OAuthSession;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct CnfClaim {
pub jkt: String,
}
impl CnfClaim {
#[must_use]
pub fn new(jkt: impl Into<String>) -> Self {
Self { jkt: jkt.into() }
}
#[must_use]
pub fn jkt(&self) -> &str {
&self.jkt
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct JwtAccessTokenClaims {
pub iss: String,
pub sub: String,
pub client_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub aud: Option<serde_json::Value>,
pub exp: u64,
#[serde(skip_serializing_if = "Option::is_none")]
pub nbf: Option<u64>,
pub iat: u64,
pub jti: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub scope: Option<String>,
pub cnf: CnfClaim,
}
impl JwtAccessTokenClaims {
#[must_use]
pub fn new(
iss: impl Into<String>,
sub: impl Into<String>,
exp: u64,
dpop_thumbprint: impl Into<String>,
) -> Self {
let iat = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
let mut raw = [0u8; 16];
rand::thread_rng().fill_bytes(&mut raw);
Self {
iss: iss.into(),
sub: sub.into(),
client_id: "https://app.example.com/client-metadata.json".to_string(),
aud: None,
exp,
nbf: None,
iat,
jti: crate::crypto::base64url_encode(&raw),
scope: None,
cnf: CnfClaim::new(dpop_thumbprint),
}
}
#[must_use]
pub fn with_audience(mut self, aud: impl Into<String>) -> Self {
self.aud = Some(serde_json::Value::String(aud.into()));
self
}
#[must_use]
pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
self.scope = Some(scope.into());
self
}
#[must_use]
pub fn with_nbf(mut self, nbf: u64) -> Self {
self.nbf = Some(nbf);
self
}
pub fn sign_jwt(
&self,
signing_key: &p256::ecdsa::SigningKey,
kid: Option<&str>,
) -> Result<String, CryptoError> {
let mut header_map = serde_json::Map::new();
header_map.insert(
"alg".to_string(),
serde_json::Value::String("ES256".to_string()),
);
header_map.insert(
"typ".to_string(),
serde_json::Value::String("at+jwt".to_string()),
);
if let Some(k) = kid {
header_map.insert("kid".to_string(), serde_json::Value::String(k.to_string()));
}
let header_str = serde_json::to_string(&serde_json::Value::Object(header_map))
.map_err(|e| CryptoError::Json(e.to_string()))?;
let payload_str =
serde_json::to_string(self).map_err(|e| CryptoError::Json(e.to_string()))?;
let header_b64 = base64url_encode(header_str.as_bytes());
let payload_b64 = base64url_encode(payload_str.as_bytes());
let signing_input = format!("{header_b64}.{payload_b64}");
let signature_bytes = sign_p256_raw(signing_key, signing_input.as_bytes())?;
let sig_b64 = base64url_encode(&signature_bytes);
Ok(format!("{signing_input}.{sig_b64}"))
}
}
pub trait AccessTokenValidator: Send + Sync + 'static {
fn validate_access_token(
&self,
token: &str,
dpop_thumbprint: &str,
) -> Pin<Box<dyn Future<Output = Result<AuthenticatedUser, IntegrationError>> + Send>>;
}
impl<T: AccessTokenValidator + ?Sized> AccessTokenValidator for Arc<T> {
fn validate_access_token(
&self,
token: &str,
dpop_thumbprint: &str,
) -> Pin<Box<dyn Future<Output = Result<AuthenticatedUser, IntegrationError>> + Send>> {
(**self).validate_access_token(token, dpop_thumbprint)
}
}
#[derive(Debug, Clone)]
pub struct JwtAccessTokenValidator {
trusted_keys: HashMap<String, VerifyingKey>,
default_key: Option<VerifyingKey>,
expected_issuer: Option<String>,
expected_audience: Option<String>,
expected_subject: Option<String>,
required_scopes: Vec<String>,
clock_skew_leeway: Duration,
}
impl Default for JwtAccessTokenValidator {
fn default() -> Self {
Self::new()
}
}
impl JwtAccessTokenValidator {
#[must_use]
pub fn new() -> Self {
Self {
trusted_keys: HashMap::new(),
default_key: None,
expected_issuer: None,
expected_audience: None,
expected_subject: None,
required_scopes: Vec::new(),
clock_skew_leeway: Duration::from_secs(60),
}
}
#[must_use]
pub fn with_trusted_key(mut self, kid: impl Into<String>, key: VerifyingKey) -> Self {
self.trusted_keys.insert(kid.into(), key);
self
}
#[must_use]
pub fn with_verifying_key(mut self, key: VerifyingKey) -> Self {
self.default_key = Some(key);
self
}
#[must_use]
pub fn with_expected_issuer(mut self, issuer: impl Into<String>) -> Self {
self.expected_issuer = Some(issuer.into());
self
}
#[must_use]
pub fn with_expected_audience(mut self, audience: impl Into<String>) -> Self {
self.expected_audience = Some(audience.into());
self
}
#[must_use]
pub fn with_expected_subject(mut self, subject: impl Into<String>) -> Self {
self.expected_subject = Some(subject.into());
self
}
#[must_use]
pub fn with_required_scope(mut self, scope: impl Into<String>) -> Self {
self.required_scopes.push(scope.into());
self
}
#[must_use]
pub fn with_clock_skew(mut self, leeway: Duration) -> Self {
self.clock_skew_leeway = leeway;
self
}
pub fn verify_token_sync(
&self,
token: &str,
dpop_thumbprint: &str,
) -> Result<AuthenticatedUser, IntegrationError> {
let parts: Vec<&str> = token.trim().split('.').collect();
if parts.len() != 3 {
return Err(IntegrationError::Token(TokenError::MalformedToken(
format!("Expected 3 parts in compact JWT, got {}", parts.len()),
)));
}
let header_bytes = base64url_decode(parts[0])
.map_err(|e| IntegrationError::Token(TokenError::MalformedToken(e.to_string())))?;
let payload_bytes = base64url_decode(parts[1])
.map_err(|e| IntegrationError::Token(TokenError::MalformedToken(e.to_string())))?;
let signature_bytes = base64url_decode(parts[2])
.map_err(|e| IntegrationError::Token(TokenError::MalformedToken(e.to_string())))?;
let header_val: serde_json::Value = serde_json::from_slice(&header_bytes)
.map_err(|e| IntegrationError::Token(TokenError::MalformedToken(e.to_string())))?;
let alg = header_val
.get("alg")
.and_then(|v| v.as_str())
.ok_or_else(|| {
IntegrationError::Token(TokenError::MalformedToken(
"Missing alg header".to_string(),
))
})?;
if alg != "ES256" {
return Err(IntegrationError::Token(TokenError::MalformedToken(
format!("Unsupported alg in access token: expected 'ES256', got '{alg}'"),
)));
}
let typ = header_val.get("typ").and_then(|v| v.as_str());
if !typ
.map(|t| {
t.eq_ignore_ascii_case("at+jwt") || t.eq_ignore_ascii_case("application/at+jwt")
})
.unwrap_or(false)
{
return Err(IntegrationError::Token(TokenError::MalformedToken(
format!(
"Unsupported typ in access token: expected 'at+jwt', got {:?}",
typ
),
)));
}
let kid = header_val.get("kid").and_then(|v| v.as_str());
let verifying_key = if let Some(k) = kid {
match self.trusted_keys.get(k) {
Some(key) => Some(key),
None if self.trusted_keys.is_empty() => self.default_key.as_ref(),
None => None,
}
} else if let Some(ref default_k) = self.default_key {
Some(default_k)
} else if self.trusted_keys.len() == 1 {
self.trusted_keys.values().next()
} else {
None
}
.ok_or_else(|| {
IntegrationError::AuthFailed(
"No matching trusted verifying key found for access token".to_string(),
)
})?;
let signing_input = format!("{}.{}", parts[0], parts[1]);
verify_p256_raw(verifying_key, signing_input.as_bytes(), &signature_bytes)
.map_err(|_| IntegrationError::Token(TokenError::InvalidSignature))?;
let claims: JwtAccessTokenClaims = serde_json::from_slice(&payload_bytes)
.map_err(|e| IntegrationError::Token(TokenError::MalformedToken(e.to_string())))?;
let now_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_err(|e| IntegrationError::Internal(e.to_string()))?
.as_secs();
let leeway_secs = self.clock_skew_leeway.as_secs();
if now_secs.saturating_add(leeway_secs) >= claims.exp {
return Err(IntegrationError::Token(TokenError::Expired {
exp: claims.exp,
now: now_secs,
}));
}
if let Some(nbf) = claims.nbf {
if nbf > now_secs.saturating_add(leeway_secs) {
return Err(IntegrationError::Token(TokenError::NotYetValid {
nbf,
now: now_secs,
}));
}
}
if claims.iss.trim().is_empty() {
return Err(IntegrationError::Token(TokenError::MissingIssuer));
}
let Some(ref exp_iss) = self.expected_issuer else {
return Err(IntegrationError::AuthFailed(
"Validator misconfigured: an expected issuer is required before any \
access token can be accepted (RFC 9068 § 4 / RFC 9449 § 7.2)"
.to_string(),
));
};
if claims.iss != *exp_iss {
return Err(IntegrationError::Token(TokenError::IssuerMismatch {
expected: exp_iss.clone(),
actual: claims.iss.clone(),
}));
}
if claims.aud.is_none() {
return Err(IntegrationError::Token(TokenError::MissingAudience));
}
let Some(ref exp_aud) = self.expected_audience else {
return Err(IntegrationError::AuthFailed(
"Validator misconfigured: an expected audience is required before any \
access token can be accepted (RFC 9068 § 4)"
.to_string(),
));
};
{
let matches_aud = match &claims.aud {
Some(serde_json::Value::String(s)) => s == exp_aud,
Some(serde_json::Value::Array(arr)) => arr
.iter()
.any(|item| item.as_str().map(|s| s == exp_aud).unwrap_or(false)),
_ => false,
};
if !matches_aud {
return Err(IntegrationError::Token(TokenError::AudienceMismatch {
expected: exp_aud.clone(),
actual: claims
.aud
.as_ref()
.map(ToString::to_string)
.unwrap_or_else(|| "none".to_string()),
}));
}
}
if claims.sub.trim().is_empty() {
return Err(IntegrationError::Token(TokenError::MissingDid));
}
if let Some(ref exp_sub) = self.expected_subject {
if &claims.sub != exp_sub {
return Err(IntegrationError::Token(TokenError::SubMismatch {
expected: exp_sub.clone(),
actual: claims.sub.clone(),
}));
}
}
if !self.required_scopes.is_empty() {
let granted_scopes: Vec<&str> = claims
.scope
.as_deref()
.unwrap_or("")
.split_whitespace()
.collect();
for required in &self.required_scopes {
if !granted_scopes.contains(&required.as_str()) {
return Err(IntegrationError::Token(TokenError::MissingAtprotoScope(
claims.scope.unwrap_or_default(),
)));
}
}
}
if !constant_time_eq(claims.cnf.jkt.as_bytes(), dpop_thumbprint.as_bytes()) {
return Err(IntegrationError::Token(TokenError::CnfThumbprintMismatch {
expected_jkt: claims.cnf.jkt,
actual_jkt: dpop_thumbprint.to_string(),
}));
}
Ok(AuthenticatedUser {
did: claims.sub,
access_token: token.to_string(),
dpop_thumbprint: dpop_thumbprint.to_string(),
scope: claims.scope,
})
}
}
impl AccessTokenValidator for JwtAccessTokenValidator {
fn validate_access_token(
&self,
token: &str,
dpop_thumbprint: &str,
) -> Pin<Box<dyn Future<Output = Result<AuthenticatedUser, IntegrationError>> + Send>> {
let res = self.verify_token_sync(token, dpop_thumbprint);
Box::pin(std::future::ready(res))
}
}
#[derive(Debug, Clone, Default)]
pub struct InMemoryTokenValidator {
tokens: Arc<RwLock<HashMap<[u8; 32], RegisteredToken>>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RegisteredToken {
pub did: String,
pub dpop_thumbprint: String,
pub scope: Option<String>,
pub expires_at: Option<SystemTime>,
}
impl InMemoryTokenValidator {
#[must_use]
pub fn new() -> Self {
Self {
tokens: Arc::new(RwLock::new(HashMap::new())),
}
}
pub fn register_token(
&self,
token: impl AsRef<[u8]>,
did: impl Into<String>,
dpop_thumbprint: impl Into<String>,
scope: Option<String>,
expires_at: Option<SystemTime>,
) {
let digest = sha256_digest(token.as_ref());
let mut guard = self.tokens.write();
guard.insert(
digest,
RegisteredToken {
did: did.into(),
dpop_thumbprint: dpop_thumbprint.into(),
scope,
expires_at,
},
);
}
pub fn register_session(&self, session: &OAuthSession) {
self.register_token(
&session.access_token,
&session.sub,
session.dpop_key.jwk_thumbprint(),
session.scope.clone(),
session.expires_at,
);
}
pub fn revoke_token(&self, token: &str) {
let digest = sha256_digest(token.as_bytes());
let mut guard = self.tokens.write();
guard.remove(&digest);
}
pub fn prune_expired(&self) -> usize {
let now = SystemTime::now();
let mut guard = self.tokens.write();
let before = guard.len();
guard.retain(|_, entry| match entry.expires_at {
Some(exp) => exp > now,
None => true,
});
before - guard.len()
}
#[must_use]
pub fn len(&self) -> usize {
self.tokens.read().len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn validate_sync(
&self,
token: &str,
dpop_thumbprint: &str,
) -> Result<AuthenticatedUser, IntegrationError> {
let digest = sha256_digest(token.as_bytes());
let guard = self.tokens.read();
let entry = guard.get(&digest).ok_or_else(|| {
IntegrationError::Token(TokenError::MalformedToken(
"Access token is not registered or has been revoked".to_string(),
))
})?;
if let Some(exp) = entry.expires_at {
let now = SystemTime::now();
if now > exp {
let exp_secs = exp
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
let now_secs = now
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
return Err(IntegrationError::Token(TokenError::Expired {
exp: exp_secs,
now: now_secs,
}));
}
}
if !constant_time_eq(entry.dpop_thumbprint.as_bytes(), dpop_thumbprint.as_bytes()) {
return Err(IntegrationError::Token(TokenError::CnfThumbprintMismatch {
expected_jkt: entry.dpop_thumbprint.clone(),
actual_jkt: dpop_thumbprint.to_string(),
}));
}
Ok(AuthenticatedUser {
did: entry.did.clone(),
access_token: token.to_string(),
dpop_thumbprint: dpop_thumbprint.to_string(),
scope: entry.scope.clone(),
})
}
}
impl AccessTokenValidator for InMemoryTokenValidator {
fn validate_access_token(
&self,
token: &str,
dpop_thumbprint: &str,
) -> Pin<Box<dyn Future<Output = Result<AuthenticatedUser, IntegrationError>> + Send>> {
let res = self.validate_sync(token, dpop_thumbprint);
Box::pin(std::future::ready(res))
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod tests {
use super::*;
use crate::dpop::DPoPKey;
use p256::ecdsa::SigningKey;
use rand::thread_rng;
#[test]
fn test_jwt_access_token_signing_and_validation_roundtrip() {
let auth_key = SigningKey::random(&mut thread_rng());
let auth_verifying_key = *auth_key.verifying_key();
let client_dpop_key = DPoPKey::generate();
let client_jkt = client_dpop_key.jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&client_jkt,
)
.with_audience("https://pds.example.com")
.with_scope("atproto transition:generic");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(auth_verifying_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com")
.with_required_scope("atproto");
let user = validator.verify_token_sync(&jwt, &client_jkt).unwrap();
assert_eq!(user.did, "did:plc:alice123");
assert_eq!(user.access_token, jwt);
assert_eq!(user.dpop_thumbprint, client_jkt);
assert_eq!(user.scope.as_deref(), Some("atproto transition:generic"));
}
#[test]
fn test_jwt_access_token_cnf_jkt_mismatch_fails() {
let auth_key = SigningKey::random(&mut thread_rng());
let auth_verifying_key = *auth_key.verifying_key();
let alice_dpop_key = DPoPKey::generate();
let attacker_dpop_key = DPoPKey::generate();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
alice_dpop_key.jwk_thumbprint(),
)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(auth_verifying_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let err = validator
.verify_token_sync(&jwt, &attacker_dpop_key.jwk_thumbprint())
.unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::CnfThumbprintMismatch { .. })
));
}
#[test]
fn test_jwt_access_token_expired_fails() {
let auth_key = SigningKey::random(&mut thread_rng());
let auth_verifying_key = *auth_key.verifying_key();
let client_dpop_key = DPoPKey::generate();
let jkt = client_dpop_key.jwk_thumbprint();
let claims =
JwtAccessTokenClaims::new("https://auth.example.com", "did:plc:alice123", 1000, &jkt)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(auth_verifying_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com")
.with_clock_skew(Duration::ZERO);
let err = validator.verify_token_sync(&jwt, &jkt).unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::Expired { .. })
));
}
#[test]
fn test_jwt_access_token_issuer_mismatch_fails() {
let auth_key = SigningKey::random(&mut thread_rng());
let auth_verifying_key = *auth_key.verifying_key();
let client_dpop_key = DPoPKey::generate();
let jkt = client_dpop_key.jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://malicious-issuer.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(auth_verifying_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let err = validator.verify_token_sync(&jwt, &jkt).unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::IssuerMismatch { .. })
));
}
#[test]
fn test_jwt_access_token_audience_mismatch_fails() {
let auth_key = SigningKey::random(&mut thread_rng());
let auth_verifying_key = *auth_key.verifying_key();
let client_dpop_key = DPoPKey::generate();
let jkt = client_dpop_key.jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
)
.with_audience("https://other-resource.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(auth_verifying_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let err = validator.verify_token_sync(&jwt, &jkt).unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::AudienceMismatch { .. })
));
}
#[test]
fn test_in_memory_token_validator_lifecycle() {
let validator = InMemoryTokenValidator::new();
let token = "active_session_token_xyz";
let did = "did:plc:carol789";
let dpop_key = DPoPKey::generate();
let jkt = dpop_key.jwk_thumbprint();
validator.register_token(token, did, &jkt, Some("atproto".to_string()), None);
let user = validator.validate_sync(token, &jkt).unwrap();
assert_eq!(user.did, did);
assert_eq!(user.dpop_thumbprint, jkt);
let wrong_key = DPoPKey::generate();
let err = validator
.validate_sync(token, &wrong_key.jwk_thumbprint())
.unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::CnfThumbprintMismatch { .. })
));
validator.revoke_token(token);
let err_revoked = validator.validate_sync(token, &jkt).unwrap_err();
assert!(matches!(
err_revoked,
IntegrationError::Token(TokenError::MalformedToken(_))
));
}
#[test]
fn test_jwt_access_token_wrong_typ_rejected() {
let auth_key = SigningKey::random(&mut thread_rng());
let client_jkt = DPoPKey::generate().jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:typ-test",
now + 3600,
&client_jkt,
)
.with_audience("https://pds.example.com");
let payload_str = serde_json::to_string(&claims).unwrap();
let header_typ_jwt = r#"{"alg":"ES256","typ":"JWT"}"#;
let header_b64 = base64url_encode(header_typ_jwt.as_bytes());
let payload_b64 = base64url_encode(payload_str.as_bytes());
let signing_input = format!("{header_b64}.{payload_b64}");
let sig = sign_p256_raw(&auth_key, signing_input.as_bytes()).unwrap();
let jwt_wrong_typ = format!("{signing_input}.{}", base64url_encode(&sig));
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(*auth_key.verifying_key())
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let err = validator
.verify_token_sync(&jwt_wrong_typ, &client_jkt)
.unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::MalformedToken(ref msg))
if msg.contains("typ") || msg.contains("at+jwt")
));
}
#[test]
fn test_jwt_access_token_unknown_kid_fail_closed() {
let auth_key = SigningKey::random(&mut thread_rng());
let client_jkt = DPoPKey::generate().jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:kid-test",
now + 3600,
&client_jkt,
)
.with_audience("https://pds.example.com");
let payload_str = serde_json::to_string(&claims).unwrap();
let header_unknown_kid = r#"{"alg":"ES256","typ":"at+jwt","kid":"key-1999"}"#;
let header_b64 = base64url_encode(header_unknown_kid.as_bytes());
let payload_b64 = base64url_encode(payload_str.as_bytes());
let signing_input = format!("{header_b64}.{payload_b64}");
let sig = sign_p256_raw(&auth_key, signing_input.as_bytes()).unwrap();
let jwt = format!("{signing_input}.{}", base64url_encode(&sig));
let validator = JwtAccessTokenValidator::new()
.with_trusted_key("key-2024", *auth_key.verifying_key())
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let err = validator.verify_token_sync(&jwt, &client_jkt).unwrap_err();
assert!(matches!(err, IntegrationError::AuthFailed(_)));
}
#[test]
fn test_jwt_access_token_missing_audience_rejected() {
let auth_key = SigningKey::random(&mut thread_rng());
let client_jkt = DPoPKey::generate().jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims_no_aud = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:aud-test",
now + 3600,
&client_jkt,
);
let jwt_no_aud = claims_no_aud.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(*auth_key.verifying_key())
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let err = validator
.verify_token_sync(&jwt_no_aud, &client_jkt)
.unwrap_err();
assert!(matches!(
err,
IntegrationError::Token(TokenError::MissingAudience)
));
}
#[test]
fn test_jwt_access_token_unconfigured_audience_fails_closed() {
let auth_key = SigningKey::random(&mut thread_rng());
let client_jkt = DPoPKey::generate().jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:aud-closed-test",
now + 3600,
&client_jkt,
)
.with_audience("https://some-other-resource.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(*auth_key.verifying_key())
.with_expected_issuer("https://auth.example.com");
let err = validator.verify_token_sync(&jwt, &client_jkt).unwrap_err();
assert!(
matches!(err, IntegrationError::AuthFailed(ref msg) if msg.contains("audience")),
"expected misconfiguration (missing expected audience) rejection, got {err:?}"
);
}
#[test]
fn test_jwt_access_token_unconfigured_issuer_fails_closed() {
let auth_key = SigningKey::random(&mut thread_rng());
let client_jkt = DPoPKey::generate().jwk_thumbprint();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:iss-closed-test",
now + 3600,
&client_jkt,
)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(*auth_key.verifying_key())
.with_expected_audience("https://pds.example.com");
let err = validator.verify_token_sync(&jwt, &client_jkt).unwrap_err();
assert!(
matches!(err, IntegrationError::AuthFailed(ref msg) if msg.contains("issuer")),
"expected misconfiguration (missing expected issuer) rejection, got {err:?}"
);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod m5_enforcement_tests {
use super::*;
use p256::ecdsa::SigningKey;
use rand::thread_rng;
fn signed_token(claims: &JwtAccessTokenClaims, key: &SigningKey) -> String {
claims.sign_jwt(key, None).unwrap()
}
#[test]
fn test_m5_exp_boundary_is_exact() {
let key = SigningKey::random(&mut thread_rng());
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now, "dpop_jkt",
);
let token = signed_token(&claims, &key);
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(*key.verifying_key())
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let res = validator.verify_token_sync(&token, "dpop_jkt");
assert!(
matches!(
res,
Err(IntegrationError::Token(TokenError::Expired { .. }))
),
"token at the exp boundary must be expired"
);
}
#[test]
fn test_m5_issuer_comparison_is_exact() {
let key = SigningKey::random(&mut thread_rng());
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com/", "did:plc:alice123",
now + 3600,
"dpop_jkt",
);
let token = signed_token(&claims, &key);
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(*key.verifying_key())
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
let res = validator.verify_token_sync(&token, "dpop_jkt");
assert!(
matches!(
res,
Err(IntegrationError::Token(TokenError::IssuerMismatch { .. }))
),
"issuer comparison must be exact"
);
}
#[test]
fn test_m5_wire_rejects_missing_mandatory_claims() {
let raw_without_mandatory = r#"{
"iss": "https://auth.example.com",
"sub": "did:plc:alice123",
"exp": 9999999999,
"scope": "atproto",
"cnf": {"jkt": "dpop_jkt"}
}"#;
let claims: Result<JwtAccessTokenClaims, _> = serde_json::from_str(raw_without_mandatory);
assert!(
claims.is_err(),
"missing client_id/iat/jti must fail deserialization (RFC 9068 § 2.2 mandatory claims)"
);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic, missing_docs)]
mod mutation_killer_tests {
use super::*;
use crate::dpop::DPoPKey;
use p256::ecdsa::SigningKey;
use rand::thread_rng;
fn now_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs()
}
fn minted_validator(now: u64) -> (JwtAccessTokenValidator, String, String) {
let auth_key = SigningKey::random(&mut thread_rng());
let auth_verifying_key = *auth_key.verifying_key();
let client_jkt = DPoPKey::generate().jwk_thumbprint();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&client_jkt,
)
.with_audience("https://pds.example.com")
.with_scope("atproto transition:generic");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(auth_verifying_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com")
.with_required_scope("atproto");
(validator, jwt, client_jkt)
}
#[test]
fn killer_cnf_claim_jkt_accessor() {
let cnf = CnfClaim::new("thumbprint-abc");
assert_eq!(
cnf.jkt(),
"thumbprint-abc",
"jkt() must return the stored value"
);
assert_ne!(cnf.jkt(), "");
let round: CnfClaim = serde_json::from_str(&serde_json::to_string(&cnf).unwrap()).unwrap();
assert_eq!(round.jkt(), "thumbprint-abc");
}
#[test]
fn killer_with_expected_subject_is_enforced() {
let now = now_secs();
let (validator, jwt, jkt) = minted_validator(now);
let ok = validator.clone().with_expected_subject("did:plc:alice123");
assert!(ok.verify_token_sync(&jwt, &jkt).is_ok());
let wrong = validator.with_expected_subject("did:plc:bob");
assert!(
wrong.verify_token_sync(&jwt, &jkt).is_err(),
"with_expected_subject must reject a different subject"
);
}
#[test]
fn killer_with_trusted_key_kid_routing() {
let now = now_secs();
let auth_key = SigningKey::random(&mut thread_rng());
let verifying = *auth_key.verifying_key();
let jkt = DPoPKey::generate().jwk_thumbprint();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, Some("key-2026")).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_trusted_key("key-2026", verifying)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
assert!(validator.verify_token_sync(&jwt, &jkt).is_ok());
let validator2 = JwtAccessTokenValidator::new()
.with_trusted_key("other-kid", verifying)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
assert!(
validator2.verify_token_sync(&jwt, &jkt).is_err(),
"unknown kid with non-empty trusted_keys must not fall back"
);
}
#[test]
fn killer_nbf_and_exp_boundaries_are_strict() {
let auth_key = SigningKey::random(&mut thread_rng());
let verifying = *auth_key.verifying_key();
let jkt = DPoPKey::generate().jwk_thumbprint();
let now = now_secs();
let future_nbf = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
)
.with_audience("https://pds.example.com")
.with_nbf(now + 3600);
let jwt_future = future_nbf.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(verifying)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
assert!(
validator.verify_token_sync(&jwt_future, &jkt).is_err(),
"nbf far in the future must be rejected"
);
let zero_leeway = validator.with_clock_skew(Duration::ZERO);
let exp_now =
JwtAccessTokenClaims::new("https://auth.example.com", "did:plc:alice123", now, &jkt)
.with_audience("https://pds.example.com");
let jwt_exp_now = exp_now.sign_jwt(&auth_key, None).unwrap();
assert!(
zero_leeway.verify_token_sync(&jwt_exp_now, &jkt).is_err(),
"exp == now with zero leeway must be expired"
);
let live = minted_validator(now);
assert!(
live.0.verify_token_sync(&live.1, &live.2).is_ok(),
"exp > now must validate"
);
}
#[test]
fn killer_audience_must_match_exactly() {
let now = now_secs();
let auth_key = SigningKey::random(&mut thread_rng());
let verifying = *auth_key.verifying_key();
let jkt = DPoPKey::generate().jwk_thumbprint();
let mut claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
);
claims.aud = Some(serde_json::Value::Array(vec![
serde_json::Value::String("https://other.example.com".to_string()),
serde_json::Value::String("https://pds.example.com".to_string()),
]));
let jwt_arr = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(verifying)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
assert!(
validator.verify_token_sync(&jwt_arr, &jkt).is_ok(),
"array audience containing the RS must be accepted"
);
let mut wrong_claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
);
wrong_claims.aud = Some(serde_json::Value::Array(vec![serde_json::Value::String(
"https://elsewhere.example.com".to_string(),
)]));
let jwt_wrong = wrong_claims.sign_jwt(&auth_key, None).unwrap();
assert!(
validator.verify_token_sync(&jwt_wrong, &jkt).is_err(),
"array audience without the RS must be rejected"
);
let mut num_claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
);
num_claims.aud = Some(serde_json::Value::Array(vec![serde_json::Value::from(
42u64,
)]));
let jwt_num = num_claims.sign_jwt(&auth_key, None).unwrap();
assert!(
validator.verify_token_sync(&jwt_num, &jkt).is_err(),
"non-string audience entries must never match"
);
}
#[test]
fn killer_required_scope_must_be_present() {
let now = now_secs();
let auth_key = SigningKey::random(&mut thread_rng());
let verifying = *auth_key.verifying_key();
let jkt = DPoPKey::generate().jwk_thumbprint();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_verifying_key(verifying)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com")
.with_required_scope("atproto");
assert!(
validator.verify_token_sync(&jwt, &jkt).is_err(),
"token without scope claim must fail a required-scope check"
);
}
#[test]
fn kid_routing_none_when_multiple_trusted_keys_and_no_default() {
let now = now_secs();
let auth_key = SigningKey::random(&mut thread_rng());
let verifying = *auth_key.verifying_key();
let other_key = *SigningKey::random(&mut thread_rng()).verifying_key();
let jkt = DPoPKey::generate().jwk_thumbprint();
let claims = JwtAccessTokenClaims::new(
"https://auth.example.com",
"did:plc:alice123",
now + 3600,
&jkt,
)
.with_audience("https://pds.example.com");
let jwt = claims.sign_jwt(&auth_key, None).unwrap();
let validator = JwtAccessTokenValidator::new()
.with_trusted_key("kid-a", verifying)
.with_trusted_key("kid-b", other_key)
.with_expected_issuer("https://auth.example.com")
.with_expected_audience("https://pds.example.com");
assert!(
validator.verify_token_sync(&jwt, &jkt).is_err(),
"kid-less token with multiple trusted keys and no default must fail"
);
}
#[test]
fn killer_in_memory_validator_full_lifecycle() {
let store = InMemoryTokenValidator::new();
assert!(store.is_empty(), "fresh store must be empty");
assert_eq!(store.len(), 0, "fresh store len must be 0");
store.register_token(
"token-a",
"did:plc:alice123",
"jkt-a",
Some("atproto".to_string()),
Some(SystemTime::now() + Duration::from_secs(3600)),
);
assert_eq!(store.len(), 1, "len must count registered tokens");
assert!(!store.is_empty());
let user = store.validate_sync("token-a", "jkt-a").unwrap();
assert_eq!(user.did, "did:plc:alice123");
assert_eq!(user.access_token, "token-a");
assert_eq!(user.dpop_thumbprint, "jkt-a");
assert_eq!(user.scope.as_deref(), Some("atproto"));
assert!(
store.validate_sync("token-a", "jkt-b").is_err(),
"mismatched thumbprint must be rejected"
);
assert!(
store.validate_sync("never-registered", "jkt-a").is_err(),
"unknown token must be rejected"
);
store.register_token(
"token-b",
"did:plc:bob",
"jkt-b",
None,
Some(SystemTime::now() + Duration::from_secs(3600)),
);
assert_eq!(store.len(), 2);
store.revoke_token("token-a");
assert_eq!(store.len(), 1, "revoke must remove exactly one entry");
assert!(store.validate_sync("token-a", "jkt-a").is_err());
assert!(store.validate_sync("token-b", "jkt-b").is_ok());
store.register_token(
"token-expired",
"did:plc:expired",
"jkt-e",
None,
Some(SystemTime::now() - Duration::from_secs(1)),
);
assert_eq!(store.len(), 2);
let pruned = store.prune_expired();
assert_eq!(
pruned, 1,
"prune_expired must evict exactly the expired entry"
);
assert_eq!(store.len(), 1);
assert!(store.validate_sync("token-expired", "jkt-e").is_err());
assert!(store.validate_sync("token-b", "jkt-b").is_ok());
store.register_token(
"token-now",
"did:plc:now",
"jkt-n",
None,
Some(SystemTime::now()),
);
assert!(
store.validate_sync("token-now", "jkt-n").is_err(),
"expires_at == now must be expired"
);
}
#[test]
fn killer_register_session_binds_session_thumbprint() {
let key = DPoPKey::generate();
let session = crate::session::OAuthSession::new(
"did:plc:alice123",
"at_session_token",
None,
"DPoP",
Some("atproto".to_string()),
Some(3600),
key.clone(),
None,
None,
None,
)
.unwrap();
let expected_jkt = key.jwk_thumbprint();
let store = InMemoryTokenValidator::new();
store.register_session(&session);
assert_eq!(store.len(), 1);
let user = store
.validate_sync("at_session_token", &expected_jkt)
.unwrap();
assert_eq!(user.did, "did:plc:alice123");
assert_eq!(user.scope.as_deref(), Some("atproto"));
assert!(store
.validate_sync("at_session_token", "some-other-jkt")
.is_err());
}
}