#[cfg(feature = "verifier")]
use crate::{
cache::Cache,
discovery::discover,
errors::{Error, Result},
http::HttpClient,
id_token::fetch_jwks,
jwks::Jwks,
types::VerifiedClaims,
};
#[cfg(feature = "verifier")]
use base64::{engine::general_purpose, Engine as _};
#[cfg(feature = "verifier")]
use josekit::{
jws::RS256,
jwt::{self, JwtPayload},
};
#[cfg(feature = "verifier")]
use std::collections::HashMap;
#[cfg(feature = "verifier")]
use std::time::{SystemTime, UNIX_EPOCH};
#[cfg(feature = "verifier")]
pub struct JwtVerifier<C: Cache<String, Jwks>, H: HttpClient> {
pub issuer_map: HashMap<String, String>,
pub audience: String,
pub http: std::sync::Arc<H>,
pub cache: std::sync::Arc<C>,
pub clock_skew_sec: i64,
pub default_issuer: Option<String>,
}
#[cfg(feature = "verifier")]
impl<C: Cache<String, Jwks>, H: HttpClient> JwtVerifier<C, H> {
pub fn new(
issuer_map: HashMap<String, String>,
audience: String,
http: std::sync::Arc<H>,
cache: std::sync::Arc<C>,
) -> Self {
Self { issuer_map, audience, http, cache, clock_skew_sec: 60, default_issuer: None }
}
pub fn builder() -> JwtVerifierBuilder<C, H> {
JwtVerifierBuilder::default()
}
pub async fn verify(&self, bearer: &str) -> Result<VerifiedClaims> {
let token = bearer.strip_prefix("Bearer ").unwrap_or(bearer);
let unverified_issuer = extract_unverified_issuer(token)?;
let expected_issuer = self.resolve_issuer(&unverified_issuer)?;
let metadata_cache = crate::cache::NoOpCache;
let metadata = discover(&expected_issuer, self.http.as_ref(), &metadata_cache).await?;
let jwks = fetch_jwks(&metadata.jwks_uri, self.http.as_ref(), self.cache.as_ref()).await?;
let kid = extract_kid(token)?.ok_or_else(|| Error::Jwt("Token missing kid".into()))?;
let jwk = jwks
.find_key(&kid)
.ok_or_else(|| Error::Jwt(format!("Key with kid '{}' not found", kid)))?;
let payload = verify_token_signature(token, jwk)?;
let claims = extract_and_validate_access_token_claims(
payload,
&expected_issuer,
&self.audience,
self.clock_skew_sec,
)?;
Ok(claims)
}
fn resolve_issuer(&self, token_issuer: &str) -> Result<String> {
if self.issuer_map.values().any(|v| v == token_issuer) {
return Ok(token_issuer.to_string());
}
if self.issuer_map.contains_key(token_issuer) {
return Ok(self.issuer_map[token_issuer].clone());
}
if let Some(default) = &self.default_issuer {
return Ok(default.clone());
}
Err(Error::Verification(format!("Issuer '{}' not in allowed list", token_issuer)))
}
pub fn resolve_issuer_with_tenant(&self, tenant: &str) -> Result<String> {
if let Some(issuer) = self.issuer_map.get(tenant) {
return Ok(issuer.clone());
}
if let Some(default) = &self.default_issuer {
return Ok(default.clone());
}
Err(Error::Verification(format!("No issuer configured for tenant '{}'", tenant)))
}
}
#[cfg(feature = "verifier")]
pub struct JwtVerifierBuilder<C: Cache<String, Jwks>, H: HttpClient> {
issuer_map: Option<HashMap<String, String>>,
audience: Option<String>,
http: Option<std::sync::Arc<H>>,
cache: Option<std::sync::Arc<C>>,
clock_skew_sec: Option<i64>,
default_issuer: Option<String>,
}
#[cfg(feature = "verifier")]
impl<C: Cache<String, Jwks>, H: HttpClient> Default for JwtVerifierBuilder<C, H> {
fn default() -> Self {
Self {
issuer_map: None,
audience: None,
http: None,
cache: None,
clock_skew_sec: None,
default_issuer: None,
}
}
}
#[cfg(feature = "verifier")]
impl<C: Cache<String, Jwks>, H: HttpClient> JwtVerifierBuilder<C, H> {
pub fn issuer_map(mut self, map: HashMap<String, String>) -> Self {
self.issuer_map = Some(map);
self
}
pub fn audience(mut self, audience: impl Into<String>) -> Self {
self.audience = Some(audience.into());
self
}
pub fn http(mut self, http: std::sync::Arc<H>) -> Self {
self.http = Some(http);
self
}
pub fn cache(mut self, cache: std::sync::Arc<C>) -> Self {
self.cache = Some(cache);
self
}
pub fn clock_skew(mut self, seconds: i64) -> Self {
self.clock_skew_sec = Some(seconds);
self
}
pub fn default_issuer(mut self, issuer: impl Into<String>) -> Self {
self.default_issuer = Some(issuer.into());
self
}
pub fn build(self) -> Result<JwtVerifier<C, H>> {
Ok(JwtVerifier {
issuer_map: self.issuer_map.unwrap_or_default(),
audience: self.audience.ok_or(Error::MissingConfig("audience"))?,
http: self.http.ok_or(Error::MissingConfig("http client"))?,
cache: self.cache.ok_or(Error::MissingConfig("cache"))?,
clock_skew_sec: self.clock_skew_sec.unwrap_or(60),
default_issuer: self.default_issuer,
})
}
}
#[cfg(feature = "verifier")]
fn extract_kid(jwt: &str) -> Result<Option<String>> {
let parts: Vec<&str> = jwt.split('.').collect();
if parts.len() != 3 {
return Err(Error::Jwt("Invalid JWT format".into()));
}
let header_bytes = general_purpose::URL_SAFE_NO_PAD
.decode(parts[0])
.map_err(|e| Error::Base64(format!("Failed to decode header: {}", e)))?;
let header_value: serde_json::Value = serde_json::from_slice(&header_bytes)
.map_err(|e| Error::Jwt(format!("Failed to parse header JSON: {}", e)))?;
Ok(header_value.get("kid").and_then(|v| v.as_str()).map(|s| s.to_string()))
}
#[cfg(feature = "verifier")]
fn extract_unverified_issuer(token: &str) -> Result<String> {
let parts: Vec<&str> = token.split('.').collect();
if parts.len() != 3 {
return Err(Error::Jwt("Invalid JWT format".into()));
}
let payload_json = general_purpose::URL_SAFE_NO_PAD
.decode(parts[1])
.map_err(|e| Error::Base64(e.to_string()))?;
let payload: serde_json::Value = serde_json::from_slice(&payload_json)?;
payload["iss"]
.as_str()
.ok_or_else(|| Error::Jwt("Token missing issuer".into()))
.map(|s| s.to_string())
}
#[cfg(feature = "verifier")]
fn verify_token_signature(token: &str, jwk: &crate::jwks::Jwk) -> Result<JwtPayload> {
let key = josekit::jwk::Jwk::from_map(serde_json::to_value(jwk)?.as_object().unwrap().clone())
.map_err(|e| Error::Jwt(format!("Invalid JWK: {}", e)))?;
let alg = jwk.alg.as_deref().unwrap_or("RS256");
let verifier = match alg {
"RS256" => RS256.verifier_from_jwk(&key),
alg => return Err(Error::Jwt(format!("Unsupported algorithm: {}", alg))),
}
.map_err(|e| Error::Jwt(format!("Failed to create verifier: {}", e)))?;
let (payload, _header) = jwt::decode_with_verifier(token, &verifier)
.map_err(|e| Error::Jwt(format!("Token verification failed: {}", e)))?;
Ok(payload)
}
#[cfg(feature = "verifier")]
fn extract_and_validate_access_token_claims(
payload: JwtPayload,
expected_issuer: &str,
expected_audience: &str,
clock_skew: i64,
) -> Result<VerifiedClaims> {
let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_secs() as i64;
let iss = payload.issuer().ok_or_else(|| Error::Verification("Missing iss claim".into()))?;
let sub = payload.subject().ok_or_else(|| Error::Verification("Missing sub claim".into()))?;
let exp = payload
.expires_at()
.ok_or_else(|| Error::Verification("Missing exp claim".into()))?
.duration_since(UNIX_EPOCH)
.map_err(|_| Error::Verification("Invalid exp time".into()))?
.as_secs() as i64;
let iat = payload
.issued_at()
.ok_or_else(|| Error::Verification("Missing iat claim".into()))?
.duration_since(UNIX_EPOCH)
.map_err(|_| Error::Verification("Invalid iat time".into()))?
.as_secs() as i64;
if iss != expected_issuer {
return Err(Error::Verification(format!(
"Invalid issuer: expected '{}', got '{}'",
expected_issuer, iss
)));
}
let aud = if let Some(audiences) = payload.audience() {
if !audiences.iter().any(|a| *a == expected_audience) {
return Err(Error::Verification(format!(
"Invalid audience: expected '{}'",
expected_audience
)));
}
expected_audience.to_string()
} else {
return Err(Error::Verification("Missing aud claim".into()));
};
if exp < now - clock_skew {
return Err(Error::Verification("Token expired".into()));
}
if iat > now + clock_skew {
return Err(Error::Verification("Token issued in the future".into()));
}
let claims_map = payload.claims_set();
let jti = claims_map.get("jti").and_then(|v| v.as_str()).unwrap_or("").to_string();
let scope = claims_map.get("scope").and_then(|v| v.as_str()).map(|s| s.to_string());
let xjp_admin = claims_map.get("xjp_admin").and_then(|v| v.as_bool());
let amr = claims_map.get("amr").and_then(|v| {
v.as_array()?
.iter()
.map(|item| item.as_str().map(|s| s.to_string()))
.collect::<Option<Vec<String>>>()
});
let auth_time = claims_map.get("auth_time").and_then(|v| v.as_i64());
Ok(VerifiedClaims {
iss: iss.to_string(),
sub: sub.to_string(),
aud: aud.to_string(),
exp,
iat,
jti,
scope,
xjp_admin,
amr,
auth_time,
})
}
#[cfg(all(test, feature = "verifier"))]
mod tests {
use super::*;
#[test]
fn test_extract_unverified_issuer() {
let token = "eyJhbGciOiJSUzI1NiIsImtpZCI6InRlc3Qta2V5In0.eyJpc3MiOiJodHRwczovL2F1dGguZXhhbXBsZS5jb20ifQ.dummy";
let issuer = extract_unverified_issuer(token).unwrap();
assert_eq!(issuer, "https://auth.example.com");
}
}