const JWKS_FETCH_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
pub(crate) const MAX_JWKS_RESPONSE_BYTES: usize = 1024 * 1024;
use std::{
collections::HashMap,
sync::RwLock,
time::{Duration, Instant},
};
use jsonwebtoken::DecodingKey;
use serde::Deserialize;
use tracing::debug;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum JwksError {
#[error("Invalid jwks_uri '{uri}': {source}")]
InvalidUrl {
uri: String,
source: url::ParseError,
},
#[error("Invalid jwks_uri scheme '{scheme}': must be https (or http on localhost)")]
InvalidScheme {
scheme: String,
},
#[error("Failed to build HTTP client: {0}")]
HttpClient(#[from] reqwest::Error),
}
#[derive(Debug, Deserialize)]
struct JwksDocument {
keys: Vec<JwkKey>,
}
#[derive(Debug, Deserialize)]
struct JwkKey {
kid: Option<String>,
kty: String,
n: Option<String>,
e: Option<String>,
x: Option<String>,
y: Option<String>,
}
pub(crate) fn is_ssrf_blocked_ip(ip: &std::net::IpAddr) -> bool {
match ip {
std::net::IpAddr::V4(v4) => {
let o = v4.octets();
o[0] == 127
|| o[0] == 10
|| (o[0] == 172 && (16..=31).contains(&o[1]))
|| (o[0] == 192 && o[1] == 168)
|| (o[0] == 169 && o[1] == 254)
|| (o[0] == 100 && (o[1] & 0b1100_0000) == 0b0100_0000)
|| o[0] == 0
},
std::net::IpAddr::V6(v6) => {
let s = v6.segments();
*v6 == std::net::Ipv6Addr::LOCALHOST
|| *v6 == std::net::Ipv6Addr::UNSPECIFIED
|| (s[0] == 0 && s[1] == 0 && s[2] == 0 && s[3] == 0 && s[4] == 0 && s[5] == 0xffff)
|| (s[0] & 0xfe00) == 0xfc00
|| (s[0] & 0xffc0) == 0xfe80
},
}
}
pub(crate) async fn dns_resolve_and_check(
host: &str,
port: u16,
) -> Result<Vec<std::net::SocketAddr>, String> {
let addrs: Vec<std::net::SocketAddr> = tokio::net::lookup_host((host, port))
.await
.map_err(|e| format!("DNS resolution failed for JWKS host '{host}': {e}"))?
.collect();
if addrs.is_empty() {
return Err(format!("DNS resolved to no addresses for JWKS host '{host}'"));
}
for addr in &addrs {
if is_ssrf_blocked_ip(&addr.ip()) {
return Err(format!(
"DNS rebinding attack blocked: JWKS host '{host}' resolved to private/reserved IP {}",
addr.ip()
));
}
}
Ok(addrs)
}
pub(crate) fn build_pinned_client(
host: &str,
addrs: &[std::net::SocketAddr],
timeout: Duration,
) -> Result<reqwest::Client, reqwest::Error> {
reqwest::Client::builder()
.timeout(timeout)
.redirect(reqwest::redirect::Policy::none())
.resolve_to_addrs(host, addrs)
.build()
}
pub struct JwksCache {
keys: RwLock<HashMap<String, DecodingKey>>,
jwks_uri: String,
last_fetched: RwLock<Option<Instant>>,
ttl: Duration,
client: reqwest::Client,
}
impl JwksCache {
pub fn new(jwks_uri: &str, ttl: Duration) -> Result<Self, JwksError> {
let parsed = reqwest::Url::parse(jwks_uri).map_err(|e| JwksError::InvalidUrl {
uri: jwks_uri.to_string(),
source: e,
})?;
let allowed = match parsed.scheme() {
"https" => true,
"http" => parsed.host_str().is_some_and(|h| {
h == "localhost" || h == "127.0.0.1" || h == "[::1]" || h == "::1"
}),
_ => false,
};
if !allowed {
return Err(JwksError::InvalidScheme {
scheme: parsed.scheme().to_string(),
});
}
let client = reqwest::Client::builder().timeout(JWKS_FETCH_TIMEOUT).build()?;
Ok(Self {
keys: RwLock::new(HashMap::new()),
jwks_uri: jwks_uri.to_string(),
last_fetched: RwLock::new(None),
ttl,
client,
})
}
pub async fn get_key(&self, kid: &str) -> Result<Option<DecodingKey>, String> {
if let Some(key) = self.get_key_from_cache(kid) {
if !self.is_stale() {
return Ok(Some(key));
}
}
self.fetch_keys().await?;
Ok(self.get_key_from_cache(kid))
}
pub fn get_key_from_cache(&self, kid: &str) -> Option<DecodingKey> {
self.keys.read().ok()?.get(kid).cloned()
}
pub async fn force_refresh(&self) -> Result<(), String> {
self.fetch_keys().await
}
fn is_stale(&self) -> bool {
self.last_fetched
.read()
.ok()
.and_then(|guard| *guard)
.is_none_or(|t| t.elapsed() > self.ttl)
}
async fn request_client(&self) -> Result<reqwest::Client, String> {
let Ok(parsed) = reqwest::Url::parse(&self.jwks_uri) else {
return Ok(self.client.clone());
};
let Some(host) = parsed.host_str() else {
return Ok(self.client.clone());
};
let is_localhost = {
let h = host.to_ascii_lowercase();
h == "localhost" || h == "127.0.0.1" || h == "[::1]" || h == "::1"
};
if is_localhost {
return Ok(self.client.clone());
}
let port = parsed.port_or_known_default().unwrap_or(443);
let validated = dns_resolve_and_check(host, port).await?;
build_pinned_client(host, &validated, JWKS_FETCH_TIMEOUT)
.map_err(|e| format!("Failed to build pinned JWKS client: {e}"))
}
async fn fetch_keys(&self) -> Result<(), String> {
debug!(uri = %self.jwks_uri, "Fetching JWKS keys");
let client = self.request_client().await?;
let body = client
.get(&self.jwks_uri)
.send()
.await
.map_err(|e| format!("JWKS fetch failed: {e}"))?
.bytes()
.await
.map_err(|e| format!("JWKS read failed: {e}"))?;
if body.len() > MAX_JWKS_RESPONSE_BYTES {
return Err(format!(
"JWKS response too large ({} bytes, max {MAX_JWKS_RESPONSE_BYTES})",
body.len()
));
}
let jwks: JwksDocument =
serde_json::from_slice(&body).map_err(|e| format!("JWKS parse failed: {e}"))?;
let mut cache = self.keys.write().map_err(|e| format!("JWKS lock poisoned: {e}"))?;
cache.clear();
for key in &jwks.keys {
if let Some(kid) = &key.kid {
if let Some(decoding_key) = Self::convert_jwk(key) {
cache.insert(kid.clone(), decoding_key);
}
}
}
if let Ok(mut last) = self.last_fetched.write() {
*last = Some(Instant::now());
}
debug!(key_count = cache.len(), "JWKS cache refreshed");
Ok(())
}
fn convert_jwk(jwk: &JwkKey) -> Option<DecodingKey> {
match jwk.kty.as_str() {
"RSA" => {
let n = jwk.n.as_ref()?;
let e = jwk.e.as_ref()?;
DecodingKey::from_rsa_components(n, e).ok()
},
"EC" => {
let x = jwk.x.as_ref()?;
let y = jwk.y.as_ref()?;
DecodingKey::from_ec_components(x, y).ok()
},
_ => None,
}
}
}
#[allow(clippy::missing_fields_in_debug)] impl std::fmt::Debug for JwksCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let key_count = self.keys.read().map(|k| k.len()).unwrap_or(0);
f.debug_struct("JwksCache")
.field("jwks_uri", &self.jwks_uri)
.field("ttl", &self.ttl)
.field("cached_keys", &key_count)
.finish_non_exhaustive()
}
}