use std::{
sync::Arc,
time::{Duration, Instant},
};
use axum_security_oauth2::HttpClient;
use jsonwebtoken::{Algorithm, jwk::JwkSet};
use tokio::sync::Mutex;
use url::Url;
use crate::{
error::VerifyError,
verifier::{IdTokenVerifier, VerifiedIdToken},
};
pub const DEFAULT_MIN_REFETCH_INTERVAL: Duration = Duration::from_secs(60);
pub struct JwksCache {
issuer: String,
audience: String,
algorithms: Option<Vec<Algorithm>>,
leeway_secs: Option<u64>,
jwks_url: Url,
http: HttpClient,
min_refetch_interval: Duration,
state: Mutex<State>,
}
#[derive(Default)]
struct State {
verifier: Option<Arc<IdTokenVerifier>>,
last_refetch: Option<Instant>,
}
impl JwksCache {
pub fn new(
issuer: impl Into<String>,
audience: impl Into<String>,
jwks_url: Url,
http: HttpClient,
) -> Self {
Self {
issuer: issuer.into(),
audience: audience.into(),
algorithms: None,
leeway_secs: None,
jwks_url,
http,
min_refetch_interval: DEFAULT_MIN_REFETCH_INTERVAL,
state: Mutex::new(State::default()),
}
}
pub fn algorithms(mut self, algorithms: &[Algorithm]) -> Self {
self.algorithms = Some(algorithms.to_vec());
self
}
pub fn leeway_secs(mut self, leeway_secs: u64) -> Self {
self.leeway_secs = Some(leeway_secs);
self
}
pub fn min_refetch_interval(mut self, interval: Duration) -> Self {
self.min_refetch_interval = interval;
self
}
pub async fn verify(
&self,
id_token: &str,
nonce: &str,
) -> Result<VerifiedIdToken, VerifyError> {
let verifier = self.current_verifier().await?;
match verifier.verify(id_token, nonce) {
Err(VerifyError::UnknownKey) => self.refetch(&verifier).await?.verify(id_token, nonce),
result => result,
}
}
pub async fn warm(&self) -> Result<(), VerifyError> {
self.current_verifier().await.map(|_| ())
}
async fn current_verifier(&self) -> Result<Arc<IdTokenVerifier>, VerifyError> {
let mut state = self.state.lock().await;
if let Some(verifier) = &state.verifier {
return Ok(verifier.clone());
}
let verifier = Arc::new(self.build_verifier(self.fetch_jwks().await?));
state.verifier = Some(verifier.clone());
state.last_refetch = Some(Instant::now());
Ok(verifier)
}
async fn refetch(
&self,
tried: &Arc<IdTokenVerifier>,
) -> Result<Arc<IdTokenVerifier>, VerifyError> {
let mut state = self.state.lock().await;
if let Some(current) = &state.verifier
&& !Arc::ptr_eq(current, tried)
{
return Ok(current.clone());
}
if let Some(last) = state.last_refetch
&& last.elapsed() < self.min_refetch_interval
{
return Ok(tried.clone());
}
state.last_refetch = Some(Instant::now());
let verifier = Arc::new(self.build_verifier(self.fetch_jwks().await?));
state.verifier = Some(verifier.clone());
Ok(verifier)
}
async fn fetch_jwks(&self) -> Result<JwkSet, VerifyError> {
let response = self
.http
.get(&self.jwks_url)
.await
.map_err(|_| VerifyError::JwksUnavailable)?;
if !response.is_success() {
return Err(VerifyError::JwksUnavailable);
}
serde_json::from_slice(&response.body).map_err(|_| VerifyError::JwksUnavailable)
}
fn build_verifier(&self, jwks: JwkSet) -> IdTokenVerifier {
let mut verifier = IdTokenVerifier::new(self.issuer.clone(), self.audience.clone(), jwks);
if let Some(algorithms) = &self.algorithms {
verifier = verifier.algorithms(algorithms);
}
if let Some(leeway_secs) = self.leeway_secs {
verifier = verifier.leeway_secs(leeway_secs);
}
verifier
}
}