use std::collections::HashSet;
use std::sync::Arc;
use std::time::{Duration, Instant};
use jsonwebtoken::DecodingKey;
use jsonwebtoken::jwk::{Jwk, KeyOperations, PublicKeyUse};
use serde_json::Value;
use tokio::sync::{Mutex, OwnedMutexGuard, RwLock};
use tracing::{debug, info, warn};
use crate::algorithms::{Algorithm, key_algorithms, signing_algorithm};
use crate::config::{KeyNamingBuf, ResolvedOAuthConfig};
use crate::token::{TokenRejection, describe_kid, for_log};
use crate::validator::plain_http_non_loopback;
pub(crate) const JWKS_MIN_REFETCH_INTERVAL: Duration = Duration::from_secs(60);
const JWKS_FETCH_TIMEOUT: Duration = Duration::from_secs(10);
pub(crate) const JWKS_BACKGROUND_REFRESH_INTERVAL: Duration = Duration::from_secs(3600);
pub(crate) fn background_retry_delay(failures: u32) -> Duration {
let doublings = failures.saturating_sub(1).min(16);
JWKS_MIN_REFETCH_INTERVAL
.saturating_mul(1 << doublings)
.min(JWKS_BACKGROUND_REFRESH_INTERVAL)
}
pub(crate) const MAX_FETCH_BYTES: usize = 256 * 1024;
const MAX_JWKS_KEYS: usize = 64;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error("{message}")]
pub struct RefreshError {
message: String,
}
fn error_chain(err: &dyn std::error::Error) -> String {
let mut out = err.to_string();
let mut source = err.source();
while let Some(cause) = source {
out.push_str(": ");
out.push_str(&cause.to_string());
source = cause.source();
}
out
}
fn context(context: &str, err: &dyn std::error::Error) -> String {
format!("{context}: {}", error_chain(err))
}
#[derive(Debug, PartialEq, Eq)]
enum Hop {
Follow,
FollowInsecure,
Refuse(String),
}
fn judge_redirect(
next: &reqwest::Url,
previous: &[reqwest::Url],
allow_insecure_http: bool,
opt_in_key: &str,
) -> Hop {
if next.scheme() != "https" && previous.iter().any(|u| u.scheme() == "https") {
return Hop::Refuse("redirect from https to a non-https URL refused".to_string());
}
let insecure = plain_http_non_loopback(next.as_str());
if insecure && !allow_insecure_http {
return Hop::Refuse(format!(
"redirect to plain http on a non-loopback host ({}) refused — set {opt_in_key} \
to permit it",
for_log(next.as_str())
));
}
if previous.len() > 3 {
Hop::Refuse("too many redirects".to_string())
} else if insecure {
Hop::FollowInsecure
} else {
Hop::Follow
}
}
pub(crate) fn http_client(
allow_insecure_http: bool,
opt_in_key: String,
) -> Result<reqwest::Client, reqwest::Error> {
reqwest::Client::builder()
.timeout(JWKS_FETCH_TIMEOUT)
.redirect(reqwest::redirect::Policy::custom(
move |attempt| match judge_redirect(
attempt.url(),
attempt.previous(),
allow_insecure_http,
&opt_in_key,
) {
Hop::Follow => attempt.follow(),
Hop::FollowInsecure => {
warn!(
url = %for_log(attempt.url().as_str()),
"OAuth: following a redirect to plain http on a non-loopback host \
({opt_in_key} is set) — signing keys fetched over it can be \
substituted by anyone on the path"
);
attempt.follow()
}
Hop::Refuse(reason) => attempt.error(reason),
},
))
.build()
}
struct CachedKey {
kid: Option<String>,
key: DecodingKey,
algorithms: Vec<Algorithm>,
ambiguous: bool,
}
#[derive(Default)]
struct JwksCache {
jwks_uri: Option<String>,
keys: Vec<CachedKey>,
last_attempt: Option<Instant>,
}
pub(crate) struct JwksStore {
issuer: String,
allow_insecure_http: bool,
algorithms: Vec<Algorithm>,
naming: KeyNamingBuf,
http: reqwest::Client,
jwks: RwLock<JwksCache>,
refresh_lock: Arc<Mutex<()>>,
min_refetch_interval: Duration,
}
impl JwksStore {
pub(crate) fn new(
config: &ResolvedOAuthConfig,
http: reqwest::Client,
min_refetch_interval: Duration,
) -> Self {
Self {
issuer: config.issuer.clone(),
allow_insecure_http: config.allow_insecure_http,
algorithms: config.algorithms.clone(),
naming: config.key_naming.clone(),
http,
jwks: RwLock::new(JwksCache {
jwks_uri: config.jwks_uri.clone().filter(|uri| !uri.trim().is_empty()),
..JwksCache::default()
}),
refresh_lock: Arc::new(Mutex::new(())),
min_refetch_interval,
}
}
pub(crate) async fn refresh_now(self: &Arc<Self>) -> Result<usize, RefreshError> {
let guard = Arc::clone(&self.refresh_lock).lock_owned().await;
self.refresh_detached(guard)
.await
.map_err(|message| RefreshError { message })
}
async fn refresh_detached(
self: &Arc<Self>,
guard: OwnedMutexGuard<()>,
) -> Result<usize, String> {
let store = Arc::clone(self);
let task = tokio::spawn(async move {
let _refreshing = guard;
store.refresh().await
});
task.await
.unwrap_or_else(|e| Err(format!("the key refresh task did not finish: {e}")))
}
pub(crate) async fn cached_decoding_key(
&self,
kid: Option<&str>,
alg: Algorithm,
) -> Option<DecodingKey> {
lookup(&self.jwks.read().await.keys, kid, alg)
}
pub(crate) async fn decoding_key(
self: &Arc<Self>,
kid: Option<&str>,
alg: Algorithm,
) -> Result<DecodingKey, TokenRejection> {
if let Some(key) = lookup(&self.jwks.read().await.keys, kid, alg) {
return Ok(key);
}
let refreshing = Arc::clone(&self.refresh_lock).lock_owned().await;
let last_attempt = {
let cache = self.jwks.read().await;
if let Some(key) = lookup(&cache.keys, kid, alg) {
return Ok(key);
}
cache.last_attempt
};
if let Some(last) = last_attempt
&& last.elapsed() < self.min_refetch_interval
{
return Err(TokenRejection::Invalid(format!(
"no {alg} key for kid {} and the JWKS was refetched less than {}s ago",
describe_kid(kid),
self.min_refetch_interval.as_secs()
)));
}
if let Err(e) = self.refresh_detached(refreshing).await {
warn!(
issuer = %self.issuer,
error = %e,
"JWKS refresh failed — tokens signed by a key we do not already hold will \
be rejected until the next attempt"
);
return Err(TokenRejection::Invalid(format!("JWKS refresh failed: {e}")));
}
lookup(&self.jwks.read().await.keys, kid, alg).ok_or_else(|| {
TokenRejection::Invalid(format!(
"no {alg} key for kid {} in the fetched JWKS",
describe_kid(kid)
))
})
}
async fn refresh(&self) -> Result<usize, String> {
let known_uri = {
let mut cache = self.jwks.write().await;
cache.last_attempt = Some(Instant::now());
cache.jwks_uri.clone()
};
let jwks_uri = match known_uri {
Some(uri) => uri,
None => {
let uri = self.discover_jwks_uri().await?;
info!(
issuer = %self.issuer,
jwks_uri = %uri,
"OAuth: discovered the JWKS URI from the issuer's metadata"
);
if plain_http_non_loopback(&uri) {
warn!(
jwks_uri = %uri,
"the discovered JWKS URI uses plain http on a non-loopback host \
({} is set) — signing keys fetched over it can be substituted by \
anyone on the path. Use https.",
self.naming.key("allow_insecure_http")
);
}
self.jwks.write().await.jwks_uri = Some(uri.clone());
uri
}
};
let keys = self
.fetch_jwks(&jwks_uri)
.await
.map_err(|e| format!("fetching the JWKS from {jwks_uri}: {e}"))?;
let count = keys.len();
debug!(count, jwks_uri = %jwks_uri, "Fetched JWKS");
let previous = std::mem::replace(&mut self.jwks.write().await.keys, keys);
self.warn_about_new_ambiguous_keys(&previous).await;
Ok(count)
}
async fn warn_about_new_ambiguous_keys(&self, previous: &[CachedKey]) {
let already: HashSet<Option<&str>> = previous
.iter()
.filter(|k| k.ambiguous)
.map(|k| k.kid.as_deref())
.collect();
let cache = self.jwks.read().await;
for key in cache.keys.iter().filter(|k| k.ambiguous) {
if already.contains(&key.kid.as_deref()) {
continue;
}
let algorithms: Vec<&str> = key.algorithms.iter().map(|a| a.as_str()).collect();
warn!(
kid = %describe_kid(key.kid.as_deref()),
algorithms = %algorithms.join(" "),
"JWKS key declares no alg, so it may verify any of {} — RFC 8725 §3.1 binds a \
key to one algorithm. Narrow {} to the algorithm the authorization server \
signs with.",
algorithms.join(", "),
self.naming.key("algorithms")
);
}
}
async fn discover_jwks_uri(&self) -> Result<String, String> {
let issuer_key = self.naming.key("issuer");
let mut errors = Vec::new();
for url in discovery_urls(&self.issuer) {
match self.fetch_json(&url).await {
Ok(doc) => match jwks_uri_from_metadata(
&doc,
&self.issuer,
&issuer_key,
self.allow_insecure_http,
&self.naming.key("allow_insecure_http"),
) {
Ok(uri) => return Ok(uri),
Err(e) => errors.push(format!("{url}: {e}")),
},
Err(e) => errors.push(format!("{url}: {e}")),
}
}
Err(format!(
"could not discover a jwks_uri for {issuer_key} {:?} — set {} explicitly or fix \
the issuer. Tried: {}",
self.issuer,
self.naming.key("jwks_uri"),
errors.join("; ")
))
}
async fn fetch_jwks(&self, uri: &str) -> Result<Vec<CachedKey>, String> {
let doc = self.fetch_json(uri).await?;
let entries = doc
.get("keys")
.and_then(Value::as_array)
.ok_or_else(|| "response is not a JWK Set (no \"keys\" array)".to_string())?;
if entries.len() > MAX_JWKS_KEYS {
warn!(
published = entries.len(),
used = MAX_JWKS_KEYS,
"JWK Set has more keys than this server will consider; the rest are ignored"
);
}
let mut keys = Vec::new();
for entry in entries.iter().take(MAX_JWKS_KEYS) {
let jwk: Jwk = match serde_json::from_value(entry.clone()) {
Ok(jwk) => jwk,
Err(e) => {
debug!(error = %e, "Skipping a JWKS entry this server cannot parse");
continue;
}
};
if let Some(key) = cached_key(&jwk, &self.algorithms) {
keys.push(key);
}
}
if keys.is_empty() {
return Err(format!(
"the JWK Set contained no usable signature keys for {} {:?}",
self.naming.key("algorithms"),
self.algorithms
));
}
Ok(keys)
}
async fn fetch_json(&self, url: &str) -> Result<Value, String> {
let mut resp = self
.http
.get(url)
.header(reqwest::header::ACCEPT, "application/json")
.send()
.await
.map_err(|e| context("request failed", &e))?
.error_for_status()
.map_err(|e| context("non-success status", &e))?;
if let Some(len) = resp.content_length()
&& len > MAX_FETCH_BYTES as u64
{
return Err(format!(
"response is {len} bytes, over the {MAX_FETCH_BYTES}-byte cap"
));
}
let mut body = Vec::new();
while let Some(chunk) = resp
.chunk()
.await
.map_err(|e| context("reading the response body", &e))?
{
if body.len() + chunk.len() > MAX_FETCH_BYTES {
return Err(format!("response exceeds the {MAX_FETCH_BYTES}-byte cap"));
}
body.extend_from_slice(&chunk);
}
serde_json::from_slice(&body).map_err(|e| context("response was not JSON", &e))
}
}
fn cached_key(jwk: &Jwk, allowed: &[Algorithm]) -> Option<CachedKey> {
match &jwk.common.public_key_use {
None | Some(PublicKeyUse::Signature) => {}
Some(_) => return None,
}
if let Some(ops) = &jwk.common.key_operations
&& !ops.contains(&KeyOperations::Verify)
{
return None;
}
let mut algorithms = key_algorithms(&jwk.algorithm)?;
if let Some(declared) = &jwk.common.key_algorithm {
match signing_algorithm(declared) {
Some(alg) if algorithms.contains(&alg) => algorithms = vec![alg],
_ => return None,
}
}
algorithms.retain(|alg| allowed.contains(alg));
if algorithms.is_empty() {
return None;
}
let ambiguous = jwk.common.key_algorithm.is_none() && algorithms.len() > 1;
match DecodingKey::from_jwk(jwk) {
Ok(key) => Some(CachedKey {
kid: jwk.common.key_id.clone(),
key,
algorithms,
ambiguous,
}),
Err(e) => {
warn!(
kid = ?jwk.common.key_id.as_deref().map(for_log),
error = %e,
"Skipping unusable JWKS entry"
);
None
}
}
}
fn lookup(keys: &[CachedKey], kid: Option<&str>, alg: Algorithm) -> Option<DecodingKey> {
let mut candidates = keys.iter().filter(|k| k.algorithms.contains(&alg));
match kid {
Some(kid) => candidates
.find(|k| k.kid.as_deref() == Some(kid))
.map(|k| k.key.clone()),
None => {
let only = candidates.next()?;
candidates.next().is_none().then(|| only.key.clone())
}
}
}
fn discovery_urls(issuer: &str) -> Vec<String> {
let trimmed = issuer.trim_end_matches('/');
let mut urls = vec![format!("{trimmed}/.well-known/openid-configuration")];
if let Some((scheme, rest)) = trimmed.split_once("://") {
let (authority, path) = match rest.find('/') {
Some(i) => (&rest[..i], &rest[i..]),
None => (rest, ""),
};
let rfc8414 =
format!("{scheme}://{authority}/.well-known/oauth-authorization-server{path}");
if !urls.contains(&rfc8414) {
urls.push(rfc8414);
}
}
urls
}
fn jwks_uri_from_metadata(
doc: &Value,
issuer: &str,
issuer_key: &str,
allow_insecure_http: bool,
opt_in_key: &str,
) -> Result<String, String> {
let found = doc.get("issuer").and_then(Value::as_str);
if found != Some(issuer) {
return Err(format!(
"metadata issuer {} does not match {issuer_key} {issuer:?} byte-for-byte \
(RFC 8414 §3.3 / OIDC Discovery §4.3: such a document must not be used)",
found.map_or_else(|| "(absent)".to_string(), |f| format!("{:?}", for_log(f)))
));
}
let uri = doc
.get("jwks_uri")
.and_then(Value::as_str)
.ok_or_else(|| "metadata has no jwks_uri".to_string())?;
let parsed =
reqwest::Url::parse(uri).map_err(|e| context("jwks_uri is not an absolute URL", &e))?;
match parsed.scheme() {
"https" => {}
"http" if issuer.starts_with("http://") => {
if !allow_insecure_http && plain_http_non_loopback(uri) {
return Err(format!(
"jwks_uri {:?} uses plain http on a non-loopback host — refused \
(RFC 8414 §2) unless {opt_in_key} is set",
for_log(uri)
));
}
}
other => {
return Err(format!(
"jwks_uri scheme {other:?} is not allowed for issuer {issuer:?}"
));
}
}
Ok(uri.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
const OPT_IN: &str = "mcp.oauth.allow_insecure_http";
#[test]
fn discovery_urls_follow_oidc_then_rfc_8414() {
assert_eq!(
discovery_urls("https://auth.example.com/application/o/wiki/"),
[
"https://auth.example.com/application/o/wiki/.well-known/openid-configuration",
"https://auth.example.com/.well-known/oauth-authorization-server/application/o/wiki",
]
);
assert_eq!(
discovery_urls("https://auth.example.com"),
[
"https://auth.example.com/.well-known/openid-configuration",
"https://auth.example.com/.well-known/oauth-authorization-server",
]
);
}
#[test]
fn discovered_jwks_uri_must_not_downgrade_transport() {
let key = "mcp.oauth.issuer";
let doc = serde_json::json!({
"issuer": "https://auth.example.com",
"jwks_uri": "http://auth.example.com/jwks",
});
assert!(
jwks_uri_from_metadata(&doc, "https://auth.example.com", key, false, OPT_IN).is_err()
);
let doc = serde_json::json!({
"issuer": "https://auth.example.com",
"jwks_uri": "file:///etc/passwd",
});
assert!(
jwks_uri_from_metadata(&doc, "https://auth.example.com", key, false, OPT_IN).is_err()
);
let doc = serde_json::json!({"issuer": "https://auth.example.com"});
assert!(
jwks_uri_from_metadata(&doc, "https://auth.example.com", key, false, OPT_IN).is_err()
);
let doc = serde_json::json!({
"issuer": "https://auth.example.com",
"jwks_uri": "http://auth.example.com/jwks",
});
assert!(
jwks_uri_from_metadata(&doc, "https://auth.example.com", key, true, OPT_IN).is_err()
);
}
#[test]
fn a_loopback_issuer_cannot_discover_a_cleartext_non_loopback_jwks_uri() {
let key = "mcp.oauth.issuer";
let issuer = "http://localhost:9000/app/";
let doc =
serde_json::json!({"issuer": issuer, "jwks_uri": "http://idp.internal.test/jwks"});
let err = jwks_uri_from_metadata(&doc, issuer, key, false, OPT_IN).unwrap_err();
assert!(err.contains("plain http on a non-loopback host"), "{err}");
assert!(err.contains(OPT_IN), "{err}");
assert_eq!(
jwks_uri_from_metadata(&doc, issuer, key, true, OPT_IN).unwrap(),
"http://idp.internal.test/jwks"
);
let doc = serde_json::json!({"issuer": issuer, "jwks_uri": "http://127.0.0.1:9000/jwks"});
assert!(jwks_uri_from_metadata(&doc, issuer, key, false, OPT_IN).is_ok());
}
#[test]
fn redirects_are_held_to_the_insecure_http_policy() {
let url = |s: &str| reqwest::Url::parse(s).unwrap();
let loopback = [url("http://127.0.0.1:9000/jwks")];
let https = [url("https://auth.example.com/jwks")];
let Hop::Refuse(reason) = judge_redirect(
&url("http://idp.internal.test/jwks"),
&loopback,
false,
OPT_IN,
) else {
panic!("a cleartext non-loopback hop must be refused without the opt-in");
};
assert!(reason.contains(OPT_IN), "{reason}");
assert_eq!(
judge_redirect(
&url("http://idp.internal.test/jwks"),
&loopback,
true,
OPT_IN
),
Hop::FollowInsecure
);
assert_eq!(
judge_redirect(&url("http://localhost:9000/keys"), &loopback, false, OPT_IN),
Hop::Follow
);
assert_eq!(
judge_redirect(
&url("https://idp.example.com/keys"),
&loopback,
false,
OPT_IN
),
Hop::Follow
);
for target in ["http://idp.internal.test/jwks", "http://127.0.0.1/jwks"] {
assert!(matches!(
judge_redirect(&url(target), &https, true, OPT_IN),
Hop::Refuse(_)
));
}
let many = [
url("https://a.example.com/"),
url("https://b.example.com/"),
url("https://c.example.com/"),
url("https://d.example.com/"),
];
assert_eq!(
judge_redirect(&url("https://e.example.com/"), &many, false, OPT_IN),
Hop::Refuse("too many redirects".to_string())
);
}
#[test]
fn error_chain_matches_the_context_colon_cause_shape() {
#[derive(Debug, thiserror::Error)]
#[error("outer")]
struct Outer(#[source] Inner);
#[derive(Debug, thiserror::Error)]
#[error("inner")]
struct Inner;
assert_eq!(context("fetching", &Outer(Inner)), "fetching: outer: inner");
}
fn rsa_jwk(extra: Value) -> Jwk {
let mut jwk = serde_json::json!({
"kty": "RSA", "kid": "k", "n": crate::testing::N_A, "e": "AQAB",
});
for (k, v) in extra.as_object().unwrap() {
jwk[k] = v.clone();
}
serde_json::from_value(jwk).unwrap()
}
#[test]
fn key_ops_without_verify_make_a_key_unusable() {
let all: Vec<Algorithm> = crate::DEFAULT_ALGORITHMS
.iter()
.map(|a| crate::parse_algorithm(a).unwrap())
.collect();
for ops in [
serde_json::json!(["encrypt"]),
serde_json::json!(["encrypt", "wrapKey"]),
serde_json::json!(["sign"]),
serde_json::json!([]),
serde_json::json!(["some-future-op"]),
] {
assert!(
cached_key(&rsa_jwk(serde_json::json!({ "key_ops": ops })), &all).is_none(),
"key_ops {ops} must not verify"
);
}
for ops in [
serde_json::json!(["verify"]),
serde_json::json!(["sign", "verify"]),
] {
assert!(
cached_key(&rsa_jwk(serde_json::json!({ "key_ops": ops })), &all).is_some(),
"key_ops {ops} may verify"
);
}
assert!(cached_key(&rsa_jwk(serde_json::json!({})), &all).is_some());
}
#[test]
fn an_alg_less_key_usable_under_several_algorithms_is_flagged_ambiguous() {
let all: Vec<Algorithm> = crate::DEFAULT_ALGORITHMS
.iter()
.map(|a| crate::parse_algorithm(a).unwrap())
.collect();
let key = cached_key(&rsa_jwk(serde_json::json!({})), &all).unwrap();
assert!(key.ambiguous);
assert_eq!(key.algorithms.len(), 6);
let key = cached_key(&rsa_jwk(serde_json::json!({"alg": "PS256"})), &all).unwrap();
assert!(!key.ambiguous);
assert_eq!(key.algorithms, [Algorithm::PS256]);
let key = cached_key(&rsa_jwk(serde_json::json!({})), &[Algorithm::RS256]).unwrap();
assert!(!key.ambiguous);
assert_eq!(key.algorithms, [Algorithm::RS256]);
}
#[test]
fn background_retries_back_off_from_a_minute_to_an_hour() {
assert_eq!(background_retry_delay(1), Duration::from_secs(60));
assert_eq!(background_retry_delay(2), Duration::from_secs(120));
assert_eq!(background_retry_delay(3), Duration::from_secs(240));
assert_eq!(background_retry_delay(6), Duration::from_secs(1920));
assert_eq!(background_retry_delay(7), Duration::from_secs(3600));
assert_eq!(background_retry_delay(u32::MAX), Duration::from_secs(3600));
}
#[test]
fn the_http_client_builds_with_the_enabled_tls_backend() {
http_client(false, OPT_IN.to_string()).expect("the JWKS HTTP client must build");
}
}