use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use http::header::{AUTHORIZATION, WWW_AUTHENTICATE};
use http::request::Parts;
use http::{HeaderMap, HeaderName, HeaderValue, Request, Response, StatusCode};
use tracing::{debug, error, info, warn};
use zeroize::Zeroizing;
use crate::authenticate::{
Credential, StaticTokenMatch, StaticTokens, authenticate_with_static_tokens,
};
use crate::config::is_scope_token;
use crate::observe::{
self, Mechanism, Outcome, REASON_MISCONFIGURED, REASON_NONE, Stage, count_request,
};
use crate::policy::StaticTokenDecision;
use crate::refusal::{BARE_INSUFFICIENT_SCOPE_CHALLENGE, DEFAULT_STATIC_CHALLENGE, select};
use crate::token::{AuthorizedToken, TokenRejection, missing_scopes};
use crate::validator::OAuthValidator;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum CredentialSource {
Bearer(HeaderName),
Raw(HeaderName),
}
impl CredentialSource {
pub fn authorization_bearer() -> Self {
Self::Bearer(AUTHORIZATION)
}
pub(crate) fn candidate<'h>(&self, headers: &'h HeaderMap) -> Option<&'h str> {
headers.get(self.header_name()).and_then(|v| self.parse(v))
}
fn parse<'h>(&self, value: &'h HeaderValue) -> Option<&'h str> {
let value = value.to_str().ok()?;
Some(match self {
Self::Bearer(_) => bearer_credential(value),
Self::Raw(_) => value,
})
}
pub(crate) fn presents_nothing(&self, headers: &HeaderMap) -> bool {
headers.get_all(self.header_name()).iter().all(|v| {
let Ok(value) = v.to_str() else {
return false;
};
match self {
Self::Raw(_) => value.trim().is_empty(),
Self::Bearer(_) => {
!names_a_token(value) && bearer_credential(value).trim().is_empty()
}
}
})
}
pub(crate) fn header_name(&self) -> &HeaderName {
match self {
Self::Bearer(name) | Self::Raw(name) => name,
}
}
}
pub(crate) fn names_a_token(value: &str) -> bool {
let value = value.trim_start_matches([' ', '\t']);
let (scheme, rest) = value
.find([' ', '\t'])
.map_or((value, ""), |i| value.split_at(i));
scheme.eq_ignore_ascii_case("dpop")
|| (scheme.eq_ignore_ascii_case("bearer") && !rest.trim().is_empty())
}
pub(crate) fn bearer_credential(header: &str) -> &str {
match header.split_once(' ') {
Some((scheme, token)) if scheme.eq_ignore_ascii_case("bearer") => token.trim(),
_ => "",
}
}
#[non_exhaustive]
pub struct RejectContext<'a> {
pub rejection: &'a TokenRejection,
pub status: StatusCode,
pub request: &'a Parts,
}
impl std::fmt::Debug for RejectContext<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let header_names: Vec<&str> = self
.request
.headers
.keys()
.map(HeaderName::as_str)
.collect();
f.debug_struct("RejectContext")
.field("rejection", self.rejection)
.field("status", &self.status)
.field("method", &self.request.method)
.field("uri", &redacted_request_uri(&self.request.uri))
.field("version", &self.request.version)
.field("header_names", &header_names)
.finish_non_exhaustive()
}
}
fn redacted_request_uri(uri: &http::Uri) -> String {
match uri.query() {
Some(_) => format!("{}?***", uri.path()),
None => uri.path().to_string(),
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum AuthLayerError {
#[error(
"no credential is configured: give a static token and/or an OAuth validator \
(the layer's allow_unauthenticated constructor is the explicit opt-out)"
)]
NoCredential,
#[error("no credential source is configured: a request could never present a credential")]
NoSources,
#[error(
"the static-token decision was made with OAuth enabled, but no OAuth validator \
was given"
)]
DecisionNeedsOAuth,
#[error(
"the static-token decision was made with OAuth disabled, but an OAuth validator \
was given"
)]
DecisionWithoutOAuth,
#[error(
"the OAuth WWW-Authenticate challenge is not a valid HTTP header value — check \
the resource URL and scopes for control or non-ASCII characters"
)]
InvalidChallenge,
#[error(
"static tokens were given, but the static-token decision was made without a static \
token: pass the current static token to static_token_policy as well"
)]
DecisionWithoutStaticToken,
#[error(
"a required scope is not a valid scope (printable ASCII with no space, '\"' or '\\', \
RFC 6749 §3.3)"
)]
InvalidScope,
#[error(
"required scopes were given, but there is no OAuth validator and static tokens do not \
bypass scopes: no request could ever pass"
)]
ScopesNeedOAuth,
#[error(
"required scopes were given, but the static-token decision allows unauthenticated \
requests, which checks no scope"
)]
ScopesWithoutAuthentication,
}
pub(crate) struct Gate {
pub(crate) static_tokens: Option<StaticTokens>,
pub(crate) oauth: Option<Arc<OAuthValidator>>,
pub(crate) sources: Vec<CredentialSource>,
oauth_challenges: Option<(HeaderValue, HeaderValue)>,
pub(crate) static_challenge: Option<HeaderValue>,
pub(crate) optional: bool,
pub(crate) required_scopes: Vec<String>,
pub(crate) static_bypasses_scopes: bool,
scope_floor: Vec<String>,
}
pub(crate) enum Admission {
Static,
OAuth(AuthorizedToken),
PassedThrough,
Refused(TokenRejection, Mechanism),
}
macro_rules! log_static_accepted {
($parts:expr) => {{
let parts: &http::request::Parts = $parts;
crate::observe::count_request(
crate::observe::Stage::Layer,
crate::observe::Outcome::Accepted,
crate::observe::Mechanism::Static,
crate::observe::REASON_NONE,
);
match parts
.extensions
.get::<crate::authenticate::StaticTokenMatch>()
.and_then(crate::authenticate::StaticTokenMatch::label)
{
Some(label) => tracing::debug!(
path = %parts.uri.path(),
auth.outcome = crate::observe::Outcome::Accepted.as_str(),
auth.mechanism = crate::observe::Mechanism::Static.as_str(),
auth.static_label = label,
"Static bearer auth accepted"
),
None => tracing::debug!(
path = %parts.uri.path(),
auth.outcome = crate::observe::Outcome::Accepted.as_str(),
auth.mechanism = crate::observe::Mechanism::Static.as_str(),
"Static bearer auth accepted"
),
}
}};
}
pub(crate) use log_static_accepted;
macro_rules! log_oauth_accepted {
($parts:expr, $token:expr) => {{
let parts: &http::request::Parts = $parts;
let token: &crate::token::AuthorizedToken = $token;
crate::observe::count_request(
crate::observe::Stage::Layer,
crate::observe::Outcome::Accepted,
crate::observe::Mechanism::OAuth,
crate::observe::REASON_NONE,
);
tracing::debug!(
path = %parts.uri.path(),
principal = ?token.principal.as_deref().map(crate::token::for_log),
subject = ?token.subject.as_deref().map(crate::token::for_log),
scopes = ?crate::token::scopes_for_log(&token.scopes),
auth.outcome = crate::observe::Outcome::Accepted.as_str(),
auth.mechanism = crate::observe::Mechanism::OAuth.as_str(),
"OAuth bearer auth accepted"
);
}};
}
pub(crate) use log_oauth_accepted;
macro_rules! log_passed_through {
($parts:expr) => {{
let parts: &http::request::Parts = $parts;
crate::observe::count_request(
crate::observe::Stage::Layer,
crate::observe::Outcome::PassedThrough,
crate::observe::Mechanism::None,
crate::observe::REASON_NONE,
);
tracing::debug!(
path = %parts.uri.path(),
auth.outcome = crate::observe::Outcome::PassedThrough.as_str(),
auth.mechanism = crate::observe::Mechanism::None.as_str(),
"No credential presented; optional auth passes the request through"
);
}};
}
pub(crate) use log_passed_through;
macro_rules! log_layer_refusal {
($gate:expr, $parts:expr, $rejection:expr, $mechanism:expr, $stage:expr) => {{
let gate: &crate::http_layer::Gate = $gate;
let parts: &http::request::Parts = $parts;
let rejection: &crate::token::TokenRejection = $rejection;
let mechanism: crate::observe::Mechanism = $mechanism;
let path = parts.uri.path();
let reason = crate::observe::reason(rejection);
let status = crate::observe::status(rejection);
crate::observe::count_request($stage, crate::observe::Outcome::Rejected, mechanism, reason);
match (&gate.oauth, rejection) {
(None, _) => tracing::warn!(
path = %path,
auth.outcome = crate::observe::Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"Bearer auth rejected"
),
(Some(_), crate::token::TokenRejection::Missing) => {
tracing::debug!(
path = %path,
auth.outcome = crate::observe::Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"No bearer credential presented"
);
}
(Some(_), _) => {
tracing::warn!(
path = %path,
reason = ?rejection,
auth.outcome = crate::observe::Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"OAuth bearer auth rejected"
);
}
}
}};
}
pub(crate) use log_layer_refusal;
impl Gate {
#[allow(clippy::too_many_arguments)] pub(crate) fn build(
static_token: Option<Zeroizing<String>>,
static_tokens: Option<StaticTokens>,
oauth: Option<Arc<OAuthValidator>>,
sources: Option<Vec<CredentialSource>>,
static_challenge: Option<Option<HeaderValue>>,
optional: bool,
required_scopes: Vec<String>,
static_bypasses_scopes: bool,
) -> Result<Self, AuthLayerError> {
let static_tokens = StaticTokens::merged(static_tokens, static_token);
if static_tokens.is_none() && oauth.is_none() {
return Err(AuthLayerError::NoCredential);
}
let sources = sources.unwrap_or_else(|| vec![CredentialSource::authorization_bearer()]);
if sources.is_empty() {
return Err(AuthLayerError::NoSources);
}
let required_scopes =
checked_scopes(required_scopes).map_err(|_| AuthLayerError::InvalidScope)?;
if !required_scopes.is_empty() && oauth.is_none() && !static_bypasses_scopes {
return Err(AuthLayerError::ScopesNeedOAuth);
}
let header = |challenge: String| {
HeaderValue::from_str(&challenge).map_err(|_| AuthLayerError::InvalidChallenge)
};
if oauth.as_ref().is_some_and(|v| v.challenge_fell_back()) {
return Err(AuthLayerError::InvalidChallenge);
}
let required: Vec<&str> = required_scopes.iter().map(String::as_str).collect();
let scope_floor: Vec<String> = match &oauth {
Some(v) => v
.scopes_with_floor(&required)
.into_iter()
.map(str::to_owned)
.collect(),
None => required_scopes.clone(),
};
let oauth_challenges = match &oauth {
Some(v) => Some((
header(v.invalid_token_challenge())?,
header(if required_scopes.is_empty() {
v.insufficient_scope_challenge()
} else {
let floor: Vec<&str> = scope_floor.iter().map(String::as_str).collect();
v.insufficient_scope_challenge_for(&floor, None)
})?,
)),
None => None,
};
let static_challenge = static_challenge
.unwrap_or_else(|| Some(HeaderValue::from_static(DEFAULT_STATIC_CHALLENGE)));
Ok(Self {
static_tokens,
oauth,
sources,
oauth_challenges,
static_challenge,
optional,
required_scopes,
static_bypasses_scopes,
scope_floor,
})
}
pub(crate) fn check_decision(
decision: &StaticTokenDecision,
has_oauth: bool,
) -> Result<(), AuthLayerError> {
match (decision.oauth_enabled(), has_oauth) {
(true, false) => Err(AuthLayerError::DecisionNeedsOAuth),
(false, true) => Err(AuthLayerError::DecisionWithoutOAuth),
_ => Ok(()),
}
}
pub(crate) fn decision_tokens(
decision: StaticTokenDecision,
static_tokens: Option<StaticTokens>,
) -> Result<(Option<Zeroizing<String>>, Option<StaticTokens>), AuthLayerError> {
let static_tokens = static_tokens.filter(|s| !s.is_empty());
match decision {
StaticTokenDecision::StaticOnly(t) | StaticTokenDecision::StaticAndOAuth(t) => {
Ok((Some(Zeroizing::new(t)), static_tokens))
}
StaticTokenDecision::StaticIgnored => Ok((None, None)),
_ if static_tokens.is_some() => Err(AuthLayerError::DecisionWithoutStaticToken),
_ => Ok((None, None)),
}
}
pub(crate) async fn admit(&self, parts: &mut Parts) -> Admission {
if self.optional {
parts.extensions.remove::<Credential>();
parts.extensions.remove::<AuthorizedToken>();
parts.extensions.remove::<StaticTokenMatch>();
}
for (name, value) in parts.headers.iter_mut() {
if self.sources.iter().any(|s| s.header_name() == name) {
value.set_sensitive(true);
}
}
let result = {
let headers = &parts.headers;
let candidates = self.sources.iter().filter_map(|s| s.candidate(headers));
authenticate_with_static_tokens(
candidates,
self.static_tokens.as_ref(),
self.oauth.as_deref(),
)
.await
};
match result {
Ok((Credential::StaticToken, _))
if !self.required_scopes.is_empty() && !self.static_bypasses_scopes =>
{
Admission::Refused(TokenRejection::InsufficientScope, Mechanism::Static)
}
Ok((Credential::OAuth(token), _))
if !missing_scopes(
&token.scopes,
self.required_scopes.iter().map(String::as_str),
)
.is_empty() =>
{
Admission::Refused(TokenRejection::InsufficientScope, Mechanism::OAuth)
}
Ok((Credential::StaticToken, matched)) => {
parts.extensions.insert(Credential::StaticToken);
parts
.extensions
.insert(matched.unwrap_or_else(StaticTokenMatch::unlabeled));
Admission::Static
}
Ok((Credential::OAuth(token), _)) => {
parts.extensions.remove::<StaticTokenMatch>();
parts.extensions.insert(token.clone());
parts.extensions.insert(Credential::OAuth(token.clone()));
Admission::OAuth(token)
}
Err(TokenRejection::Missing)
if self.optional
&& self
.sources
.iter()
.all(|s| s.presents_nothing(&parts.headers)) =>
{
Admission::PassedThrough
}
Err(rejection) => {
let mechanism = Mechanism::of_rejection(&rejection);
Admission::Refused(rejection, mechanism)
}
}
}
pub(crate) fn status_and_challenge(
&self,
rejection: &TokenRejection,
) -> (StatusCode, Option<&HeaderValue>) {
self.status_and_challenge_with(rejection, None)
}
pub(crate) fn status_and_challenge_with<'a>(
&'a self,
rejection: &TokenRejection,
insufficient: Option<&'a HeaderValue>,
) -> (StatusCode, Option<&'a HeaderValue>) {
let (status, challenge) = select(
rejection,
self.oauth_challenges
.as_ref()
.map(|(i, s)| (i, insufficient.unwrap_or(s))),
self.static_challenge.as_ref(),
);
let challenge = match (rejection, &self.oauth_challenges, insufficient) {
(TokenRejection::InsufficientScope, None, Some(bare)) => Some(bare),
_ => challenge,
};
let status = StatusCode::from_u16(status).unwrap_or(StatusCode::UNAUTHORIZED);
(status, challenge)
}
pub(crate) fn scope_challenge(&self, extra: &[String]) -> Option<HeaderValue> {
let Some(validator) = self.oauth.as_ref() else {
return Some(HeaderValue::from_static(BARE_INSUFFICIENT_SCOPE_CHALLENGE));
};
let mut all: Vec<&str> = self.scope_floor.iter().map(String::as_str).collect();
for scope in extra {
if !all.contains(&scope.as_str()) {
all.push(scope);
}
}
HeaderValue::from_str(&validator.insufficient_scope_challenge_for(&all, None)).ok()
}
pub(crate) fn finish<B>(
&self,
rejection: &TokenRejection,
response: Response<B>,
) -> Response<B> {
self.finish_with(rejection, None, response)
}
pub(crate) fn finish_with<B>(
&self,
rejection: &TokenRejection,
insufficient: Option<&HeaderValue>,
mut response: Response<B>,
) -> Response<B> {
let (status, challenge) = self.status_and_challenge_with(rejection, insufficient);
*response.status_mut() = status;
if let Some(value) = challenge {
response
.headers_mut()
.insert(WWW_AUTHENTICATE, value.clone());
}
response
}
}
pub trait RefusalResponse<B>: sealed::Sealed<B> {
fn refusal_response(&self, cx: RejectContext<'_>) -> Response<B>;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct EmptyRefusal;
mod sealed {
pub trait Sealed<B> {}
impl<B: Default> Sealed<B> for super::EmptyRefusal {}
impl<B, F> Sealed<B> for F where F: Fn(super::RejectContext<'_>) -> http::Response<B> {}
}
impl<B: Default> RefusalResponse<B> for EmptyRefusal {
fn refusal_response(&self, _cx: RejectContext<'_>) -> Response<B> {
Response::new(B::default())
}
}
impl<B, F> RefusalResponse<B> for F
where
F: Fn(RejectContext<'_>) -> Response<B>,
{
fn refusal_response(&self, cx: RejectContext<'_>) -> Response<B> {
self(cx)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error(
"{scope:?} is not a valid scope (printable ASCII with no space, '\"' or '\\', RFC 6749 §3.3)"
)]
#[non_exhaustive]
pub struct InvalidScope {
scope: String,
}
impl InvalidScope {
pub fn scope(&self) -> &str {
&self.scope
}
}
pub(crate) fn checked_scopes(
scopes: impl IntoIterator<Item = impl Into<String>>,
) -> Result<Vec<String>, InvalidScope> {
let mut out: Vec<String> = Vec::new();
for scope in scopes {
let scope = scope.into();
if !is_scope_token(&scope) {
return Err(InvalidScope { scope });
}
if !out.contains(&scope) {
out.push(scope);
}
}
Ok(out)
}
#[cfg(feature = "axum")]
pub(crate) const fn all_scope_tokens(scopes: &[&str]) -> bool {
let mut i = 0;
while i < scopes.len() {
let bytes = scopes[i].as_bytes();
if bytes.is_empty() {
return false;
}
let mut j = 0;
while j < bytes.len() {
let b = bytes[j];
if !(b == 0x21 || (b >= 0x23 && b <= 0x5B) || (b >= 0x5D && b <= 0x7E)) {
return false;
}
j += 1;
}
i += 1;
}
true
}
pub(crate) fn mark_authorization_sensitive(headers: &mut http::HeaderMap) {
for (name, value) in headers.iter_mut() {
if name == http::header::AUTHORIZATION {
value.set_sensitive(true);
}
}
}
#[derive(Clone)]
pub(crate) struct GateRan(pub(crate) Option<Arc<Gate>>);
pub(crate) struct RefusalBody<B> {
pub(crate) source: Arc<dyn std::any::Any + Send + Sync>,
pub(crate) build:
fn(&(dyn std::any::Any + Send + Sync), RejectContext<'_>) -> Option<Response<B>>,
}
impl<B> Clone for RefusalBody<B> {
fn clone(&self) -> Self {
Self {
source: Arc::clone(&self.source),
build: self.build,
}
}
}
fn http_refusal_body<R, B>(
source: &(dyn std::any::Any + Send + Sync),
cx: RejectContext<'_>,
) -> Option<Response<B>>
where
R: RefusalResponse<B> + 'static,
{
source.downcast_ref::<R>().map(|r| r.refusal_response(cx))
}
pub(crate) enum ScopeVerdict {
Pass,
Refuse(TokenRejection),
NoLayer,
}
pub(crate) fn judge_scopes(
parts: &Parts,
required: &[String],
static_bypasses: bool,
) -> ScopeVerdict {
if parts.extensions.get::<GateRan>().is_none() {
return ScopeVerdict::NoLayer;
}
if required.is_empty() {
return ScopeVerdict::Pass;
}
match parts.extensions.get::<Credential>() {
Some(Credential::OAuth(token)) => {
if missing_scopes(&token.scopes, required.iter().map(String::as_str)).is_empty() {
ScopeVerdict::Pass
} else {
ScopeVerdict::Refuse(TokenRejection::InsufficientScope)
}
}
Some(Credential::StaticToken) if static_bypasses => ScopeVerdict::Pass,
Some(_) => ScopeVerdict::Refuse(TokenRejection::InsufficientScope),
None => ScopeVerdict::Refuse(TokenRejection::Missing),
}
}
pub(crate) fn scope_refusal<B: Default + 'static>(
parts: &Parts,
rejection: &TokenRejection,
required: &[String],
what: &'static str,
) -> Response<B> {
let path = parts.uri.path();
let mechanism = Mechanism::of_request(parts.extensions.get::<Credential>(), rejection);
let stage = if matches!(
what,
"Scoped" | "AuthorizedToken" | "Credential" | "StaticTokenMatch"
) {
Stage::Handler
} else {
Stage::Route
};
let Some(GateRan(gate)) = parts.extensions.get::<GateRan>() else {
error!(
path = %path,
what,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = 500u16,
"Server misconfiguration: a scope requirement ran on a route no authentication \
layer covers; refusing the request"
);
count_request(stage, Outcome::Rejected, mechanism, REASON_MISCONFIGURED);
let mut response = Response::new(B::default());
*response.status_mut() = StatusCode::INTERNAL_SERVER_ERROR;
return response;
};
let Some(gate) = gate else {
error!(
path = %path,
what,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = REASON_MISCONFIGURED,
auth.status = 401u16,
"Server misconfiguration: a scope requirement needs a credential, but its \
authentication layer allows unauthenticated requests; refusing the request"
);
count_request(stage, Outcome::Rejected, mechanism, REASON_MISCONFIGURED);
let mut response = Response::new(B::default());
*response.status_mut() = StatusCode::UNAUTHORIZED;
response.headers_mut().insert(
WWW_AUTHENTICATE,
HeaderValue::from_static(DEFAULT_STATIC_CHALLENGE),
);
return response;
};
let status = observe::status(rejection);
let reason = match (rejection, &gate.oauth) {
(TokenRejection::InsufficientScope, None) => REASON_MISCONFIGURED,
_ => observe::reason(rejection),
};
count_request(stage, Outcome::Rejected, mechanism, reason);
match (rejection, &gate.oauth) {
(TokenRejection::InsufficientScope, None) => error!(
path = %path,
what,
required = ?required,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"Server misconfiguration: the route requires scopes, but its authentication layer \
has no OAuth validator, so no credential can carry them; refusing the request"
),
(TokenRejection::InsufficientScope, Some(_)) => {
let present = match parts.extensions.get::<Credential>() {
Some(Credential::OAuth(token)) => token.scopes.clone(),
_ => Vec::new(),
};
info!(
path = %path,
what,
required = ?required,
present = ?crate::token::scopes_for_log(&present),
static_token = matches!(parts.extensions.get::<Credential>(), Some(Credential::StaticToken)),
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"The credential lacks the scopes this route requires"
);
}
(TokenRejection::Missing, Some(_)) => {
debug!(
path = %path,
what,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"No bearer credential presented"
);
}
_ => warn!(
path = %path,
what,
reason = ?rejection,
auth.outcome = Outcome::Rejected.as_str(),
auth.mechanism = mechanism.as_str(),
auth.reason = reason,
auth.status = status,
"Bearer auth rejected"
),
}
let insufficient = match rejection {
TokenRejection::InsufficientScope => gate.scope_challenge(required),
_ => None,
};
let (status, _) = gate.status_and_challenge_with(rejection, insufficient.as_ref());
let response = parts
.extensions
.get::<RefusalBody<B>>()
.and_then(|body| {
(body.build)(
&*body.source,
RejectContext {
rejection,
status,
request: parts,
},
)
})
.unwrap_or_else(|| Response::new(B::default()));
gate.finish_with(rejection, insufficient.as_ref(), response)
}
#[derive(Clone, Debug)]
pub struct RequireScopes {
scopes: Arc<[String]>,
static_bypasses: bool,
}
impl RequireScopes {
pub fn new(scopes: impl IntoIterator<Item = impl Into<String>>) -> Self {
Self::try_new(scopes).unwrap_or_else(|e| panic!("RequireScopes::new: {e}"))
}
pub fn try_new(
scopes: impl IntoIterator<Item = impl Into<String>>,
) -> Result<Self, InvalidScope> {
Ok(Self {
scopes: checked_scopes(scopes)?.into(),
static_bypasses: false,
})
}
pub fn static_token_bypasses_scopes(mut self) -> Self {
self.static_bypasses = true;
self
}
pub fn scopes(&self) -> &[String] {
&self.scopes
}
}
impl<S> tower_layer::Layer<S> for RequireScopes {
type Service = RequireScopesService<S>;
fn layer(&self, inner: S) -> Self::Service {
RequireScopesService {
require: self.clone(),
inner,
}
}
}
#[derive(Clone, Debug)]
pub struct RequireScopesService<S> {
require: RequireScopes,
inner: S,
}
impl<S, ReqBody, ResBody> tower_service::Service<Request<ReqBody>> for RequireScopesService<S>
where
S: tower_service::Service<Request<ReqBody>, Response = Response<ResBody>>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
ReqBody: Send + 'static,
ResBody: Default + 'static,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future =
Pin<Box<dyn Future<Output = Result<Response<ResBody>, 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<ReqBody>) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let require = self.require.clone();
Box::pin(async move {
let (parts, body) = request.into_parts();
let verdict = judge_scopes(&parts, &require.scopes, require.static_bypasses);
let rejection = match verdict {
ScopeVerdict::Pass => return inner.call(Request::from_parts(parts, body)).await,
ScopeVerdict::Refuse(rejection) => rejection,
ScopeVerdict::NoLayer => TokenRejection::Missing,
};
Ok(scope_refusal(
&parts,
&rejection,
&require.scopes,
"RequireScopes",
))
})
}
}
pub struct HttpAuthLayer<R = EmptyRefusal> {
mode: Arc<HttpMode>,
on_reject: Arc<R>,
}
enum HttpMode {
Enforce(Arc<Gate>),
AllowUnauthenticated,
}
impl<R> Clone for HttpAuthLayer<R> {
fn clone(&self) -> Self {
Self {
mode: Arc::clone(&self.mode),
on_reject: Arc::clone(&self.on_reject),
}
}
}
impl<R> std::fmt::Debug for HttpAuthLayer<R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &*self.mode {
HttpMode::AllowUnauthenticated => f
.debug_struct("HttpAuthLayer")
.field("allow_unauthenticated", &true)
.finish(),
HttpMode::Enforce(g) => f
.debug_struct("HttpAuthLayer")
.field("static_tokens", &g.static_tokens)
.field("oauth", &g.oauth)
.field("sources", &g.sources)
.field("on_reject", &std::any::type_name::<R>())
.field("static_challenge", &g.static_challenge)
.field("optional", &g.optional)
.finish(),
}
}
}
impl HttpAuthLayer {
pub fn builder() -> HttpAuthLayerBuilder {
HttpAuthLayerBuilder::default()
}
pub fn allow_unauthenticated() -> Self {
Self {
mode: Arc::new(HttpMode::AllowUnauthenticated),
on_reject: Arc::new(EmptyRefusal),
}
}
pub fn from_decision(
decision: StaticTokenDecision,
oauth: Option<Arc<OAuthValidator>>,
) -> Result<Self, AuthLayerError> {
Self::builder()
.optional_oauth(oauth)
.build_with_decision(decision)
}
}
impl<R> HttpAuthLayer<R> {
pub fn allows_unauthenticated(&self) -> bool {
matches!(*self.mode, HttpMode::AllowUnauthenticated)
}
pub fn oauth(&self) -> Option<&Arc<OAuthValidator>> {
match &*self.mode {
HttpMode::Enforce(g) => g.oauth.as_ref(),
HttpMode::AllowUnauthenticated => None,
}
}
async fn check<ReqBody, ResBody>(
&self,
request: Request<ReqBody>,
) -> Result<Request<ReqBody>, Response<ResBody>>
where
R: RefusalResponse<ResBody> + Send + Sync + 'static,
ResBody: 'static,
{
let gate = match &*self.mode {
HttpMode::AllowUnauthenticated => {
count_request(
Stage::Layer,
Outcome::PassedThrough,
Mechanism::None,
REASON_NONE,
);
let mut request = request;
mark_authorization_sensitive(request.headers_mut());
self.mark::<ResBody>(request.extensions_mut(), None);
return Ok(request);
}
HttpMode::Enforce(gate) => gate,
};
let (mut parts, body) = request.into_parts();
match 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) => {
log_layer_refusal!(gate, &parts, &rejection, mechanism, Stage::Layer);
let (status, _) = gate.status_and_challenge(&rejection);
let response = self.on_reject.refusal_response(RejectContext {
rejection: &rejection,
status,
request: &parts,
});
return Err(gate.finish(&rejection, response));
}
}
self.mark::<ResBody>(&mut parts.extensions, Some(Arc::clone(gate)));
Ok(Request::from_parts(parts, body))
}
fn mark<ResBody>(&self, extensions: &mut http::Extensions, gate: Option<Arc<Gate>>)
where
R: RefusalResponse<ResBody> + Send + Sync + 'static,
ResBody: 'static,
{
if gate.is_none() && extensions.get::<GateRan>().is_some() {
return;
}
extensions.insert(GateRan(gate));
extensions.insert(RefusalBody::<ResBody> {
source: Arc::clone(&self.on_reject) as Arc<dyn std::any::Any + Send + Sync>,
build: http_refusal_body::<R, ResBody>,
});
}
}
pub struct HttpAuthLayerBuilder<R = EmptyRefusal> {
static_token: Option<Zeroizing<String>>,
static_tokens: Option<StaticTokens>,
oauth: Option<Arc<OAuthValidator>>,
sources: Option<Vec<CredentialSource>>,
static_challenge: Option<Option<HeaderValue>>,
optional: bool,
required_scopes: Vec<String>,
static_bypasses_scopes: bool,
on_reject: R,
}
impl Default for HttpAuthLayerBuilder {
fn default() -> Self {
Self {
static_token: None,
static_tokens: None,
oauth: None,
sources: None,
static_challenge: None,
optional: false,
required_scopes: Vec::new(),
static_bypasses_scopes: false,
on_reject: EmptyRefusal,
}
}
}
impl<R> std::fmt::Debug for HttpAuthLayerBuilder<R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HttpAuthLayerBuilder")
.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", &std::any::type_name::<R>())
.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<R> HttpAuthLayerBuilder<R> {
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<B, F>(self, f: F) -> HttpAuthLayerBuilder<F>
where
F: Fn(RejectContext<'_>) -> Response<B> + Send + Sync + 'static,
{
HttpAuthLayerBuilder {
static_token: self.static_token,
static_tokens: self.static_tokens,
oauth: self.oauth,
sources: self.sources,
static_challenge: self.static_challenge,
optional: self.optional,
required_scopes: self.required_scopes,
static_bypasses_scopes: self.static_bypasses_scopes,
on_reject: f,
}
}
pub fn build_with_decision(
mut self,
decision: StaticTokenDecision,
) -> Result<HttpAuthLayer<R>, 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(HttpAuthLayer {
mode: Arc::new(HttpMode::AllowUnauthenticated),
on_reject: Arc::new(self.on_reject),
});
}
self.static_token = token;
self.static_tokens = tokens;
self.build()
}
pub fn build(self) -> Result<HttpAuthLayer<R>, 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(HttpAuthLayer {
mode: Arc::new(HttpMode::Enforce(Arc::new(gate))),
on_reject: Arc::new(self.on_reject),
})
}
}
impl<S, R> tower_layer::Layer<S> for HttpAuthLayer<R> {
type Service = HttpAuthService<S, R>;
fn layer(&self, inner: S) -> Self::Service {
HttpAuthService {
layer: self.clone(),
inner,
}
}
}
pub struct HttpAuthService<S, R = EmptyRefusal> {
layer: HttpAuthLayer<R>,
inner: S,
}
impl<S: Clone, R> Clone for HttpAuthService<S, R> {
fn clone(&self) -> Self {
Self {
layer: self.layer.clone(),
inner: self.inner.clone(),
}
}
}
impl<S: std::fmt::Debug, R> std::fmt::Debug for HttpAuthService<S, R> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HttpAuthService")
.field("layer", &self.layer)
.field("inner", &self.inner)
.finish()
}
}
impl<S, R, ReqBody, ResBody> tower_service::Service<Request<ReqBody>> for HttpAuthService<S, R>
where
S: tower_service::Service<Request<ReqBody>, Response = Response<ResBody>>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
R: RefusalResponse<ResBody> + Send + Sync + 'static,
ReqBody: Send + 'static,
ResBody: 'static,
{
type Response = Response<ResBody>;
type Error = S::Error;
type Future =
Pin<Box<dyn Future<Output = Result<Response<ResBody>, 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<ReqBody>) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let layer = self.layer.clone();
Box::pin(async move {
let request = match layer.check(request).await {
Ok(request) => request,
Err(refusal) => return Ok(refusal),
};
inner.call(request).await
})
}
}
#[cfg(test)]
mod tests {
use ::tower::{ServiceExt, service_fn};
use super::*;
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 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",
}),
)
}
async fn inner(request: Request<String>) -> Result<Response<String>, std::convert::Infallible> {
for name in ["authorization", "x-api-key"] {
for value in request.headers().get_all(name) {
assert!(
value.is_sensitive(),
"{name} must reach the service sensitive"
);
}
}
let token = request.extensions().get::<AuthorizedToken>();
let credential = request.extensions().get::<Credential>();
let body = match (credential, token) {
(Some(Credential::OAuth(_)), Some(t)) => format!("oauth {:?}", t.subject),
(Some(Credential::StaticToken), None) => "static".to_string(),
(None, None) => "anonymous".to_string(),
other => panic!("unexpected extensions: {other:?}"),
};
Ok(Response::new(body))
}
async fn send<R>(layer: &HttpAuthLayer<R>, headers: &[(&str, &str)]) -> Response<String>
where
R: RefusalResponse<String> + Send + Sync + 'static,
{
let mut request = Request::builder().uri("/test");
for (name, value) in headers {
request = request.header(*name, *value);
}
let service = tower_layer::Layer::layer(layer, service_fn(inner));
service
.oneshot(request.body(String::new()).unwrap())
.await
.unwrap()
}
fn challenge<B>(response: &Response<B>) -> Option<&str> {
response
.headers()
.get(WWW_AUTHENTICATE)
.map(|v| v.to_str().unwrap())
}
#[test]
fn the_builder_fails_closed() {
assert_eq!(
HttpAuthLayer::builder().build().unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
HttpAuthLayer::builder()
.static_token("")
.build()
.unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
HttpAuthLayer::builder().optional().build().unwrap_err(),
AuthLayerError::NoCredential
);
assert_eq!(
HttpAuthLayer::builder()
.static_token(STATIC)
.sources([])
.build()
.unwrap_err(),
AuthLayerError::NoSources
);
assert_eq!(
HttpAuthLayer::builder()
.build_with_decision(StaticTokenDecision::OAuthOnly)
.unwrap_err(),
AuthLayerError::DecisionNeedsOAuth
);
let built = HttpAuthLayer::builder()
.static_token(STATIC)
.optional()
.build()
.unwrap();
assert!(!built.allows_unauthenticated() && built.oauth().is_none());
assert!(
HttpAuthLayer::builder()
.build_with_decision(StaticTokenDecision::Unauthenticated)
.unwrap()
.allows_unauthenticated()
);
}
#[test]
fn a_validator_on_its_fallback_challenges_fails_the_build_as_for_axum() {
let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
cfg.resource = "https://api.example.test/v1\r\nX-Injected: 1".into();
let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
assert_eq!(
HttpAuthLayer::builder()
.static_token(STATIC)
.oauth(Arc::clone(&v))
.build()
.unwrap_err(),
AuthLayerError::InvalidChallenge
);
#[cfg(feature = "axum")]
assert_eq!(
crate::axum::AuthLayer::builder()
.oauth(v)
.build()
.unwrap_err(),
AuthLayerError::InvalidChallenge
);
}
#[tokio::test]
async fn only_the_explicit_opt_out_passes_everything() {
let layer = HttpAuthLayer::allow_unauthenticated();
let service = tower_layer::Layer::layer(
&layer,
service_fn(|request: Request<String>| async move {
let inserted = request.extensions().get::<Credential>().is_some()
|| request.extensions().get::<AuthorizedToken>().is_some();
assert!(!inserted, "a pass-through inserts nothing");
for value in request.headers().get_all("authorization") {
assert!(value.is_sensitive(), "Authorization must be sensitive");
}
Ok::<_, std::convert::Infallible>(Response::new(String::from("anonymous")))
}),
);
for authorization in [None, Some("Bearer junk")] {
let mut request = Request::builder().uri("/test");
if let Some(value) = authorization {
request = request.header("authorization", value);
}
let response = service
.clone()
.oneshot(request.body(String::new()).unwrap())
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.body(), "anonymous");
}
}
#[tokio::test]
async fn oauth_accepts_a_valid_token_and_refuses_the_rest_with_the_validators_challenge() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let layer = HttpAuthLayer::builder()
.oauth(Arc::clone(&v))
.build()
.unwrap();
let bearer = format!("Bearer {}", testing::valid_token());
let ok = send(&layer, &[("authorization", &bearer)]).await;
assert_eq!(ok.status(), StatusCode::OK);
assert!(ok.body().starts_with("oauth "), "{}", ok.body());
assert_eq!(challenge(&ok), None);
let cases = [
(None, 401, v.invalid_token_challenge()),
(
Some(format!("Bearer {}", expired_token())),
401,
v.invalid_token_challenge(),
),
(
Some("Bearer not-a-jwt".to_string()),
401,
v.invalid_token_challenge(),
),
(
Some(format!("Bearer {}", unscoped_token())),
403,
v.insufficient_scope_challenge(),
),
];
for (header, status, expected) in cases {
let headers: Vec<(&str, &str)> = header
.iter()
.map(|h| ("authorization", h.as_str()))
.collect();
let response = send(&layer, &headers).await;
assert_eq!(response.status().as_u16(), status, "{header:?}");
assert_eq!(challenge(&response), Some(expected.as_str()), "{header:?}");
assert_eq!(
response.body(),
"",
"the default body is ResBody::default()"
);
}
}
#[tokio::test]
async fn static_only_sends_the_static_challenge_unless_opted_out() {
let default = HttpAuthLayer::builder()
.static_token(STATIC)
.build()
.unwrap();
let ok = send(&default, &[("authorization", "Bearer secret")]).await;
assert_eq!(
(ok.status(), ok.body().as_str()),
(StatusCode::OK, "static")
);
for headers in [&[][..], &[("authorization", "Bearer wrong")][..]] {
let refused = send(&default, headers).await;
assert_eq!(refused.status(), StatusCode::UNAUTHORIZED);
assert_eq!(challenge(&refused), Some(DEFAULT_STATIC_CHALLENGE));
}
let custom = HttpAuthLayer::builder()
.static_token(STATIC)
.static_challenge(Some(HeaderValue::from_static("ApiKey realm=\"x\"")))
.build()
.unwrap();
assert_eq!(
challenge(&send(&custom, &[]).await),
Some("ApiKey realm=\"x\"")
);
let none = HttpAuthLayer::builder()
.static_token(STATIC)
.static_challenge(None)
.build()
.unwrap();
let refused = send(&none, &[]).await;
assert_eq!(refused.status(), StatusCode::UNAUTHORIZED);
assert_eq!(challenge(&refused), None);
}
#[tokio::test]
async fn every_source_header_is_sensitive_for_the_callback_and_the_service() {
let layer = HttpAuthLayer::builder()
.static_token(STATIC)
.sources([
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
])
.on_reject(|cx: RejectContext<'_>| {
for name in ["authorization", "x-api-key"] {
assert!(cx.request.headers[name].is_sensitive(), "{name}");
}
assert!(!format!("{cx:?}").contains("wrong"));
Response::new(String::from("refused"))
})
.build()
.unwrap();
let ok = send(
&layer,
&[("authorization", "Bearer wrong"), ("x-api-key", STATIC)],
)
.await;
assert_eq!(ok.status(), StatusCode::OK);
let refused = send(
&layer,
&[("authorization", "Bearer wrong"), ("x-api-key", "wrong")],
)
.await;
assert_eq!(refused.status(), StatusCode::UNAUTHORIZED);
assert_eq!(refused.body(), "refused");
}
#[tokio::test]
async fn on_reject_shapes_the_body_but_not_the_status_or_the_challenge() {
let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
let v = validator(&jwks.url);
let layer = HttpAuthLayer::builder()
.oauth(Arc::clone(&v))
.on_reject(|cx: RejectContext<'_>| {
Response::builder()
.status(StatusCode::IM_A_TEAPOT)
.header(WWW_AUTHENTICATE, "Basic realm=\"nope\"")
.header("content-type", "application/json")
.body(format!("{{\"status\":{}}}", cx.status.as_u16()))
.unwrap()
})
.build()
.unwrap();
let bearer = format!("Bearer {}", unscoped_token());
let refused = send(&layer, &[("authorization", &bearer)]).await;
assert_eq!(refused.status(), StatusCode::FORBIDDEN);
let challenges: Vec<_> = refused.headers().get_all(WWW_AUTHENTICATE).iter().collect();
assert_eq!(challenges, [v.insufficient_scope_challenge().as_str()]);
assert_eq!(refused.headers()["content-type"], "application/json");
assert_eq!(refused.body(), "{\"status\":403}");
}
#[tokio::test]
async fn an_optional_layer_passes_only_a_request_presenting_nothing() {
let layer = HttpAuthLayer::builder()
.static_token(STATIC)
.optional()
.build()
.unwrap();
for headers in [
&[][..],
&[("authorization", "Bearer ")][..],
&[("authorization", "Basic x")][..],
] {
let response = send(&layer, headers).await;
assert_eq!(
(response.status(), response.body().as_str()),
(StatusCode::OK, "anonymous")
);
}
for value in ["Bearer wrong", "DPoP x", "Bearer\tx"] {
let response = send(&layer, &[("authorization", value)]).await;
assert_eq!(response.status(), StatusCode::UNAUTHORIZED, "{value:?}");
assert_eq!(challenge(&response), Some(DEFAULT_STATIC_CHALLENGE));
}
}
#[test]
fn debug_never_prints_the_static_token() {
let builder = HttpAuthLayer::builder().static_token("hunter2");
assert!(!format!("{builder:?}").contains("hunter2"));
let layer = builder.build().unwrap();
let rendered = format!("{layer:?}");
assert!(
!rendered.contains("hunter2") && rendered.contains("<redacted>"),
"{rendered}"
);
let service = tower_layer::Layer::layer(&layer, "inner");
assert!(!format!("{service:?}").contains("hunter2"));
}
fn rotation() -> StaticTokens {
StaticTokens::new()
.with(Some("current"), "key-current")
.and_then(|t| t.with(Some("next"), "key-next"))
.unwrap()
}
type Seen = (u16, Vec<(String, String)>, String);
async fn seen<R>(layer: &HttpAuthLayer<R>, headers: &[(&str, &str)]) -> Seen
where
R: RefusalResponse<String> + Send + Sync + 'static,
{
let mut request = Request::builder().uri("/test");
for (name, value) in headers {
request = request.header(*name, *value);
}
let service = tower_layer::Layer::layer(
layer,
service_fn(|request: Request<String>| async move {
let body = format!(
"{:?} {:?}",
request.extensions().get::<Credential>(),
request.extensions().get::<StaticTokenMatch>(),
);
Ok::<_, std::convert::Infallible>(Response::new(body))
}),
);
let response = service
.oneshot(request.body(String::new()).unwrap())
.await
.unwrap();
let headers = response
.headers()
.iter()
.map(|(k, v)| (k.to_string(), v.to_str().unwrap().to_string()))
.collect();
(response.status().as_u16(), headers, response.into_body())
}
#[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 unscoped = format!("Bearer {}", unscoped_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![("x-api-key", "secret")],
vec![("authorization", "Bearer wrong"), ("x-api-key", "secret")],
vec![("authorization", valid.as_str())],
vec![("authorization", unscoped.as_str())],
];
for (oauth, optional) in [(None, false), (Some(&v), false), (None, true)] {
let sources = [
CredentialSource::authorization_bearer(),
CredentialSource::Raw(HeaderName::from_static("x-api-key")),
];
let mut old = HttpAuthLayer::builder()
.static_token(STATIC)
.optional_oauth(oauth.cloned())
.sources(sources.clone());
let mut new = HttpAuthLayer::builder()
.static_tokens(StaticTokens::single(STATIC).unwrap())
.optional_oauth(oauth.cloned())
.sources(sources);
if optional {
old = old.optional();
new = new.optional();
}
let (old, new) = (old.build().unwrap(), new.build().unwrap());
for headers in &requests {
let a = seen(&old, headers).await;
let b = seen(&new, headers).await;
assert_eq!(a, b, "oauth={} {headers:.40?}", oauth.is_some());
if a.2.starts_with("Some(StaticToken)") {
assert!(
a.2.ends_with("Some(StaticTokenMatch { label: None })"),
"{a:?}"
);
}
}
}
}
#[tokio::test]
async fn every_token_in_a_set_is_accepted_with_its_label() {
let layer = HttpAuthLayer::builder()
.static_tokens(rotation())
.build()
.unwrap();
for (secret, label) in [("key-current", "current"), ("key-next", "next")] {
let (status, _, body) =
seen(&layer, &[("authorization", &format!("Bearer {secret}"))]).await;
assert_eq!(status, 200);
assert_eq!(
body,
format!("Some(StaticToken) Some(StaticTokenMatch {{ label: Some({label:?}) }})")
);
}
let (status, headers, _) = seen(&layer, &[("authorization", "Bearer key-old")]).await;
assert_eq!(status, 401);
assert!(headers.contains(&(
"www-authenticate".to_string(),
DEFAULT_STATIC_CHALLENGE.to_string()
)));
}
#[tokio::test]
async fn static_token_and_static_tokens_are_merged() {
let layer = HttpAuthLayer::builder()
.static_token("key-next")
.static_tokens(rotation())
.build()
.unwrap();
let (_, _, body) = seen(&layer, &[("authorization", "Bearer key-next")]).await;
assert!(body.ends_with("label: Some(\"next\") })"), "{body}");
let layer = HttpAuthLayer::builder()
.static_token("key-extra")
.static_tokens(rotation())
.build()
.unwrap();
for secret in ["key-current", "key-next", "key-extra"] {
let (status, _, _) =
seen(&layer, &[("authorization", &format!("Bearer {secret}"))]).await;
assert_eq!(status, 200, "{secret}");
}
let (_, _, body) = seen(&layer, &[("authorization", "Bearer key-extra")]).await;
assert!(body.ends_with("label: None })"), "{body}");
}
#[tokio::test]
async fn a_whitespace_static_token_is_no_credential() {
for blank in [" ", "\t", " \r\n "] {
for builder in [
HttpAuthLayer::builder().static_token(blank),
HttpAuthLayer::builder().static_token(blank).optional(),
HttpAuthLayer::builder()
.static_token(blank)
.static_tokens(StaticTokens::new()),
] {
assert_eq!(
builder.build().unwrap_err(),
AuthLayerError::NoCredential,
"{blank:?}"
);
}
assert_eq!(
HttpAuthLayer::builder()
.build_with_decision(StaticTokenDecision::StaticOnly(blank.into()))
.unwrap_err(),
AuthLayerError::NoCredential
);
}
let v = validator("http://127.0.0.1:1/jwks");
let layer = HttpAuthLayer::builder()
.oauth(v)
.static_token(" ")
.build()
.unwrap();
let refused = send(&layer, &[("authorization", "Bearer ")]).await;
assert_eq!(refused.status(), StatusCode::UNAUTHORIZED);
}
#[test]
fn an_empty_set_is_no_credential() {
for builder in [
HttpAuthLayer::builder().static_tokens(StaticTokens::new()),
HttpAuthLayer::builder()
.static_tokens(StaticTokens::new())
.static_token(""),
HttpAuthLayer::builder()
.static_tokens(StaticTokens::new())
.optional(),
HttpAuthLayer::builder()
.static_tokens(rotation())
.optional_static_tokens(None),
] {
assert_eq!(builder.build().unwrap_err(), AuthLayerError::NoCredential);
}
}
#[tokio::test]
async fn static_tokens_follow_the_decision() {
let v = validator("http://127.0.0.1:1/jwks");
let accepts = |layer: HttpAuthLayer| async move {
let mut accepted = Vec::new();
for secret in ["key-current", "key-next", "decided"] {
let (status, _, body) =
seen(&layer, &[("authorization", &format!("Bearer {secret}"))]).await;
if status == 200 {
accepted.push(format!(
"{secret}={}",
body.rsplit("label: ").next().unwrap()
));
}
}
accepted
};
for (decision, oauth) in [
(StaticTokenDecision::StaticOnly("decided".into()), None),
(
StaticTokenDecision::StaticAndOAuth("decided".into()),
Some(Arc::clone(&v)),
),
] {
let layer = HttpAuthLayer::builder()
.optional_oauth(oauth)
.static_tokens(rotation())
.build_with_decision(decision)
.unwrap();
assert_eq!(
accepts(layer).await,
[
"key-current=Some(\"current\") })",
"key-next=Some(\"next\") })",
"decided=None })"
]
);
}
let layer = HttpAuthLayer::builder()
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::StaticOnly("key-current".into()))
.unwrap();
assert_eq!(
accepts(layer).await,
[
"key-current=Some(\"current\") })",
"key-next=Some(\"next\") })"
]
);
let layer = HttpAuthLayer::builder()
.oauth(Arc::clone(&v))
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::StaticIgnored)
.unwrap();
assert!(accepts(layer).await.is_empty());
assert_eq!(
HttpAuthLayer::builder()
.oauth(Arc::clone(&v))
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::OAuthOnly)
.unwrap_err(),
AuthLayerError::DecisionWithoutStaticToken
);
assert_eq!(
HttpAuthLayer::builder()
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::Unauthenticated)
.unwrap_err(),
AuthLayerError::DecisionWithoutStaticToken
);
assert_eq!(
HttpAuthLayer::builder()
.static_tokens(rotation())
.build_with_decision(StaticTokenDecision::OAuthOnly)
.unwrap_err(),
AuthLayerError::DecisionNeedsOAuth
);
assert!(
HttpAuthLayer::builder()
.static_tokens(StaticTokens::new())
.build_with_decision(StaticTokenDecision::Unauthenticated)
.unwrap()
.allows_unauthenticated()
);
assert!(
HttpAuthLayer::builder()
.oauth(Arc::clone(&v))
.static_tokens(StaticTokens::new())
.build_with_decision(StaticTokenDecision::OAuthOnly)
.is_ok()
);
}
#[tokio::test]
async fn an_optional_layer_with_several_tokens() {
let layer = HttpAuthLayer::builder()
.static_tokens(rotation())
.optional()
.build()
.unwrap();
let (status, _, body) = seen(&layer, &[]).await;
assert_eq!((status, body.as_str()), (200, "None None"));
for (secret, label) in [("key-current", "current"), ("key-next", "next")] {
let (status, _, body) =
seen(&layer, &[("authorization", &format!("Bearer {secret}"))]).await;
assert_eq!(status, 200);
assert!(body.ends_with(&format!("Some({label:?}) }})")), "{body}");
}
let (status, _, _) = seen(&layer, &[("authorization", "Bearer key-old")]).await;
assert_eq!(status, 401);
}
#[test]
fn debug_never_prints_a_token_from_a_set() {
let builder = HttpAuthLayer::builder()
.static_token("hunter2-single")
.static_tokens(
StaticTokens::new()
.with(Some("current"), "hunter2-current")
.unwrap(),
);
let rendered = format!("{builder:?}");
assert!(
!rendered.contains("hunter2") && rendered.contains("current"),
"{rendered}"
);
let layer = builder.build().unwrap();
let rendered = format!("{layer:?}");
assert!(
!rendered.contains("hunter2") && rendered.contains("len: 2"),
"{rendered}"
);
let service = tower_layer::Layer::layer(&layer, "inner");
assert!(!format!("{service:?}").contains("hunter2"));
}
}