use std::collections::{HashMap, HashSet};
use std::fmt;
use std::sync::Arc;
use std::time::{Duration, SystemTime};
use async_trait::async_trait;
use jsonwebtoken::{decode, decode_header, Algorithm, DecodingKey, Validation};
use moka::future::Cache;
use rvoip_core_traits::identity::IdentityAssurance;
use rvoip_core_traits::ids::IdentityId;
use serde::Deserialize;
use tracing::{debug, warn};
use url::Url;
use crate::bearer::{
unix_time_from_seconds, validate_optional_token_id, AuthenticatedPrincipal,
AuthenticationMethod, BearerAuthError, BearerValidator, ValidatedBearer,
};
use crate::providers::{
CredentialAuthError, TokenRevocationChecker, TokenRevocationContext, TokenRevocationStatus,
};
pub const DEFAULT_JWKS_CACHE_TTL: Duration = Duration::from_secs(3600);
const JWKS_CACHE_MAX_CAPACITY: u64 = 64;
#[derive(Deserialize)]
struct JwksDocument {
keys: Vec<JwksKey>,
}
impl fmt::Debug for JwksDocument {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("JwksDocument")
.field("key_count", &self.keys.len())
.finish()
}
}
#[derive(Deserialize)]
struct JwksKey {
kty: String,
kid: Option<String>,
n: Option<String>,
e: Option<String>,
#[allow(dead_code)] crv: Option<String>,
x: Option<String>,
y: Option<String>,
}
impl fmt::Debug for JwksKey {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("JwksKey")
.field("key_type_class", &jwk_key_type_class(&self.kty))
.field("key_id_present", &self.kid.is_some())
.field("key_id_bytes", &self.kid.as_ref().map_or(0, String::len))
.field("rsa_modulus_present", &self.n.is_some())
.field("rsa_modulus_bytes", &self.n.as_ref().map_or(0, String::len))
.field("rsa_exponent_present", &self.e.is_some())
.field(
"rsa_exponent_bytes",
&self.e.as_ref().map_or(0, String::len),
)
.field("curve_present", &self.crv.is_some())
.field("curve_bytes", &self.crv.as_ref().map_or(0, String::len))
.field("x_coordinate_present", &self.x.is_some())
.field(
"x_coordinate_bytes",
&self.x.as_ref().map_or(0, String::len),
)
.field("y_coordinate_present", &self.y.is_some())
.field(
"y_coordinate_bytes",
&self.y.as_ref().map_or(0, String::len),
)
.finish()
}
}
fn jwk_key_type_class(key_type: &str) -> &'static str {
match key_type {
"RSA" => "rsa",
"EC" => "ec",
"oct" => "symmetric",
_ => "other",
}
}
#[derive(Deserialize)]
struct TokenClaims {
sub: String,
#[serde(default)]
iss: Option<String>,
#[serde(default)]
iat: Option<u64>,
#[serde(default)]
exp: Option<u64>,
#[serde(default)]
jti: Option<String>,
#[serde(default)]
scope: Option<String>,
#[serde(default)]
scopes: Option<Vec<String>>,
#[serde(default)]
roles: Option<Vec<String>>,
#[serde(default)]
realm_access: Option<RoleAccess>,
#[serde(default)]
resource_access: Option<HashMap<String, RoleAccess>>,
#[serde(default, alias = "tenant", alias = "tid")]
tenant_id: Option<String>,
}
impl fmt::Debug for TokenClaims {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("TokenClaims")
.field("subject_present", &!self.sub.is_empty())
.field("issuer_present", &self.iss.is_some())
.field("issued_at_present", &self.iat.is_some())
.field("expires_at_present", &self.exp.is_some())
.field("token_id_present", &self.jti.is_some())
.field("scope_present", &self.scope.is_some())
.field("scope_bytes", &self.scope.as_ref().map_or(0, String::len))
.field(
"scope_list_count",
&self.scopes.as_ref().map_or(0, Vec::len),
)
.field("role_count", &self.roles.as_ref().map_or(0, Vec::len))
.field("realm_access_present", &self.realm_access.is_some())
.field(
"resource_access_count",
&self.resource_access.as_ref().map_or(0, HashMap::len),
)
.field("tenant_present", &self.tenant_id.is_some())
.finish()
}
}
#[derive(Deserialize)]
struct RoleAccess {
#[serde(default)]
roles: Vec<String>,
}
impl fmt::Debug for RoleAccess {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RoleAccess")
.field("role_count", &self.roles.len())
.finish()
}
}
#[derive(Clone)]
pub struct JwksJwtValidator {
inner: Arc<Inner>,
}
struct Inner {
jwks_url: Url,
client: reqwest::Client,
cache: Cache<String, DecodingKey>,
validation: Validation,
revocation_checker: Option<Arc<dyn TokenRevocationChecker>>,
require_jti: bool,
}
impl JwksJwtValidator {
pub fn new(jwks_url: Url) -> Self {
let mut validation = Validation::new(Algorithm::RS256);
validation.validate_aud = false;
Self::new_with_validation(jwks_url, validation)
}
pub fn new_with_validation(jwks_url: Url, validation: Validation) -> Self {
let client = reqwest::Client::builder()
.user_agent("rvoip-auth-core/0.1 (jwks)")
.timeout(Duration::from_secs(10))
.build()
.expect("reqwest::Client::builder default config never fails");
Self {
inner: Arc::new(Inner {
jwks_url,
client,
cache: Cache::builder()
.max_capacity(JWKS_CACHE_MAX_CAPACITY)
.time_to_live(DEFAULT_JWKS_CACHE_TTL)
.build(),
validation,
revocation_checker: None,
require_jti: false,
}),
}
}
pub fn with_cache_ttl(self, ttl: Duration) -> Self {
let inner = &*self.inner;
let new_cache = Cache::builder()
.max_capacity(JWKS_CACHE_MAX_CAPACITY)
.time_to_live(ttl)
.build();
Self {
inner: Arc::new(Inner {
jwks_url: inner.jwks_url.clone(),
client: inner.client.clone(),
cache: new_cache,
validation: inner.validation.clone(),
revocation_checker: inner.revocation_checker.clone(),
require_jti: inner.require_jti,
}),
}
}
pub fn with_audience<I, S>(self, audiences: I) -> Self
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let inner = &*self.inner;
let mut validation = inner.validation.clone();
let auds: HashSet<String> = audiences
.into_iter()
.map(|s| s.as_ref().to_string())
.collect();
validation.set_audience(&auds.into_iter().collect::<Vec<_>>());
validation.validate_aud = true;
Self {
inner: Arc::new(Inner {
jwks_url: inner.jwks_url.clone(),
client: inner.client.clone(),
cache: inner.cache.clone(),
validation,
revocation_checker: inner.revocation_checker.clone(),
require_jti: inner.require_jti,
}),
}
}
pub fn with_issuer<I, S>(self, issuers: I) -> Self
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let inner = &*self.inner;
let mut validation = inner.validation.clone();
validation.set_issuer(
&issuers
.into_iter()
.map(|s| s.as_ref().to_string())
.collect::<Vec<_>>(),
);
Self {
inner: Arc::new(Inner {
jwks_url: inner.jwks_url.clone(),
client: inner.client.clone(),
cache: inner.cache.clone(),
validation,
revocation_checker: inner.revocation_checker.clone(),
require_jti: inner.require_jti,
}),
}
}
pub fn with_algorithms(self, algorithms: Vec<Algorithm>) -> Self {
let inner = &*self.inner;
let mut validation = inner.validation.clone();
validation.algorithms = algorithms;
Self {
inner: Arc::new(Inner {
jwks_url: inner.jwks_url.clone(),
client: inner.client.clone(),
cache: inner.cache.clone(),
validation,
revocation_checker: inner.revocation_checker.clone(),
require_jti: inner.require_jti,
}),
}
}
pub fn with_revocation_checker(self, checker: Arc<dyn TokenRevocationChecker>) -> Self {
let inner = &*self.inner;
Self {
inner: Arc::new(Inner {
jwks_url: inner.jwks_url.clone(),
client: inner.client.clone(),
cache: inner.cache.clone(),
validation: inner.validation.clone(),
revocation_checker: Some(checker),
require_jti: inner.require_jti,
}),
}
}
pub fn with_required_jti(self) -> Self {
let inner = &*self.inner;
Self {
inner: Arc::new(Inner {
jwks_url: inner.jwks_url.clone(),
client: inner.client.clone(),
cache: inner.cache.clone(),
validation: inner.validation.clone(),
revocation_checker: inner.revocation_checker.clone(),
require_jti: true,
}),
}
}
pub fn into_arc(self) -> Arc<dyn BearerValidator> {
Arc::new(self)
}
async fn resolve_key(&self, kid: &str) -> Result<DecodingKey, BearerAuthError> {
if let Some(key) = self.inner.cache.get(kid).await {
return Ok(key);
}
debug!(
key_id_present = !kid.is_empty(),
"jwks: cache miss, refetching"
);
let doc = self.fetch_jwks().await?;
for jwk in doc.keys {
let Some(jwk_kid) = jwk.kid.clone() else {
continue;
};
match decoding_key_from_jwk(&jwk) {
Ok(key) => {
self.inner.cache.insert(jwk_kid, key).await;
}
Err(_) => {
warn!(
key_id_present = !jwk_kid.is_empty(),
error_class = "invalid-jwk",
"jwks: skipping unparseable key"
);
}
}
}
self.inner
.cache
.get(kid)
.await
.ok_or_else(|| BearerAuthError::Invalid(format!("no signing key for kid={}", kid)))
}
async fn fetch_jwks(&self) -> Result<JwksDocument, BearerAuthError> {
let resp = self
.inner
.client
.get(self.inner.jwks_url.clone())
.send()
.await
.map_err(|e| BearerAuthError::Unavailable(format!("JWKS fetch: {e}")))?;
if !resp.status().is_success() {
return Err(BearerAuthError::Unavailable(format!(
"JWKS endpoint returned {}",
resp.status()
)));
}
resp.json::<JwksDocument>()
.await
.map_err(|e| BearerAuthError::Unavailable(format!("JWKS parse: {e}")))
}
}
fn decoding_key_from_jwk(jwk: &JwksKey) -> Result<DecodingKey, BearerAuthError> {
match jwk.kty.as_str() {
"RSA" => {
let n = jwk
.n
.as_deref()
.ok_or_else(|| BearerAuthError::Invalid("RSA jwk missing n".into()))?;
let e = jwk
.e
.as_deref()
.ok_or_else(|| BearerAuthError::Invalid("RSA jwk missing e".into()))?;
DecodingKey::from_rsa_components(n, e)
.map_err(|err| BearerAuthError::Invalid(format!("RSA jwk: {err}")))
}
"EC" => {
let x = jwk
.x
.as_deref()
.ok_or_else(|| BearerAuthError::Invalid("EC jwk missing x".into()))?;
let y = jwk
.y
.as_deref()
.ok_or_else(|| BearerAuthError::Invalid("EC jwk missing y".into()))?;
let _ = jwk.crv.as_deref().unwrap_or("P-256");
DecodingKey::from_ec_components(x, y)
.map_err(|err| BearerAuthError::Invalid(format!("EC jwk: {err}")))
}
"oct" => {
Err(BearerAuthError::Invalid(
"oct (symmetric) keys in JWKS not supported; use HMAC JwtValidator directly".into(),
))
}
other => Err(BearerAuthError::Invalid(format!("unsupported kty={other}"))),
}
}
#[async_trait]
impl BearerValidator for JwksJwtValidator {
async fn validate(&self, token: &str) -> Result<IdentityAssurance, BearerAuthError> {
Ok(self.validate_credential(token).await?.principal.assurance)
}
async fn validate_principal(
&self,
token: &str,
) -> Result<AuthenticatedPrincipal, BearerAuthError> {
Ok(self.validate_credential(token).await?.principal)
}
async fn validate_credential(&self, token: &str) -> Result<ValidatedBearer, BearerAuthError> {
if token.is_empty() {
return Err(BearerAuthError::Empty);
}
let header =
decode_header(token).map_err(|e| BearerAuthError::Invalid(format!("header: {e}")))?;
let kid = header
.kid
.as_ref()
.ok_or_else(|| BearerAuthError::Invalid("token header missing kid".into()))?;
let key = self.resolve_key(kid).await?;
let data = decode::<TokenClaims>(token, &key, &self.inner.validation)
.map_err(|e| BearerAuthError::Invalid(e.to_string()))?;
let claims = data.claims;
let token_id = validate_optional_token_id(claims.jti.clone())?;
if self.inner.require_jti && token_id.is_none() {
return Err(BearerAuthError::Invalid(
"token missing required jti".into(),
));
}
let issued_at = claims
.iat
.map(|iat| unix_time_from_seconds(iat, "iat"))
.transpose()?;
let expires_at_system = claims
.exp
.map(|exp| unix_time_from_seconds(exp, "exp"))
.transpose()?;
let revocation_context = revocation_context_from_claims(
&claims,
token_id.as_deref(),
issued_at,
expires_at_system,
);
check_revocation(
self.inner.revocation_checker.as_ref(),
revocation_context.as_ref(),
)
.await?;
let subject = claims.sub.clone();
let expires_at = claims.exp.map(expiration_from_unix).transpose()?;
let identity = IdentityId::from_string(subject.clone());
let scopes = scopes_from_claims(
claims.scope,
claims.scopes,
claims.roles,
claims.realm_access,
claims.resource_access,
);
let assurance = IdentityAssurance::UserAuthorized {
identity: identity.clone(),
user_id: identity,
scopes: scopes.clone(),
};
ValidatedBearer::new(
AuthenticatedPrincipal {
subject,
tenant: claims.tenant_id,
scopes,
issuer: claims.iss,
expires_at,
method: AuthenticationMethod::Oidc,
assurance,
},
token_id,
issued_at,
)
}
}
fn expiration_from_unix(seconds: u64) -> Result<chrono::DateTime<chrono::Utc>, BearerAuthError> {
i64::try_from(seconds)
.ok()
.and_then(|seconds| chrono::DateTime::from_timestamp(seconds, 0))
.ok_or_else(|| BearerAuthError::Invalid("token exp is outside the supported range".into()))
}
async fn check_revocation(
checker: Option<&Arc<dyn TokenRevocationChecker>>,
context: Option<&TokenRevocationContext>,
) -> Result<(), BearerAuthError> {
let Some(checker) = checker else {
return Ok(());
};
let Some(context) = context else {
return Err(BearerAuthError::Invalid(
"token missing jti for revocation check".into(),
));
};
match checker.check_token(context).await {
Ok(TokenRevocationStatus::Active) => Ok(()),
Ok(TokenRevocationStatus::Revoked) => Err(BearerAuthError::Invalid("token revoked".into())),
Err(CredentialAuthError::Invalid) | Err(CredentialAuthError::PolicyRejected(_)) => Err(
BearerAuthError::Invalid("revocation check rejected token".into()),
),
Err(CredentialAuthError::Unavailable(err)) => Err(BearerAuthError::Unavailable(err)),
}
}
fn revocation_context_from_claims(
claims: &TokenClaims,
token_id: Option<&str>,
issued_at: Option<SystemTime>,
expires_at: Option<SystemTime>,
) -> Option<TokenRevocationContext> {
let mut context = TokenRevocationContext::new(token_id?).with_subject(claims.sub.clone());
if let Some(issuer) = claims.iss.clone() {
context = context.with_issuer(issuer);
}
context = context.with_times(issued_at, expires_at);
Some(context)
}
fn scopes_from_claims(
scope: Option<String>,
scopes: Option<Vec<String>>,
roles: Option<Vec<String>>,
realm_access: Option<RoleAccess>,
resource_access: Option<HashMap<String, RoleAccess>>,
) -> Vec<String> {
let mut values = Vec::new();
if let Some(scope) = scope {
values.extend(scope.split_whitespace().map(str::to_string));
}
if let Some(scopes) = scopes {
for scope in scopes {
push_unique(&mut values, scope);
}
}
if let Some(roles) = roles {
for role in roles {
push_unique(&mut values, format!("role:{role}"));
}
}
if let Some(realm_access) = realm_access {
for role in realm_access.roles {
push_unique(&mut values, format!("realm:{role}"));
}
}
if let Some(resource_access) = resource_access {
for (client, access) in resource_access {
for role in access.roles {
push_unique(&mut values, format!("{client}:{role}"));
}
}
}
values
}
fn push_unique(values: &mut Vec<String>, value: String) {
if !values.contains(&value) {
values.push(value);
}
}
#[cfg(test)]
mod diagnostic_tests {
use super::*;
const CANARY: &str = "jwks-claims-malicious-canary\r\nAuthorization: exposed";
#[test]
fn decoded_jwks_keys_keep_exact_values_out_of_debug() {
let document: JwksDocument = serde_json::from_value(serde_json::json!({
"keys": [{
"kty": CANARY,
"kid": CANARY,
"n": CANARY,
"e": CANARY,
"crv": CANARY,
"x": CANARY,
"y": CANARY,
}]
}))
.unwrap();
let key = &document.keys[0];
for rendered in [format!("{document:?}"), format!("{key:?}")] {
assert!(!rendered.contains(CANARY), "JWKS value leaked: {rendered}");
}
assert_eq!(key.kty, CANARY);
assert_eq!(key.kid.as_deref(), Some(CANARY));
assert_eq!(key.n.as_deref(), Some(CANARY));
assert_eq!(key.e.as_deref(), Some(CANARY));
assert_eq!(key.crv.as_deref(), Some(CANARY));
assert_eq!(key.x.as_deref(), Some(CANARY));
assert_eq!(key.y.as_deref(), Some(CANARY));
assert_eq!(jwk_key_type_class(&key.kty), "other");
}
#[test]
fn decoded_claims_keep_values_out_of_debug() {
let claims: TokenClaims = serde_json::from_value(serde_json::json!({
"sub": CANARY,
"iss": CANARY,
"iat": 1,
"exp": 2,
"jti": CANARY,
"scope": CANARY,
"scopes": [CANARY],
"roles": [CANARY],
"realm_access": { "roles": [CANARY] },
"resource_access": { (CANARY): { "roles": [CANARY] } },
"tenant_id": CANARY,
}))
.unwrap();
for rendered in [
format!("{claims:?}"),
format!("{:?}", claims.realm_access.as_ref().unwrap()),
format!(
"{:?}",
claims
.resource_access
.as_ref()
.unwrap()
.get(CANARY)
.unwrap()
),
] {
assert!(!rendered.contains(CANARY), "claim leaked: {rendered}");
}
assert_eq!(claims.sub, CANARY);
assert_eq!(claims.iss.as_deref(), Some(CANARY));
assert_eq!(claims.jti.as_deref(), Some(CANARY));
assert_eq!(claims.scope.as_deref(), Some(CANARY));
assert_eq!(claims.scopes.as_deref(), Some(&[CANARY.to_string()][..]));
assert_eq!(claims.roles.as_deref(), Some(&[CANARY.to_string()][..]));
assert_eq!(claims.tenant_id.as_deref(), Some(CANARY));
assert_eq!(
claims.realm_access.as_ref().unwrap().roles,
[CANARY.to_string()]
);
assert_eq!(
claims.resource_access.as_ref().unwrap()[CANARY].roles,
[CANARY.to_string()]
);
}
}