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::{FromRequestParts, OptionalFromRequestParts, Request, State};
use ::axum::middleware::Next;
use ::axum::response::{IntoResponse, Response};
use ::axum::routing::{any, get};
use http::header::WWW_AUTHENTICATE;
use http::request::Parts;
use http::{HeaderValue, Method, StatusCode};
use tracing::{error, info};
use zeroize::Zeroizing;
use crate::authenticate::{Credential, StaticTokenMatch, StaticTokens};
use crate::challenge::PROTECTED_RESOURCE_METADATA_PREFIX;
use crate::policy::StaticTokenDecision;
use crate::token::{AuthorizedToken, InvalidTokenKind, TokenRejection};
use crate::validator::OAuthValidator;
use crate::http_layer::{
Admission, Gate, GateRan, RefusalBody, log_layer_refusal, log_oauth_accepted,
log_passed_through, log_static_accepted,
};
#[doc(inline)]
pub use crate::http_layer::{
AuthLayerError, CredentialSource, InvalidScope, RejectContext, RequireScopes,
RequireScopesService,
};
use crate::observe::{
self, Mechanism, Outcome, REASON_MISCONFIGURED, REASON_NONE, Stage, count_request,
};
pub use crate::refusal::DEFAULT_STATIC_CHALLENGE;
#[cfg(test)]
use crate::http_layer::{bearer_credential, names_a_token};
#[cfg(test)]
use http::{HeaderMap, HeaderName};
pub type RejectFn = Arc<dyn Fn(RejectContext<'_>) -> Response + Send + Sync>;
#[derive(Clone)]
pub struct AuthLayer {
inner: Arc<Mode>,
}
enum Mode {
Enforce(Enforce),
AllowUnauthenticated,
}
struct Enforce {
gate: Arc<Gate>,
on_reject: Option<RejectFn>,
}
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_tokens", &e.gate.static_tokens)
.field("oauth", &e.gate.oauth)
.field("sources", &e.gate.sources)
.field("on_reject", &e.on_reject.as_ref().map(|_| "<fn>"))
.field("static_challenge", &e.gate.static_challenge)
.field("optional", &e.gate.optional)
.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.gate.oauth.as_ref(),
Mode::AllowUnauthenticated => None,
}
}
}
#[derive(Default)]
pub struct AuthLayerBuilder {
static_token: Option<Zeroizing<String>>,
static_tokens: Option<StaticTokens>,
oauth: Option<Arc<OAuthValidator>>,
sources: Option<Vec<CredentialSource>>,
on_reject: Option<RejectFn>,
static_challenge: Option<Option<HeaderValue>>,
optional: bool,
required_scopes: Vec<String>,
static_bypasses_scopes: bool,
}
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("static_tokens", &self.static_tokens)
.field("oauth", &self.oauth)
.field("sources", &self.sources)
.field("on_reject", &self.on_reject.as_ref().map(|_| "<fn>"))
.field("static_challenge", &self.static_challenge)
.field("optional", &self.optional)
.field("required_scopes", &self.required_scopes)
.field("static_bypasses_scopes", &self.static_bypasses_scopes)
.finish()
}
}
impl AuthLayerBuilder {
pub fn static_token(mut self, token: impl Into<String>) -> Self {
self.static_token = Some(Zeroizing::new(token.into()));
self
}
pub fn optional_static_token(mut self, token: Option<String>) -> Self {
self.static_token = token.map(Zeroizing::new);
self
}
pub fn static_tokens(mut self, tokens: StaticTokens) -> Self {
self.static_tokens = Some(tokens);
self
}
pub fn optional_static_tokens(mut self, tokens: Option<StaticTokens>) -> Self {
self.static_tokens = tokens;
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 optional(mut self) -> Self {
self.optional = true;
self
}
pub fn require_scopes(mut self, scopes: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.required_scopes = scopes.into_iter().map(Into::into).collect();
self
}
pub fn static_token_bypasses_scopes(mut self) -> Self {
self.static_bypasses_scopes = true;
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> {
Gate::check_decision(&decision, self.oauth.is_some())?;
let unauthenticated = decision == StaticTokenDecision::Unauthenticated;
let (token, tokens) = Gate::decision_tokens(decision, self.static_tokens.take())?;
if unauthenticated && !self.required_scopes.is_empty() {
return Err(AuthLayerError::ScopesWithoutAuthentication);
}
if unauthenticated {
return Ok(AuthLayer::allow_unauthenticated());
}
self.static_token = token;
self.static_tokens = tokens;
self.build()
}
pub fn build(self) -> Result<AuthLayer, AuthLayerError> {
let gate = Gate::build(
self.static_token,
self.static_tokens,
self.oauth,
self.sources,
self.static_challenge,
self.optional,
self.required_scopes,
self.static_bypasses_scopes,
)?;
Ok(AuthLayer {
inner: Arc::new(Mode::Enforce(Enforce {
gate: Arc::new(gate),
on_reject: self.on_reject,
})),
})
}
}
impl Enforce {
fn reject(&self, rejection: &TokenRejection, request: &Parts) -> Response {
self.reject_with(rejection, request, None)
}
fn reject_with(
&self,
rejection: &TokenRejection,
request: &Parts,
insufficient: Option<&HeaderValue>,
) -> Response {
let (status, _) = self.gate.status_and_challenge_with(rejection, insufficient);
let response = match &self.on_reject {
Some(f) => f(RejectContext {
rejection,
status,
request,
}),
None => Response::new(Body::empty()),
};
self.gate.finish_with(rejection, insufficient, response)
}
fn refuse_scoped(&self, request: &Parts, required: &[String]) -> Response {
let path = request.uri.path();
let rejection = TokenRejection::InsufficientScope;
let mechanism = Mechanism::of_request(request.extensions.get::<Credential>(), &rejection);
if self.gate.oauth.is_none() {
count_request(
Stage::Handler,
Outcome::Rejected,
mechanism,
REASON_MISCONFIGURED,
);
error!(
path = %path,
required = ?required,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = 403u16,
"Server misconfiguration: the handler requires scopes, but its AuthLayer has no \
OAuth validator, so no credential can carry them; refusing the request"
);
} else {
count_request(
Stage::Handler,
Outcome::Rejected,
mechanism,
observe::reason(&rejection),
);
let present = match request.extensions.get::<Credential>() {
Some(Credential::OAuth(token)) => token.scopes.clone(),
_ => Vec::new(),
};
info!(
path = %path,
required = ?required,
present = ?crate::token::scopes_for_log(&present),
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = observe::reason(&rejection),
auth.status = 403u16,
"The credential lacks the scopes this handler requires"
);
}
let insufficient = self.gate.scope_challenge(required);
self.reject_with(&rejection, request, insufficient.as_ref())
}
fn refuse(
&self,
rejection: &TokenRejection,
request: &Parts,
mechanism: Mechanism,
stage: Stage,
) -> Response {
log_layer_refusal!(&self.gate, request, rejection, mechanism, stage);
self.reject(rejection, request)
}
}
#[derive(Clone)]
struct LayerRan(AuthLayer);
fn axum_refusal_body(
source: &(dyn std::any::Any + Send + Sync),
cx: RejectContext<'_>,
) -> Option<Response> {
match source.downcast_ref::<Mode>()? {
Mode::Enforce(Enforce {
on_reject: Some(f), ..
}) => Some(f(cx)),
_ => None,
}
}
impl AuthLayer {
fn mark(&self, extensions: &mut http::Extensions) {
let gate = match &*self.inner {
Mode::Enforce(enforce) => Some(Arc::clone(&enforce.gate)),
Mode::AllowUnauthenticated => None,
};
if gate.is_none() && extensions.get::<GateRan>().is_some() {
return;
}
extensions.insert(GateRan(gate));
extensions.insert(RefusalBody::<Body> {
source: Arc::clone(&self.inner) as Arc<dyn std::any::Any + Send + Sync>,
build: axum_refusal_body,
});
}
}
impl AuthLayer {
async fn check(&self, mut request: Request) -> Result<Request, Response> {
let enforce = match &*self.inner {
Mode::AllowUnauthenticated => {
count_request(
Stage::Layer,
Outcome::PassedThrough,
Mechanism::None,
REASON_NONE,
);
if request.extensions().get::<LayerRan>().is_none() {
request.extensions_mut().insert(LayerRan(self.clone()));
}
crate::http_layer::mark_authorization_sensitive(request.headers_mut());
self.mark(request.extensions_mut());
return Ok(request);
}
Mode::Enforce(enforce) => enforce,
};
let (mut parts, body) = request.into_parts();
match enforce.gate.admit(&mut parts).await {
Admission::Static => log_static_accepted!(&parts),
Admission::OAuth(token) => log_oauth_accepted!(&parts, &token),
Admission::PassedThrough => log_passed_through!(&parts),
Admission::Refused(rejection, mechanism) => {
return Err(enforce.refuse(&rejection, &parts, mechanism, Stage::Layer));
}
}
parts.extensions.insert(LayerRan(self.clone()));
self.mark(&mut parts.extensions);
Ok(Request::from_parts(parts, body))
}
fn refuse_extraction(
&self,
rejection: &TokenRejection,
parts: &Parts,
wants: Wants,
) -> Response {
let mechanism = Mechanism::of_request(parts.extensions.get::<Credential>(), rejection);
let misconfigured = || {
count_request(
Stage::Handler,
Outcome::Rejected,
mechanism,
REASON_MISCONFIGURED,
)
};
match &*self.inner {
Mode::Enforce(enforce)
if wants == Wants::OAuthToken && enforce.gate.oauth.is_none() =>
{
misconfigured();
error!(
path = %parts.uri.path(),
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = observe::status(rejection),
"Server misconfiguration: the handler requires an OAuth access token, but \
its AuthLayer has no OAuth validator; refusing the request"
);
enforce.reject(rejection, parts)
}
Mode::Enforce(enforce)
if wants == Wants::StaticToken && enforce.gate.static_tokens.is_none() =>
{
misconfigured();
error!(
path = %parts.uri.path(),
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = observe::status(rejection),
"Server misconfiguration: the handler requires a static token, but its \
AuthLayer has no static token; refusing the request"
);
enforce.reject(rejection, parts)
}
Mode::Enforce(enforce) => enforce.refuse(rejection, parts, mechanism, Stage::Handler),
Mode::AllowUnauthenticated
if matches!(parts.extensions.get::<GateRan>(), Some(GateRan(Some(_)))) =>
{
crate::http_layer::scope_refusal::<Body>(parts, rejection, &[], wants.extractor())
}
Mode::AllowUnauthenticated => {
misconfigured();
error!(
path = %parts.uri.path(),
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = 401u16,
"Server misconfiguration: the handler requires a credential, but its \
AuthLayer allows unauthenticated requests; refusing the request"
);
(
StatusCode::UNAUTHORIZED,
[(
WWW_AUTHENTICATE,
HeaderValue::from_static(DEFAULT_STATIC_CHALLENGE),
)],
)
.into_response()
}
}
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum Wants {
AnyCredential,
OAuthToken,
StaticToken,
}
impl Wants {
fn extractor(self) -> &'static str {
match self {
Self::AnyCredential => "Credential",
Self::OAuthToken => "AuthorizedToken",
Self::StaticToken => "StaticTokenMatch",
}
}
}
enum Found<T> {
Present(T),
Absent(AuthLayer),
NoLayer,
}
fn find<T: Clone + Send + Sync + 'static>(parts: &Parts) -> Found<T> {
match (
parts.extensions.get::<T>(),
parts.extensions.get::<LayerRan>(),
) {
(Some(value), _) => Found::Present(value.clone()),
(None, Some(LayerRan(layer))) => Found::Absent(layer.clone()),
(None, None) => Found::NoLayer,
}
}
fn no_layer(parts: &Parts, extractor: &'static str) -> Response {
count_request(
Stage::Handler,
Outcome::Rejected,
Mechanism::None,
REASON_MISCONFIGURED,
);
error!(
path = %parts.uri.path(),
extractor,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = Mechanism::None.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = 500u16,
"Server misconfiguration: an authentication extractor ran on a route no AuthLayer \
covers; refusing the request"
);
StatusCode::INTERNAL_SERVER_ERROR.into_response()
}
fn refuse_absent(layer: &AuthLayer, parts: &Parts, wants: Wants) -> Response {
let rejection = match (parts.extensions.get::<Credential>(), wants) {
(Some(_), Wants::StaticToken) => TokenRejection::invalid(
InvalidTokenKind::StaticTokenRequired,
"a credential was accepted, but the handler requires a static token",
),
(Some(_), _) => TokenRejection::invalid(
InvalidTokenKind::OAuthTokenRequired,
"a credential was accepted, but the handler requires an OAuth access token",
),
(None, _) => TokenRejection::Missing,
};
layer.refuse_extraction(&rejection, parts, wants)
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: Send + Sync> FromRequestParts<S> for AuthorizedToken {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Response> {
match find::<AuthorizedToken>(parts) {
Found::Present(token) => Ok(token),
Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::OAuthToken)),
Found::NoLayer => Err(no_layer(parts, "AuthorizedToken")),
}
}
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: Send + Sync> OptionalFromRequestParts<S> for AuthorizedToken {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Option<Self>, Response> {
match find::<AuthorizedToken>(parts) {
Found::Present(token) => Ok(Some(token)),
Found::Absent(_) => Ok(None),
Found::NoLayer => Err(no_layer(parts, "Option<AuthorizedToken>")),
}
}
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: Send + Sync> FromRequestParts<S> for Credential {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Response> {
match find::<Credential>(parts) {
Found::Present(credential) => Ok(credential),
Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::AnyCredential)),
Found::NoLayer => Err(no_layer(parts, "Credential")),
}
}
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: Send + Sync> OptionalFromRequestParts<S> for Credential {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Option<Self>, Response> {
match find::<Credential>(parts) {
Found::Present(credential) => Ok(Some(credential)),
Found::Absent(_) => Ok(None),
Found::NoLayer => Err(no_layer(parts, "Option<Credential>")),
}
}
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: Send + Sync> FromRequestParts<S> for StaticTokenMatch {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Response> {
match find::<StaticTokenMatch>(parts) {
Found::Present(matched) => Ok(matched),
Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::StaticToken)),
Found::NoLayer => Err(no_layer(parts, "StaticTokenMatch")),
}
}
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: Send + Sync> OptionalFromRequestParts<S> for StaticTokenMatch {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Option<Self>, Response> {
match find::<StaticTokenMatch>(parts) {
Found::Present(matched) => Ok(Some(matched)),
Found::Absent(_) => Ok(None),
Found::NoLayer => Err(no_layer(parts, "Option<StaticTokenMatch>")),
}
}
}
pub trait ScopeSet: Send + Sync + 'static {
const SCOPES: &'static [&'static str];
}
pub struct Scoped<S: ScopeSet> {
token: AuthorizedToken,
_scopes: std::marker::PhantomData<fn() -> S>,
}
impl<S: ScopeSet> Scoped<S> {
pub fn token(&self) -> &AuthorizedToken {
&self.token
}
pub fn into_token(self) -> AuthorizedToken {
self.token
}
}
impl<S: ScopeSet> std::ops::Deref for Scoped<S> {
type Target = AuthorizedToken;
fn deref(&self) -> &AuthorizedToken {
&self.token
}
}
impl<S: ScopeSet> Clone for Scoped<S> {
fn clone(&self) -> Self {
Self {
token: self.token.clone(),
_scopes: std::marker::PhantomData,
}
}
}
impl<S: ScopeSet> std::fmt::Debug for Scoped<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Scoped")
.field("required", &S::SCOPES)
.field("token", &self.token)
.finish()
}
}
#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
impl<S: ScopeSet, St: Send + Sync> FromRequestParts<St> for Scoped<S> {
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, _state: &St) -> Result<Self, Response> {
let () = ValidScopeSet::<S>::CHECKED;
let Ok(required) = crate::http_layer::checked_scopes(S::SCOPES.iter().copied()) else {
count_request(
Stage::Handler,
Outcome::Rejected,
Mechanism::None,
REASON_MISCONFIGURED,
);
error!(
path = %parts.uri.path(),
extractor = std::any::type_name::<S>(),
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = Mechanism::None.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = 500u16,
"Server misconfiguration: a ScopeSet holds an entry that is not a valid scope, \
which no token can carry; refusing the request"
);
return Err(StatusCode::INTERNAL_SERVER_ERROR.into_response());
};
match find::<Credential>(parts) {
Found::Present(Credential::OAuth(token)) if token.require_scopes(S::SCOPES).is_ok() => {
Ok(Self {
token,
_scopes: std::marker::PhantomData,
})
}
Found::Present(_) => {
let gate = match parts.extensions.get::<GateRan>() {
Some(GateRan(Some(gate))) => Some(Arc::clone(gate)),
_ => None,
};
let layer = parts.extensions.get::<LayerRan>().map(|l| &*l.0.inner);
Err(match (gate, layer) {
(Some(gate), Some(Mode::Enforce(enforce)))
if Arc::ptr_eq(&gate, &enforce.gate) =>
{
enforce.refuse_scoped(parts, &required)
}
(None, Some(Mode::Enforce(enforce))) => enforce.refuse_scoped(parts, &required),
_ => crate::http_layer::scope_refusal::<Body>(
parts,
&TokenRejection::InsufficientScope,
&required,
"Scoped",
),
})
}
Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::OAuthToken)),
Found::NoLayer if parts.extensions.get::<GateRan>().is_some() => {
Err(crate::http_layer::scope_refusal::<Body>(
parts,
&TokenRejection::Missing,
&required,
"Scoped",
))
}
Found::NoLayer => Err(no_layer(parts, "Scoped")),
}
}
}
struct ValidScopeSet<S>(std::marker::PhantomData<S>);
impl<S: ScopeSet> ValidScopeSet<S> {
const CHECKED: () = assert!(
crate::http_layer::all_scope_tokens(S::SCOPES),
"every ScopeSet::SCOPES entry must be an RFC 6749 §3.3 scope-token: printable ASCII \
with no space, '\"' or '\\'"
);
}
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 scope_tests;
#[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
);
for blank in [" ", "\t", " \n "] {
assert_eq!(
AuthLayer::builder()
.static_token(blank)
.build()
.unwrap_err(),
AuthLayerError::NoCredential,
"{blank:?}"
);
assert_eq!(
AuthLayer::builder()
.static_token(blank)
.optional()
.build()
.unwrap_err(),
AuthLayerError::NoCredential,
"{blank:?}"
);
assert_eq!(
AuthLayer::builder()
.build_with_decision(StaticTokenDecision::StaticOnly(blank.into()))
.unwrap_err(),
AuthLayerError::NoCredential,
"{blank:?}"
);
}
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>>,
headers: http::HeaderMap| async move {
assert!(c.is_none() && t.is_none(), "a pass-through inserts nothing");
for value in headers.get_all(http::header::AUTHORIZATION) {
assert!(value.is_sensitive(), "Authorization must be sensitive");
}
"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
);
}
async fn observed(resp: Response) -> (StatusCode, Vec<(String, Vec<u8>)>, Vec<u8>) {
let status = resp.status();
let headers = resp
.headers()
.iter()
.map(|(k, v)| (k.to_string(), v.as_bytes().to_vec()))
.collect();
(status, headers, body_bytes(resp).await)
}
fn json_reject(cx: RejectContext<'_>) -> Response {
Response::new(Body::from(format!("refused {}", cx.status.as_u16())))
}
fn optional_extractor_app(
layer: AuthLayer,
runs: Arc<std::sync::atomic::AtomicUsize>,
) -> Router {
let handler = move |credential: Option<Credential>, token: Option<AuthorizedToken>| {
let runs = Arc::clone(&runs);
async move {
runs.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
match (credential, token) {
(None, None) => "none".to_string(),
(Some(Credential::StaticToken), None) => "static".to_string(),
(Some(Credential::OAuth(c)), Some(t)) => {
assert_eq!(c, t);
format!("oauth {}", t.subject.as_deref().unwrap_or_default())
}
other => panic!("inconsistent extraction: {other:?}"),
}
}
};
Router::new()
.route("/test", get(handler))
.route_layer(layer)
}
type Headers<'a> = Vec<(&'a str, &'a [u8])>;
async fn send_raw(app: &Router, headers: &[(&str, &[u8])]) -> Response {
let mut req = Request::builder().uri("/test");
for (name, value) in headers {
req = req.header(*name, HeaderValue::from_bytes(value).unwrap());
}
app.clone()
.oneshot(req.body(Body::empty()).unwrap())
.await
.unwrap()
}
#[tokio::test]
async fn an_oauth_extractor_refusing_a_static_token_names_its_kind() {
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<'_>| {
let label = match cx.rejection {
TokenRejection::Invalid(invalid) => invalid.kind().as_str(),
_ => "not invalid",
};
Response::new(Body::from(label))
})
.build()
.unwrap();
let app = Router::new()
.route("/test", get(|_: AuthorizedToken| async { "ok" }))
.route("/static", get(|_: StaticTokenMatch| async { "ok" }))
.route_layer(layer);
let resp = get_with_auth(&app, Some("Bearer secret")).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
assert_eq!(body_bytes(resp).await, b"oauth_token_required");
let resp = get_with_auth(&app, Some("Bearer wrong")).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(body_bytes(resp).await, b"not_jwt");
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/static")
.header(
"authorization",
format!("Bearer {}", testing::valid_token()),
)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
assert_eq!(body_bytes(resp).await, b"static_token_required");
}
#[tokio::test]
async fn the_extractors_read_a_valid_token_and_the_static_token() {
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))
.build()
.unwrap();
let app = Router::new()
.route(
"/credential",
get(|credential: Credential| async move {
match credential {
Credential::StaticToken => "static".to_string(),
Credential::OAuth(t) => format!("oauth {}", t.subject.unwrap_or_default()),
}
}),
)
.route(
"/token",
get(|token: AuthorizedToken| async move {
format!(
"{} {}",
token.subject.as_deref().unwrap_or_default(),
token.has_scope("mcp:read")
)
}),
)
.route_layer(layer);
let get_at = |path: &'static str, header: String| {
let app = app.clone();
async move {
app.oneshot(
Request::builder()
.uri(path)
.header("authorization", header)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap()
}
};
let valid = format!("Bearer {}", testing::valid_token());
let resp = get_at("/credential", valid.clone()).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"oauth user-1");
let resp = get_at("/token", valid).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"user-1 true");
let resp = get_at("/credential", "Bearer secret".into()).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"static");
let resp = get_at("/token", "Bearer secret".into()).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
}
#[tokio::test]
async fn an_extractor_outside_every_layer_fails_closed_with_500() {
let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let count = |runs: &Arc<std::sync::atomic::AtomicUsize>| {
runs.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
};
let (r1, r2, r3, r4) = (
Arc::clone(&runs),
Arc::clone(&runs),
Arc::clone(&runs),
Arc::clone(&runs),
);
let app = Router::new()
.route(
"/credential",
get(move |_: Credential| async move { count(&r1) }),
)
.route(
"/token",
get(move |_: AuthorizedToken| async move { count(&r2) }),
)
.route(
"/opt-credential",
get(move |_: Option<Credential>| async move { count(&r3) }),
)
.route(
"/opt-token",
get(move |_: Option<AuthorizedToken>| async move { count(&r4) }),
);
for path in ["/credential", "/token", "/opt-credential", "/opt-token"] {
for header in [None, Some("Bearer secret")] {
let mut req = Request::builder().uri(path);
if let Some(h) = header {
req = req.header("authorization", h);
}
let resp = app
.clone()
.oneshot(req.body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(
resp.status(),
StatusCode::INTERNAL_SERVER_ERROR,
"{path} {header:?}"
);
assert!(resp.headers().get(WWW_AUTHENTICATE).is_none());
assert!(body_bytes(resp).await.is_empty(), "{path}");
}
}
assert_eq!(runs.load(std::sync::atomic::Ordering::SeqCst), 0);
}
#[test]
fn optional_still_needs_a_credential_to_build() {
assert_eq!(
AuthLayer::builder().optional().build().unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
AuthLayer::builder()
.static_token("")
.optional()
.build()
.unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
AuthLayer::builder()
.static_token(STATIC)
.optional()
.sources([])
.build()
.unwrap_err(),
AuthLayerError::NoSources
);
let layer = AuthLayer::builder()
.static_token("hunter2")
.optional()
.build()
.unwrap();
assert!(!layer.allows_unauthenticated());
let rendered = format!("{layer:?}");
assert!(!rendered.contains("hunter2") && rendered.contains("optional: true"));
}
#[tokio::test]
async fn an_optional_layer_passes_only_a_request_with_no_credential() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let builder = || {
AuthLayer::builder()
.static_token(STATIC)
.oauth(Arc::clone(&v))
.sources([
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
])
.on_reject(json_reject)
};
let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let optional =
optional_extractor_app(builder().optional().build().unwrap(), Arc::clone(&runs));
let strict =
optional_extractor_app(builder().build().unwrap(), Arc::new(Default::default()));
let ran = || runs.load(std::sync::atomic::Ordering::SeqCst);
let blanks: &[&[(&str, &[u8])]] = &[
&[],
&[("authorization", b"")],
&[("authorization", b"Bearer ")],
&[("authorization", b"bearer ")],
&[("authorization", b"Basic c2VjcmV0")],
&[("x-api-key", b" ")],
&[("authorization", b"Bearer "), ("x-api-key", b"")],
&[("authorization", b"Bearer "), ("authorization", b" ")],
];
for headers in blanks {
let before = ran();
let resp = send_raw(&optional, headers).await;
assert_eq!(resp.status(), StatusCode::OK, "{headers:?}");
assert!(resp.headers().get(WWW_AUTHENTICATE).is_none());
assert_eq!(body_bytes(resp).await, b"none", "{headers:?}");
assert_eq!(ran(), before + 1);
let resp = send_raw(&strict, headers).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED, "{headers:?}");
assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
}
let invalid = format!("Bearer {}", testing::valid_token().replace('.', "x."));
let expired = format!("Bearer {}", expired_token());
let unscoped = format!("Bearer {}", unscoped_token());
let refused: Vec<(Headers<'_>, StatusCode, String)> = vec![
(
vec![("authorization", b"Bearer not-a-jwt")],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![("authorization", invalid.as_bytes())],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![("authorization", expired.as_bytes())],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![("x-api-key", b"wrong-key")],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![("authorization", b"Bearer "), ("x-api-key", b"wrong-key")],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![
("authorization", b"Bearer "),
("authorization", b"Bearer junk"),
],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![("authorization", b"Bearer \xff")],
StatusCode::UNAUTHORIZED,
v.invalid_token_challenge(),
),
(
vec![("authorization", unscoped.as_bytes())],
StatusCode::FORBIDDEN,
v.insufficient_scope_challenge(),
),
];
for (headers, status, challenge) in &refused {
let before = ran();
let resp = send_raw(&optional, headers).await;
assert_eq!(resp.status(), *status, "{headers:?}");
assert_eq!(&www_authenticate(&resp), challenge, "{headers:?}");
let got = observed(resp).await;
assert_eq!(got.2, format!("refused {}", status.as_u16()).as_bytes());
assert_eq!(
got,
observed(send_raw(&strict, headers).await).await,
"{headers:?}"
);
assert_eq!(ran(), before, "the handler ran for {headers:?}");
}
let valid = format!("Bearer {}", testing::valid_token());
let resp = send_raw(&optional, &[("authorization", valid.as_bytes())]).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"oauth user-1");
let resp = send_raw(&optional, &[("x-api-key", STATIC.as_bytes())]).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"static");
}
#[tokio::test]
async fn a_required_extractor_behind_an_optional_layer_gets_the_layers_own_refusal() {
let v = unreachable_validator();
let builder = || {
AuthLayer::builder()
.static_token(STATIC)
.oauth(Arc::clone(&v))
.on_reject(json_reject)
};
let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let make = |layer: AuthLayer| {
let (r1, r2) = (Arc::clone(&runs), Arc::clone(&runs));
Router::new()
.route(
"/test",
get(move |_: Credential| async move {
r1.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}),
)
.route(
"/token",
get(move |_: AuthorizedToken| async move {
r2.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}),
)
.route_layer(layer)
};
let optional = make(builder().optional().build().unwrap());
let strict = make(builder().build().unwrap());
for path in ["/test", "/token"] {
let request = || Request::builder().uri(path).body(Body::empty()).unwrap();
let got = observed(optional.clone().oneshot(request()).await.unwrap()).await;
let want = observed(strict.clone().oneshot(request()).await.unwrap()).await;
assert_eq!(got.0, StatusCode::UNAUTHORIZED, "{path}");
assert_eq!(got, want, "{path}");
assert!(
got.1.iter().any(|(k, val)| k == "www-authenticate"
&& val == v.invalid_token_challenge().as_bytes()),
"{path}"
);
}
assert_eq!(runs.load(std::sync::atomic::Ordering::SeqCst), 0);
let optional = make(
AuthLayer::builder()
.static_token(STATIC)
.optional()
.build()
.unwrap(),
);
let resp = optional
.oneshot(Request::builder().uri("/test").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
}
#[tokio::test]
async fn the_extractors_under_allow_unauthenticated() {
let app = Router::new()
.route(
"/test",
get(
|c: Option<Credential>, t: Option<AuthorizedToken>| async move {
assert!(c.is_none() && t.is_none());
"none"
},
),
)
.route("/required", get(|_: Credential| async { "unreachable" }))
.route_layer(AuthLayer::allow_unauthenticated());
let resp = get_with_auth(&app, Some("Bearer anything")).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"none");
let resp = app
.oneshot(
Request::builder()
.uri("/required")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
}
#[tokio::test]
async fn a_non_optional_layer_with_extractors_matches_extension_handlers() {
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 via_extension = Router::new()
.route(
"/test",
get(|Extension(c): Extension<Credential>| async move { format!("{c:?}") }),
)
.route_layer(layer.clone());
let via_extractor = Router::new()
.route(
"/test",
get(|c: Credential| async move { format!("{c:?}") }),
)
.route_layer(layer);
let valid = format!("Bearer {}", testing::valid_token());
let unscoped = format!("Bearer {}", unscoped_token());
let expired = format!("Bearer {}", expired_token());
for header in [
None,
Some("Bearer "),
Some("Bearer secret"),
Some("Bearer wrong"),
Some(valid.as_str()),
Some(unscoped.as_str()),
Some(expired.as_str()),
] {
assert_eq!(
observed(get_with_auth(&via_extractor, header).await).await,
observed(get_with_auth(&via_extension, header).await).await,
"{header:?}"
);
}
}
fn claims_with(extra: serde_json::Value) -> serde_json::Value {
let mut claims = serde_json::json!({
"iss": testing::ISSUER, "aud": testing::AUDIENCE,
"exp": testing::now() + 3600, "scope": "mcp:read mcp:write", "sub": "user-1",
});
for (k, v) in extra.as_object().unwrap() {
claims[k] = v.clone();
}
claims
}
#[test]
fn names_a_token_only_for_dpop_and_tab_separated_bearer() {
for value in [
"DPoP x",
"dpop x",
"DPoP",
"Bearer\tx",
"bearer\t x",
" Bearer x",
"\tBEARER\tx",
] {
assert!(names_a_token(value), "{value:?}");
}
for value in [
"",
"Bearer",
"Bearer ",
"Bearer\t",
"Bearer \t ",
"Basic x",
"x",
] {
assert!(!names_a_token(value), "{value:?}");
}
assert_eq!(bearer_credential("Bearer\tx"), "");
assert_eq!(bearer_credential("DPoP x"), "");
}
#[tokio::test]
async fn an_optional_layer_refuses_every_presented_shape_like_the_strict_layer() {
use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let builder = || {
AuthLayer::builder()
.static_token(STATIC)
.oauth(Arc::clone(&v))
.sources([
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
])
.on_reject(json_reject)
};
let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let optional =
optional_extractor_app(builder().optional().build().unwrap(), Arc::clone(&runs));
let strict =
optional_extractor_app(builder().build().unwrap(), Arc::new(Default::default()));
let forged = testing::mint(
testing::KEY_B_PEM,
testing::KID_A,
&claims_with(serde_json::json!({})),
);
let wrong_aud = testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&claims_with(serde_json::json!({ "aud": "some-other-client" })),
);
let cnf = testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&claims_with(serde_json::json!({ "cnf": { "jkt": "abc" } })),
);
let valid = testing::valid_token();
let crit = {
let mut parts: Vec<String> = valid.split('.').map(str::to_string).collect();
parts[0] =
URL_SAFE_NO_PAD.encode(br#"{"alg":"RS256","kid":"test-key-a","crit":["exp"]}"#);
parts.join(".")
};
let cases: Vec<(&str, Vec<(&str, String)>)> = vec![
(
"forged",
vec![("authorization", format!("Bearer {forged}"))],
),
(
"wrong aud",
vec![("authorization", format!("Bearer {wrong_aud}"))],
),
(
"cnf bearer",
vec![("authorization", format!("Bearer {cnf}"))],
),
("crit", vec![("authorization", format!("Bearer {crit}"))]),
(
"BEARER forged",
vec![("authorization", format!("BEARER {forged}"))],
),
(
"bearer bad",
vec![("authorization", "bearer not-a-jwt".to_string())],
),
(
"Bearer<TAB>forged",
vec![("authorization", format!("Bearer\t{forged}"))],
),
(
"Bearer<TAB>valid",
vec![("authorization", format!("Bearer\t{valid}"))],
),
(
" Bearer forged",
vec![("authorization", format!(" Bearer {forged}"))],
),
("DPoP cnf", vec![("authorization", format!("DPoP {cnf}"))]),
(
"DPoP forged",
vec![("authorization", format!("DPoP {forged}"))],
),
(
"Basic, then Bearer forged",
vec![
("authorization", "Basic x".to_string()),
("authorization", format!("Bearer {forged}")),
],
),
(
"blank Bearer, then Bearer forged",
vec![
("authorization", "Bearer ".to_string()),
("authorization", format!("Bearer {forged}")),
],
),
];
for (name, headers) in &cases {
let headers: Vec<(&str, &[u8])> =
headers.iter().map(|(n, v)| (*n, v.as_bytes())).collect();
let before = runs.load(std::sync::atomic::Ordering::SeqCst);
let got = observed(send_raw(&optional, &headers).await).await;
let want = observed(send_raw(&strict, &headers).await).await;
assert_eq!(got.0, StatusCode::UNAUTHORIZED, "{name}");
assert!(
got.1.iter().any(|(k, val)| k == "www-authenticate"
&& val == v.invalid_token_challenge().as_bytes()),
"{name}"
);
assert_eq!(got, want, "{name}");
assert_eq!(
runs.load(std::sync::atomic::Ordering::SeqCst),
before,
"the handler ran for {name}"
);
}
}
#[tokio::test]
async fn an_inner_optional_layer_does_not_leak_an_outer_layers_credential() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let outer = AuthLayer::builder()
.oauth(validator(&jwks.url))
.build()
.unwrap();
let inner = AuthLayer::builder()
.static_token("inner-key")
.sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
.optional()
.build()
.unwrap();
let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let app = optional_extractor_app(inner, Arc::clone(&runs)).layer(outer);
let valid = format!("Bearer {}", testing::valid_token());
let resp = send_raw(&app, &[("authorization", valid.as_bytes())]).await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"none");
let resp = send_raw(
&app,
&[
("authorization", valid.as_bytes()),
("x-inner", b"inner-key"),
],
)
.await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"static");
assert_eq!(runs.load(std::sync::atomic::Ordering::SeqCst), 2);
}
#[tokio::test]
async fn strict_nested_layers_accumulate_extensions_as_documented() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let outer = AuthLayer::builder()
.oauth(validator(&jwks.url))
.build()
.unwrap();
let inner = AuthLayer::builder()
.static_token("inner-key")
.sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
.build()
.unwrap();
let app = Router::new()
.route(
"/test",
get(|c: Credential, t: AuthorizedToken| async move {
format!(
"{} {}",
matches!(c, Credential::StaticToken),
t.subject.unwrap_or_default()
)
}),
)
.route_layer(inner)
.layer(outer);
let valid = format!("Bearer {}", testing::valid_token());
let resp = send_raw(
&app,
&[
("authorization", valid.as_bytes()),
("x-inner", b"inner-key"),
],
)
.await;
assert_eq!(resp.status(), StatusCode::OK);
assert_eq!(body_bytes(resp).await, b"true user-1");
let resp = send_raw(&app, &[("authorization", valid.as_bytes())]).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
}
}
#[cfg(test)]
mod shared_refusal_tests {
use ::tower::{ServiceExt, service_fn};
use super::*;
use crate::http_layer::HttpAuthLayer;
use crate::testing;
use crate::{Refusal, refusal, refusal_with_static_challenge};
const STATIC: &str = "secret";
const CUSTOM: &str = "ApiKey realm=\"example\"";
fn validator(jwks_uri: &str) -> Arc<OAuthValidator> {
Arc::new(OAuthValidator::new(&testing::resolved_config(jwks_uri)).unwrap())
}
#[derive(Clone, Copy, Debug)]
enum Static {
Unset,
Custom,
Off,
}
fn what_the_axum_layer_sends(
oauth: Option<Arc<OAuthValidator>>,
setting: Static,
rejection: &TokenRejection,
) -> (u16, Vec<String>) {
let mut builder = AuthLayer::builder()
.static_token(STATIC)
.optional_oauth(oauth)
.on_reject(|_| {
(StatusCode::IM_A_TEAPOT, [(WWW_AUTHENTICATE, "Callback x")]).into_response()
});
builder = match setting {
Static::Unset => builder,
Static::Custom => builder.static_challenge(Some(HeaderValue::from_static(CUSTOM))),
Static::Off => builder.static_challenge(None),
};
let layer = builder.build().unwrap();
let Mode::Enforce(enforce) = &*layer.inner else {
unreachable!("an enforcing layer was built")
};
let (parts, ()) = Request::builder()
.uri("/test")
.body(())
.unwrap()
.into_parts();
let response = enforce.reject(rejection, &parts);
let challenges = response
.headers()
.get_all(WWW_AUTHENTICATE)
.iter()
.map(|v| v.to_str().unwrap().to_string())
.collect();
(response.status().as_u16(), challenges)
}
#[test]
fn refusal_gives_exactly_what_enforce_reject_gives() {
let v = validator("http://127.0.0.1:1/jwks");
let rejections = [
TokenRejection::Missing,
TokenRejection::Invalid("any reason".into()),
TokenRejection::InsufficientScope,
];
let mut rows = 0;
for oauth in [None, Some(Arc::clone(&v))] {
for setting in [Static::Unset, Static::Custom, Static::Off] {
for rejection in &rejections {
let static_str = match setting {
Static::Unset => Some(DEFAULT_STATIC_CHALLENGE),
Static::Custom => Some(CUSTOM),
Static::Off => None,
};
let ours =
refusal_with_static_challenge(rejection, oauth.as_deref(), static_str);
if let Static::Unset = setting {
assert_eq!(ours, refusal(rejection, oauth.as_deref()));
}
let (status, challenges) =
what_the_axum_layer_sends(oauth.clone(), setting, rejection);
let context = format!("oauth={} {setting:?} {rejection:?}", oauth.is_some());
assert_eq!(ours.status, status, "{context}");
let expected = match &ours {
Refusal {
www_authenticate: Some(c),
..
} => vec![c.clone()],
_ => vec!["Callback x".to_string()],
};
assert_eq!(challenges, expected, "{context}");
rows += 1;
}
}
}
assert_eq!(rows, 18);
}
#[test]
fn the_status_and_challenge_are_what_rfc_6750_asks_for() {
let v = validator("http://127.0.0.1:1/jwks");
let r = refusal(&TokenRejection::Missing, Some(&v));
assert_eq!(r.status, 401);
assert_eq!(r.www_authenticate, Some(v.invalid_token_challenge()));
let r = refusal(&TokenRejection::Invalid("x".into()), Some(&v));
assert_eq!(r.status, 401);
assert_eq!(r.www_authenticate, Some(v.invalid_token_challenge()));
let r = refusal(&TokenRejection::InsufficientScope, Some(&v));
assert_eq!(r.status, 403);
assert_eq!(r.www_authenticate, Some(v.insufficient_scope_challenge()));
assert_eq!(
refusal_with_static_challenge(&TokenRejection::Missing, Some(&v), None),
refusal(&TokenRejection::Missing, Some(&v))
);
}
type Seen = (u16, Vec<String>, String);
async fn through_axum(layer: AuthLayer, headers: &[(&str, &str)]) -> Seen {
let app: Router = Router::new()
.route(
"/test",
get(|credential: Option<Credential>| async move { format!("{credential:?}") }),
)
.route_layer(layer);
let mut request = Request::builder().uri("/test");
for (name, value) in headers {
request = request.header(*name, *value);
}
let response = app
.oneshot(request.body(Body::empty()).unwrap())
.await
.unwrap();
let status = response.status().as_u16();
let challenges = response
.headers()
.get_all(WWW_AUTHENTICATE)
.iter()
.map(|v| v.to_str().unwrap().to_string())
.collect();
let body = ::axum::body::to_bytes(response.into_body(), 64 * 1024)
.await
.unwrap();
(
status,
challenges,
String::from_utf8(body.to_vec()).unwrap(),
)
}
async fn through_tower(layer: HttpAuthLayer, headers: &[(&str, &str)]) -> Seen {
let service = tower_layer::Layer::layer(
&layer,
service_fn(|request: http::Request<String>| async move {
let credential = request.extensions().get::<Credential>().cloned();
Ok::<_, std::convert::Infallible>(http::Response::new(format!("{credential:?}")))
}),
);
let mut request = http::Request::builder().uri("/test");
for (name, value) in headers {
request = request.header(*name, *value);
}
let response = service
.oneshot(request.body(String::new()).unwrap())
.await
.unwrap();
let challenges = response
.headers()
.get_all(WWW_AUTHENTICATE)
.iter()
.map(|v| v.to_str().unwrap().to_string())
.collect();
(response.status().as_u16(), challenges, response.into_body())
}
#[tokio::test]
async fn the_axum_and_tower_layers_answer_every_request_identically() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let mint = |scope: &str, exp_offset: i64| {
testing::mint(
testing::KEY_A_PEM,
testing::KID_A,
&serde_json::json!({
"iss": testing::ISSUER, "aud": testing::AUDIENCE, "sub": "user-1",
"exp": testing::now() as i64 + exp_offset, "scope": scope,
}),
)
};
let valid = format!("Bearer {}", mint("mcp:read", 3600));
let expired = format!("Bearer {}", mint("mcp:read", -3600));
let unscoped = format!("Bearer {}", mint("openid", 3600));
let requests: Vec<Vec<(&str, &str)>> = vec![
vec![],
vec![("authorization", valid.as_str())],
vec![("authorization", expired.as_str())],
vec![("authorization", unscoped.as_str())],
vec![("authorization", "Bearer not-a-jwt")],
vec![("authorization", "Bearer secret")],
vec![("authorization", "bearer secret")],
vec![("authorization", "Bearer ")],
vec![("authorization", "Basic abc")],
vec![("authorization", "DPoP abc")],
vec![("x-api-key", "secret")],
vec![("authorization", "Bearer wrong"), ("x-api-key", "secret")],
];
type Config = (
Option<&'static str>,
bool,
Option<Option<&'static str>>,
bool,
bool,
);
let configs: [Config; 7] = [
(None, true, None, false, false),
(Some(STATIC), true, None, false, true),
(Some(STATIC), false, None, false, false),
(Some(STATIC), false, Some(None), false, false),
(Some(STATIC), false, Some(Some(CUSTOM)), false, true),
(Some(STATIC), true, None, true, false),
(Some(STATIC), false, None, true, true),
];
for (static_token, with_oauth, static_challenge, optional, api_key) in configs {
let oauth = with_oauth.then(|| Arc::clone(&v));
let sources = if api_key {
vec![
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
]
} else {
vec![CredentialSource::authorization_bearer()]
};
let challenge = static_challenge.map(|c| c.map(HeaderValue::from_static));
let mut axum_builder = AuthLayer::builder()
.optional_static_token(static_token.map(str::to_string))
.optional_oauth(oauth.clone())
.sources(sources.clone());
let mut tower_builder = HttpAuthLayer::builder()
.optional_static_token(static_token.map(str::to_string))
.optional_oauth(oauth.clone())
.sources(sources);
if let Some(c) = challenge {
axum_builder = axum_builder.static_challenge(c.clone());
tower_builder = tower_builder.static_challenge(c);
}
if optional {
axum_builder = axum_builder.optional();
tower_builder = tower_builder.optional();
}
let axum_layer = axum_builder.build().unwrap();
let tower_layer = tower_builder.build().unwrap();
for headers in &requests {
let a = through_axum(axum_layer.clone(), headers).await;
let t = through_tower(tower_layer.clone(), headers).await;
assert_eq!(
a, t,
"config {static_token:?} oauth={with_oauth} {static_challenge:?} \
optional={optional} api_key={api_key}, request {headers:?}"
);
}
}
}
type SeenFull = (u16, Vec<String>, Option<String>, String);
fn seen_parts(headers: &HeaderMap, status: u16, body: String) -> SeenFull {
let challenges = headers
.get_all(WWW_AUTHENTICATE)
.iter()
.map(|v| v.to_str().unwrap().to_string())
.collect();
let content_type = headers
.get(http::header::CONTENT_TYPE)
.map(|v| v.to_str().unwrap().to_string());
(status, challenges, content_type, body)
}
fn describe(extensions: &http::Extensions) -> String {
format!(
"{:?} token={} {:?}",
extensions.get::<Credential>(),
extensions.get::<AuthorizedToken>().is_some(),
extensions.get::<StaticTokenMatch>()
)
}
type RawHeaders = Vec<(&'static str, HeaderValue)>;
async fn axum_full(app: Router, headers: &RawHeaders) -> SeenFull {
let mut request = Request::builder().uri("/test");
for (name, value) in headers {
request = request.header(*name, value.clone());
}
let response = app
.oneshot(request.body(Body::empty()).unwrap())
.await
.unwrap();
let (parts, body) = response.into_parts();
let body = ::axum::body::to_bytes(body, 64 * 1024).await.unwrap();
seen_parts(
&parts.headers,
parts.status.as_u16(),
String::from_utf8(body.to_vec()).unwrap(),
)
}
async fn tower_full<S>(service: S, headers: &RawHeaders) -> SeenFull
where
S: tower_service::Service<
http::Request<String>,
Response = http::Response<String>,
Error = std::convert::Infallible,
>,
{
let mut request = http::Request::builder().uri("/test");
for (name, value) in headers {
request = request.header(*name, value.clone());
}
let response = service
.oneshot(request.body(String::new()).unwrap())
.await
.unwrap();
let (parts, body) = response.into_parts();
seen_parts(&parts.headers, parts.status.as_u16(), body)
}
fn axum_app(layer: AuthLayer) -> Router {
Router::new()
.route(
"/test",
get(|request: Request| async move { describe(request.extensions()) }),
)
.route_layer(layer)
}
fn tower_handler(
request: http::Request<String>,
) -> std::future::Ready<Result<http::Response<String>, std::convert::Infallible>> {
let mut response = http::Response::new(describe(request.extensions()));
response.headers_mut().insert(
http::header::CONTENT_TYPE,
HeaderValue::from_static("text/plain; charset=utf-8"),
);
std::future::ready(Ok(response))
}
#[tokio::test]
async fn the_layers_agree_on_callbacks_repeated_and_unreadable_headers() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let valid = HeaderValue::from_str(&format!("Bearer {}", testing::valid_token())).unwrap();
let requests: Vec<RawHeaders> = vec![
vec![],
vec![("authorization", valid.clone())],
vec![("authorization", HeaderValue::from_static("Bearer wrong"))],
vec![
("authorization", valid.clone()),
("authorization", HeaderValue::from_static("Bearer wrong")),
],
vec![
("authorization", HeaderValue::from_static("Bearer ")),
("authorization", HeaderValue::from_static("Bearer secret")),
],
vec![
("x-api-key", HeaderValue::from_static("wrong")),
("x-api-key", HeaderValue::from_static("secret")),
],
vec![(
"authorization",
HeaderValue::from_bytes(b"Bearer s\xe9cret").unwrap(),
)],
vec![("x-api-key", HeaderValue::from_bytes(b"\xff").unwrap())],
vec![("x-api-key", HeaderValue::from_static("key-next"))],
];
let next = || {
crate::StaticTokens::new()
.with(Some("next"), "key-next")
.unwrap()
};
for with_oauth in [false, true] {
for optional in [false, true] {
for static_challenge in [None, Some(None)] {
let sources = [
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
];
let oauth = with_oauth.then(|| Arc::clone(&v));
let mut a = AuthLayer::builder()
.static_token(STATIC)
.static_tokens(next())
.optional_oauth(oauth.clone())
.sources(sources.clone())
.on_reject(|cx| {
(
StatusCode::IM_A_TEAPOT,
[
(WWW_AUTHENTICATE, "Callback x"),
(http::header::CONTENT_TYPE, "application/json"),
],
format!("{{\"status\":{}}}", cx.status.as_u16()),
)
.into_response()
});
let mut t = HttpAuthLayer::builder()
.static_token(STATIC)
.static_tokens(next())
.optional_oauth(oauth)
.sources(sources)
.on_reject(|cx: RejectContext<'_>| {
http::Response::builder()
.status(StatusCode::IM_A_TEAPOT)
.header(WWW_AUTHENTICATE, "Callback x")
.header(http::header::CONTENT_TYPE, "application/json")
.body(format!("{{\"status\":{}}}", cx.status.as_u16()))
.unwrap()
});
if let Some(c) = &static_challenge {
a = a.static_challenge(c.clone());
t = t.static_challenge(c.clone());
}
if optional {
a = a.optional();
t = t.optional();
}
let (a, t) = (a.build().unwrap(), t.build().unwrap());
for headers in &requests {
let service = tower_layer::Layer::layer(&t, service_fn(tower_handler));
assert_eq!(
axum_full(axum_app(a.clone()), headers).await,
tower_full(service, headers).await,
"oauth={with_oauth} optional={optional} \
static_challenge={static_challenge:?} {headers:?}"
);
}
}
}
}
}
#[tokio::test]
async fn the_layers_agree_that_optional_clears_an_outer_layers_credential() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let bearer = HeaderValue::from_str(&format!("Bearer {}", testing::valid_token())).unwrap();
let outer_sources = [CredentialSource::Bearer(HeaderName::from_static("x-outer"))];
let requests: Vec<RawHeaders> = vec![
vec![("x-outer", bearer.clone())],
vec![
("x-outer", bearer.clone()),
("authorization", HeaderValue::from_static("Bearer secret")),
],
vec![
("x-outer", bearer.clone()),
("authorization", HeaderValue::from_static("Bearer wrong")),
],
];
let axum_outer = AuthLayer::builder()
.oauth(Arc::clone(&v))
.sources(outer_sources.clone())
.build()
.unwrap();
let axum_inner = AuthLayer::builder()
.static_token(STATIC)
.optional()
.build()
.unwrap();
let tower_outer = HttpAuthLayer::builder()
.oauth(Arc::clone(&v))
.sources(outer_sources)
.build()
.unwrap();
let tower_inner = HttpAuthLayer::builder()
.static_token(STATIC)
.optional()
.build()
.unwrap();
let mut outcomes = Vec::new();
for headers in &requests {
let app = axum_app(axum_inner.clone()).layer(axum_outer.clone());
let service = ::tower::ServiceBuilder::new()
.layer(tower_outer.clone())
.layer(tower_inner.clone())
.service(service_fn(tower_handler));
let a = axum_full(app, headers).await;
assert_eq!(a, tower_full(service, headers).await, "{headers:?}");
outcomes.push(a);
}
assert_eq!(outcomes[0].3, "None token=false None");
assert_eq!(
outcomes[1].3,
"Some(StaticToken) token=false Some(StaticTokenMatch { label: None })"
);
assert_eq!(outcomes[2].0, 401);
}
}
#[cfg(test)]
mod static_tokens_tests {
use ::tower::ServiceExt;
use super::*;
use crate::testing;
use http::HeaderName;
const STATIC: &str = "secret";
fn rotation() -> StaticTokens {
StaticTokens::new()
.with(Some("current"), "key-current")
.and_then(|t| t.with(Some("next"), "key-next"))
.unwrap()
}
fn validator(jwks_uri: &str) -> Arc<OAuthValidator> {
Arc::new(OAuthValidator::new(&testing::resolved_config(jwks_uri)).unwrap())
}
async fn report(credential: Option<Credential>, matched: Option<StaticTokenMatch>) -> String {
format!("{credential:?} {matched:?}")
}
fn app(layer: AuthLayer) -> Router {
Router::new()
.route("/test", get(report))
.route(
"/required",
get(|m: StaticTokenMatch| async move { format!("{:?}", m.label()) }),
)
.route_layer(layer)
}
async fn send(app: &Router, path: &str, headers: &[(&str, &str)]) -> Response {
let mut request = Request::builder().uri(path);
for (name, value) in headers {
request = request.header(*name, *value);
}
app.clone()
.oneshot(request.body(Body::empty()).unwrap())
.await
.unwrap()
}
async fn observed(resp: Response) -> (u16, Vec<(String, Vec<u8>)>, String) {
let status = resp.status().as_u16();
let headers = resp
.headers()
.iter()
.map(|(k, v)| (k.to_string(), v.as_bytes().to_vec()))
.collect();
let body = ::axum::body::to_bytes(resp.into_body(), 64 * 1024)
.await
.unwrap();
(status, headers, String::from_utf8(body.to_vec()).unwrap())
}
fn json_reject(cx: RejectContext<'_>) -> Response {
(
cx.status,
Json(serde_json::json!({ "status": cx.status.as_u16() })),
)
.into_response()
}
#[tokio::test]
async fn a_one_entry_set_answers_exactly_like_static_token() {
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 requests: Vec<Vec<(&str, &str)>> = vec![
vec![],
vec![("authorization", "Bearer secret")],
vec![("authorization", "bearer secret")],
vec![("authorization", "Bearer wrong")],
vec![("authorization", "Bearer ")],
vec![("authorization", "DPoP secret")],
vec![("x-api-key", "secret")],
vec![("authorization", "Bearer wrong"), ("x-api-key", "secret")],
vec![("authorization", valid.as_str())],
];
let configs = [
(false, false, false, false),
(false, false, true, false),
(true, false, false, true),
(false, true, false, true),
(true, true, false, false),
];
for (with_oauth, optional, no_challenge, callback) in configs {
let build = |use_set: bool| {
let mut b = AuthLayer::builder()
.optional_oauth(with_oauth.then(|| Arc::clone(&v)))
.sources([
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
]);
b = if use_set {
b.static_tokens(StaticTokens::single(STATIC).unwrap())
} else {
b.static_token(STATIC)
};
if optional {
b = b.optional();
}
if no_challenge {
b = b.static_challenge(None);
}
if callback {
b = b.on_reject(json_reject);
}
app(b.build().unwrap())
};
let (old, new) = (build(false), build(true));
for headers in &requests {
for path in ["/test", "/required"] {
let a = observed(send(&old, path, headers).await).await;
let b = observed(send(&new, path, headers).await).await;
assert_eq!(
a,
b,
"config {:?} {path} {headers:.40?}",
(with_oauth, optional, no_challenge, callback)
);
}
}
}
}
#[tokio::test]
async fn the_extractors_name_the_matching_key() {
let app = app(AuthLayer::builder()
.static_tokens(rotation())
.build()
.unwrap());
for (secret, label) in [("key-current", "current"), ("key-next", "next")] {
let bearer = format!("Bearer {secret}");
let (status, _, body) =
observed(send(&app, "/test", &[("authorization", &bearer)]).await).await;
assert_eq!(status, 200);
assert_eq!(
body,
format!("Some(StaticToken) Some(StaticTokenMatch {{ label: Some({label:?}) }})")
);
let (status, _, body) =
observed(send(&app, "/required", &[("authorization", &bearer)]).await).await;
assert_eq!((status, body), (200, format!("Some({label:?})")));
}
let (status, headers, _) =
observed(send(&app, "/test", &[("authorization", "Bearer key-old")]).await).await;
assert_eq!(status, 401);
assert!(headers.contains(&(
"www-authenticate".to_string(),
DEFAULT_STATIC_CHALLENGE.as_bytes().to_vec()
)));
}
#[tokio::test]
async fn the_extractor_refuses_an_oauth_request_and_a_route_outside_every_layer() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let layer = AuthLayer::builder()
.oauth(Arc::clone(&v))
.static_tokens(rotation())
.build()
.unwrap();
let router = app(layer);
let valid = format!("Bearer {}", testing::valid_token());
let (status, _, body) =
observed(send(&router, "/test", &[("authorization", &valid)]).await).await;
assert_eq!(status, 200);
assert!(
body.starts_with("Some(OAuth(") && body.ends_with(" None"),
"{body}"
);
let resp = send(&router, "/required", &[("authorization", &valid)]).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
assert_eq!(
resp.headers()[WWW_AUTHENTICATE],
v.invalid_token_challenge().as_str()
);
let oauth_only = app(AuthLayer::builder().oauth(Arc::clone(&v)).build().unwrap());
let resp = send(&oauth_only, "/required", &[("authorization", &valid)]).await;
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
let bare: Router = Router::new().route("/test", get(report)).route(
"/required",
get(|_: StaticTokenMatch| async { "unreachable" }),
);
for path in ["/test", "/required"] {
let resp = send(&bare, path, &[("authorization", "Bearer key-current")]).await;
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR, "{path}");
}
}
#[tokio::test]
async fn an_optional_layer_with_several_tokens() {
let app = app(AuthLayer::builder()
.static_tokens(rotation())
.optional()
.build()
.unwrap());
let (status, _, body) = observed(send(&app, "/test", &[]).await).await;
assert_eq!((status, body.as_str()), (200, "None None"));
let (status, _, _) = observed(send(&app, "/required", &[]).await).await;
assert_eq!(status, 401);
for (secret, label) in [("key-current", "current"), ("key-next", "next")] {
let (status, _, body) = observed(
send(
&app,
"/required",
&[("authorization", &format!("Bearer {secret}"))],
)
.await,
)
.await;
assert_eq!((status, body), (200, format!("Some({label:?})")));
}
let (status, _, _) =
observed(send(&app, "/test", &[("authorization", "Bearer key-old")]).await).await;
assert_eq!(status, 401);
}
#[tokio::test]
async fn nested_layers_keep_the_match_paired_with_the_credential() {
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 outer_static = AuthLayer::builder()
.static_tokens(rotation())
.sources([CredentialSource::Raw(HeaderName::from_static("x-outer"))])
.build()
.unwrap();
let inner_oauth = AuthLayer::builder().oauth(Arc::clone(&v)).build().unwrap();
let router = app(inner_oauth).layer(outer_static.clone());
let (status, _, body) = observed(
send(
&router,
"/test",
&[("x-outer", "key-next"), ("authorization", &valid)],
)
.await,
)
.await;
assert_eq!(status, 200);
assert!(
body.starts_with("Some(OAuth(") && body.ends_with(" None"),
"{body}"
);
let inner_optional = AuthLayer::builder()
.static_token("inner-key")
.sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
.optional()
.build()
.unwrap();
let router = app(inner_optional).layer(outer_static.clone());
let (_, _, body) = observed(send(&router, "/test", &[("x-outer", "key-next")]).await).await;
assert_eq!(body, "None None");
let inner_static = AuthLayer::builder()
.static_token("inner-key")
.sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
.build()
.unwrap();
let router = app(inner_static).layer(outer_static);
let (_, _, body) = observed(
send(
&router,
"/test",
&[("x-outer", "key-next"), ("x-inner", "inner-key")],
)
.await,
)
.await;
assert_eq!(
body,
"Some(StaticToken) Some(StaticTokenMatch { label: None })"
);
}
#[test]
fn build_with_decision_combines_the_set_as_documented() {
let v = validator("http://127.0.0.1:1/jwks");
assert!(
AuthLayer::builder()
.optional_oauth(Some(Arc::clone(&v)))
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::StaticAndOAuth("x".into()))
.is_ok()
);
let dropped = AuthLayer::builder()
.oauth(Arc::clone(&v))
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::StaticIgnored)
.unwrap();
assert!(
format!("{dropped:?}").contains("static_tokens: None"),
"{dropped:?}"
);
for (decision, oauth) in [
(StaticTokenDecision::OAuthOnly, Some(Arc::clone(&v))),
(StaticTokenDecision::Unauthenticated, None),
] {
assert_eq!(
AuthLayer::builder()
.optional_oauth(oauth)
.static_tokens(rotation())
.build_with_decision(decision)
.unwrap_err(),
AuthLayerError::DecisionWithoutStaticToken
);
}
assert!(
AuthLayer::from_decision(StaticTokenDecision::Unauthenticated, None)
.unwrap()
.allows_unauthenticated()
);
assert_eq!(
AuthLayer::builder()
.static_tokens(StaticTokens::new())
.build()
.unwrap_err(),
AuthLayerError::NoCredential
);
}
#[test]
fn debug_never_prints_a_token_from_a_set() {
let builder = AuthLayer::builder()
.static_token("hunter2-single")
.static_tokens(
StaticTokens::new()
.with(Some("next"), "hunter2-next")
.unwrap(),
);
let rendered = format!("{builder:?}");
assert!(
!rendered.contains("hunter2") && rendered.contains("next"),
"{rendered}"
);
let layer = builder.build().unwrap();
let rendered = format!("{layer:?}");
assert!(
!rendered.contains("hunter2") && rendered.contains("<redacted>"),
"{rendered}"
);
}
}