use std::collections::HashSet;
use std::sync::{Arc, PoisonError};
use std::time::{Duration, Instant, SystemTime};
use base64::Engine as _;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use jsonwebtoken::DecodingKey;
use jsonwebtoken::jwk::{AlgorithmParameters, Jwk, KeyOperations, PublicKeyUse};
use serde_json::Value;
use tokio::sync::{Mutex, OwnedMutexGuard, RwLock};
use tracing::{Instrument, Span, debug, info, warn};
use crate::algorithms::{Algorithm, key_algorithms, signing_algorithm};
use crate::config::{KeyNamingBuf, ResolvedOAuthConfig};
use crate::observe::record_field;
use crate::token::{InvalidTokenKind, TokenRejection, describe_kid, for_log};
use crate::validator::{
is_canonical_url, is_loopback_url, parsed_plain_http_non_loopback, plain_http_non_loopback,
url_is_loopback,
};
pub(crate) const JWKS_MIN_REFETCH_INTERVAL: Duration = Duration::from_secs(60);
pub const DEFAULT_FETCH_TIMEOUT: Duration = Duration::from_secs(10);
pub const MIN_FETCH_TIMEOUT: Duration = Duration::from_secs(1);
pub const MAX_FETCH_TIMEOUT: Duration = Duration::from_secs(60);
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 KEYLESS_RETRY_FLOOR: Duration = Duration::from_secs(5);
pub(crate) const KEYLESS_RETRY_CAP: Duration = Duration::from_secs(300);
pub(crate) fn keyless_retry_delay(failures: u32) -> Duration {
let doublings = failures.saturating_sub(1).min(16);
KEYLESS_RETRY_FLOOR
.saturating_mul(1 << doublings)
.min(KEYLESS_RETRY_CAP)
}
pub(crate) const MAX_FETCH_BYTES: usize = 256 * 1024;
pub(crate) const MAX_JWKS_KEYS: usize = 64;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error("{message}")]
pub struct RefreshError {
kind: RefreshErrorKind,
message: String,
}
impl RefreshError {
fn new(kind: RefreshErrorKind, message: impl Into<String>) -> Self {
Self {
kind,
message: message.into(),
}
}
fn context(self, context: impl std::fmt::Display) -> Self {
Self {
kind: self.kind,
message: format!("{context}: {}", self.message),
}
}
pub fn kind(&self) -> RefreshErrorKind {
self.kind
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum RefreshErrorKind {
Discovery,
Fetch,
Parse,
NoUsableKeys,
}
impl RefreshErrorKind {
pub fn as_str(self) -> &'static str {
match self {
Self::Discovery => "discovery",
Self::Fetch => "fetch",
Self::Parse => "parse",
Self::NoUsableKeys => "no_usable_keys",
}
}
}
impl std::fmt::Display for RefreshErrorKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct KeySetStatus {
pub keys: usize,
pub jwks_uri: Option<String>,
pub last_attempt: Option<SystemTime>,
pub last_success: Option<SystemTime>,
pub last_error: Option<RefreshError>,
}
impl KeySetStatus {
pub fn is_ready(&self) -> bool {
self.keys > 0
}
}
pub(crate) fn redact_url(raw: &str) -> String {
try_redact_url(raw).unwrap_or_else(|| "<unparseable URL, redacted>".to_string())
}
pub(crate) fn jwks_host(raw: &str) -> String {
try_redact_url(raw)
.and_then(|shown| {
reqwest::Url::parse(&shown)
.ok()?
.host_str()
.map(str::to_owned)
})
.unwrap_or_else(|| redact_url(""))
}
pub(crate) fn debug_url(raw: &str) -> String {
if raw.trim().is_empty() {
raw.to_string()
} else {
redact_url(raw)
}
}
pub(crate) fn try_redact_url(raw: &str) -> Option<String> {
let mut url = reqwest::Url::parse(raw).ok()?;
url.host_str()?;
if url.path().contains('@') {
return None;
}
let userinfo = !url.username().is_empty() || url.password().is_some();
if !userinfo && url.query().is_none() && url.fragment().is_none() {
return Some(raw.to_string());
}
if userinfo && (url.set_password(None).is_err() || url.set_username("***").is_err()) {
return None;
}
if url.query().is_some() {
url.set_query(Some("***"));
}
if url.fragment().is_some() {
url.set_fragment(Some("***"));
}
Some(url.to_string())
}
pub(crate) 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());
}
if previous.first().is_some_and(is_loopback_url) && !is_loopback_url(next) {
return Hop::Refuse(format!(
"redirect from a loopback URL to a non-loopback host ({}) refused — a fetch that \
starts on loopback stays on loopback",
for_log(&redact_url(next.as_str()))
));
}
let insecure = parsed_plain_http_non_loopback(next);
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(&redact_url(next.as_str()))
));
}
if previous.len() > 3 {
Hop::Refuse("too many redirects".to_string())
} else if insecure {
Hop::FollowInsecure
} else {
Hop::Follow
}
}
pub(crate) const PROXY_BYPASS: &str = "localhost, 127.0.0.0/8, ::1";
pub(crate) struct FetchSettings {
pub(crate) timeout: Duration,
pub(crate) roots: Vec<reqwest::Certificate>,
pub(crate) proxy: Option<reqwest::Url>,
}
impl Default for FetchSettings {
fn default() -> Self {
Self {
timeout: DEFAULT_FETCH_TIMEOUT,
roots: Vec::new(),
proxy: None,
}
}
}
pub(crate) struct HttpClients {
pub(crate) normal: reqwest::Client,
pub(crate) loopback: reqwest::Client,
}
impl HttpClients {
pub(crate) fn for_url(&self, url: &str) -> &reqwest::Client {
if url_is_loopback(url) {
&self.loopback
} else {
&self.normal
}
}
}
pub(crate) fn http_clients(
allow_insecure_http: bool,
opt_in_key: &str,
settings: &FetchSettings,
) -> Result<HttpClients, reqwest::Error> {
let mut normal = client_builder(allow_insecure_http, opt_in_key, settings);
if let Some(proxy) = &settings.proxy {
normal = normal.no_proxy().proxy(
reqwest::Proxy::all(proxy.clone())?
.no_proxy(reqwest::NoProxy::from_string(PROXY_BYPASS)),
);
}
Ok(HttpClients {
normal: normal.build()?,
loopback: client_builder(allow_insecure_http, opt_in_key, settings)
.no_proxy()
.dns_resolver(Arc::new(LoopbackResolver))
.build()?,
})
}
struct LoopbackResolver;
impl reqwest::dns::Resolve for LoopbackResolver {
fn resolve(&self, _name: reqwest::dns::Name) -> reqwest::dns::Resolving {
let addrs: Vec<std::net::SocketAddr> = vec![
(std::net::Ipv6Addr::LOCALHOST, 0).into(),
(std::net::Ipv4Addr::LOCALHOST, 0).into(),
];
Box::pin(std::future::ready(Ok(
Box::new(addrs.into_iter()) as reqwest::dns::Addrs
)))
}
}
fn client_builder(
allow_insecure_http: bool,
opt_in_key: &str,
settings: &FetchSettings,
) -> reqwest::ClientBuilder {
let opt_in_key = opt_in_key.to_string();
let mut builder = reqwest::Client::builder().timeout(settings.timeout);
for root in &settings.roots {
builder = builder.add_root_certificate(root.clone());
}
builder.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(&redact_url(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),
},
))
}
pub(crate) struct CachedKey {
kid: Option<String>,
key: DecodingKey,
pub(crate) algorithms: Vec<Algorithm>,
pub(crate) ambiguous: bool,
}
#[derive(Default)]
struct JwksCache {
keys: Vec<CachedKey>,
last_attempt: Option<Instant>,
}
#[derive(Default)]
struct Tracked {
public: KeySetStatus,
jwks_uri: Option<String>,
attempts: u64,
}
impl Tracked {
fn set_jwks_uri(&mut self, uri: Option<String>) {
self.public.jwks_uri = uri.as_deref().map(redact_url);
self.jwks_uri = uri;
}
}
pub(crate) struct JwksStore {
issuer: String,
issuer_host: String,
allow_insecure_http: bool,
discovers_jwks_uri: bool,
algorithms: Vec<Algorithm>,
naming: KeyNamingBuf,
http: HttpClients,
jwks: RwLock<JwksCache>,
status: std::sync::Mutex<Tracked>,
refresh_lock: Arc<Mutex<()>>,
min_refetch_interval: Duration,
}
impl JwksStore {
pub(crate) fn new(
config: &ResolvedOAuthConfig,
http: HttpClients,
min_refetch_interval: Duration,
seed: Vec<CachedKey>,
) -> Self {
let mut tracked = Tracked::default();
let configured_uri = config.jwks_uri.clone().filter(|uri| !uri.trim().is_empty());
let discovers_jwks_uri = configured_uri.is_none();
tracked.set_jwks_uri(configured_uri);
tracked.public.keys = seed.len();
warn_about_new_ambiguous_keys(&[], &seed, &config.key_naming);
let issuer_host = jwks_host(&config.issuer);
crate::observe::set_keys(&issuer_host, seed.len());
Self {
issuer: config.issuer.clone(),
issuer_host,
allow_insecure_http: config.allow_insecure_http,
discovers_jwks_uri,
algorithms: config.algorithms.clone(),
naming: config.key_naming.clone(),
http,
jwks: RwLock::new(JwksCache {
keys: seed,
last_attempt: None,
}),
status: std::sync::Mutex::new(tracked),
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
}
fn status_fields(&self) -> std::sync::MutexGuard<'_, Tracked> {
self.status.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn status(&self) -> KeySetStatus {
self.status_fields().public.clone()
}
pub(crate) fn has_keys(&self) -> bool {
self.status_fields().public.keys > 0
}
#[cfg(test)]
pub(crate) fn refresh_in_flight(&self) -> bool {
self.refresh_lock.try_lock().is_err()
}
async fn refresh_detached(
self: &Arc<Self>,
guard: OwnedMutexGuard<()>,
) -> Result<usize, RefreshError> {
let ours = self.status_fields().attempts + 1;
let store = Arc::clone(self);
let task = tokio::spawn(
async move {
let _refreshing = guard;
store.refresh().await
}
.in_current_span(),
);
task.await.unwrap_or_else(|e| {
let err = RefreshError::new(
RefreshErrorKind::Fetch,
format!("the key refresh task did not finish: {e}"),
);
if e.is_panic() {
let mut status = self.status_fields();
if status.attempts <= ours {
status.public.last_error = Some(err.clone());
}
}
Err(err)
})
}
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
{
let (outage, detail_suffix) = {
let status = self.status_fields();
if status.public.last_error.is_some() {
(true, " (the last JWKS refresh failed)")
} else if status.public.keys == 0 {
(true, " (no signing key is held)")
} else {
(false, "")
}
};
return Err(TokenRejection::invalid(
if outage {
InvalidTokenKind::KeySetUnavailable
} else {
InvalidTokenKind::KeyNotFound
},
format!(
"no {alg} key for kid {} and the JWKS was refetched less than {}s \
ago{detail_suffix}",
describe_kid(kid),
self.min_refetch_interval.as_secs()
),
));
}
if let Err(e) = self.refresh_detached(refreshing).await {
warn!(
issuer = %redact_url(&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(
InvalidTokenKind::KeySetUnavailable,
format!("JWKS refresh failed: {e}"),
));
}
lookup(&self.jwks.read().await.keys, kid, alg).ok_or_else(|| {
TokenRejection::invalid(
InvalidTokenKind::KeyNotFound,
format!(
"no {alg} key for kid {} in the fetched JWKS",
describe_kid(kid)
),
)
})
}
async fn refresh(&self) -> Result<usize, RefreshError> {
let span = tracing::info_span!(
"oauth_rs.jwks_refresh",
jwks.host = tracing::field::Empty,
result = tracing::field::Empty,
keys = tracing::field::Empty,
);
let result = self.refresh_in(&span).instrument(span.clone()).await;
let label = match &result {
Ok(_) => "success",
Err(e) => e.kind().as_str(),
};
let held = self.status_fields().public.keys;
record_field(&span, "result", label);
record_field(&span, "keys", held);
crate::observe::count_refresh(&self.issuer_host, label, held);
result
}
async fn refresh_in(&self, span: &Span) -> Result<usize, RefreshError> {
self.jwks.write().await.last_attempt = Some(Instant::now());
let known_uri = {
let mut status = self.status_fields();
status.attempts += 1;
status.public.last_attempt = Some(SystemTime::now());
status.jwks_uri.clone()
};
if !span.is_disabled()
&& let Some(uri) = known_uri.as_deref()
{
record_field(span, "jwks.host", jwks_host(uri).as_str());
}
let result = self.load(known_uri, span).await;
let status = &mut self.status_fields().public;
match &result {
Ok(count) => {
status.keys = *count;
status.last_success = Some(SystemTime::now());
status.last_error = None;
}
Err(e) => status.last_error = Some(e.clone()),
}
result
}
async fn load(&self, known_uri: Option<String>, span: &Span) -> Result<usize, RefreshError> {
let jwks_uri = match known_uri {
Some(uri) => uri,
None => {
let uri = self.discover_jwks_uri().await?;
if !span.is_disabled() {
record_field(span, "jwks.host", jwks_host(&uri).as_str());
}
info!(
issuer = %redact_url(&self.issuer),
jwks_uri = %redact_url(&uri),
"OAuth: discovered the JWKS URI from the issuer's metadata"
);
if plain_http_non_loopback(&uri) {
warn!(
jwks_uri = %redact_url(&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.status_fields().set_jwks_uri(Some(uri.clone()));
uri
}
};
let shown = redact_url(&jwks_uri);
let keys = match self.fetch_jwks(&jwks_uri).await {
Ok(keys) => keys,
Err(e) => {
if self.discovers_jwks_uri {
self.status_fields().jwks_uri = None;
}
return Err(e.context(format_args!("fetching the JWKS from {shown}")));
}
};
let count = keys.len();
debug!(count, jwks_uri = %shown, "Fetched JWKS");
let previous = std::mem::replace(&mut self.jwks.write().await.keys, keys);
let cache = self.jwks.read().await;
warn_about_new_ambiguous_keys(&previous, &cache.keys, &self.naming);
Ok(count)
}
async fn discover_jwks_uri(&self) -> Result<String, RefreshError> {
let span = tracing::info_span!(
"oauth_rs.jwks_discovery",
issuer.host = tracing::field::Empty,
jwks.host = tracing::field::Empty,
result = tracing::field::Empty,
);
if !span.is_disabled() {
record_field(&span, "issuer.host", self.issuer_host.as_str());
}
let result = self.discover_in().instrument(span.clone()).await;
match &result {
Ok(uri) => {
if !span.is_disabled() {
record_field(&span, "jwks.host", jwks_host(uri).as_str());
}
record_field(&span, "result", "success");
}
Err(e) => {
record_field(&span, "result", e.kind().as_str());
}
}
result
}
async fn discover_in(&self) -> Result<String, RefreshError> {
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!("{}: {e}", redact_url(&url))),
},
Err(e) => errors.push(format!("{}: {e}", redact_url(&url))),
}
}
Err(RefreshError::new(
RefreshErrorKind::Discovery,
format!(
"could not discover a jwks_uri for {issuer_key} {:?} — set {} explicitly or \
fix the issuer. Tried: {}",
redact_url(&self.issuer),
self.naming.key("jwks_uri"),
errors.join("; ")
),
))
}
async fn fetch_jwks(&self, uri: &str) -> Result<Vec<CachedKey>, RefreshError> {
let doc = self.fetch_json(uri).await?;
keys_from_jwk_set(&doc, &self.algorithms, &self.naming)
}
async fn fetch_json(&self, url: &str) -> Result<Value, RefreshError> {
let fetch = |message: String| RefreshError::new(RefreshErrorKind::Fetch, message);
let mut resp = self
.http
.for_url(url)
.get(url)
.header(reqwest::header::ACCEPT, "application/json")
.send()
.await
.map_err(|e| fetch(context("request failed", &e.without_url())))?;
let status = resp.status();
if !status.is_success() {
return Err(fetch(format!("non-success status: {status}")));
}
if let Some(len) = resp.content_length()
&& len > MAX_FETCH_BYTES as u64
{
return Err(fetch(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| fetch(context("reading the response body", &e.without_url())))?
{
if body.len() + chunk.len() > MAX_FETCH_BYTES {
return Err(fetch(format!(
"response exceeds the {MAX_FETCH_BYTES}-byte cap"
)));
}
body.extend_from_slice(&chunk);
}
serde_json::from_slice(&body).map_err(|e| {
RefreshError::new(
RefreshErrorKind::Parse,
context("response was not JSON", &e),
)
})
}
}
fn warn_about_new_ambiguous_keys(
previous: &[CachedKey],
current: &[CachedKey],
naming: &KeyNamingBuf,
) {
let already: HashSet<Option<&str>> = previous
.iter()
.filter(|k| k.ambiguous)
.map(|k| k.kid.as_deref())
.collect();
for key in current.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(", "),
naming.key("algorithms")
);
}
}
pub(crate) fn keys_from_jwk_set(
doc: &Value,
allowed: &[Algorithm],
naming: &KeyNamingBuf,
) -> Result<Vec<CachedKey>, RefreshError> {
let entries = doc.get("keys").and_then(Value::as_array).ok_or_else(|| {
RefreshError::new(RefreshErrorKind::Parse, "not a JWK Set (no \"keys\" array)")
})?;
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) {
if let Some(key) = parse_jwks_entry(entry, allowed) {
keys.push(key);
}
}
if keys.is_empty() {
return Err(RefreshError::new(
RefreshErrorKind::NoUsableKeys,
format!(
"the JWK Set contained no usable signature keys for {} {:?}",
naming.key("algorithms"),
allowed
),
));
}
Ok(keys)
}
pub(crate) fn keys_from_jwk_set_json(
json: &str,
allowed: &[Algorithm],
naming: &KeyNamingBuf,
) -> Result<Vec<CachedKey>, RefreshError> {
if json.len() > MAX_FETCH_BYTES {
return Err(RefreshError::new(
RefreshErrorKind::Parse,
format!(
"it is {} bytes, over the {MAX_FETCH_BYTES}-byte cap",
json.len()
),
));
}
let doc: Value = serde_json::from_str(json)
.map_err(|e| RefreshError::new(RefreshErrorKind::Parse, context("not JSON", &e)))?;
keys_from_jwk_set(&doc, allowed, naming)
}
pub(crate) fn parse_jwks_entry(entry: &Value, allowed: &[Algorithm]) -> Option<CachedKey> {
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");
return None;
}
};
cached_key(&jwk, allowed)
}
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;
if let AlgorithmParameters::RSA(rsa) = &jwk.algorithm
&& !rsa_components_can_verify(&rsa.n, &rsa.e)
{
warn!(
kid = ?jwk.common.key_id.as_deref().map(for_log),
"Skipping an RSA JWKS entry that cannot verify any signature (its modulus is not \
2048 to 8192 bits, or its exponent is not an odd number from 3 to 2^33 - 1, \
minimally encoded)"
);
return None;
}
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 rsa_components_can_verify(n: &str, e: &str) -> bool {
let (Ok(n), Ok(e)) = (URL_SAFE_NO_PAD.decode(n), URL_SAFE_NO_PAD.decode(e)) else {
return false;
};
let modulus_bits = match n.first() {
Some(&top) if top != 0 => (n.len() - 1) * 8 + (8 - top.leading_zeros() as usize),
_ => 0,
};
let modulus_ok = (2048..=8192).contains(&modulus_bits) && n.last().is_some_and(|b| b & 1 == 1);
let exponent_ok = (1..=5).contains(&e.len())
&& e[0] != 0
&& e[e.len() - 1] & 1 == 1
&& (3..=(1u64 << 33) - 1)
.contains(&e.iter().fold(0u64, |acc, b| (acc << 8) | u64::from(*b)));
modulus_ok && exponent_ok
}
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())
}
}
}
pub(crate) fn discovery_urls(issuer: &str) -> Vec<String> {
const OIDC: &str = "/.well-known/openid-configuration";
if reqwest::Url::parse(issuer.trim()).is_err() {
return vec![OIDC.to_string()];
}
let trimmed = issuer.trim_end_matches('/');
let mut urls = vec![format!("{trimmed}{OIDC}")];
if is_canonical_url(issuer)
&& 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 shown_issuer = redact_url(issuer);
let found = doc.get("issuer").and_then(Value::as_str);
if found != Some(issuer) {
return Err(format!(
"metadata issuer {} does not match {issuer_key} {shown_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(&redact_url(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.trim())
.map_err(|e| context("jwks_uri is not an absolute URL", &e))?;
let issuer_is_http =
reqwest::Url::parse(issuer.trim()).is_ok_and(|issuer| issuer.scheme() == "http");
match parsed.scheme() {
"https" => {}
"http" if issuer_is_http => {
if !allow_insecure_http && parsed_plain_http_non_loopback(&parsed) {
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(&redact_url(uri))
));
}
}
other => {
return Err(format!(
"jwks_uri scheme {other:?} is not allowed for issuer {shown_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 a_non_canonical_cleartext_jwks_uri_is_judged_by_where_it_really_goes() {
let key = "mcp.oauth.issuer";
let issuer = "http://localhost:9000/app/";
for uri in [
"http:/idp.internal.test/jwks",
"http:idp.internal.test/jwks",
"HTTP:\\\\idp.internal.test\\jwks",
" http://idp.internal.test/jwks",
] {
let doc = serde_json::json!({"issuer": issuer, "jwks_uri": uri});
let err = jwks_uri_from_metadata(&doc, issuer, key, false, OPT_IN)
.expect_err(&format!("{uri:?} must be refused"));
assert!(
err.contains("plain http on a non-loopback host"),
"{uri:?}: {err}"
);
}
let issuer = "https://auth.example.com";
for uri in [
"http:/idp.internal.test/jwks",
"HTTP:idp.internal.test/jwks",
] {
let doc = serde_json::json!({"issuer": issuer, "jwks_uri": uri});
assert!(
jwks_uri_from_metadata(&doc, issuer, key, true, OPT_IN).is_err(),
"{uri:?}"
);
}
let issuer = "HTTP://localhost:9000/app/";
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 plain = [url("http://idp.internal.test/start")];
let https = [url("https://auth.example.com/jwks")];
let Hop::Refuse(reason) =
judge_redirect(&url("http://idp.internal.test/jwks"), &plain, 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"), &plain, true, OPT_IN),
Hop::FollowInsecure
);
for target in [
"http://idp.internal.test/jwks",
"https://idp.example.com/keys",
"http://203.0.113.1/keys",
"http://localhost.example.test/keys",
] {
for opt_in in [false, true] {
let Hop::Refuse(reason) = judge_redirect(&url(target), &loopback, opt_in, OPT_IN)
else {
panic!("{target} after a loopback start must be refused");
};
assert!(
reason.contains("loopback URL to a non-loopback host"),
"{reason}"
);
}
let chain = [loopback[0].clone(), url("http://localhost:9000/moved")];
assert!(matches!(
judge_redirect(&url(target), &chain, true, OPT_IN),
Hop::Refuse(_)
));
}
for target in [
"http://localhost:9000/keys",
"http://127.0.0.2:9000/keys",
"http://[::1]:9000/keys",
"http://app.localhost:9000/keys",
] {
assert_eq!(
judge_redirect(&url(target), &loopback, false, OPT_IN),
Hop::Follow,
"{target}"
);
}
assert_eq!(
judge_redirect(&url("https://idp.example.com/keys"), &plain, 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());
}
fn all_algorithms() -> Vec<Algorithm> {
crate::DEFAULT_ALGORITHMS
.iter()
.map(|a| crate::parse_algorithm(a).unwrap())
.collect()
}
fn rsa_entry(extra: Value) -> Value {
let mut entry = serde_json::json!({
"kty": "RSA", "kid": "k", "n": crate::testing::N_A, "e": "AQAB",
});
for (k, v) in extra.as_object().unwrap() {
entry[k] = v.clone();
}
entry
}
#[test]
fn a_plain_signature_entry_is_parsed_into_a_key() {
let all = all_algorithms();
let key = parse_jwks_entry(&rsa_entry(serde_json::json!({})), &all).unwrap();
assert_eq!(key.kid.as_deref(), Some("k"));
let key = parse_jwks_entry(&rsa_entry(serde_json::json!({"use": "sig"})), &all).unwrap();
assert_eq!(key.kid.as_deref(), Some("k"));
}
#[test]
fn an_encryption_use_entry_is_skipped() {
let all = all_algorithms();
for usage in ["enc", "something-else"] {
assert!(
parse_jwks_entry(&rsa_entry(serde_json::json!({ "use": usage })), &all).is_none(),
"use {usage} must not verify"
);
}
}
#[test]
fn a_key_ops_entry_without_verify_is_skipped_at_entry_level_too() {
let all = all_algorithms();
let entry = rsa_entry(serde_json::json!({ "key_ops": ["encrypt"] }));
assert!(parse_jwks_entry(&entry, &all).is_none());
}
#[test]
fn an_hmac_entry_is_skipped() {
let all = all_algorithms();
let entry = serde_json::json!({"kty": "oct", "kid": "hmac", "k": "c2VjcmV0"});
assert!(parse_jwks_entry(&entry, &all).is_none());
}
#[test]
fn an_unparseable_entry_is_skipped_not_fatal() {
let all = all_algorithms();
for entry in [
serde_json::json!({"kty": "OKP", "crv": "X25519", "kid": "x", "x": "AA"}),
serde_json::json!({"kty": "no-such-type"}),
serde_json::json!({"kid": "no kty at all"}),
serde_json::json!("not an object"),
serde_json::json!(null),
serde_json::json!(42),
serde_json::json!([]),
] {
assert!(parse_jwks_entry(&entry, &all).is_none(), "{entry}");
}
assert!(parse_jwks_entry(&rsa_entry(serde_json::json!({})), &all).is_some());
}
#[test]
fn an_entry_outside_the_allowlist_is_skipped() {
let entry = rsa_entry(serde_json::json!({}));
assert!(parse_jwks_entry(&entry, &[Algorithm::ES256]).is_none());
assert!(parse_jwks_entry(&entry, &[]).is_none());
}
#[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 an_rsa_key_that_cannot_verify_is_skipped() {
let all = all_algorithms();
let b64 = |bytes: &[u8]| URL_SAFE_NO_PAD.encode(bytes);
let mut modulus = vec![0xc5_u8; 256];
modulus[255] = 0x01;
let odd_2048 = b64(&modulus);
assert!(rsa_components_can_verify(crate::testing::N_A, "AQAB"));
assert!(rsa_components_can_verify(&odd_2048, "Aw"));
assert!(rsa_components_can_verify(&b64(&[0xff; 1024]), "AQAB"));
for (what, n, e) in [
("empty e", odd_2048.clone(), String::new()),
("e = 1", odd_2048.clone(), b64(&[1])),
("even e", odd_2048.clone(), b64(&[0x01, 0x00, 0x00])),
(
"e with a leading zero",
odd_2048.clone(),
b64(&[0, 1, 0, 1]),
),
(
"e over 2^33 - 1",
odd_2048.clone(),
b64(&[0x02, 0, 0, 0, 1]),
),
("e over 5 bytes", odd_2048.clone(), b64(&[1, 0, 0, 0, 0, 1])),
("empty n", String::new(), "AQAB".into()),
("n under 2048 bits", b64(&[0xc5; 255]), "AQAB".into()),
(
"n of 2047 bits",
b64(&[&[0x7f][..], &modulus[1..]].concat()),
"AQAB".into(),
),
(
"n of 2041 bits",
b64(&[&[0x01][..], &modulus[1..]].concat()),
"AQAB".into(),
),
("n over 8192 bits", b64(&[0xc5; 1025]), "AQAB".into()),
("even n", b64(&[0xc4; 256]), "AQAB".into()),
(
"n with a leading zero",
b64(&[&[0][..], &modulus[..]].concat()),
"AQAB".into(),
),
("n not base64url", "!!".repeat(200), "AQAB".into()),
] {
assert!(!rsa_components_can_verify(&n, &e), "{what}");
let entry = serde_json::json!({"kty": "RSA", "kid": "k", "n": n, "e": e});
assert!(parse_jwks_entry(&entry, &all).is_none(), "{what}");
}
let naming = KeyNamingBuf::Dotted("oauth".into());
let bad = serde_json::json!({"kty": "RSA", "kid": "bad", "n": odd_2048, "e": ""});
let mut entries = vec![bad; MAX_JWKS_KEYS];
entries.push(rsa_entry(serde_json::json!({})));
let doc = serde_json::json!({ "keys": entries });
let err = keys_from_jwk_set(&doc, &all, &naming).err().unwrap();
assert_eq!(err.kind(), RefreshErrorKind::NoUsableKeys);
}
#[tokio::test]
async fn the_loopback_client_resolves_a_name_no_dns_would() {
let jwks = crate::testing::spawn_jwks_server("200 OK", crate::testing::jwks_body()).await;
let url = jwks.url.replace("127.0.0.1", "nonexistent-name.invalid");
let clients = http_clients(false, OPT_IN, &FetchSettings::default()).unwrap();
let fetched = clients.loopback.get(&url).send().await.unwrap();
assert!(fetched.status().is_success(), "{}", fetched.status());
assert_eq!(
jwks.hits.load(std::sync::atomic::Ordering::SeqCst),
1,
"the URL's port was kept"
);
assert!(
clients.normal.get(&url).send().await.is_err(),
"the system resolver must not resolve {url}"
);
assert_eq!(jwks.hits.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[tokio::test]
async fn the_loopback_resolver_answers_every_name_with_loopback() {
use reqwest::dns::Resolve;
for name in ["localhost", "idp.localhost", "example.test", "a.b.c"] {
let addrs: Vec<std::net::SocketAddr> = LoopbackResolver
.resolve(name.parse().unwrap())
.await
.unwrap()
.collect();
assert_eq!(addrs.len(), 2, "{name}");
assert!(
addrs.iter().all(|a| a.ip().is_loopback() && a.port() == 0),
"{name}: {addrs:?}"
);
}
}
#[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 redact_url_masks_userinfo_query_and_fragment_only() {
for (raw, shown) in [
(
"https://alice:s3cret@idp.example.com:8443/jwks",
"https://***@idp.example.com:8443/jwks",
),
(
"https://alice@idp.example.com/jwks",
"https://***@idp.example.com/jwks",
),
(
"https://:s3cret@idp.example.com/jwks",
"https://***@idp.example.com/jwks",
),
(
"https://idp.example.com/jwks?key=t0ken",
"https://idp.example.com/jwks?***",
),
(
"https://alice:s3cret@idp.example.com/o/app/jwks?key=t0ken#frag",
"https://***@idp.example.com/o/app/jwks?***#***",
),
(
"http://alice:s3cret@[::1]:9000/jwks?key=t0ken",
"http://***@[::1]:9000/jwks?***",
),
] {
let redacted = redact_url(raw);
assert_eq!(redacted, shown, "{raw}");
for secret in ["alice", "s3cret", "t0ken", "frag"] {
assert!(!redacted.contains(secret), "{raw} -> {redacted}");
}
}
for raw in [
"https://idp.example.com/app/",
"https://IDP.example.com",
"http://[::1]:9000/jwks",
] {
assert_eq!(redact_url(raw), raw);
}
assert_eq!(
redact_url("https://idp.example.test/jwks?u=alice:s3cret@x#y@z"),
"https://idp.example.test/jwks?***#***"
);
for raw in [
"http://alice:1234/s3cret@proxy.example.test:3128",
"https://idp.example.test/alice:s3cret@x/jwks",
"",
"not a url",
"alice:s3cret@idp.example.com/jwks",
"http://[::1",
"https://alice:s3cret@",
] {
let redacted = redact_url(raw);
assert!(!redacted.contains("s3cret"), "{raw} -> {redacted}");
assert_eq!(redacted, "<unparseable URL, redacted>", "{raw}");
}
}
#[test]
fn keyless_retries_back_off_from_five_seconds_to_five_minutes() {
let secs: Vec<u64> = (1..=9).map(|n| keyless_retry_delay(n).as_secs()).collect();
assert_eq!(secs, [5, 10, 20, 40, 80, 160, 300, 300, 300]);
assert_eq!(keyless_retry_delay(0), KEYLESS_RETRY_FLOOR);
assert_eq!(keyless_retry_delay(u32::MAX), KEYLESS_RETRY_CAP);
for n in 0..64 {
assert!(
keyless_retry_delay(n) >= Duration::from_secs(5),
"never under the floor"
);
}
}
#[test]
fn refresh_error_kinds_have_stable_labels() {
let labels: Vec<&str> = [
RefreshErrorKind::Discovery,
RefreshErrorKind::Fetch,
RefreshErrorKind::Parse,
RefreshErrorKind::NoUsableKeys,
]
.iter()
.map(|k| k.as_str())
.collect();
assert_eq!(labels, ["discovery", "fetch", "parse", "no_usable_keys"]);
let err = RefreshError::new(RefreshErrorKind::Parse, "inner").context("outer");
assert_eq!(err.to_string(), "outer: inner");
assert_eq!(err.kind(), RefreshErrorKind::Parse);
}
#[test]
fn the_http_client_builds_with_the_enabled_tls_backend() {
http_clients(false, OPT_IN, &FetchSettings::default())
.expect("the JWKS HTTP clients must build");
}
}