use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
const DEFAULT_TTL: Duration = Duration::from_secs(3600);
const MIN_TTL: Duration = Duration::from_secs(60);
const MAX_TTL: Duration = Duration::from_secs(86_400);
const MAX_DOC_BYTES: usize = 65_536;
const FETCH_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Discovered {
pub authorize_url: String,
pub token_url: String,
pub jwks_url: String,
pub userinfo_url: Option<String>,
pub issuer: String,
}
#[derive(serde::Deserialize)]
struct DiscoveryDoc {
issuer: String,
authorization_endpoint: String,
token_endpoint: String,
jwks_uri: String,
#[serde(default)]
userinfo_endpoint: Option<String>,
}
struct Entry {
doc: Arc<Discovered>,
fetched_at: Instant,
ttl: Duration,
}
pub struct DiscoveryCache {
entries: tokio::sync::RwLock<HashMap<String, Arc<Entry>>>,
flights: tokio::sync::Mutex<HashMap<String, Arc<tokio::sync::Mutex<()>>>>,
client: reqwest::Client,
allow_private: bool,
}
impl DiscoveryCache {
pub fn new(client: reqwest::Client, allow_private: bool) -> Self {
Self {
entries: tokio::sync::RwLock::new(HashMap::new()),
flights: tokio::sync::Mutex::new(HashMap::new()),
client,
allow_private,
}
}
async fn fresh(&self, issuer: &str) -> Option<Arc<Discovered>> {
self.entries
.read()
.await
.get(issuer)
.filter(|e| e.fetched_at.elapsed() < e.ttl)
.map(|e| Arc::clone(&e.doc))
}
pub async fn resolve(&self, issuer: &str) -> Result<Arc<Discovered>, String> {
if let Some(doc) = self.fresh(issuer).await {
return Ok(doc);
}
let flight = {
let mut flights = self.flights.lock().await;
if let Some(flight) = flights.get(issuer) {
Arc::clone(flight)
} else {
let flight = Arc::new(tokio::sync::Mutex::new(()));
flights.insert(issuer.to_string(), Arc::clone(&flight));
flight
}
};
let _guard = flight.lock().await;
if let Some(doc) = self.fresh(issuer).await {
return Ok(doc);
}
match self.fetch(issuer).await {
Ok((doc, ttl)) => {
let doc = Arc::new(doc);
self.entries.write().await.insert(
issuer.to_string(),
Arc::new(Entry {
doc: Arc::clone(&doc),
fetched_at: Instant::now(),
ttl,
}),
);
Ok(doc)
}
Err(e) => {
if let Some(entry) = self.entries.read().await.get(issuer).cloned() {
tracing::warn!(
issuer = %issuer,
error = %e,
"OIDC discovery refresh failed; serving the last good document"
);
return Ok(Arc::clone(&entry.doc));
}
Err(e)
}
}
}
async fn fetch(&self, issuer: &str) -> Result<(Discovered, Duration), String> {
let url = well_known_url(issuer)?;
if !self.allow_private {
crate::validation::validate_url_not_private(&url).await?;
}
let response = self
.client
.get(&url)
.timeout(FETCH_TIMEOUT)
.send()
.await
.map_err(|e| format!("discovery fetch failed: {e}"))?;
if !response.status().is_success() {
return Err(format!("discovery HTTP {}", response.status()));
}
let ttl = ttl_from_cache_control(
response
.headers()
.get("cache-control")
.and_then(|v| v.to_str().ok()),
);
let body = crate::http_body::read_bounded(response, MAX_DOC_BYTES)
.await
.map_err(|e| format!("discovery document {e}"))?;
let doc: DiscoveryDoc = serde_json::from_slice(&body)
.map_err(|e| format!("discovery document is not valid metadata: {e}"))?;
if doc.issuer != issuer {
return Err(format!(
"discovery document issuer '{}' does not match the configured issuer '{issuer}'",
doc.issuer
));
}
require_https("authorization_endpoint", &doc.authorization_endpoint)?;
require_https("token_endpoint", &doc.token_endpoint)?;
require_https("jwks_uri", &doc.jwks_uri)?;
if let Some(ref u) = doc.userinfo_endpoint {
require_https("userinfo_endpoint", u)?;
}
Ok((
Discovered {
authorize_url: doc.authorization_endpoint,
token_url: doc.token_endpoint,
jwks_url: doc.jwks_uri,
userinfo_url: doc.userinfo_endpoint,
issuer: doc.issuer,
},
ttl,
))
}
}
fn well_known_url(issuer: &str) -> Result<String, String> {
let trimmed = issuer.trim_end_matches('/');
if trimmed.is_empty() {
return Err("issuer must not be empty".to_string());
}
Ok(format!("{trimmed}/.well-known/openid-configuration"))
}
fn require_https(field: &str, value: &str) -> Result<(), String> {
let url = url::Url::parse(value)
.map_err(|e| format!("discovery {field} '{value}' is not a URL: {e}"))?;
if url.scheme() == "https"
|| (url.scheme() == "http" && super::oauth2_login::is_loopback_host(&url))
{
return Ok(());
}
Err(format!(
"discovery {field} must be https — '{value}' is {}",
url.scheme()
))
}
fn ttl_from_cache_control(header: Option<&str>) -> Duration {
let Some(header) = header else {
return DEFAULT_TTL;
};
header
.split(',')
.filter_map(|d| {
let d = d.trim();
d.strip_prefix("max-age=")
.and_then(|v| v.parse::<u64>().ok())
.map(Duration::from_secs)
})
.next()
.map(|d| d.clamp(MIN_TTL, MAX_TTL))
.unwrap_or(DEFAULT_TTL)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn well_known_url_appends_and_preserves_the_issuer_path() {
assert_eq!(
well_known_url("https://issuer.example").expect("ok"),
"https://issuer.example/.well-known/openid-configuration"
);
assert_eq!(
well_known_url("https://issuer.example/").expect("ok"),
"https://issuer.example/.well-known/openid-configuration"
);
assert_eq!(
well_known_url("https://login.example/tenant-123/v2.0").expect("ok"),
"https://login.example/tenant-123/v2.0/.well-known/openid-configuration"
);
}
#[test]
fn require_https_matches_the_hand_typed_rule() {
assert!(require_https("token_endpoint", "https://idp.example/token").is_ok());
assert!(require_https("token_endpoint", "http://127.0.0.1:9/token").is_ok());
assert!(require_https("token_endpoint", "http://idp.example/token").is_err());
}
#[test]
fn ttl_is_clamped() {
assert_eq!(ttl_from_cache_control(None), DEFAULT_TTL);
assert_eq!(
ttl_from_cache_control(Some("max-age=30")),
MIN_TTL,
"below the floor clamps up"
);
assert_eq!(
ttl_from_cache_control(Some("public, max-age=7200")),
Duration::from_secs(7200)
);
assert_eq!(
ttl_from_cache_control(Some("max-age=999999")),
MAX_TTL,
"above the ceiling clamps down"
);
}
}