use std::sync::Arc;
use jsonwebtoken::{jwk::JwkSet, Algorithm, DecodingKey, Validation};
use serde::de::DeserializeOwned;
use crate::{
cache::{JwkCacheKey, JwksCache},
config::AuthenticationConfigProvider,
error::OidcError,
};
pub fn validate_token_with_jwk<T: DeserializeOwned>(
token: &str,
jwk: &jsonwebtoken::jwk::Jwk,
alg: Algorithm,
expected_issuer: Option<&str>,
expected_audience: Option<&[String]>,
) -> Result<T, String> {
let mut validation = Validation::new(alg);
validation.validate_aud = expected_audience.is_some();
if let Some(audience) = expected_audience {
validation.set_audience(audience);
}
if let Some(issuer) = expected_issuer {
validation.set_issuer(&[issuer]);
}
let mut required_claims = vec!["exp"];
if expected_issuer.is_some() {
required_claims.push("iss");
}
if expected_audience.is_some() {
required_claims.push("aud");
}
validation.set_required_spec_claims(&required_claims);
let decoding_key = DecodingKey::from_jwk(jwk)
.map_err(|e| format!("Failed to create DecodingKey from JWK: {e}"))?;
let token_data = jsonwebtoken::decode::<T>(token, &decoding_key, &validation)
.map_err(|e| format!("JWT validation failed: {e}"))?;
Ok(token_data.claims)
}
pub async fn fetch_and_cache_jwks(
cache: &Arc<dyn JwksCache>,
jwks_uri: &str,
config: &(impl AuthenticationConfigProvider + Send + Sync),
requested_kid: &str,
) -> Result<jsonwebtoken::jwk::Jwk, OidcError> {
let response = reqwest::get(jwks_uri)
.await
.map_err(|e| OidcError::JwksError(format!("Failed to fetch JWKS: {e}")))?;
let jwks: JwkSet = response
.json()
.await
.map_err(|e| OidcError::JwksError(format!("Failed to parse JWKS: {e}")))?;
let jwks_cache_ttl = config.get_jwks_cache_ttl();
for jwk in &jwks.keys {
if let Some(kid) = &jwk.common.key_id {
let jwk_cache_key = JwkCacheKey::from_jwks_uri_and_kid(jwks_uri, kid);
if let Ok(serialized) = serde_json::to_string(jwk) {
cache.set(jwk_cache_key.as_str(), serialized, jwks_cache_ttl);
}
}
}
jwks.find(requested_kid)
.cloned()
.ok_or_else(|| OidcError::JwksError(format!("No JWK found for kid: {requested_kid}")))
}
#[cfg(test)]
mod tests {
use super::*;
use jsonwebtoken::Algorithm;
use serde::{Deserialize, Serialize};
#[derive(Debug, Serialize, Deserialize, PartialEq)]
struct TestClaims {
sub: String,
exp: usize,
}
#[test]
fn test_validate_token_with_invalid_jwk() {
let token = "invalid.token.here";
let jwk: jsonwebtoken::jwk::Jwk = serde_json::from_str(
r#"{"kty":"oct","k":"AAECAwQFBgcICQoLDA0ODw"}"#,
)
.expect("valid JWK JSON");
let result = validate_token_with_jwk::<TestClaims>(
token,
&jwk,
Algorithm::HS256,
None,
None,
);
assert!(result.is_err());
}
}