use std::{
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
thread::{self, JoinHandle},
time::{Duration, Instant},
};
use jsonwebtoken::jwk::JwkSet;
use parking_lot::RwLock;
use thiserror::Error;
const OPENID_CONFIGURATION_ADDRESS_FORMAT: &str =
"{authority}/common/.well-known/openid-configuration";
const REFRESH_INTERVAL: Duration = Duration::from_secs(24 * 60 * 60);
#[derive(Debug, Error)]
pub enum AadError {
#[error("document fetch failed: {0}")]
Fetch(String),
#[error("document parse failed: {0}")]
Parse(String),
}
pub trait DocumentFetch: Send + Sync {
fn get(&self, url: &str) -> Result<Vec<u8>, AadError>;
}
#[cfg(test)]
struct NoopFetch;
#[cfg(test)]
impl DocumentFetch for NoopFetch {
fn get(&self, _url: &str) -> Result<Vec<u8>, AadError> {
Err(AadError::Fetch("noop".into()))
}
}
#[derive(serde::Deserialize)]
struct OpenIdConfiguration {
jwks_uri: String,
}
pub struct IssuerSigningTokenProvider {
authority: String,
signing_tokens: RwLock<Arc<JwkSet>>,
refresh: RwLock<Option<JoinHandle<()>>>,
stopped: Arc<AtomicBool>,
fetch: Arc<dyn DocumentFetch>,
}
impl IssuerSigningTokenProvider {
pub fn create(authority: &str, fetch: Arc<dyn DocumentFetch>) -> Result<Arc<Self>, AadError> {
if authority.is_empty() {
return Err(AadError::Fetch("Authority cannot be null".into()));
}
let signing_tokens = Self::retrieve_signing_tokens(authority, fetch.as_ref())?;
let provider = Arc::new(Self {
authority: authority.to_string(),
signing_tokens: RwLock::new(Arc::new(signing_tokens)),
refresh: RwLock::new(None),
stopped: Arc::new(AtomicBool::new(false)),
fetch,
});
let for_thread = Arc::clone(&provider);
let thread_stopped = Arc::clone(&provider.stopped);
if let Ok(handle) = thread::Builder::new()
.name("aad-signing-tokens".into())
.spawn(move || {
let mut last = Instant::now();
while !thread_stopped.load(Ordering::Relaxed) {
thread::sleep(Duration::from_secs(1));
if last.elapsed() >= REFRESH_INTERVAL {
last = Instant::now();
for_thread.refresh_signing_tokens();
}
}
})
{
*provider.refresh.write() = Some(handle);
}
Ok(provider)
}
#[cfg(test)]
pub fn create_for_test(signing_tokens: JwkSet) -> Arc<Self> {
Arc::new(Self {
authority: String::new(),
signing_tokens: RwLock::new(Arc::new(signing_tokens)),
refresh: RwLock::new(None),
stopped: Arc::new(AtomicBool::new(true)),
fetch: Arc::new(NoopFetch),
})
}
pub fn signing_tokens(&self) -> Arc<JwkSet> {
Arc::clone(&self.signing_tokens.read())
}
pub fn refresh_signing_tokens(&self) {
match Self::retrieve_signing_tokens(&self.authority, self.fetch.as_ref()) {
Ok(tokens) => *self.signing_tokens.write() = Arc::new(tokens),
Err(err) => {
log::error!("Failed to retrieve issuer signing tokens: {err}");
}
}
}
pub fn retrieve_signing_tokens(
authority: &str,
fetch: &dyn DocumentFetch,
) -> Result<JwkSet, AadError> {
let config_url = OPENID_CONFIGURATION_ADDRESS_FORMAT.replace("{authority}", authority);
let config_doc = fetch.get(&config_url)?;
let config: OpenIdConfiguration =
sonic_rs::from_slice(&config_doc).map_err(|e| AadError::Parse(e.to_string()))?;
let jwks_doc = fetch.get(&config.jwks_uri)?;
sonic_rs::from_slice(&jwks_doc).map_err(|e| AadError::Parse(e.to_string()))
}
}
impl Drop for IssuerSigningTokenProvider {
fn drop(&mut self) {
self.stopped.store(true, Ordering::Relaxed);
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::AtomicUsize;
use super::*;
struct MemFetch {
openid: Vec<u8>,
jwks: Vec<u8>,
hits: AtomicUsize,
}
impl DocumentFetch for MemFetch {
fn get(&self, url: &str) -> Result<Vec<u8>, AadError> {
self.hits.fetch_add(1, Ordering::Relaxed);
if url.ends_with("/.well-known/openid-configuration") {
Ok(self.openid.clone())
} else if url == "https://keys.example/keys" {
Ok(self.jwks.clone())
} else {
Err(AadError::Fetch(format!("unexpected url {url}")))
}
}
}
fn fixtures() -> (Vec<u8>, Vec<u8>) {
let openid = br#"{"jwks_uri":"https://keys.example/keys"}"#.to_vec();
let jwks = br#"{"keys":[{"kty":"RSA","kid":"k1","n":"xYuC","e":"AQAB"}]}"#.to_vec();
(openid, jwks)
}
#[test]
fn retrieve_signing_tokens_follows_openid_flow() {
let (openid, jwks) = fixtures();
let fetch = MemFetch {
openid,
jwks,
hits: AtomicUsize::new(0),
};
let set =
IssuerSigningTokenProvider::retrieve_signing_tokens("https://login.example", &fetch).unwrap();
assert_eq!(set.keys.len(), 1);
assert_eq!(set.keys[0].common.key_id.as_deref(), Some("k1"));
assert_eq!(fetch.hits.load(Ordering::Relaxed), 2);
}
#[test]
fn retrieve_signing_tokens_parse_errors() {
let fetch = MemFetch {
openid: b"not json".to_vec(),
jwks: Vec::new(),
hits: AtomicUsize::new(0),
};
assert!(matches!(
IssuerSigningTokenProvider::retrieve_signing_tokens("https://login.example", &fetch),
Err(AadError::Parse(_))
));
let fetch = MemFetch {
openid: br#"{"jwks_uri":"https://keys.example/keys"}"#.to_vec(),
jwks: b"{}".to_vec(),
hits: AtomicUsize::new(0),
};
assert!(matches!(
IssuerSigningTokenProvider::retrieve_signing_tokens("https://login.example", &fetch),
Err(AadError::Parse(_))
));
}
#[test]
fn create_rejects_empty_authority() {
let (openid, jwks) = fixtures();
let fetch: Arc<MemFetch> = Arc::new(MemFetch {
openid,
jwks,
hits: AtomicUsize::new(0),
});
assert!(IssuerSigningTokenProvider::create("", fetch).is_err());
}
#[test]
fn create_and_refresh_swaps_keys() {
let (openid, jwks) = fixtures();
let fetch: Arc<MemFetch> = Arc::new(MemFetch {
openid,
jwks,
hits: AtomicUsize::new(0),
});
let provider =
IssuerSigningTokenProvider::create("https://login.example", fetch.clone()).unwrap();
assert_eq!(fetch.hits.load(Ordering::Relaxed), 2); assert_eq!(provider.signing_tokens().keys.len(), 1);
let jwk = provider.signing_tokens();
provider.refresh_signing_tokens();
let jwk2 = provider.signing_tokens();
assert!(!Arc::ptr_eq(&jwk, &jwk2));
assert_eq!(jwk2.keys.len(), 1);
assert_eq!(fetch.hits.load(Ordering::Relaxed), 4);
drop(provider);
}
}