use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use ::axum::Json;
use ::axum::Router;
use ::axum::body::Body;
use ::axum::extract::{Request, State};
use ::axum::middleware::Next;
use ::axum::response::{IntoResponse, Response};
use ::axum::routing::{any, get};
use http::header::{AUTHORIZATION, WWW_AUTHENTICATE};
use http::request::Parts;
use http::{HeaderMap, HeaderName, HeaderValue, Method, StatusCode};
use tracing::{debug, warn};
use crate::authenticate::{Credential, authenticate};
use crate::challenge::PROTECTED_RESOURCE_METADATA_PREFIX;
use crate::policy::StaticTokenDecision;
use crate::token::{TokenRejection, for_log};
use crate::validator::OAuthValidator;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum CredentialSource {
Bearer(HeaderName),
Raw(HeaderName),
}
impl CredentialSource {
pub fn authorization_bearer() -> Self {
Self::Bearer(AUTHORIZATION)
}
fn candidate<'h>(&self, headers: &'h HeaderMap) -> Option<&'h str> {
match self {
Self::Bearer(name) => headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(bearer_credential),
Self::Raw(name) => headers.get(name).and_then(|v| v.to_str().ok()),
}
}
fn header_name(&self) -> &HeaderName {
match self {
Self::Bearer(name) | Self::Raw(name) => name,
}
}
}
fn bearer_credential(header: &str) -> &str {
match header.split_once(' ') {
Some((scheme, token)) if scheme.eq_ignore_ascii_case("bearer") => token.trim(),
_ => "",
}
}
#[non_exhaustive]
pub struct RejectContext<'a> {
pub rejection: &'a TokenRejection,
pub status: StatusCode,
pub request: &'a Parts,
}
impl std::fmt::Debug for RejectContext<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let header_names: Vec<&str> = self
.request
.headers
.keys()
.map(HeaderName::as_str)
.collect();
f.debug_struct("RejectContext")
.field("rejection", self.rejection)
.field("status", &self.status)
.field("method", &self.request.method)
.field("uri", &self.request.uri)
.field("version", &self.request.version)
.field("header_names", &header_names)
.finish_non_exhaustive()
}
}
pub type RejectFn = Arc<dyn Fn(RejectContext<'_>) -> Response + Send + Sync>;
pub const DEFAULT_STATIC_CHALLENGE: &str = "Bearer error=\"invalid_token\"";
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum AuthLayerError {
#[error(
"no credential is configured: give a static token and/or an OAuth validator \
(AuthLayer::allow_unauthenticated is the explicit opt-out)"
)]
NoCredential,
#[error("no credential source is configured: a request could never present a credential")]
NoSources,
#[error(
"the static-token decision was made with OAuth enabled, but no OAuth validator \
was given"
)]
DecisionNeedsOAuth,
#[error(
"the static-token decision was made with OAuth disabled, but an OAuth validator \
was given"
)]
DecisionWithoutOAuth,
#[error(
"the OAuth WWW-Authenticate challenge is not a valid HTTP header value — check \
the resource URL and scopes for control or non-ASCII characters"
)]
InvalidChallenge,
}
#[derive(Clone)]
pub struct AuthLayer {
inner: Arc<Mode>,
}
enum Mode {
Enforce(Enforce),
AllowUnauthenticated,
}
struct Enforce {
static_token: Option<String>,
oauth: Option<Arc<OAuthValidator>>,
sources: Vec<CredentialSource>,
on_reject: Option<RejectFn>,
oauth_challenges: Option<(HeaderValue, HeaderValue)>,
static_challenge: Option<HeaderValue>,
}
impl std::fmt::Debug for AuthLayer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &*self.inner {
Mode::AllowUnauthenticated => f
.debug_struct("AuthLayer")
.field("allow_unauthenticated", &true)
.finish(),
Mode::Enforce(e) => f
.debug_struct("AuthLayer")
.field(
"static_token",
&e.static_token.as_ref().map(|_| "<redacted>"),
)
.field("oauth", &e.oauth)
.field("sources", &e.sources)
.field("on_reject", &e.on_reject.as_ref().map(|_| "<fn>"))
.field("static_challenge", &e.static_challenge)
.finish(),
}
}
}
impl AuthLayer {
pub fn builder() -> AuthLayerBuilder {
AuthLayerBuilder::default()
}
pub fn allow_unauthenticated() -> Self {
Self {
inner: Arc::new(Mode::AllowUnauthenticated),
}
}
pub fn from_decision(
decision: StaticTokenDecision,
oauth: Option<Arc<OAuthValidator>>,
) -> Result<Self, AuthLayerError> {
Self::builder()
.optional_oauth(oauth)
.build_with_decision(decision)
}
pub fn allows_unauthenticated(&self) -> bool {
matches!(*self.inner, Mode::AllowUnauthenticated)
}
pub fn oauth(&self) -> Option<&Arc<OAuthValidator>> {
match &*self.inner {
Mode::Enforce(e) => e.oauth.as_ref(),
Mode::AllowUnauthenticated => None,
}
}
}
#[derive(Default)]
pub struct AuthLayerBuilder {
static_token: Option<String>,
oauth: Option<Arc<OAuthValidator>>,
sources: Option<Vec<CredentialSource>>,
on_reject: Option<RejectFn>,
static_challenge: Option<Option<HeaderValue>>,
}
impl std::fmt::Debug for AuthLayerBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AuthLayerBuilder")
.field(
"static_token",
&self.static_token.as_ref().map(|_| "<redacted>"),
)
.field("oauth", &self.oauth)
.field("sources", &self.sources)
.field("on_reject", &self.on_reject.as_ref().map(|_| "<fn>"))
.field("static_challenge", &self.static_challenge)
.finish()
}
}
impl AuthLayerBuilder {
pub fn static_token(mut self, token: impl Into<String>) -> Self {
self.static_token = Some(token.into());
self
}
pub fn optional_static_token(mut self, token: Option<String>) -> Self {
self.static_token = token;
self
}
pub fn oauth(mut self, validator: Arc<OAuthValidator>) -> Self {
self.oauth = Some(validator);
self
}
pub fn optional_oauth(mut self, validator: Option<Arc<OAuthValidator>>) -> Self {
self.oauth = validator;
self
}
pub fn sources(mut self, sources: impl IntoIterator<Item = CredentialSource>) -> Self {
self.sources = Some(sources.into_iter().collect());
self
}
pub fn static_challenge(mut self, challenge: Option<HeaderValue>) -> Self {
self.static_challenge = Some(challenge);
self
}
pub fn on_reject(
mut self,
f: impl Fn(RejectContext<'_>) -> Response + Send + Sync + 'static,
) -> Self {
self.on_reject = Some(Arc::new(f));
self
}
pub fn build_with_decision(
mut self,
decision: StaticTokenDecision,
) -> Result<AuthLayer, AuthLayerError> {
match (decision.oauth_enabled(), self.oauth.is_some()) {
(true, false) => return Err(AuthLayerError::DecisionNeedsOAuth),
(false, true) => return Err(AuthLayerError::DecisionWithoutOAuth),
_ => {}
}
if decision == StaticTokenDecision::Unauthenticated {
return Ok(AuthLayer::allow_unauthenticated());
}
self.static_token = decision.into_static_token();
self.build()
}
pub fn build(self) -> Result<AuthLayer, AuthLayerError> {
let static_token = self.static_token.filter(|t| !t.is_empty());
if static_token.is_none() && self.oauth.is_none() {
return Err(AuthLayerError::NoCredential);
}
let sources = self
.sources
.unwrap_or_else(|| vec![CredentialSource::authorization_bearer()]);
if sources.is_empty() {
return Err(AuthLayerError::NoSources);
}
let header = |challenge: String| {
HeaderValue::from_str(&challenge).map_err(|_| AuthLayerError::InvalidChallenge)
};
let oauth_challenges = match &self.oauth {
Some(v) => Some((
header(v.invalid_token_challenge())?,
header(v.insufficient_scope_challenge())?,
)),
None => None,
};
let static_challenge = self
.static_challenge
.unwrap_or_else(|| Some(HeaderValue::from_static(DEFAULT_STATIC_CHALLENGE)));
Ok(AuthLayer {
inner: Arc::new(Mode::Enforce(Enforce {
static_token,
oauth: self.oauth,
sources,
on_reject: self.on_reject,
oauth_challenges,
static_challenge,
})),
})
}
}
impl Enforce {
fn reject(&self, rejection: &TokenRejection, request: &Parts) -> Response {
let status = match rejection {
TokenRejection::InsufficientScope => StatusCode::FORBIDDEN,
TokenRejection::Invalid(_) | TokenRejection::Missing => StatusCode::UNAUTHORIZED,
};
let challenge = match (&self.oauth_challenges, rejection) {
(Some((_, insufficient)), TokenRejection::InsufficientScope) => Some(insufficient),
(Some((invalid, _)), _) => Some(invalid),
(None, _) => self.static_challenge.as_ref(),
};
let mut response = match &self.on_reject {
Some(f) => f(RejectContext {
rejection,
status,
request,
}),
None => {
let mut response = Response::new(Body::empty());
*response.status_mut() = status;
response
}
};
*response.status_mut() = status;
if let Some(value) = challenge {
response
.headers_mut()
.insert(WWW_AUTHENTICATE, value.clone());
}
response
}
}
impl AuthLayer {
async fn check(&self, request: Request) -> Result<Request, Response> {
let enforce = match &*self.inner {
Mode::AllowUnauthenticated => return Ok(request),
Mode::Enforce(enforce) => enforce,
};
let (mut parts, body) = request.into_parts();
for (name, value) in parts.headers.iter_mut() {
if enforce.sources.iter().any(|s| s.header_name() == name) {
value.set_sensitive(true);
}
}
let result = {
let headers = &parts.headers;
let candidates = enforce.sources.iter().filter_map(|s| s.candidate(headers));
authenticate(
candidates,
enforce.static_token.as_deref(),
enforce.oauth.as_deref(),
)
.await
};
match result {
Ok(Credential::StaticToken) => {
parts.extensions.insert(Credential::StaticToken);
}
Ok(Credential::OAuth(token)) => {
debug!(
path = %parts.uri.path(),
principal = ?token.principal.as_deref().map(for_log),
subject = ?token.subject.as_deref().map(for_log),
scopes = ?token.scopes,
"OAuth bearer auth accepted"
);
parts.extensions.insert(token.clone());
parts.extensions.insert(Credential::OAuth(token));
}
Err(rejection) => {
let path = parts.uri.path();
match (&enforce.oauth, &rejection) {
(None, _) => warn!(path = %path, "Bearer auth rejected"),
(Some(_), TokenRejection::Missing) => {
debug!(path = %path, "No bearer credential presented");
}
(Some(_), _) => {
warn!(path = %path, reason = ?rejection, "OAuth bearer auth rejected");
}
}
return Err(enforce.reject(&rejection, &parts));
}
}
Ok(Request::from_parts(parts, body))
}
}
pub async fn require_auth(State(auth): State<AuthLayer>, request: Request, next: Next) -> Response {
match auth.check(request).await {
Ok(request) => next.run(request).await,
Err(refusal) => refusal,
}
}
impl<S> tower_layer::Layer<S> for AuthLayer {
type Service = AuthService<S>;
fn layer(&self, inner: S) -> Self::Service {
AuthService {
auth: self.clone(),
inner,
}
}
}
#[derive(Clone)]
pub struct AuthService<S> {
auth: AuthLayer,
inner: S,
}
impl<S: std::fmt::Debug> std::fmt::Debug for AuthService<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AuthService")
.field("auth", &self.auth)
.field("inner", &self.inner)
.finish()
}
}
impl<S> tower_service::Service<Request> for AuthService<S>
where
S: tower_service::Service<Request, Response = Response> + Clone + Send + 'static,
S::Future: Send + 'static,
{
type Response = Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Response, S::Error>> + Send + 'static>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let auth = self.auth.clone();
Box::pin(async move {
match auth.check(request).await {
Ok(request) => inner.call(request).await,
Err(refusal) => Ok(refusal),
}
})
}
}
pub fn metadata_router<S>(oauth: Option<Arc<OAuthValidator>>) -> Router<S>
where
S: Clone + Send + Sync + 'static,
{
async fn not_found() -> StatusCode {
StatusCode::NOT_FOUND
}
let catch_all = format!("{PROTECTED_RESOURCE_METADATA_PREFIX}/{{*rest}}");
let Some(validator) = oauth else {
return Router::new()
.route(&catch_all, any(not_found))
.route(PROTECTED_RESOURCE_METADATA_PREFIX, get(not_found));
};
let serve = {
let validator = Arc::clone(&validator);
move || {
let validator = Arc::clone(&validator);
async move { Json(validator.metadata()).into_response() }
}
};
let path: Arc<str> = validator.metadata_path().into();
let suffix = move |request: Request| {
let validator = Arc::clone(&validator);
let path = Arc::clone(&path);
async move {
if request.uri().path() != &*path {
return StatusCode::NOT_FOUND.into_response();
}
match *request.method() {
Method::GET | Method::HEAD => Json(validator.metadata()).into_response(),
_ => method_not_allowed(),
}
}
};
Router::new()
.route(&catch_all, any(suffix))
.route(PROTECTED_RESOURCE_METADATA_PREFIX, get(serve))
}
fn method_not_allowed() -> Response {
(
StatusCode::METHOD_NOT_ALLOWED,
[(http::header::ALLOW, HeaderValue::from_static("GET,HEAD"))],
)
.into_response()
}
#[cfg(test)]
mod tests {
use ::axum::Extension;
use ::axum::middleware;
use tower::ServiceExt;
use super::*;
use crate::AuthorizedToken;
use crate::testing;
const STATIC: &str = "secret";
fn validator(jwks_uri: &str) -> Arc<OAuthValidator> {
Arc::new(OAuthValidator::new(&testing::resolved_config(jwks_uri)).unwrap())
}
fn unreachable_validator() -> Arc<OAuthValidator> {
validator("http://127.0.0.1:1/jwks")
}
fn app(auth: AuthLayer) -> Router {
Router::new()
.route("/test", get(|| async { "ok" }))
.route_layer(middleware::from_fn_with_state(auth, require_auth))
}
fn wiki_app(static_token: Option<&str>, oauth: Option<Arc<OAuthValidator>>) -> Router {
app(AuthLayer::builder()
.optional_static_token(static_token.map(str::to_string))
.optional_oauth(oauth)
.static_challenge(None)
.build()
.unwrap())
}
async fn send(app: &Router, headers: &[(&str, &str)]) -> Response {
let mut req = Request::builder().uri("/test");
for (name, value) in headers {
req = req.header(*name, *value);
}
app.clone()
.oneshot(req.body(Body::empty()).unwrap())
.await
.unwrap()
}
async fn get_with_auth(app: &Router, header: Option<&str>) -> Response {
match header {
Some(h) => send(app, &[("authorization", h)]).await,
None => send(app, &[]).await,
}
}
fn www_authenticate(resp: &Response) -> String {
resp.headers()
.get(WWW_AUTHENTICATE)
.expect("a refusal with OAuth configured must carry WWW-Authenticate")
.to_str()
.unwrap()
.to_string()
}
async fn body_bytes(resp: Response) -> Vec<u8> {
::axum::body::to_bytes(resp.into_body(), 64 * 1024)
.await
.unwrap()
.to_vec()
}
fn unscoped_token() -> String {
testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&serde_json::json!({
"iss": testing::ISSUER, "aud": testing::AUDIENCE,
"exp": testing::now() + 3600, "scope": "openid profile",
}),
)
}
fn expired_token() -> String {
testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&serde_json::json!({
"iss": testing::ISSUER, "aud": testing::AUDIENCE,
"exp": testing::now() - 3600, "scope": "mcp:read",
}),
)
}
#[test]
fn the_builder_refuses_to_build_a_pass_through() {
assert_eq!(
AuthLayer::builder().build().unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
AuthLayer::builder().static_token("").build().unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
AuthLayer::builder()
.optional_static_token(None)
.optional_oauth(None)
.build()
.unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
AuthLayer::builder()
.static_token(STATIC)
.sources([])
.build()
.unwrap_err(),
AuthLayerError::NoSources
);
let built = AuthLayer::builder().static_token(STATIC).build().unwrap();
assert!(!built.allows_unauthenticated());
assert!(built.oauth().is_none());
}
#[tokio::test]
async fn only_the_explicit_opt_out_passes_requests_through() {
let layer = AuthLayer::allow_unauthenticated();
assert!(layer.allows_unauthenticated());
let app =
Router::new()
.route(
"/test",
get(
|c: Option<Extension<Credential>>,
t: Option<Extension<AuthorizedToken>>| async move {
assert!(c.is_none() && t.is_none(), "a pass-through inserts nothing");
"ok"
},
),
)
.route_layer(middleware::from_fn_with_state(layer, require_auth));
assert_eq!(get_with_auth(&app, None).await.status(), StatusCode::OK);
assert_eq!(
get_with_auth(&app, Some("Bearer anything")).await.status(),
StatusCode::OK
);
}
#[test]
fn debug_never_prints_the_static_token() {
let layer = AuthLayer::builder()
.static_token("hunter2")
.build()
.unwrap();
let rendered = format!("{layer:?}");
assert!(!rendered.contains("hunter2"), "{rendered}");
let builder = AuthLayer::builder().static_token("hunter2");
let rendered = format!("{builder:?}");
assert!(!rendered.contains("hunter2"), "{rendered}");
}
#[tokio::test]
async fn static_token_only() {
let app = wiki_app(Some(STATIC), None);
for (header, status) in [
(Some("Bearer secret"), StatusCode::OK),
(Some("Bearer wrong-token"), StatusCode::UNAUTHORIZED),
(None, StatusCode::UNAUTHORIZED),
(Some("Basic c2VjcmV0LXRva2Vu"), StatusCode::UNAUTHORIZED),
] {
let resp = get_with_auth(&app, header).await;
assert_eq!(resp.status(), status, "{header:?}");
assert!(resp.headers().get(WWW_AUTHENTICATE).is_none(), "{header:?}");
}
}
#[tokio::test]
async fn a_static_only_401_carries_a_bearer_challenge_by_default() {
let app = app(AuthLayer::builder().static_token(STATIC).build().unwrap());
for header in [None, Some("Bearer wrong-token"), Some("Basic abc")] {
let resp = get_with_auth(&app, header).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED, "{header:?}");
assert_eq!(
resp.headers()[WWW_AUTHENTICATE],
DEFAULT_STATIC_CHALLENGE,
"{header:?}"
);
}
assert_eq!(
get_with_auth(&app, Some("Bearer secret")).await.status(),
StatusCode::OK
);
let app = super::tests::app(
AuthLayer::from_decision(StaticTokenDecision::StaticOnly(STATIC.into()), None).unwrap(),
);
assert_eq!(
get_with_auth(&app, None).await.headers()[WWW_AUTHENTICATE],
DEFAULT_STATIC_CHALLENGE
);
let custom = HeaderValue::from_static("Bearer realm=\"my-api\"");
let app = super::tests::app(
AuthLayer::builder()
.static_token(STATIC)
.static_challenge(Some(custom.clone()))
.build()
.unwrap(),
);
assert_eq!(
get_with_auth(&app, None).await.headers()[WWW_AUTHENTICATE],
custom
);
}
#[tokio::test]
async fn the_static_challenge_is_ignored_when_oauth_is_configured() {
let v = unreachable_validator();
let layer = AuthLayer::builder()
.static_token(STATIC)
.oauth(Arc::clone(&v))
.static_challenge(Some(HeaderValue::from_static("Bearer realm=\"x\"")))
.build()
.unwrap();
let resp = get_with_auth(&app(layer), None).await;
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
}
#[test]
fn an_oauth_challenge_that_is_not_a_header_value_fails_the_build() {
let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
cfg.resource = "https://kb.example.test/m\ncp".into();
let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
assert_eq!(
AuthLayer::builder().oauth(v).build().unwrap_err(),
AuthLayerError::InvalidChallenge
);
let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
cfg.required_scopes = vec!["a\u{1}b".into()];
let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
assert_eq!(
AuthLayer::builder()
.static_token(STATIC)
.oauth(v)
.build()
.unwrap_err(),
AuthLayerError::InvalidChallenge
);
}
#[tokio::test]
async fn static_bearer_token_still_works_with_oauth_enabled() {
let app = wiki_app(Some(STATIC), Some(unreachable_validator()));
assert_eq!(
get_with_auth(&app, Some("Bearer secret")).await.status(),
StatusCode::OK
);
}
#[tokio::test]
async fn an_oauth_token_is_accepted_alongside_the_static_token() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = wiki_app(Some(STATIC), Some(validator(&jwks.url)));
let header = format!("Bearer {}", testing::valid_token());
assert_eq!(
get_with_auth(&app, Some(&header)).await.status(),
StatusCode::OK
);
assert_eq!(
get_with_auth(&app, Some("Bearer secret")).await.status(),
StatusCode::OK
);
}
#[tokio::test]
async fn a_missing_credential_gets_401_with_a_well_formed_challenge() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = wiki_app(Some(STATIC), Some(validator(&jwks.url)));
for header in [None, Some("Bearer not-the-secret"), Some("Basic abc")] {
let resp = get_with_auth(&app, header).await;
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"header: {header:?}"
);
assert_eq!(
www_authenticate(&resp),
"Bearer error=\"invalid_token\", \
resource_metadata=\"https://kb.example.test\
/.well-known/oauth-protected-resource/mcp\", \
scope=\"mcp:read mcp:write\""
);
}
}
#[tokio::test]
async fn an_invalid_token_gets_401_and_an_insufficient_scope_token_gets_403() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = wiki_app(None, Some(validator(&jwks.url)));
let resp = get_with_auth(&app, Some(&format!("Bearer {}", expired_token()))).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert!(www_authenticate(&resp).contains("error=\"invalid_token\""));
let resp = get_with_auth(&app, Some(&format!("Bearer {}", unscoped_token()))).await;
assert_eq!(
resp.status(),
StatusCode::FORBIDDEN,
"a valid token missing the scope is 403, not 401"
);
assert_eq!(
www_authenticate(&resp),
"Bearer error=\"insufficient_scope\", scope=\"mcp:read\", \
resource_metadata=\"https://kb.example.test\
/.well-known/oauth-protected-resource/mcp\""
);
}
#[tokio::test]
async fn an_authelia_style_scp_token_is_accepted_through_the_middleware() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = wiki_app(None, Some(validator(&jwks.url)));
let token = testing::mint_with(
crate::Algorithm::RS256,
Some(testing::KID_A),
Some("at+jwt"),
&serde_json::json!({
"iss": testing::ISSUER, "aud": [testing::AUDIENCE],
"exp": testing::now() + 3600, "nbf": testing::now(),
"sub": "44726d41-0000-4000-8000-000000000000",
"scp": ["mcp:read", "mcp:write"],
}),
);
let resp = get_with_auth(&app, Some(&format!("Bearer {token}"))).await;
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn the_bearer_scheme_is_case_insensitive_for_both_credentials() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = wiki_app(Some(STATIC), Some(validator(&jwks.url)));
for header in [
"bearer secret".to_string(),
"BEARER secret".to_string(),
"Bearer secret ".to_string(),
format!("bearer {}", testing::valid_token()),
] {
assert_eq!(
get_with_auth(&app, Some(&header)).await.status(),
StatusCode::OK,
"{header:.20}"
);
}
for header in ["Basic secret", "Bearersecret", "secret", "Bearer\tsecret"] {
assert_eq!(
get_with_auth(&app, Some(header)).await.status(),
StatusCode::UNAUTHORIZED,
"{header}"
);
}
}
#[test]
fn bearer_credential_parsing() {
assert_eq!(bearer_credential("Bearer abc"), "abc");
assert_eq!(bearer_credential("bEaReR abc "), "abc");
assert_eq!(bearer_credential("Bearer "), "");
assert_eq!(bearer_credential("Bearer"), "");
assert_eq!(bearer_credential("Basic abc"), "");
assert_eq!(bearer_credential(""), "");
}
mod wiki {
use subtle::ConstantTimeEq;
use super::super::*;
#[derive(Clone)]
pub(super) struct AuthState {
pub(super) bearer_token: Option<String>,
pub(super) oauth: Option<Arc<OAuthValidator>>,
}
impl AuthState {
fn challenge(&self, rejection: &TokenRejection) -> Option<String> {
let oauth = self.oauth.as_ref()?;
Some(match rejection {
TokenRejection::InsufficientScope => oauth.insufficient_scope_challenge(),
TokenRejection::Invalid(_) | TokenRejection::Missing => {
oauth.invalid_token_challenge()
}
})
}
}
fn auth_rejection(auth: &AuthState, rejection: TokenRejection) -> Response {
let status = match rejection {
TokenRejection::InsufficientScope => StatusCode::FORBIDDEN,
TokenRejection::Invalid(_) | TokenRejection::Missing => StatusCode::UNAUTHORIZED,
};
let mut response = Response::builder().status(status);
if let Some(challenge) = auth.challenge(&rejection)
&& let Ok(value) = HeaderValue::from_str(&challenge)
{
response = response.header(WWW_AUTHENTICATE, value);
}
response
.body(Body::empty())
.expect("a status-and-header-only response is always constructible")
}
pub(super) async fn bearer_auth(
State(auth): State<AuthState>,
headers: HeaderMap,
request: Request,
next: Next,
) -> Response {
if auth.bearer_token.is_none() && auth.oauth.is_none() {
return next.run(request).await;
}
let auth_header = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let token = bearer_credential(auth_header);
if let Some(ref expected_token) = auth.bearer_token
&& !token.is_empty()
&& token.as_bytes().ct_eq(expected_token.as_bytes()).into()
{
return next.run(request).await;
}
let Some(ref oauth) = auth.oauth else {
return auth_rejection(
&auth,
TokenRejection::Invalid("static token mismatch".into()),
);
};
match oauth.validate(token).await {
Ok(claims) => {
let mut request = request;
request.extensions_mut().insert(claims);
next.run(request).await
}
Err(TokenRejection::Missing) => auth_rejection(&auth, TokenRejection::Missing),
Err(rejection) => auth_rejection(&auth, rejection),
}
}
fn bearer_credential(header: &str) -> &str {
match header.split_once(' ') {
Some((scheme, token)) if scheme.eq_ignore_ascii_case("bearer") => token.trim(),
_ => "",
}
}
}
fn wiki_oracle_app(static_token: Option<&str>, oauth: Option<Arc<OAuthValidator>>) -> Router {
let auth_state = wiki::AuthState {
bearer_token: static_token.map(str::to_string),
oauth,
};
Router::new()
.route("/test", get(|| async { "ok" }))
.route_layer(middleware::from_fn_with_state(
auth_state,
wiki::bearer_auth,
))
}
async fn assert_same_response(actual: Response, expected: Response, what: &str) {
assert_eq!(actual.status(), expected.status(), "{what}");
assert_eq!(actual.version(), expected.version(), "{what}");
assert_eq!(actual.headers(), expected.headers(), "{what}");
assert_eq!(
body_bytes(actual).await,
body_bytes(expected).await,
"{what}"
);
}
#[tokio::test]
async fn default_responses_are_byte_identical_to_the_wiki_middleware() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let valid = format!("Bearer {}", testing::valid_token());
let lower = format!("bearer {}", testing::valid_token());
let expired = format!("Bearer {}", expired_token());
let unscoped = format!("Bearer {}", unscoped_token());
let headers: Vec<Option<&str>> = vec![
None,
Some(""),
Some("Bearer"),
Some("Bearer "),
Some("Bearer secret"),
Some("bearer secret"),
Some("Bearer secret "),
Some("Bearer wrong"),
Some("Bearersecret"),
Some("Bearer\tsecret"),
Some("Basic c2VjcmV0"),
Some("secret"),
Some(&valid),
Some(&lower),
Some(&expired),
Some(&unscoped),
];
for (static_token, oauth) in [
(Some(STATIC), None),
(Some(STATIC), Some(Arc::clone(&v))),
(None, Some(Arc::clone(&v))),
] {
let ours = wiki_app(static_token, oauth.clone());
let theirs = wiki_oracle_app(static_token, oauth.clone());
for header in &headers {
assert_same_response(
get_with_auth(&ours, *header).await,
get_with_auth(&theirs, *header).await,
&format!(
"static={static_token:?} oauth={} header={header:.30?}",
oauth.is_some()
),
)
.await;
}
}
}
#[tokio::test]
async fn on_reject_shapes_the_body_but_not_the_status_or_challenge() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let layer = AuthLayer::builder()
.static_token(STATIC)
.oauth(Arc::clone(&v))
.on_reject(|cx: RejectContext<'_>| {
assert_eq!(cx.request.uri.path(), "/test");
let status = cx.status;
let body = match cx.rejection {
TokenRejection::InsufficientScope => r#"{"error":"insufficient_scope"}"#,
_ => r#"{"error":"unauthorized"}"#,
};
Response::builder()
.status(StatusCode::OK)
.header("content-type", "application/json")
.header("x-seen-status", status.as_str())
.header(WWW_AUTHENTICATE, "Basic realm=\"nope\"")
.header(WWW_AUTHENTICATE, "Bearer realm=\"also-nope\"")
.body(Body::from(body))
.unwrap()
})
.build()
.unwrap();
let app = app(layer);
let resp = get_with_auth(&app, None).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(resp.headers()["x-seen-status"], "401");
assert_eq!(resp.headers()["content-type"], "application/json");
assert_eq!(
resp.headers().get_all(WWW_AUTHENTICATE).iter().count(),
1,
"the callback's challenges are replaced, not added to"
);
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
assert_eq!(body_bytes(resp).await, br#"{"error":"unauthorized"}"#);
let resp = get_with_auth(&app, Some(&format!("Bearer {}", unscoped_token()))).await;
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
assert_eq!(resp.headers()["x-seen-status"], "403");
assert_eq!(www_authenticate(&resp), v.insufficient_scope_challenge());
assert_eq!(body_bytes(resp).await, br#"{"error":"insufficient_scope"}"#);
}
#[tokio::test]
async fn on_reject_without_oauth_or_a_static_challenge_keeps_its_own_headers() {
let builder = || {
AuthLayer::builder()
.static_token(STATIC)
.on_reject(|cx: RejectContext<'_>| {
Response::builder()
.status(cx.status)
.header(WWW_AUTHENTICATE, "ApiKey")
.body(Body::from("nope"))
.unwrap()
})
};
let layer = builder().static_challenge(None).build().unwrap();
let resp = get_with_auth(&app(layer), Some("Bearer wrong")).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(resp.headers()[WWW_AUTHENTICATE], "ApiKey");
assert_eq!(body_bytes(resp).await, b"nope");
let resp = get_with_auth(&app(builder().build().unwrap()), Some("Bearer wrong")).await;
assert_eq!(resp.headers().get_all(WWW_AUTHENTICATE).iter().count(), 1);
assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
assert_eq!(body_bytes(resp).await, b"nope");
}
#[tokio::test]
async fn the_presented_credential_never_reaches_a_debug_rendering() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let token = unscoped_token(); let seen: Arc<std::sync::Mutex<Vec<String>>> = Arc::default();
let log = Arc::clone(&seen);
let layer = AuthLayer::builder()
.oauth(v)
.sources([
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
])
.on_reject(move |cx: RejectContext<'_>| {
log.lock().unwrap().push(format!("{cx:?}"));
log.lock().unwrap().push(format!("{:?}", cx.request));
Response::new(Body::empty())
})
.build()
.unwrap();
let bearer = format!("Bearer {token}");
let resp = send(
&app(layer),
&[
("authorization", bearer.as_str()),
("x-api-key", "raw-api-key-value"),
("accept", "application/json"),
],
)
.await;
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
let seen = seen.lock().unwrap();
assert_eq!(seen.len(), 2);
for rendered in seen.iter() {
assert!(!rendered.contains(&token), "{rendered}");
assert!(!rendered.contains("raw-api-key-value"), "{rendered}");
}
assert!(seen[0].contains("InsufficientScope"), "{}", seen[0]);
assert!(seen[0].contains("403"), "{}", seen[0]);
assert!(seen[0].contains("authorization"), "{}", seen[0]);
assert!(!seen[0].contains("application/json"), "{}", seen[0]);
}
#[tokio::test]
async fn the_inner_service_sees_the_credential_headers_marked_sensitive() {
let layer = AuthLayer::builder().static_token(STATIC).build().unwrap();
let app = Router::new()
.route(
"/test",
get(|headers: HeaderMap| async move {
assert!(headers["authorization"].is_sensitive());
assert!(!format!("{headers:?}").contains(STATIC));
assert!(!headers["accept"].is_sensitive());
"ok"
}),
)
.route_layer(layer);
let resp = send(
&app,
&[("authorization", "Bearer secret"), ("accept", "text/plain")],
)
.await;
assert_eq!(resp.status(), StatusCode::OK);
}
fn multi_source_app(v: Arc<OAuthValidator>) -> Router {
app(AuthLayer::builder()
.static_token(STATIC)
.oauth(v)
.sources([
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
])
.build()
.unwrap())
}
#[tokio::test]
async fn a_bad_authorization_header_does_not_mask_a_good_raw_header() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = multi_source_app(validator(&jwks.url));
let foreign = format!("Bearer {}", expired_token());
for authorization in [foreign.as_str(), "Bearer garbage", "Basic abc"] {
let resp = send(
&app,
&[("authorization", authorization), ("x-api-key", STATIC)],
)
.await;
assert_eq!(resp.status(), StatusCode::OK, "{authorization:.30}");
}
}
#[tokio::test]
async fn a_bad_raw_header_does_not_mask_a_good_authorization_header() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let app = multi_source_app(validator(&jwks.url));
let valid = format!("Bearer {}", testing::valid_token());
for authorization in ["Bearer secret", valid.as_str()] {
let resp = send(
&app,
&[("authorization", authorization), ("x-api-key", "garbage")],
)
.await;
assert_eq!(resp.status(), StatusCode::OK, "{authorization:.30}");
}
let resp = send(&app, &[("x-api-key", &testing::valid_token())]).await;
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn multi_source_refusals() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let app = multi_source_app(Arc::clone(&v));
let resp = send(
&app,
&[
("authorization", "Bearer garbage"),
("x-api-key", "also-garbage"),
],
)
.await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
let unscoped = unscoped_token();
let resp = send(
&app,
&[
("authorization", "Bearer garbage"),
("x-api-key", &unscoped),
],
)
.await;
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
assert_eq!(www_authenticate(&resp), v.insufficient_scope_challenge());
let resp = send(&app, &[("x-api-key", "Bearer secret")]).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
let resp = send(&app, &[]).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn a_raw_only_layer_ignores_the_authorization_header() {
let layer = AuthLayer::builder()
.static_token(STATIC)
.sources([CredentialSource::Raw(HeaderName::from_static("x-api-key"))])
.build()
.unwrap();
let app = app(layer);
assert_eq!(
send(&app, &[("authorization", "Bearer secret")])
.await
.status(),
StatusCode::UNAUTHORIZED
);
assert_eq!(
send(&app, &[("x-api-key", STATIC)]).await.status(),
StatusCode::OK
);
}
#[tokio::test]
async fn the_credential_and_oauth_token_are_inserted_into_extensions() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let layer = AuthLayer::builder()
.static_token(STATIC)
.oauth(validator(&jwks.url))
.build()
.unwrap();
let app = Router::new()
.route(
"/test",
get(
|Extension(credential): Extension<Credential>,
token: Option<Extension<AuthorizedToken>>| async move {
match (credential, token) {
(Credential::StaticToken, None) => "static".to_string(),
(Credential::OAuth(c), Some(Extension(t))) => {
assert_eq!(c, t);
format!(
"oauth {} {}",
t.subject.as_deref().unwrap_or_default(),
t.has_scope("mcp:write")
)
}
other => panic!("inconsistent extensions: {other:?}"),
}
},
),
)
.route_layer(middleware::from_fn_with_state(layer, require_auth));
let resp = get_with_auth(&app, Some("Bearer secret")).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"static");
let header = format!("Bearer {}", testing::valid_token());
let resp = get_with_auth(&app, Some(&header)).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"oauth user-1 true");
}
#[tokio::test]
async fn the_layer_behaves_exactly_like_the_middleware_function() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let layer = AuthLayer::builder()
.static_token(STATIC)
.oauth(validator(&jwks.url))
.on_reject(|cx: RejectContext<'_>| Response::new(Body::from(cx.status.to_string())))
.build()
.unwrap();
let handler = get(|c: Option<Extension<Credential>>| async move {
match c {
Some(Extension(Credential::StaticToken)) => "static",
Some(Extension(Credential::OAuth(_))) => "oauth",
None => "none",
}
});
let via_fn = Router::new()
.route("/test", handler.clone())
.route_layer(middleware::from_fn_with_state(layer.clone(), require_auth));
let via_route_layer = Router::new()
.route("/test", handler.clone())
.route_layer(layer.clone());
let via_layer = Router::new().route("/test", handler).layer(layer);
let valid = format!("Bearer {}", testing::valid_token());
let unscoped = format!("Bearer {}", unscoped_token());
for header in [
None,
Some("Bearer secret"),
Some("Bearer wrong"),
Some(valid.as_str()),
Some(unscoped.as_str()),
] {
let expected = get_with_auth(&via_fn, header).await;
let (status, challenge) = (
expected.status(),
expected.headers().get(WWW_AUTHENTICATE).cloned(),
);
let expected_body = body_bytes(expected).await;
for app in [&via_route_layer, &via_layer] {
let resp = get_with_auth(app, header).await;
assert_eq!(resp.status(), status, "{header:?}");
assert_eq!(
resp.headers().get(WWW_AUTHENTICATE),
challenge.as_ref(),
"{header:?}"
);
assert_eq!(body_bytes(resp).await, expected_body, "{header:?}");
}
}
}
#[tokio::test]
async fn from_decision_maps_every_decision_and_refuses_a_mismatch() {
use StaticTokenDecision::*;
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let valid = format!("Bearer {}", testing::valid_token());
let status = |layer: AuthLayer, header: &'static str| async move {
get_with_auth(&app(layer), Some(header)).await.status()
};
let valid: &'static str = Box::leak(valid.into_boxed_str());
let open = AuthLayer::from_decision(Unauthenticated, None).unwrap();
assert!(open.allows_unauthenticated());
let layer = AuthLayer::from_decision(StaticOnly(STATIC.into()), None).unwrap();
assert!(!layer.allows_unauthenticated());
assert_eq!(status(layer.clone(), "Bearer secret").await, StatusCode::OK);
assert_eq!(status(layer, valid).await, StatusCode::UNAUTHORIZED);
let layer =
AuthLayer::from_decision(StaticAndOAuth(STATIC.into()), Some(Arc::clone(&v))).unwrap();
assert_eq!(status(layer.clone(), "Bearer secret").await, StatusCode::OK);
assert_eq!(status(layer, valid).await, StatusCode::OK);
for decision in [OAuthOnly, StaticIgnored] {
let layer = AuthLayer::from_decision(decision, Some(Arc::clone(&v))).unwrap();
assert_eq!(
status(layer.clone(), "Bearer secret").await,
StatusCode::UNAUTHORIZED
);
assert_eq!(status(layer, valid).await, StatusCode::OK);
}
for decision in [StaticAndOAuth(STATIC.into()), OAuthOnly, StaticIgnored] {
assert_eq!(
AuthLayer::from_decision(decision, None).unwrap_err(),
AuthLayerError::DecisionNeedsOAuth
);
}
for decision in [StaticOnly(STATIC.into()), Unauthenticated] {
assert_eq!(
AuthLayer::from_decision(decision, Some(Arc::clone(&v))).unwrap_err(),
AuthLayerError::DecisionWithoutOAuth
);
}
}
#[tokio::test]
async fn build_with_decision_keeps_the_builders_sources_and_replaces_its_token() {
let layer = AuthLayer::builder()
.static_token("builder-token")
.sources([CredentialSource::Raw(HeaderName::from_static("x-api-key"))])
.build_with_decision(StaticTokenDecision::StaticOnly(STATIC.into()))
.unwrap();
let app = app(layer);
assert_eq!(
send(&app, &[("x-api-key", STATIC)]).await.status(),
StatusCode::OK
);
assert_eq!(
send(&app, &[("x-api-key", "builder-token")]).await.status(),
StatusCode::UNAUTHORIZED
);
assert_eq!(
send(&app, &[("authorization", "Bearer secret")])
.await
.status(),
StatusCode::UNAUTHORIZED
);
}
async fn get_path(app: &Router, path: &str) -> Response {
app.clone()
.oneshot(Request::builder().uri(path).body(Body::empty()).unwrap())
.await
.unwrap()
}
const MCP_METADATA_PATH: &str = "/.well-known/oauth-protected-resource/mcp";
#[tokio::test]
async fn metadata_routes_for_a_resource_with_a_path() {
let v = unreachable_validator();
assert_eq!(v.metadata_path(), MCP_METADATA_PATH);
let app: Router = metadata_router(Some(Arc::clone(&v)));
for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
let resp = get_path(&app, path).await;
assert_eq!(resp.status(), StatusCode::OK, "{path}");
assert_eq!(resp.headers()["content-type"], "application/json", "{path}");
let body = body_bytes(resp).await;
assert_eq!(body, serde_json::to_vec(&v.metadata()).unwrap(), "{path}");
let doc: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(doc["resource"], testing::RESOURCE);
assert_eq!(doc["authorization_servers"][0], testing::ISSUER);
assert_eq!(
doc["scopes_supported"],
serde_json::json!(["mcp:read", "mcp:write"])
);
assert_eq!(
doc["bearer_methods_supported"],
serde_json::json!(["header"])
);
}
for path in [
"/.well-known/oauth-protected-resource/other",
"/.well-known/oauth-protected-resource/mcp/deeper",
] {
assert_eq!(
get_path(&app, path).await.status(),
StatusCode::NOT_FOUND,
"{path}"
);
}
}
async fn post_path(app: &Router, path: &str) -> Response {
app.clone()
.oneshot(
Request::builder()
.method("POST")
.uri(path)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap()
}
#[tokio::test]
async fn metadata_routes_answer_other_methods_by_whether_the_path_is_served() {
let app: Router = metadata_router(Some(unreachable_validator()));
for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
assert_eq!(
post_path(&app, path).await.status(),
StatusCode::METHOD_NOT_ALLOWED,
"{path}"
);
}
for path in [
"/.well-known/oauth-protected-resource/other",
"/.well-known/oauth-protected-resource/mcp/deeper",
] {
let resp = post_path(&app, path).await;
assert_eq!(resp.status(), StatusCode::NOT_FOUND, "{path}");
assert!(!resp.headers().contains_key("allow"), "{path}");
}
let app: Router = metadata_router(None);
assert_eq!(
post_path(&app, MCP_METADATA_PATH).await.status(),
StatusCode::NOT_FOUND
);
}
#[tokio::test]
async fn metadata_routes_for_a_root_resource() {
let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
cfg.resource = "https://api.example.test/".to_string();
let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
assert_eq!(v.metadata_path(), PROTECTED_RESOURCE_METADATA_PREFIX);
let app: Router = metadata_router(Some(v));
let resp = get_path(&app, PROTECTED_RESOURCE_METADATA_PREFIX).await;
assert_eq!(resp.status(), StatusCode::OK);
let doc: serde_json::Value = serde_json::from_slice(&body_bytes(resp).await).unwrap();
assert_eq!(doc["resource"], "https://api.example.test/");
assert_eq!(
get_path(&app, MCP_METADATA_PATH).await.status(),
StatusCode::NOT_FOUND
);
}
#[tokio::test]
async fn metadata_routes_with_a_nested_and_a_brace_bearing_path() {
for (resource, path) in [
(
"https://api.example.test/v1/things",
"/.well-known/oauth-protected-resource/v1/things",
),
(
"https://api.example.test/v1/",
"/.well-known/oauth-protected-resource/v1/",
),
(
"https://api.example.test/a{b}",
"/.well-known/oauth-protected-resource/a{b}",
),
] {
let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
cfg.resource = resource.to_string();
let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
assert_eq!(v.metadata_path(), path);
let app: Router = metadata_router(Some(v));
let resp = get_path(&app, path).await;
assert_eq!(resp.status(), StatusCode::OK, "{resource}");
}
}
#[tokio::test]
async fn metadata_routes_for_a_path_axum_would_read_as_route_syntax() {
for (resource, path) in [
(
"https://api.example.test/a/:id",
"/.well-known/oauth-protected-resource/a/:id",
),
(
"https://api.example.test/a/*x",
"/.well-known/oauth-protected-resource/a/*x",
),
(
"https://api.example.test/:id",
"/.well-known/oauth-protected-resource/:id",
),
(
"https://api.example.test/*",
"/.well-known/oauth-protected-resource/*",
),
(
"https://api.example.test/{x}",
"/.well-known/oauth-protected-resource/{x}",
),
(
"https://api.example.test/%7Bx%7D",
"/.well-known/oauth-protected-resource/%7Bx%7D",
),
(
"https://api.example.test//mcp",
"/.well-known/oauth-protected-resource//mcp",
),
] {
let resolved = crate::OAuthConfig {
enabled: true,
issuer: testing::ISSUER.to_string(),
jwks_uri: Some("http://127.0.0.1:1/jwks".to_string()),
audience: testing::AUDIENCE.to_string(),
resource: resource.to_string(),
required_scope: Some("mcp:read".to_string()),
..crate::OAuthConfig::default()
}
.resolve(crate::KeyNaming::Dotted("oauth"))
.unwrap_or_else(|e| panic!("{resource}: {e}"))
.expect("enabled");
let v = Arc::new(OAuthValidator::new(&resolved).unwrap());
assert_eq!(v.metadata_path(), path, "{resource}");
let app: Router = metadata_router(Some(Arc::clone(&v)));
for served in [path, PROTECTED_RESOURCE_METADATA_PREFIX] {
let resp = get_path(&app, served).await;
assert_eq!(resp.status(), StatusCode::OK, "{resource} {served}");
assert_eq!(
body_bytes(resp).await,
serde_json::to_vec(&v.metadata()).unwrap(),
"{resource} {served}"
);
}
for other in [
"/.well-known/oauth-protected-resource/a/other",
"/.well-known/oauth-protected-resource/a/:id/x",
"/.well-known/oauth-protected-resource/other",
"/.well-known/oauth-protected-resource/mcp",
] {
assert_eq!(
get_path(&app, other).await.status(),
StatusCode::NOT_FOUND,
"{resource} {other}"
);
}
}
}
#[tokio::test]
async fn metadata_path_answers_methods_exactly_like_the_bare_prefix_route() {
let app: Router = metadata_router(Some(unreachable_validator()));
let send = |method: &'static str, path: &'static str| {
app.clone().oneshot(
Request::builder()
.method(method)
.uri(path)
.body(Body::empty())
.unwrap(),
)
};
for method in ["GET", "HEAD", "POST", "PUT", "DELETE", "OPTIONS", "PATCH"] {
let bare = send(method, PROTECTED_RESOURCE_METADATA_PREFIX)
.await
.unwrap();
let suffixed = send(method, MCP_METADATA_PATH).await.unwrap();
assert_eq!(bare.status(), suffixed.status(), "{method}");
assert_eq!(bare.headers(), suffixed.headers(), "{method}");
let (bare, suffixed) = (body_bytes(bare).await, body_bytes(suffixed).await);
assert_eq!(bare, suffixed, "{method}");
if method == "HEAD" {
assert!(suffixed.is_empty());
}
}
}
#[tokio::test]
async fn metadata_routes_404_when_oauth_is_not_configured() {
let app: Router = metadata_router(None).fallback(|| async { "spa shell" });
for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
let resp = get_path(&app, path).await;
assert_eq!(resp.status(), StatusCode::NOT_FOUND, "{path}");
assert!(body_bytes(resp).await.is_empty(), "{path}");
}
}
#[tokio::test]
async fn metadata_routes_are_reachable_outside_the_auth_layer() {
let v = unreachable_validator();
let layer = AuthLayer::builder().oauth(Arc::clone(&v)).build().unwrap();
let app = Router::new()
.route("/mcp", get(|| async { "ok" }))
.route_layer(middleware::from_fn_with_state(layer, require_auth))
.merge(metadata_router(Some(v)));
assert_eq!(
get_path(&app, "/mcp").await.status(),
StatusCode::UNAUTHORIZED
);
for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
assert_eq!(
get_path(&app, path).await.status(),
StatusCode::OK,
"{path}"
);
}
}
#[tokio::test]
async fn metadata_router_works_with_app_state() {
#[derive(Clone)]
struct AppState;
let app: Router = Router::new()
.route("/x", get(|State(_): State<AppState>| async { "x" }))
.merge(metadata_router(Some(unreachable_validator())))
.with_state(AppState);
assert_eq!(
get_path(&app, MCP_METADATA_PATH).await.status(),
StatusCode::OK
);
}
}