1use std::future::Future;
184use std::pin::Pin;
185use std::sync::Arc;
186use std::task::{Context, Poll};
187
188use ::axum::Json;
189use ::axum::Router;
190use ::axum::body::Body;
191use ::axum::extract::{FromRequestParts, OptionalFromRequestParts, Request, State};
192use ::axum::middleware::Next;
193use ::axum::response::{IntoResponse, Response};
194use ::axum::routing::{any, get};
195use http::header::WWW_AUTHENTICATE;
196use http::request::Parts;
197use http::{HeaderValue, Method, StatusCode};
198use tracing::{error, info};
199use zeroize::Zeroizing;
200
201use crate::authenticate::{Credential, StaticTokenMatch, StaticTokens};
202use crate::challenge::PROTECTED_RESOURCE_METADATA_PREFIX;
203use crate::policy::StaticTokenDecision;
204use crate::token::{AuthorizedToken, InvalidTokenKind, TokenRejection};
205use crate::validator::OAuthValidator;
206
207use crate::http_layer::{
208 Admission, Gate, GateRan, RefusalBody, log_layer_refusal, log_oauth_accepted,
209 log_passed_through, log_static_accepted,
210};
211#[doc(inline)]
212pub use crate::http_layer::{
213 AuthLayerError, CredentialSource, InvalidScope, RejectContext, RequireScopes,
214 RequireScopesService,
215};
216use crate::observe::{
217 self, Mechanism, Outcome, REASON_MISCONFIGURED, REASON_NONE, Stage, count_request,
218};
219pub use crate::refusal::DEFAULT_STATIC_CHALLENGE;
220#[cfg(test)]
223use crate::http_layer::{bearer_credential, names_a_token};
224#[cfg(test)]
225use http::{HeaderMap, HeaderName};
226
227pub type RejectFn = Arc<dyn Fn(RejectContext<'_>) -> Response + Send + Sync>;
229
230#[derive(Clone)]
251pub struct AuthLayer {
252 inner: Arc<Mode>,
253}
254
255enum Mode {
256 Enforce(Enforce),
257 AllowUnauthenticated,
258}
259
260struct Enforce {
265 gate: Arc<Gate>,
266 on_reject: Option<RejectFn>,
267}
268
269impl std::fmt::Debug for AuthLayer {
271 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
272 match &*self.inner {
273 Mode::AllowUnauthenticated => f
274 .debug_struct("AuthLayer")
275 .field("allow_unauthenticated", &true)
276 .finish(),
277 Mode::Enforce(e) => f
278 .debug_struct("AuthLayer")
279 .field("static_tokens", &e.gate.static_tokens)
280 .field("oauth", &e.gate.oauth)
281 .field("sources", &e.gate.sources)
282 .field("on_reject", &e.on_reject.as_ref().map(|_| "<fn>"))
283 .field("static_challenge", &e.gate.static_challenge)
284 .field("optional", &e.gate.optional)
285 .finish(),
286 }
287 }
288}
289
290impl AuthLayer {
291 pub fn builder() -> AuthLayerBuilder {
326 AuthLayerBuilder::default()
327 }
328
329 pub fn allow_unauthenticated() -> Self {
350 Self {
351 inner: Arc::new(Mode::AllowUnauthenticated),
352 }
353 }
354
355 pub fn from_decision(
365 decision: StaticTokenDecision,
366 oauth: Option<Arc<OAuthValidator>>,
367 ) -> Result<Self, AuthLayerError> {
368 Self::builder()
369 .optional_oauth(oauth)
370 .build_with_decision(decision)
371 }
372
373 pub fn allows_unauthenticated(&self) -> bool {
375 matches!(*self.inner, Mode::AllowUnauthenticated)
376 }
377
378 pub fn oauth(&self) -> Option<&Arc<OAuthValidator>> {
380 match &*self.inner {
381 Mode::Enforce(e) => e.gate.oauth.as_ref(),
382 Mode::AllowUnauthenticated => None,
383 }
384 }
385}
386
387#[derive(Default)]
389pub struct AuthLayerBuilder {
390 static_token: Option<Zeroizing<String>>,
391 static_tokens: Option<StaticTokens>,
392 oauth: Option<Arc<OAuthValidator>>,
393 sources: Option<Vec<CredentialSource>>,
394 on_reject: Option<RejectFn>,
395 static_challenge: Option<Option<HeaderValue>>,
397 optional: bool,
398 required_scopes: Vec<String>,
399 static_bypasses_scopes: bool,
400}
401
402impl std::fmt::Debug for AuthLayerBuilder {
403 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
404 f.debug_struct("AuthLayerBuilder")
405 .field(
406 "static_token",
407 &self.static_token.as_ref().map(|_| "<redacted>"),
408 )
409 .field("static_tokens", &self.static_tokens)
410 .field("oauth", &self.oauth)
411 .field("sources", &self.sources)
412 .field("on_reject", &self.on_reject.as_ref().map(|_| "<fn>"))
413 .field("static_challenge", &self.static_challenge)
414 .field("optional", &self.optional)
415 .field("required_scopes", &self.required_scopes)
416 .field("static_bypasses_scopes", &self.static_bypasses_scopes)
417 .finish()
418 }
419}
420
421impl AuthLayerBuilder {
422 pub fn static_token(mut self, token: impl Into<String>) -> Self {
432 self.static_token = Some(Zeroizing::new(token.into()));
433 self
434 }
435
436 pub fn optional_static_token(mut self, token: Option<String>) -> Self {
439 self.static_token = token.map(Zeroizing::new);
440 self
441 }
442
443 pub fn static_tokens(mut self, tokens: StaticTokens) -> Self {
498 self.static_tokens = Some(tokens);
499 self
500 }
501
502 pub fn optional_static_tokens(mut self, tokens: Option<StaticTokens>) -> Self {
506 self.static_tokens = tokens;
507 self
508 }
509
510 pub fn oauth(mut self, validator: Arc<OAuthValidator>) -> Self {
512 self.oauth = Some(validator);
513 self
514 }
515
516 pub fn optional_oauth(mut self, validator: Option<Arc<OAuthValidator>>) -> Self {
518 self.oauth = validator;
519 self
520 }
521
522 pub fn sources(mut self, sources: impl IntoIterator<Item = CredentialSource>) -> Self {
526 self.sources = Some(sources.into_iter().collect());
527 self
528 }
529
530 pub fn static_challenge(mut self, challenge: Option<HeaderValue>) -> Self {
553 self.static_challenge = Some(challenge);
554 self
555 }
556
557 pub fn optional(mut self) -> Self {
626 self.optional = true;
627 self
628 }
629
630 pub fn require_scopes(mut self, scopes: impl IntoIterator<Item = impl Into<String>>) -> Self {
685 self.required_scopes = scopes.into_iter().map(Into::into).collect();
686 self
687 }
688
689 pub fn static_token_bypasses_scopes(mut self) -> Self {
698 self.static_bypasses_scopes = true;
699 self
700 }
701
702 pub fn on_reject(
723 mut self,
724 f: impl Fn(RejectContext<'_>) -> Response + Send + Sync + 'static,
725 ) -> Self {
726 self.on_reject = Some(Arc::new(f));
727 self
728 }
729
730 pub fn build_with_decision(
764 mut self,
765 decision: StaticTokenDecision,
766 ) -> Result<AuthLayer, AuthLayerError> {
767 Gate::check_decision(&decision, self.oauth.is_some())?;
768 let unauthenticated = decision == StaticTokenDecision::Unauthenticated;
769 let (token, tokens) = Gate::decision_tokens(decision, self.static_tokens.take())?;
770 if unauthenticated && !self.required_scopes.is_empty() {
771 return Err(AuthLayerError::ScopesWithoutAuthentication);
772 }
773 if unauthenticated {
774 return Ok(AuthLayer::allow_unauthenticated());
775 }
776 self.static_token = token;
777 self.static_tokens = tokens;
778 self.build()
779 }
780
781 pub fn build(self) -> Result<AuthLayer, AuthLayerError> {
795 let gate = Gate::build(
798 self.static_token,
799 self.static_tokens,
800 self.oauth,
801 self.sources,
802 self.static_challenge,
803 self.optional,
804 self.required_scopes,
805 self.static_bypasses_scopes,
806 )?;
807 Ok(AuthLayer {
808 inner: Arc::new(Mode::Enforce(Enforce {
809 gate: Arc::new(gate),
810 on_reject: self.on_reject,
811 })),
812 })
813 }
814}
815
816impl Enforce {
817 fn reject(&self, rejection: &TokenRejection, request: &Parts) -> Response {
838 self.reject_with(rejection, request, None)
839 }
840
841 fn reject_with(
844 &self,
845 rejection: &TokenRejection,
846 request: &Parts,
847 insufficient: Option<&HeaderValue>,
848 ) -> Response {
849 let (status, _) = self.gate.status_and_challenge_with(rejection, insufficient);
850 let response = match &self.on_reject {
851 Some(f) => f(RejectContext {
852 rejection,
853 status,
854 request,
855 }),
856 None => Response::new(Body::empty()),
857 };
858 self.gate.finish_with(rejection, insufficient, response)
861 }
862
863 fn refuse_scoped(&self, request: &Parts, required: &[String]) -> Response {
869 let path = request.uri.path();
870 let rejection = TokenRejection::InsufficientScope;
871 let mechanism = Mechanism::of_request(request.extensions.get::<Credential>(), &rejection);
872 if self.gate.oauth.is_none() {
873 count_request(
874 Stage::Handler,
875 Outcome::Rejected,
876 mechanism,
877 REASON_MISCONFIGURED,
878 );
879 error!(
880 path = %path,
881 required = ?required,
882 auth.outcome = Outcome::Rejected.as_str(),
883 auth.mechanism = mechanism.as_str(),
884 auth.reason = REASON_MISCONFIGURED,
885 auth.status = 403u16,
886 "Server misconfiguration: the handler requires scopes, but its AuthLayer has no \
887 OAuth validator, so no credential can carry them; refusing the request"
888 );
889 } else {
890 count_request(
891 Stage::Handler,
892 Outcome::Rejected,
893 mechanism,
894 observe::reason(&rejection),
895 );
896 let present = match request.extensions.get::<Credential>() {
897 Some(Credential::OAuth(token)) => token.scopes.clone(),
898 _ => Vec::new(),
899 };
900 info!(
901 path = %path,
902 required = ?required,
903 present = ?crate::token::scopes_for_log(&present),
904 auth.outcome = Outcome::Rejected.as_str(),
905 auth.mechanism = mechanism.as_str(),
906 auth.reason = observe::reason(&rejection),
907 auth.status = 403u16,
908 "The credential lacks the scopes this handler requires"
909 );
910 }
911 let insufficient = self.gate.scope_challenge(required);
913 self.reject_with(&rejection, request, insufficient.as_ref())
914 }
915
916 fn refuse(
921 &self,
922 rejection: &TokenRejection,
923 request: &Parts,
924 mechanism: Mechanism,
925 stage: Stage,
926 ) -> Response {
927 log_layer_refusal!(&self.gate, request, rejection, mechanism, stage);
931 self.reject(rejection, request)
932 }
933}
934
935#[derive(Clone)]
941struct LayerRan(AuthLayer);
942
943fn axum_refusal_body(
945 source: &(dyn std::any::Any + Send + Sync),
946 cx: RejectContext<'_>,
947) -> Option<Response> {
948 match source.downcast_ref::<Mode>()? {
949 Mode::Enforce(Enforce {
950 on_reject: Some(f), ..
951 }) => Some(f(cx)),
952 _ => None,
953 }
954}
955
956impl AuthLayer {
957 fn mark(&self, extensions: &mut http::Extensions) {
961 let gate = match &*self.inner {
962 Mode::Enforce(enforce) => Some(Arc::clone(&enforce.gate)),
963 Mode::AllowUnauthenticated => None,
964 };
965 if gate.is_none() && extensions.get::<GateRan>().is_some() {
970 return;
971 }
972 extensions.insert(GateRan(gate));
973 extensions.insert(RefusalBody::<Body> {
974 source: Arc::clone(&self.inner) as Arc<dyn std::any::Any + Send + Sync>,
975 build: axum_refusal_body,
976 });
977 }
978}
979
980impl AuthLayer {
981 async fn check(&self, mut request: Request) -> Result<Request, Response> {
986 let enforce = match &*self.inner {
987 Mode::AllowUnauthenticated => {
988 count_request(
989 Stage::Layer,
990 Outcome::PassedThrough,
991 Mechanism::None,
992 REASON_NONE,
993 );
994 if request.extensions().get::<LayerRan>().is_none() {
1000 request.extensions_mut().insert(LayerRan(self.clone()));
1001 }
1002 crate::http_layer::mark_authorization_sensitive(request.headers_mut());
1003 self.mark(request.extensions_mut());
1004 return Ok(request);
1005 }
1006 Mode::Enforce(enforce) => enforce,
1007 };
1008
1009 let (mut parts, body) = request.into_parts();
1010 match enforce.gate.admit(&mut parts).await {
1016 Admission::Static => log_static_accepted!(&parts),
1017 Admission::OAuth(token) => log_oauth_accepted!(&parts, &token),
1018 Admission::PassedThrough => log_passed_through!(&parts),
1019 Admission::Refused(rejection, mechanism) => {
1020 return Err(enforce.refuse(&rejection, &parts, mechanism, Stage::Layer));
1021 }
1022 }
1023 parts.extensions.insert(LayerRan(self.clone()));
1024 self.mark(&mut parts.extensions);
1025 Ok(Request::from_parts(parts, body))
1026 }
1027
1028 fn refuse_extraction(
1044 &self,
1045 rejection: &TokenRejection,
1046 parts: &Parts,
1047 wants: Wants,
1048 ) -> Response {
1049 let mechanism = Mechanism::of_request(parts.extensions.get::<Credential>(), rejection);
1050 let misconfigured = || {
1051 count_request(
1052 Stage::Handler,
1053 Outcome::Rejected,
1054 mechanism,
1055 REASON_MISCONFIGURED,
1056 )
1057 };
1058 match &*self.inner {
1059 Mode::Enforce(enforce)
1060 if wants == Wants::OAuthToken && enforce.gate.oauth.is_none() =>
1061 {
1062 misconfigured();
1063 error!(
1064 path = %parts.uri.path(),
1065 auth.outcome = Outcome::Rejected.as_str(),
1066 auth.mechanism = mechanism.as_str(),
1067 auth.reason = REASON_MISCONFIGURED,
1068 auth.status = observe::status(rejection),
1069 "Server misconfiguration: the handler requires an OAuth access token, but \
1070 its AuthLayer has no OAuth validator; refusing the request"
1071 );
1072 enforce.reject(rejection, parts)
1073 }
1074 Mode::Enforce(enforce)
1075 if wants == Wants::StaticToken && enforce.gate.static_tokens.is_none() =>
1076 {
1077 misconfigured();
1078 error!(
1079 path = %parts.uri.path(),
1080 auth.outcome = Outcome::Rejected.as_str(),
1081 auth.mechanism = mechanism.as_str(),
1082 auth.reason = REASON_MISCONFIGURED,
1083 auth.status = observe::status(rejection),
1084 "Server misconfiguration: the handler requires a static token, but its \
1085 AuthLayer has no static token; refusing the request"
1086 );
1087 enforce.reject(rejection, parts)
1088 }
1089 Mode::Enforce(enforce) => enforce.refuse(rejection, parts, mechanism, Stage::Handler),
1090 Mode::AllowUnauthenticated
1098 if matches!(parts.extensions.get::<GateRan>(), Some(GateRan(Some(_)))) =>
1099 {
1100 crate::http_layer::scope_refusal::<Body>(parts, rejection, &[], wants.extractor())
1101 }
1102 Mode::AllowUnauthenticated => {
1105 misconfigured();
1106 error!(
1107 path = %parts.uri.path(),
1108 auth.outcome = Outcome::Rejected.as_str(),
1109 auth.mechanism = mechanism.as_str(),
1110 auth.reason = REASON_MISCONFIGURED,
1111 auth.status = 401u16,
1112 "Server misconfiguration: the handler requires a credential, but its \
1113 AuthLayer allows unauthenticated requests; refusing the request"
1114 );
1115 (
1116 StatusCode::UNAUTHORIZED,
1117 [(
1118 WWW_AUTHENTICATE,
1119 HeaderValue::from_static(DEFAULT_STATIC_CHALLENGE),
1120 )],
1121 )
1122 .into_response()
1123 }
1124 }
1125 }
1126}
1127
1128#[derive(Clone, Copy, PartialEq, Eq)]
1130enum Wants {
1131 AnyCredential,
1133 OAuthToken,
1135 StaticToken,
1137}
1138
1139impl Wants {
1140 fn extractor(self) -> &'static str {
1142 match self {
1143 Self::AnyCredential => "Credential",
1144 Self::OAuthToken => "AuthorizedToken",
1145 Self::StaticToken => "StaticTokenMatch",
1146 }
1147 }
1148}
1149
1150enum Found<T> {
1152 Present(T),
1154 Absent(AuthLayer),
1156 NoLayer,
1158}
1159
1160fn find<T: Clone + Send + Sync + 'static>(parts: &Parts) -> Found<T> {
1161 match (
1162 parts.extensions.get::<T>(),
1163 parts.extensions.get::<LayerRan>(),
1164 ) {
1165 (Some(value), _) => Found::Present(value.clone()),
1166 (None, Some(LayerRan(layer))) => Found::Absent(layer.clone()),
1167 (None, None) => Found::NoLayer,
1168 }
1169}
1170
1171fn no_layer(parts: &Parts, extractor: &'static str) -> Response {
1174 count_request(
1175 Stage::Handler,
1176 Outcome::Rejected,
1177 Mechanism::None,
1178 REASON_MISCONFIGURED,
1179 );
1180 error!(
1181 path = %parts.uri.path(),
1182 extractor,
1183 auth.outcome = Outcome::Rejected.as_str(),
1184 auth.mechanism = Mechanism::None.as_str(),
1185 auth.reason = REASON_MISCONFIGURED,
1186 auth.status = 500u16,
1187 "Server misconfiguration: an authentication extractor ran on a route no AuthLayer \
1188 covers; refusing the request"
1189 );
1190 StatusCode::INTERNAL_SERVER_ERROR.into_response()
1191}
1192
1193fn refuse_absent(layer: &AuthLayer, parts: &Parts, wants: Wants) -> Response {
1198 let rejection = match (parts.extensions.get::<Credential>(), wants) {
1199 (Some(_), Wants::StaticToken) => TokenRejection::invalid(
1200 InvalidTokenKind::StaticTokenRequired,
1201 "a credential was accepted, but the handler requires a static token",
1202 ),
1203 (Some(_), _) => TokenRejection::invalid(
1204 InvalidTokenKind::OAuthTokenRequired,
1205 "a credential was accepted, but the handler requires an OAuth access token",
1206 ),
1207 (None, _) => TokenRejection::Missing,
1208 };
1209 layer.refuse_extraction(&rejection, parts, wants)
1210}
1211
1212#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1234impl<S: Send + Sync> FromRequestParts<S> for AuthorizedToken {
1235 type Rejection = Response;
1236
1237 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Response> {
1238 match find::<AuthorizedToken>(parts) {
1239 Found::Present(token) => Ok(token),
1240 Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::OAuthToken)),
1241 Found::NoLayer => Err(no_layer(parts, "AuthorizedToken")),
1242 }
1243 }
1244}
1245
1246#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1254impl<S: Send + Sync> OptionalFromRequestParts<S> for AuthorizedToken {
1255 type Rejection = Response;
1256
1257 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Option<Self>, Response> {
1258 match find::<AuthorizedToken>(parts) {
1259 Found::Present(token) => Ok(Some(token)),
1260 Found::Absent(_) => Ok(None),
1261 Found::NoLayer => Err(no_layer(parts, "Option<AuthorizedToken>")),
1262 }
1263 }
1264}
1265
1266#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1290impl<S: Send + Sync> FromRequestParts<S> for Credential {
1291 type Rejection = Response;
1292
1293 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Response> {
1294 match find::<Credential>(parts) {
1295 Found::Present(credential) => Ok(credential),
1296 Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::AnyCredential)),
1297 Found::NoLayer => Err(no_layer(parts, "Credential")),
1298 }
1299 }
1300}
1301
1302#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1308impl<S: Send + Sync> OptionalFromRequestParts<S> for Credential {
1309 type Rejection = Response;
1310
1311 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Option<Self>, Response> {
1312 match find::<Credential>(parts) {
1313 Found::Present(credential) => Ok(Some(credential)),
1314 Found::Absent(_) => Ok(None),
1315 Found::NoLayer => Err(no_layer(parts, "Option<Credential>")),
1316 }
1317 }
1318}
1319
1320#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1348impl<S: Send + Sync> FromRequestParts<S> for StaticTokenMatch {
1349 type Rejection = Response;
1350
1351 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Response> {
1352 match find::<StaticTokenMatch>(parts) {
1353 Found::Present(matched) => Ok(matched),
1354 Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::StaticToken)),
1355 Found::NoLayer => Err(no_layer(parts, "StaticTokenMatch")),
1356 }
1357 }
1358}
1359
1360#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1367impl<S: Send + Sync> OptionalFromRequestParts<S> for StaticTokenMatch {
1368 type Rejection = Response;
1369
1370 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Option<Self>, Response> {
1371 match find::<StaticTokenMatch>(parts) {
1372 Found::Present(matched) => Ok(Some(matched)),
1373 Found::Absent(_) => Ok(None),
1374 Found::NoLayer => Err(no_layer(parts, "Option<StaticTokenMatch>")),
1375 }
1376 }
1377}
1378
1379pub trait ScopeSet: Send + Sync + 'static {
1412 const SCOPES: &'static [&'static str];
1414}
1415
1416pub struct Scoped<S: ScopeSet> {
1460 token: AuthorizedToken,
1461 _scopes: std::marker::PhantomData<fn() -> S>,
1462}
1463
1464impl<S: ScopeSet> Scoped<S> {
1465 pub fn token(&self) -> &AuthorizedToken {
1467 &self.token
1468 }
1469
1470 pub fn into_token(self) -> AuthorizedToken {
1472 self.token
1473 }
1474}
1475
1476impl<S: ScopeSet> std::ops::Deref for Scoped<S> {
1477 type Target = AuthorizedToken;
1478
1479 fn deref(&self) -> &AuthorizedToken {
1480 &self.token
1481 }
1482}
1483
1484impl<S: ScopeSet> Clone for Scoped<S> {
1485 fn clone(&self) -> Self {
1486 Self {
1487 token: self.token.clone(),
1488 _scopes: std::marker::PhantomData,
1489 }
1490 }
1491}
1492
1493impl<S: ScopeSet> std::fmt::Debug for Scoped<S> {
1496 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1497 f.debug_struct("Scoped")
1498 .field("required", &S::SCOPES)
1499 .field("token", &self.token)
1500 .finish()
1501 }
1502}
1503
1504#[cfg_attr(docsrs, doc(cfg(feature = "axum")))]
1505impl<S: ScopeSet, St: Send + Sync> FromRequestParts<St> for Scoped<S> {
1506 type Rejection = Response;
1507
1508 async fn from_request_parts(parts: &mut Parts, _state: &St) -> Result<Self, Response> {
1509 let () = ValidScopeSet::<S>::CHECKED;
1512 let Ok(required) = crate::http_layer::checked_scopes(S::SCOPES.iter().copied()) else {
1515 count_request(
1516 Stage::Handler,
1517 Outcome::Rejected,
1518 Mechanism::None,
1519 REASON_MISCONFIGURED,
1520 );
1521 error!(
1522 path = %parts.uri.path(),
1523 extractor = std::any::type_name::<S>(),
1524 auth.outcome = Outcome::Rejected.as_str(),
1525 auth.mechanism = Mechanism::None.as_str(),
1526 auth.reason = REASON_MISCONFIGURED,
1527 auth.status = 500u16,
1528 "Server misconfiguration: a ScopeSet holds an entry that is not a valid scope, \
1529 which no token can carry; refusing the request"
1530 );
1531 return Err(StatusCode::INTERNAL_SERVER_ERROR.into_response());
1532 };
1533 match find::<Credential>(parts) {
1534 Found::Present(Credential::OAuth(token)) if token.require_scopes(S::SCOPES).is_ok() => {
1535 Ok(Self {
1536 token,
1537 _scopes: std::marker::PhantomData,
1538 })
1539 }
1540 Found::Present(_) => {
1546 let gate = match parts.extensions.get::<GateRan>() {
1547 Some(GateRan(Some(gate))) => Some(Arc::clone(gate)),
1548 _ => None,
1549 };
1550 let layer = parts.extensions.get::<LayerRan>().map(|l| &*l.0.inner);
1551 Err(match (gate, layer) {
1552 (Some(gate), Some(Mode::Enforce(enforce)))
1553 if Arc::ptr_eq(&gate, &enforce.gate) =>
1554 {
1555 enforce.refuse_scoped(parts, &required)
1556 }
1557 (None, Some(Mode::Enforce(enforce))) => enforce.refuse_scoped(parts, &required),
1558 _ => crate::http_layer::scope_refusal::<Body>(
1563 parts,
1564 &TokenRejection::InsufficientScope,
1565 &required,
1566 "Scoped",
1567 ),
1568 })
1569 }
1570 Found::Absent(layer) => Err(refuse_absent(&layer, parts, Wants::OAuthToken)),
1571 Found::NoLayer if parts.extensions.get::<GateRan>().is_some() => {
1575 Err(crate::http_layer::scope_refusal::<Body>(
1576 parts,
1577 &TokenRejection::Missing,
1578 &required,
1579 "Scoped",
1580 ))
1581 }
1582 Found::NoLayer => Err(no_layer(parts, "Scoped")),
1583 }
1584 }
1585}
1586
1587struct ValidScopeSet<S>(std::marker::PhantomData<S>);
1591
1592impl<S: ScopeSet> ValidScopeSet<S> {
1593 const CHECKED: () = assert!(
1594 crate::http_layer::all_scope_tokens(S::SCOPES),
1595 "every ScopeSet::SCOPES entry must be an RFC 6749 §3.3 scope-token: printable ASCII \
1596 with no space, '\"' or '\\'"
1597 );
1598}
1599
1600pub async fn require_auth(State(auth): State<AuthLayer>, request: Request, next: Next) -> Response {
1628 match auth.check(request).await {
1629 Ok(request) => next.run(request).await,
1630 Err(refusal) => refusal,
1631 }
1632}
1633
1634impl<S> tower_layer::Layer<S> for AuthLayer {
1635 type Service = AuthService<S>;
1636
1637 fn layer(&self, inner: S) -> Self::Service {
1638 AuthService {
1639 auth: self.clone(),
1640 inner,
1641 }
1642 }
1643}
1644
1645#[derive(Clone)]
1648pub struct AuthService<S> {
1649 auth: AuthLayer,
1650 inner: S,
1651}
1652
1653impl<S: std::fmt::Debug> std::fmt::Debug for AuthService<S> {
1654 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1655 f.debug_struct("AuthService")
1656 .field("auth", &self.auth)
1657 .field("inner", &self.inner)
1658 .finish()
1659 }
1660}
1661
1662impl<S> tower_service::Service<Request> for AuthService<S>
1663where
1664 S: tower_service::Service<Request, Response = Response> + Clone + Send + 'static,
1665 S::Future: Send + 'static,
1666{
1667 type Response = Response;
1668 type Error = S::Error;
1669 type Future = Pin<Box<dyn Future<Output = Result<Response, S::Error>> + Send + 'static>>;
1670
1671 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
1672 self.inner.poll_ready(cx)
1673 }
1674
1675 fn call(&mut self, request: Request) -> Self::Future {
1676 let clone = self.inner.clone();
1679 let mut inner = std::mem::replace(&mut self.inner, clone);
1680 let auth = self.auth.clone();
1681 Box::pin(async move {
1682 match auth.check(request).await {
1683 Ok(request) => inner.call(request).await,
1684 Err(refusal) => Ok(refusal),
1685 }
1686 })
1687 }
1688}
1689
1690pub fn metadata_router<S>(oauth: Option<Arc<OAuthValidator>>) -> Router<S>
1740where
1741 S: Clone + Send + Sync + 'static,
1742{
1743 async fn not_found() -> StatusCode {
1744 StatusCode::NOT_FOUND
1745 }
1746 let catch_all = format!("{PROTECTED_RESOURCE_METADATA_PREFIX}/{{*rest}}");
1747
1748 let Some(validator) = oauth else {
1749 return Router::new()
1751 .route(&catch_all, any(not_found))
1752 .route(PROTECTED_RESOURCE_METADATA_PREFIX, get(not_found));
1753 };
1754
1755 let serve = {
1756 let validator = Arc::clone(&validator);
1757 move || {
1758 let validator = Arc::clone(&validator);
1759 async move { Json(validator.metadata()).into_response() }
1760 }
1761 };
1762 let path: Arc<str> = validator.metadata_path().into();
1768 let suffix = move |request: Request| {
1769 let validator = Arc::clone(&validator);
1770 let path = Arc::clone(&path);
1771 async move {
1772 if request.uri().path() != &*path {
1773 return StatusCode::NOT_FOUND.into_response();
1774 }
1775 match *request.method() {
1776 Method::GET | Method::HEAD => Json(validator.metadata()).into_response(),
1777 _ => method_not_allowed(),
1778 }
1779 }
1780 };
1781 Router::new()
1782 .route(&catch_all, any(suffix))
1783 .route(PROTECTED_RESOURCE_METADATA_PREFIX, get(serve))
1784}
1785
1786fn method_not_allowed() -> Response {
1789 (
1790 StatusCode::METHOD_NOT_ALLOWED,
1791 [(http::header::ALLOW, HeaderValue::from_static("GET,HEAD"))],
1792 )
1793 .into_response()
1794}
1795
1796#[cfg(test)]
1797mod scope_tests;
1798
1799#[cfg(test)]
1800mod tests {
1801 use ::axum::Extension;
1802 use ::axum::middleware;
1803 use tower::ServiceExt;
1804
1805 use super::*;
1806 use crate::AuthorizedToken;
1807 use crate::testing;
1808
1809 const STATIC: &str = "secret";
1810
1811 fn validator(jwks_uri: &str) -> Arc<OAuthValidator> {
1812 Arc::new(OAuthValidator::new(&testing::resolved_config(jwks_uri)).unwrap())
1813 }
1814
1815 fn unreachable_validator() -> Arc<OAuthValidator> {
1816 validator("http://127.0.0.1:1/jwks")
1817 }
1818
1819 fn app(auth: AuthLayer) -> Router {
1820 Router::new()
1821 .route("/test", get(|| async { "ok" }))
1822 .route_layer(middleware::from_fn_with_state(auth, require_auth))
1823 }
1824
1825 fn wiki_app(static_token: Option<&str>, oauth: Option<Arc<OAuthValidator>>) -> Router {
1829 app(AuthLayer::builder()
1830 .optional_static_token(static_token.map(str::to_string))
1831 .optional_oauth(oauth)
1832 .static_challenge(None)
1833 .build()
1834 .unwrap())
1835 }
1836
1837 async fn send(app: &Router, headers: &[(&str, &str)]) -> Response {
1838 let mut req = Request::builder().uri("/test");
1839 for (name, value) in headers {
1840 req = req.header(*name, *value);
1841 }
1842 app.clone()
1843 .oneshot(req.body(Body::empty()).unwrap())
1844 .await
1845 .unwrap()
1846 }
1847
1848 async fn get_with_auth(app: &Router, header: Option<&str>) -> Response {
1849 match header {
1850 Some(h) => send(app, &[("authorization", h)]).await,
1851 None => send(app, &[]).await,
1852 }
1853 }
1854
1855 fn www_authenticate(resp: &Response) -> String {
1856 resp.headers()
1857 .get(WWW_AUTHENTICATE)
1858 .expect("a refusal with OAuth configured must carry WWW-Authenticate")
1859 .to_str()
1860 .unwrap()
1861 .to_string()
1862 }
1863
1864 async fn body_bytes(resp: Response) -> Vec<u8> {
1865 ::axum::body::to_bytes(resp.into_body(), 64 * 1024)
1866 .await
1867 .unwrap()
1868 .to_vec()
1869 }
1870
1871 fn unscoped_token() -> String {
1872 testing::mint(
1873 testing::KEY_A_PEM,
1874 testing::KID_A,
1875 &serde_json::json!({
1876 "iss": testing::ISSUER, "aud": testing::AUDIENCE,
1877 "exp": testing::now() + 3600, "scope": "openid profile",
1878 }),
1879 )
1880 }
1881
1882 fn expired_token() -> String {
1883 testing::mint(
1884 testing::KEY_A_PEM,
1885 testing::KID_A,
1886 &serde_json::json!({
1887 "iss": testing::ISSUER, "aud": testing::AUDIENCE,
1888 "exp": testing::now() - 3600, "scope": "mcp:read",
1889 }),
1890 )
1891 }
1892
1893 #[test]
1896 fn the_builder_refuses_to_build_a_pass_through() {
1897 assert_eq!(
1898 AuthLayer::builder().build().unwrap_err(),
1899 AuthLayerError::NoCredential
1900 );
1901 assert_eq!(
1902 AuthLayer::builder().static_token("").build().unwrap_err(),
1903 AuthLayerError::NoCredential
1904 );
1905 for blank in [" ", "\t", " \n "] {
1908 assert_eq!(
1909 AuthLayer::builder()
1910 .static_token(blank)
1911 .build()
1912 .unwrap_err(),
1913 AuthLayerError::NoCredential,
1914 "{blank:?}"
1915 );
1916 assert_eq!(
1917 AuthLayer::builder()
1918 .static_token(blank)
1919 .optional()
1920 .build()
1921 .unwrap_err(),
1922 AuthLayerError::NoCredential,
1923 "{blank:?}"
1924 );
1925 assert_eq!(
1926 AuthLayer::builder()
1927 .build_with_decision(StaticTokenDecision::StaticOnly(blank.into()))
1928 .unwrap_err(),
1929 AuthLayerError::NoCredential,
1930 "{blank:?}"
1931 );
1932 }
1933 assert_eq!(
1934 AuthLayer::builder()
1935 .optional_static_token(None)
1936 .optional_oauth(None)
1937 .build()
1938 .unwrap_err(),
1939 AuthLayerError::NoCredential
1940 );
1941 assert_eq!(
1942 AuthLayer::builder()
1943 .static_token(STATIC)
1944 .sources([])
1945 .build()
1946 .unwrap_err(),
1947 AuthLayerError::NoSources
1948 );
1949 let built = AuthLayer::builder().static_token(STATIC).build().unwrap();
1950 assert!(!built.allows_unauthenticated());
1951 assert!(built.oauth().is_none());
1952 }
1953
1954 #[tokio::test]
1955 async fn only_the_explicit_opt_out_passes_requests_through() {
1956 let layer = AuthLayer::allow_unauthenticated();
1957 assert!(layer.allows_unauthenticated());
1958 let app = Router::new()
1959 .route(
1960 "/test",
1961 get(
1962 |c: Option<Extension<Credential>>,
1963 t: Option<Extension<AuthorizedToken>>,
1964 headers: http::HeaderMap| async move {
1965 assert!(c.is_none() && t.is_none(), "a pass-through inserts nothing");
1966 for value in headers.get_all(http::header::AUTHORIZATION) {
1968 assert!(value.is_sensitive(), "Authorization must be sensitive");
1969 }
1970 "ok"
1971 },
1972 ),
1973 )
1974 .route_layer(middleware::from_fn_with_state(layer, require_auth));
1975 assert_eq!(get_with_auth(&app, None).await.status(), StatusCode::OK);
1976 assert_eq!(
1977 get_with_auth(&app, Some("Bearer anything")).await.status(),
1978 StatusCode::OK
1979 );
1980 }
1981
1982 #[test]
1983 fn debug_never_prints_the_static_token() {
1984 let layer = AuthLayer::builder()
1985 .static_token("hunter2")
1986 .build()
1987 .unwrap();
1988 let rendered = format!("{layer:?}");
1989 assert!(!rendered.contains("hunter2"), "{rendered}");
1990 let builder = AuthLayer::builder().static_token("hunter2");
1991 let rendered = format!("{builder:?}");
1992 assert!(!rendered.contains("hunter2"), "{rendered}");
1993 }
1994
1995 #[tokio::test]
1998 async fn static_token_only() {
1999 let app = wiki_app(Some(STATIC), None);
2000 for (header, status) in [
2001 (Some("Bearer secret"), StatusCode::OK),
2002 (Some("Bearer wrong-token"), StatusCode::UNAUTHORIZED),
2003 (None, StatusCode::UNAUTHORIZED),
2004 (Some("Basic c2VjcmV0LXRva2Vu"), StatusCode::UNAUTHORIZED),
2005 ] {
2006 let resp = get_with_auth(&app, header).await;
2007 assert_eq!(resp.status(), status, "{header:?}");
2008 assert!(resp.headers().get(WWW_AUTHENTICATE).is_none(), "{header:?}");
2010 }
2011 }
2012
2013 #[tokio::test]
2014 async fn a_static_only_401_carries_a_bearer_challenge_by_default() {
2015 let app = app(AuthLayer::builder().static_token(STATIC).build().unwrap());
2017 for header in [None, Some("Bearer wrong-token"), Some("Basic abc")] {
2018 let resp = get_with_auth(&app, header).await;
2019 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED, "{header:?}");
2020 assert_eq!(
2021 resp.headers()[WWW_AUTHENTICATE],
2022 DEFAULT_STATIC_CHALLENGE,
2023 "{header:?}"
2024 );
2025 }
2026 assert_eq!(
2027 get_with_auth(&app, Some("Bearer secret")).await.status(),
2028 StatusCode::OK
2029 );
2030 let app = super::tests::app(
2032 AuthLayer::from_decision(StaticTokenDecision::StaticOnly(STATIC.into()), None).unwrap(),
2033 );
2034 assert_eq!(
2035 get_with_auth(&app, None).await.headers()[WWW_AUTHENTICATE],
2036 DEFAULT_STATIC_CHALLENGE
2037 );
2038 let custom = HeaderValue::from_static("Bearer realm=\"my-api\"");
2040 let app = super::tests::app(
2041 AuthLayer::builder()
2042 .static_token(STATIC)
2043 .static_challenge(Some(custom.clone()))
2044 .build()
2045 .unwrap(),
2046 );
2047 assert_eq!(
2048 get_with_auth(&app, None).await.headers()[WWW_AUTHENTICATE],
2049 custom
2050 );
2051 }
2052
2053 #[tokio::test]
2054 async fn the_static_challenge_is_ignored_when_oauth_is_configured() {
2055 let v = unreachable_validator();
2056 let layer = AuthLayer::builder()
2057 .static_token(STATIC)
2058 .oauth(Arc::clone(&v))
2059 .static_challenge(Some(HeaderValue::from_static("Bearer realm=\"x\"")))
2060 .build()
2061 .unwrap();
2062 let resp = get_with_auth(&app(layer), None).await;
2063 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
2064 }
2065
2066 #[test]
2067 fn an_oauth_challenge_that_is_not_a_header_value_fails_the_build() {
2068 let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
2071 cfg.resource = "https://kb.example.test/m\ncp".into();
2072 let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
2073 assert_eq!(
2074 AuthLayer::builder().oauth(v).build().unwrap_err(),
2075 AuthLayerError::InvalidChallenge
2076 );
2077 let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
2078 cfg.required_scopes = vec!["a\u{1}b".into()];
2079 let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
2080 assert_eq!(
2081 AuthLayer::builder()
2082 .static_token(STATIC)
2083 .oauth(v)
2084 .build()
2085 .unwrap_err(),
2086 AuthLayerError::InvalidChallenge
2087 );
2088 }
2089
2090 #[tokio::test]
2091 async fn static_bearer_token_still_works_with_oauth_enabled() {
2092 let app = wiki_app(Some(STATIC), Some(unreachable_validator()));
2095 assert_eq!(
2096 get_with_auth(&app, Some("Bearer secret")).await.status(),
2097 StatusCode::OK
2098 );
2099 }
2100
2101 #[tokio::test]
2102 async fn an_oauth_token_is_accepted_alongside_the_static_token() {
2103 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2104 let app = wiki_app(Some(STATIC), Some(validator(&jwks.url)));
2105 let header = format!("Bearer {}", testing::valid_token());
2106 assert_eq!(
2107 get_with_auth(&app, Some(&header)).await.status(),
2108 StatusCode::OK
2109 );
2110 assert_eq!(
2111 get_with_auth(&app, Some("Bearer secret")).await.status(),
2112 StatusCode::OK
2113 );
2114 }
2115
2116 #[tokio::test]
2117 async fn a_missing_credential_gets_401_with_a_well_formed_challenge() {
2118 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2119 let app = wiki_app(Some(STATIC), Some(validator(&jwks.url)));
2120 for header in [None, Some("Bearer not-the-secret"), Some("Basic abc")] {
2124 let resp = get_with_auth(&app, header).await;
2125 assert_eq!(
2126 resp.status(),
2127 StatusCode::UNAUTHORIZED,
2128 "header: {header:?}"
2129 );
2130 assert_eq!(
2131 www_authenticate(&resp),
2132 "Bearer error=\"invalid_token\", \
2133 resource_metadata=\"https://kb.example.test\
2134 /.well-known/oauth-protected-resource/mcp\", \
2135 scope=\"mcp:read mcp:write\""
2136 );
2137 }
2138 }
2139
2140 #[tokio::test]
2141 async fn an_invalid_token_gets_401_and_an_insufficient_scope_token_gets_403() {
2142 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2143 let app = wiki_app(None, Some(validator(&jwks.url)));
2144
2145 let resp = get_with_auth(&app, Some(&format!("Bearer {}", expired_token()))).await;
2146 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
2147 assert!(www_authenticate(&resp).contains("error=\"invalid_token\""));
2148
2149 let resp = get_with_auth(&app, Some(&format!("Bearer {}", unscoped_token()))).await;
2150 assert_eq!(
2151 resp.status(),
2152 StatusCode::FORBIDDEN,
2153 "a valid token missing the scope is 403, not 401"
2154 );
2155 assert_eq!(
2156 www_authenticate(&resp),
2157 "Bearer error=\"insufficient_scope\", scope=\"mcp:read\", \
2158 resource_metadata=\"https://kb.example.test\
2159 /.well-known/oauth-protected-resource/mcp\""
2160 );
2161 }
2162
2163 #[tokio::test]
2164 async fn an_authelia_style_scp_token_is_accepted_through_the_middleware() {
2165 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2166 let app = wiki_app(None, Some(validator(&jwks.url)));
2167 let token = testing::mint_with(
2168 crate::Algorithm::RS256,
2169 Some(testing::KID_A),
2170 Some("at+jwt"),
2171 &serde_json::json!({
2172 "iss": testing::ISSUER, "aud": [testing::AUDIENCE],
2173 "exp": testing::now() + 3600, "nbf": testing::now(),
2174 "sub": "44726d41-0000-4000-8000-000000000000",
2175 "scp": ["mcp:read", "mcp:write"],
2176 }),
2177 );
2178 let resp = get_with_auth(&app, Some(&format!("Bearer {token}"))).await;
2179 assert_eq!(resp.status(), StatusCode::OK);
2180 }
2181
2182 #[tokio::test]
2183 async fn the_bearer_scheme_is_case_insensitive_for_both_credentials() {
2184 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2185 let app = wiki_app(Some(STATIC), Some(validator(&jwks.url)));
2186 for header in [
2187 "bearer secret".to_string(),
2188 "BEARER secret".to_string(),
2189 "Bearer secret ".to_string(),
2190 format!("bearer {}", testing::valid_token()),
2191 ] {
2192 assert_eq!(
2193 get_with_auth(&app, Some(&header)).await.status(),
2194 StatusCode::OK,
2195 "{header:.20}"
2196 );
2197 }
2198 for header in ["Basic secret", "Bearersecret", "secret", "Bearer\tsecret"] {
2199 assert_eq!(
2200 get_with_auth(&app, Some(header)).await.status(),
2201 StatusCode::UNAUTHORIZED,
2202 "{header}"
2203 );
2204 }
2205 }
2206
2207 #[test]
2208 fn bearer_credential_parsing() {
2209 assert_eq!(bearer_credential("Bearer abc"), "abc");
2210 assert_eq!(bearer_credential("bEaReR abc "), "abc");
2211 assert_eq!(bearer_credential("Bearer "), "");
2212 assert_eq!(bearer_credential("Bearer"), "");
2213 assert_eq!(bearer_credential("Basic abc"), "");
2214 assert_eq!(bearer_credential(""), "");
2215 }
2216
2217 mod wiki {
2224 use subtle::ConstantTimeEq;
2225
2226 use super::super::*;
2227
2228 #[derive(Clone)]
2229 pub(super) struct AuthState {
2230 pub(super) bearer_token: Option<String>,
2231 pub(super) oauth: Option<Arc<OAuthValidator>>,
2232 }
2233
2234 impl AuthState {
2235 fn challenge(&self, rejection: &TokenRejection) -> Option<String> {
2236 let oauth = self.oauth.as_ref()?;
2237 Some(match rejection {
2238 TokenRejection::InsufficientScope => oauth.insufficient_scope_challenge(),
2239 TokenRejection::Invalid(_) | TokenRejection::Missing => {
2240 oauth.invalid_token_challenge()
2241 }
2242 })
2243 }
2244 }
2245
2246 fn auth_rejection(auth: &AuthState, rejection: TokenRejection) -> Response {
2247 let status = match rejection {
2248 TokenRejection::InsufficientScope => StatusCode::FORBIDDEN,
2249 TokenRejection::Invalid(_) | TokenRejection::Missing => StatusCode::UNAUTHORIZED,
2250 };
2251 let mut response = Response::builder().status(status);
2252 if let Some(challenge) = auth.challenge(&rejection)
2253 && let Ok(value) = HeaderValue::from_str(&challenge)
2254 {
2255 response = response.header(WWW_AUTHENTICATE, value);
2256 }
2257 response
2258 .body(Body::empty())
2259 .expect("a status-and-header-only response is always constructible")
2260 }
2261
2262 pub(super) async fn bearer_auth(
2263 State(auth): State<AuthState>,
2264 headers: HeaderMap,
2265 request: Request,
2266 next: Next,
2267 ) -> Response {
2268 if auth.bearer_token.is_none() && auth.oauth.is_none() {
2269 return next.run(request).await;
2270 }
2271 let auth_header = headers
2272 .get("authorization")
2273 .and_then(|v| v.to_str().ok())
2274 .unwrap_or("");
2275 let token = bearer_credential(auth_header);
2276 if let Some(ref expected_token) = auth.bearer_token
2277 && !token.is_empty()
2278 && token.as_bytes().ct_eq(expected_token.as_bytes()).into()
2279 {
2280 return next.run(request).await;
2281 }
2282 let Some(ref oauth) = auth.oauth else {
2283 return auth_rejection(
2284 &auth,
2285 TokenRejection::Invalid("static token mismatch".into()),
2286 );
2287 };
2288 match oauth.validate(token).await {
2289 Ok(claims) => {
2290 let mut request = request;
2291 request.extensions_mut().insert(claims);
2292 next.run(request).await
2293 }
2294 Err(TokenRejection::Missing) => auth_rejection(&auth, TokenRejection::Missing),
2295 Err(rejection) => auth_rejection(&auth, rejection),
2296 }
2297 }
2298
2299 fn bearer_credential(header: &str) -> &str {
2300 match header.split_once(' ') {
2301 Some((scheme, token)) if scheme.eq_ignore_ascii_case("bearer") => token.trim(),
2302 _ => "",
2303 }
2304 }
2305 }
2306
2307 fn wiki_oracle_app(static_token: Option<&str>, oauth: Option<Arc<OAuthValidator>>) -> Router {
2308 let auth_state = wiki::AuthState {
2309 bearer_token: static_token.map(str::to_string),
2310 oauth,
2311 };
2312 Router::new()
2313 .route("/test", get(|| async { "ok" }))
2314 .route_layer(middleware::from_fn_with_state(
2315 auth_state,
2316 wiki::bearer_auth,
2317 ))
2318 }
2319
2320 async fn assert_same_response(actual: Response, expected: Response, what: &str) {
2321 assert_eq!(actual.status(), expected.status(), "{what}");
2322 assert_eq!(actual.version(), expected.version(), "{what}");
2323 assert_eq!(actual.headers(), expected.headers(), "{what}");
2324 assert_eq!(
2325 body_bytes(actual).await,
2326 body_bytes(expected).await,
2327 "{what}"
2328 );
2329 }
2330
2331 #[tokio::test]
2332 async fn default_responses_are_byte_identical_to_the_wiki_middleware() {
2333 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2334 let v = validator(&jwks.url);
2335 let valid = format!("Bearer {}", testing::valid_token());
2336 let lower = format!("bearer {}", testing::valid_token());
2337 let expired = format!("Bearer {}", expired_token());
2338 let unscoped = format!("Bearer {}", unscoped_token());
2339 let headers: Vec<Option<&str>> = vec![
2340 None,
2341 Some(""),
2342 Some("Bearer"),
2343 Some("Bearer "),
2344 Some("Bearer secret"),
2345 Some("bearer secret"),
2346 Some("Bearer secret "),
2347 Some("Bearer wrong"),
2348 Some("Bearersecret"),
2349 Some("Bearer\tsecret"),
2350 Some("Basic c2VjcmV0"),
2351 Some("secret"),
2352 Some(&valid),
2353 Some(&lower),
2354 Some(&expired),
2355 Some(&unscoped),
2356 ];
2357
2358 for (static_token, oauth) in [
2360 (Some(STATIC), None),
2361 (Some(STATIC), Some(Arc::clone(&v))),
2362 (None, Some(Arc::clone(&v))),
2363 ] {
2364 let ours = wiki_app(static_token, oauth.clone());
2365 let theirs = wiki_oracle_app(static_token, oauth.clone());
2366 for header in &headers {
2367 assert_same_response(
2368 get_with_auth(&ours, *header).await,
2369 get_with_auth(&theirs, *header).await,
2370 &format!(
2371 "static={static_token:?} oauth={} header={header:.30?}",
2372 oauth.is_some()
2373 ),
2374 )
2375 .await;
2376 }
2377 }
2378 }
2379
2380 #[tokio::test]
2383 async fn on_reject_shapes_the_body_but_not_the_status_or_challenge() {
2384 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2385 let v = validator(&jwks.url);
2386 let layer = AuthLayer::builder()
2387 .static_token(STATIC)
2388 .oauth(Arc::clone(&v))
2389 .on_reject(|cx: RejectContext<'_>| {
2390 assert_eq!(cx.request.uri.path(), "/test");
2391 let status = cx.status;
2392 let body = match cx.rejection {
2393 TokenRejection::InsufficientScope => r#"{"error":"insufficient_scope"}"#,
2394 _ => r#"{"error":"unauthorized"}"#,
2395 };
2396 Response::builder()
2397 .status(StatusCode::OK)
2400 .header("content-type", "application/json")
2401 .header("x-seen-status", status.as_str())
2402 .header(WWW_AUTHENTICATE, "Basic realm=\"nope\"")
2403 .header(WWW_AUTHENTICATE, "Bearer realm=\"also-nope\"")
2404 .body(Body::from(body))
2405 .unwrap()
2406 })
2407 .build()
2408 .unwrap();
2409 let app = app(layer);
2410
2411 let resp = get_with_auth(&app, None).await;
2412 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
2413 assert_eq!(resp.headers()["x-seen-status"], "401");
2414 assert_eq!(resp.headers()["content-type"], "application/json");
2415 assert_eq!(
2416 resp.headers().get_all(WWW_AUTHENTICATE).iter().count(),
2417 1,
2418 "the callback's challenges are replaced, not added to"
2419 );
2420 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
2421 assert_eq!(body_bytes(resp).await, br#"{"error":"unauthorized"}"#);
2422
2423 let resp = get_with_auth(&app, Some(&format!("Bearer {}", unscoped_token()))).await;
2424 assert_eq!(resp.status(), StatusCode::FORBIDDEN);
2425 assert_eq!(resp.headers()["x-seen-status"], "403");
2426 assert_eq!(www_authenticate(&resp), v.insufficient_scope_challenge());
2427 assert_eq!(body_bytes(resp).await, br#"{"error":"insufficient_scope"}"#);
2428 }
2429
2430 #[tokio::test]
2431 async fn on_reject_without_oauth_or_a_static_challenge_keeps_its_own_headers() {
2432 let builder = || {
2433 AuthLayer::builder()
2434 .static_token(STATIC)
2435 .on_reject(|cx: RejectContext<'_>| {
2436 Response::builder()
2437 .status(cx.status)
2438 .header(WWW_AUTHENTICATE, "ApiKey")
2439 .body(Body::from("nope"))
2440 .unwrap()
2441 })
2442 };
2443 let layer = builder().static_challenge(None).build().unwrap();
2444 let resp = get_with_auth(&app(layer), Some("Bearer wrong")).await;
2445 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
2446 assert_eq!(resp.headers()[WWW_AUTHENTICATE], "ApiKey");
2447 assert_eq!(body_bytes(resp).await, b"nope");
2448
2449 let resp = get_with_auth(&app(builder().build().unwrap()), Some("Bearer wrong")).await;
2452 assert_eq!(resp.headers().get_all(WWW_AUTHENTICATE).iter().count(), 1);
2453 assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
2454 assert_eq!(body_bytes(resp).await, b"nope");
2455 }
2456
2457 #[tokio::test]
2458 async fn the_presented_credential_never_reaches_a_debug_rendering() {
2459 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2460 let v = validator(&jwks.url);
2461 let token = unscoped_token(); let seen: Arc<std::sync::Mutex<Vec<String>>> = Arc::default();
2463 let log = Arc::clone(&seen);
2464 let layer = AuthLayer::builder()
2465 .oauth(v)
2466 .sources([
2467 CredentialSource::authorization_bearer(),
2468 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
2469 ])
2470 .on_reject(move |cx: RejectContext<'_>| {
2471 log.lock().unwrap().push(format!("{cx:?}"));
2472 log.lock().unwrap().push(format!("{:?}", cx.request));
2473 Response::new(Body::empty())
2474 })
2475 .build()
2476 .unwrap();
2477 let bearer = format!("Bearer {token}");
2478 let resp = send(
2479 &app(layer),
2480 &[
2481 ("authorization", bearer.as_str()),
2482 ("x-api-key", "raw-api-key-value"),
2483 ("accept", "application/json"),
2484 ],
2485 )
2486 .await;
2487 assert_eq!(resp.status(), StatusCode::FORBIDDEN);
2488 let seen = seen.lock().unwrap();
2489 assert_eq!(seen.len(), 2);
2490 for rendered in seen.iter() {
2491 assert!(!rendered.contains(&token), "{rendered}");
2492 assert!(!rendered.contains("raw-api-key-value"), "{rendered}");
2493 }
2494 assert!(seen[0].contains("InsufficientScope"), "{}", seen[0]);
2496 assert!(seen[0].contains("403"), "{}", seen[0]);
2497 assert!(seen[0].contains("authorization"), "{}", seen[0]);
2498 assert!(!seen[0].contains("application/json"), "{}", seen[0]);
2499 }
2500
2501 #[tokio::test]
2502 async fn the_inner_service_sees_the_credential_headers_marked_sensitive() {
2503 let layer = AuthLayer::builder().static_token(STATIC).build().unwrap();
2504 let app = Router::new()
2505 .route(
2506 "/test",
2507 get(|headers: HeaderMap| async move {
2508 assert!(headers["authorization"].is_sensitive());
2509 assert!(!format!("{headers:?}").contains(STATIC));
2510 assert!(!headers["accept"].is_sensitive());
2511 "ok"
2512 }),
2513 )
2514 .route_layer(layer);
2515 let resp = send(
2516 &app,
2517 &[("authorization", "Bearer secret"), ("accept", "text/plain")],
2518 )
2519 .await;
2520 assert_eq!(resp.status(), StatusCode::OK);
2521 }
2522
2523 fn multi_source_app(v: Arc<OAuthValidator>) -> Router {
2526 app(AuthLayer::builder()
2527 .static_token(STATIC)
2528 .oauth(v)
2529 .sources([
2530 CredentialSource::authorization_bearer(),
2531 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
2532 ])
2533 .build()
2534 .unwrap())
2535 }
2536
2537 #[tokio::test]
2538 async fn a_bad_authorization_header_does_not_mask_a_good_raw_header() {
2539 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2540 let app = multi_source_app(validator(&jwks.url));
2541 let foreign = format!("Bearer {}", expired_token());
2542 for authorization in [foreign.as_str(), "Bearer garbage", "Basic abc"] {
2543 let resp = send(
2544 &app,
2545 &[("authorization", authorization), ("x-api-key", STATIC)],
2546 )
2547 .await;
2548 assert_eq!(resp.status(), StatusCode::OK, "{authorization:.30}");
2549 }
2550 }
2551
2552 #[tokio::test]
2553 async fn a_bad_raw_header_does_not_mask_a_good_authorization_header() {
2554 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2555 let app = multi_source_app(validator(&jwks.url));
2556 let valid = format!("Bearer {}", testing::valid_token());
2557 for authorization in ["Bearer secret", valid.as_str()] {
2558 let resp = send(
2559 &app,
2560 &[("authorization", authorization), ("x-api-key", "garbage")],
2561 )
2562 .await;
2563 assert_eq!(resp.status(), StatusCode::OK, "{authorization:.30}");
2564 }
2565 let resp = send(&app, &[("x-api-key", &testing::valid_token())]).await;
2567 assert_eq!(resp.status(), StatusCode::OK);
2568 }
2569
2570 #[tokio::test]
2571 async fn multi_source_refusals() {
2572 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2573 let v = validator(&jwks.url);
2574 let app = multi_source_app(Arc::clone(&v));
2575
2576 let resp = send(
2577 &app,
2578 &[
2579 ("authorization", "Bearer garbage"),
2580 ("x-api-key", "also-garbage"),
2581 ],
2582 )
2583 .await;
2584 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
2585 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
2586
2587 let unscoped = unscoped_token();
2589 let resp = send(
2590 &app,
2591 &[
2592 ("authorization", "Bearer garbage"),
2593 ("x-api-key", &unscoped),
2594 ],
2595 )
2596 .await;
2597 assert_eq!(resp.status(), StatusCode::FORBIDDEN);
2598 assert_eq!(www_authenticate(&resp), v.insufficient_scope_challenge());
2599
2600 let resp = send(&app, &[("x-api-key", "Bearer secret")]).await;
2603 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
2604
2605 let resp = send(&app, &[]).await;
2606 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
2607 }
2608
2609 #[tokio::test]
2610 async fn a_raw_only_layer_ignores_the_authorization_header() {
2611 let layer = AuthLayer::builder()
2612 .static_token(STATIC)
2613 .sources([CredentialSource::Raw(HeaderName::from_static("x-api-key"))])
2614 .build()
2615 .unwrap();
2616 let app = app(layer);
2617 assert_eq!(
2618 send(&app, &[("authorization", "Bearer secret")])
2619 .await
2620 .status(),
2621 StatusCode::UNAUTHORIZED
2622 );
2623 assert_eq!(
2624 send(&app, &[("x-api-key", STATIC)]).await.status(),
2625 StatusCode::OK
2626 );
2627 }
2628
2629 #[tokio::test]
2632 async fn the_credential_and_oauth_token_are_inserted_into_extensions() {
2633 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2634 let layer = AuthLayer::builder()
2635 .static_token(STATIC)
2636 .oauth(validator(&jwks.url))
2637 .build()
2638 .unwrap();
2639 let app = Router::new()
2640 .route(
2641 "/test",
2642 get(
2643 |Extension(credential): Extension<Credential>,
2644 token: Option<Extension<AuthorizedToken>>| async move {
2645 match (credential, token) {
2646 (Credential::StaticToken, None) => "static".to_string(),
2647 (Credential::OAuth(c), Some(Extension(t))) => {
2648 assert_eq!(c, t);
2649 format!(
2650 "oauth {} {}",
2651 t.subject.as_deref().unwrap_or_default(),
2652 t.has_scope("mcp:write")
2653 )
2654 }
2655 other => panic!("inconsistent extensions: {other:?}"),
2656 }
2657 },
2658 ),
2659 )
2660 .route_layer(middleware::from_fn_with_state(layer, require_auth));
2661
2662 let resp = get_with_auth(&app, Some("Bearer secret")).await;
2663 assert_eq!(resp.status(), StatusCode::OK);
2664 assert_eq!(body_bytes(resp).await, b"static");
2665
2666 let header = format!("Bearer {}", testing::valid_token());
2667 let resp = get_with_auth(&app, Some(&header)).await;
2668 assert_eq!(resp.status(), StatusCode::OK);
2669 assert_eq!(body_bytes(resp).await, b"oauth user-1 true");
2670 }
2671
2672 #[tokio::test]
2678 async fn the_layer_behaves_exactly_like_the_middleware_function() {
2679 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2680 let layer = AuthLayer::builder()
2681 .static_token(STATIC)
2682 .oauth(validator(&jwks.url))
2683 .on_reject(|cx: RejectContext<'_>| Response::new(Body::from(cx.status.to_string())))
2684 .build()
2685 .unwrap();
2686 let handler = get(|c: Option<Extension<Credential>>| async move {
2687 match c {
2688 Some(Extension(Credential::StaticToken)) => "static",
2689 Some(Extension(Credential::OAuth(_))) => "oauth",
2690 None => "none",
2691 }
2692 });
2693 let via_fn = Router::new()
2694 .route("/test", handler.clone())
2695 .route_layer(middleware::from_fn_with_state(layer.clone(), require_auth));
2696 let via_route_layer = Router::new()
2697 .route("/test", handler.clone())
2698 .route_layer(layer.clone());
2699 let via_layer = Router::new().route("/test", handler).layer(layer);
2700
2701 let valid = format!("Bearer {}", testing::valid_token());
2702 let unscoped = format!("Bearer {}", unscoped_token());
2703 for header in [
2704 None,
2705 Some("Bearer secret"),
2706 Some("Bearer wrong"),
2707 Some(valid.as_str()),
2708 Some(unscoped.as_str()),
2709 ] {
2710 let expected = get_with_auth(&via_fn, header).await;
2711 let (status, challenge) = (
2712 expected.status(),
2713 expected.headers().get(WWW_AUTHENTICATE).cloned(),
2714 );
2715 let expected_body = body_bytes(expected).await;
2716 for app in [&via_route_layer, &via_layer] {
2717 let resp = get_with_auth(app, header).await;
2718 assert_eq!(resp.status(), status, "{header:?}");
2719 assert_eq!(
2720 resp.headers().get(WWW_AUTHENTICATE),
2721 challenge.as_ref(),
2722 "{header:?}"
2723 );
2724 assert_eq!(body_bytes(resp).await, expected_body, "{header:?}");
2725 }
2726 }
2727 }
2728
2729 #[tokio::test]
2732 async fn from_decision_maps_every_decision_and_refuses_a_mismatch() {
2733 use StaticTokenDecision::*;
2734
2735 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
2736 let v = validator(&jwks.url);
2737 let valid = format!("Bearer {}", testing::valid_token());
2738 let status = |layer: AuthLayer, header: &'static str| async move {
2739 get_with_auth(&app(layer), Some(header)).await.status()
2740 };
2741 let valid: &'static str = Box::leak(valid.into_boxed_str());
2742
2743 let open = AuthLayer::from_decision(Unauthenticated, None).unwrap();
2745 assert!(open.allows_unauthenticated());
2746
2747 let layer = AuthLayer::from_decision(StaticOnly(STATIC.into()), None).unwrap();
2749 assert!(!layer.allows_unauthenticated());
2750 assert_eq!(status(layer.clone(), "Bearer secret").await, StatusCode::OK);
2751 assert_eq!(status(layer, valid).await, StatusCode::UNAUTHORIZED);
2752
2753 let layer =
2755 AuthLayer::from_decision(StaticAndOAuth(STATIC.into()), Some(Arc::clone(&v))).unwrap();
2756 assert_eq!(status(layer.clone(), "Bearer secret").await, StatusCode::OK);
2757 assert_eq!(status(layer, valid).await, StatusCode::OK);
2758
2759 for decision in [OAuthOnly, StaticIgnored] {
2761 let layer = AuthLayer::from_decision(decision, Some(Arc::clone(&v))).unwrap();
2762 assert_eq!(
2763 status(layer.clone(), "Bearer secret").await,
2764 StatusCode::UNAUTHORIZED
2765 );
2766 assert_eq!(status(layer, valid).await, StatusCode::OK);
2767 }
2768
2769 for decision in [StaticAndOAuth(STATIC.into()), OAuthOnly, StaticIgnored] {
2771 assert_eq!(
2772 AuthLayer::from_decision(decision, None).unwrap_err(),
2773 AuthLayerError::DecisionNeedsOAuth
2774 );
2775 }
2776 for decision in [StaticOnly(STATIC.into()), Unauthenticated] {
2777 assert_eq!(
2778 AuthLayer::from_decision(decision, Some(Arc::clone(&v))).unwrap_err(),
2779 AuthLayerError::DecisionWithoutOAuth
2780 );
2781 }
2782 }
2783
2784 #[tokio::test]
2785 async fn build_with_decision_keeps_the_builders_sources_and_replaces_its_token() {
2786 let layer = AuthLayer::builder()
2787 .static_token("builder-token")
2788 .sources([CredentialSource::Raw(HeaderName::from_static("x-api-key"))])
2789 .build_with_decision(StaticTokenDecision::StaticOnly(STATIC.into()))
2790 .unwrap();
2791 let app = app(layer);
2792 assert_eq!(
2793 send(&app, &[("x-api-key", STATIC)]).await.status(),
2794 StatusCode::OK
2795 );
2796 assert_eq!(
2797 send(&app, &[("x-api-key", "builder-token")]).await.status(),
2798 StatusCode::UNAUTHORIZED
2799 );
2800 assert_eq!(
2801 send(&app, &[("authorization", "Bearer secret")])
2802 .await
2803 .status(),
2804 StatusCode::UNAUTHORIZED
2805 );
2806 }
2807
2808 async fn get_path(app: &Router, path: &str) -> Response {
2811 app.clone()
2812 .oneshot(Request::builder().uri(path).body(Body::empty()).unwrap())
2813 .await
2814 .unwrap()
2815 }
2816
2817 const MCP_METADATA_PATH: &str = "/.well-known/oauth-protected-resource/mcp";
2818
2819 #[tokio::test]
2820 async fn metadata_routes_for_a_resource_with_a_path() {
2821 let v = unreachable_validator();
2822 assert_eq!(v.metadata_path(), MCP_METADATA_PATH);
2823 let app: Router = metadata_router(Some(Arc::clone(&v)));
2824 for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
2825 let resp = get_path(&app, path).await;
2826 assert_eq!(resp.status(), StatusCode::OK, "{path}");
2827 assert_eq!(resp.headers()["content-type"], "application/json", "{path}");
2828 let body = body_bytes(resp).await;
2829 assert_eq!(body, serde_json::to_vec(&v.metadata()).unwrap(), "{path}");
2831 let doc: serde_json::Value = serde_json::from_slice(&body).unwrap();
2832 assert_eq!(doc["resource"], testing::RESOURCE);
2833 assert_eq!(doc["authorization_servers"][0], testing::ISSUER);
2834 assert_eq!(
2835 doc["scopes_supported"],
2836 serde_json::json!(["mcp:read", "mcp:write"])
2837 );
2838 assert_eq!(
2839 doc["bearer_methods_supported"],
2840 serde_json::json!(["header"])
2841 );
2842 }
2843 for path in [
2844 "/.well-known/oauth-protected-resource/other",
2845 "/.well-known/oauth-protected-resource/mcp/deeper",
2846 ] {
2847 assert_eq!(
2848 get_path(&app, path).await.status(),
2849 StatusCode::NOT_FOUND,
2850 "{path}"
2851 );
2852 }
2853 }
2854
2855 async fn post_path(app: &Router, path: &str) -> Response {
2856 app.clone()
2857 .oneshot(
2858 Request::builder()
2859 .method("POST")
2860 .uri(path)
2861 .body(Body::empty())
2862 .unwrap(),
2863 )
2864 .await
2865 .unwrap()
2866 }
2867
2868 #[tokio::test]
2869 async fn metadata_routes_answer_other_methods_by_whether_the_path_is_served() {
2870 let app: Router = metadata_router(Some(unreachable_validator()));
2873 for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
2874 assert_eq!(
2875 post_path(&app, path).await.status(),
2876 StatusCode::METHOD_NOT_ALLOWED,
2877 "{path}"
2878 );
2879 }
2880 for path in [
2881 "/.well-known/oauth-protected-resource/other",
2882 "/.well-known/oauth-protected-resource/mcp/deeper",
2883 ] {
2884 let resp = post_path(&app, path).await;
2885 assert_eq!(resp.status(), StatusCode::NOT_FOUND, "{path}");
2886 assert!(!resp.headers().contains_key("allow"), "{path}");
2887 }
2888 let app: Router = metadata_router(None);
2889 assert_eq!(
2890 post_path(&app, MCP_METADATA_PATH).await.status(),
2891 StatusCode::NOT_FOUND
2892 );
2893 }
2894
2895 #[tokio::test]
2896 async fn metadata_routes_for_a_root_resource() {
2897 let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
2898 cfg.resource = "https://api.example.test/".to_string();
2899 let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
2900 assert_eq!(v.metadata_path(), PROTECTED_RESOURCE_METADATA_PREFIX);
2901 let app: Router = metadata_router(Some(v));
2902 let resp = get_path(&app, PROTECTED_RESOURCE_METADATA_PREFIX).await;
2903 assert_eq!(resp.status(), StatusCode::OK);
2904 let doc: serde_json::Value = serde_json::from_slice(&body_bytes(resp).await).unwrap();
2905 assert_eq!(doc["resource"], "https://api.example.test/");
2906 assert_eq!(
2907 get_path(&app, MCP_METADATA_PATH).await.status(),
2908 StatusCode::NOT_FOUND
2909 );
2910 }
2911
2912 #[tokio::test]
2913 async fn metadata_routes_with_a_nested_and_a_brace_bearing_path() {
2914 for (resource, path) in [
2915 (
2916 "https://api.example.test/v1/things",
2917 "/.well-known/oauth-protected-resource/v1/things",
2918 ),
2919 (
2921 "https://api.example.test/v1/",
2922 "/.well-known/oauth-protected-resource/v1/",
2923 ),
2924 (
2925 "https://api.example.test/a{b}",
2926 "/.well-known/oauth-protected-resource/a{b}",
2927 ),
2928 ] {
2929 let mut cfg = testing::resolved_config("http://127.0.0.1:1/jwks");
2930 cfg.resource = resource.to_string();
2931 let v = Arc::new(OAuthValidator::new(&cfg).unwrap());
2932 assert_eq!(v.metadata_path(), path);
2933 let app: Router = metadata_router(Some(v));
2934 let resp = get_path(&app, path).await;
2935 assert_eq!(resp.status(), StatusCode::OK, "{resource}");
2936 }
2937 }
2938
2939 #[tokio::test]
2943 async fn metadata_routes_for_a_path_axum_would_read_as_route_syntax() {
2944 for (resource, path) in [
2945 (
2946 "https://api.example.test/a/:id",
2947 "/.well-known/oauth-protected-resource/a/:id",
2948 ),
2949 (
2950 "https://api.example.test/a/*x",
2951 "/.well-known/oauth-protected-resource/a/*x",
2952 ),
2953 (
2954 "https://api.example.test/:id",
2955 "/.well-known/oauth-protected-resource/:id",
2956 ),
2957 (
2958 "https://api.example.test/*",
2959 "/.well-known/oauth-protected-resource/*",
2960 ),
2961 (
2962 "https://api.example.test/{x}",
2963 "/.well-known/oauth-protected-resource/{x}",
2964 ),
2965 (
2966 "https://api.example.test/%7Bx%7D",
2967 "/.well-known/oauth-protected-resource/%7Bx%7D",
2968 ),
2969 (
2970 "https://api.example.test//mcp",
2971 "/.well-known/oauth-protected-resource//mcp",
2972 ),
2973 ] {
2974 let resolved = crate::OAuthConfig {
2975 enabled: true,
2976 issuer: testing::ISSUER.to_string(),
2977 jwks_uri: Some("http://127.0.0.1:1/jwks".to_string()),
2978 audience: testing::AUDIENCE.to_string(),
2979 resource: resource.to_string(),
2980 required_scope: Some("mcp:read".to_string()),
2981 ..crate::OAuthConfig::default()
2982 }
2983 .resolve(crate::KeyNaming::Dotted("oauth"))
2984 .unwrap_or_else(|e| panic!("{resource}: {e}"))
2985 .expect("enabled");
2986 let v = Arc::new(OAuthValidator::new(&resolved).unwrap());
2987 assert_eq!(v.metadata_path(), path, "{resource}");
2988 let app: Router = metadata_router(Some(Arc::clone(&v)));
2989 for served in [path, PROTECTED_RESOURCE_METADATA_PREFIX] {
2990 let resp = get_path(&app, served).await;
2991 assert_eq!(resp.status(), StatusCode::OK, "{resource} {served}");
2992 assert_eq!(
2993 body_bytes(resp).await,
2994 serde_json::to_vec(&v.metadata()).unwrap(),
2995 "{resource} {served}"
2996 );
2997 }
2998 for other in [
3000 "/.well-known/oauth-protected-resource/a/other",
3001 "/.well-known/oauth-protected-resource/a/:id/x",
3002 "/.well-known/oauth-protected-resource/other",
3003 "/.well-known/oauth-protected-resource/mcp",
3004 ] {
3005 assert_eq!(
3006 get_path(&app, other).await.status(),
3007 StatusCode::NOT_FOUND,
3008 "{resource} {other}"
3009 );
3010 }
3011 }
3012 }
3013
3014 #[tokio::test]
3018 async fn metadata_path_answers_methods_exactly_like_the_bare_prefix_route() {
3019 let app: Router = metadata_router(Some(unreachable_validator()));
3020 let send = |method: &'static str, path: &'static str| {
3021 app.clone().oneshot(
3022 Request::builder()
3023 .method(method)
3024 .uri(path)
3025 .body(Body::empty())
3026 .unwrap(),
3027 )
3028 };
3029 for method in ["GET", "HEAD", "POST", "PUT", "DELETE", "OPTIONS", "PATCH"] {
3030 let bare = send(method, PROTECTED_RESOURCE_METADATA_PREFIX)
3031 .await
3032 .unwrap();
3033 let suffixed = send(method, MCP_METADATA_PATH).await.unwrap();
3034 assert_eq!(bare.status(), suffixed.status(), "{method}");
3035 assert_eq!(bare.headers(), suffixed.headers(), "{method}");
3036 let (bare, suffixed) = (body_bytes(bare).await, body_bytes(suffixed).await);
3037 assert_eq!(bare, suffixed, "{method}");
3038 if method == "HEAD" {
3039 assert!(suffixed.is_empty());
3040 }
3041 }
3042 }
3043
3044 #[tokio::test]
3045 async fn metadata_routes_404_when_oauth_is_not_configured() {
3046 let app: Router = metadata_router(None).fallback(|| async { "spa shell" });
3048 for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
3049 let resp = get_path(&app, path).await;
3050 assert_eq!(resp.status(), StatusCode::NOT_FOUND, "{path}");
3051 assert!(body_bytes(resp).await.is_empty(), "{path}");
3052 }
3053 }
3054
3055 #[tokio::test]
3056 async fn metadata_routes_are_reachable_outside_the_auth_layer() {
3057 let v = unreachable_validator();
3058 let layer = AuthLayer::builder().oauth(Arc::clone(&v)).build().unwrap();
3059 let app = Router::new()
3060 .route("/mcp", get(|| async { "ok" }))
3061 .route_layer(middleware::from_fn_with_state(layer, require_auth))
3062 .merge(metadata_router(Some(v)));
3063 assert_eq!(
3064 get_path(&app, "/mcp").await.status(),
3065 StatusCode::UNAUTHORIZED
3066 );
3067 for path in [MCP_METADATA_PATH, PROTECTED_RESOURCE_METADATA_PREFIX] {
3068 assert_eq!(
3069 get_path(&app, path).await.status(),
3070 StatusCode::OK,
3071 "{path}"
3072 );
3073 }
3074 }
3075
3076 #[tokio::test]
3077 async fn metadata_router_works_with_app_state() {
3078 #[derive(Clone)]
3079 struct AppState;
3080 let app: Router = Router::new()
3081 .route("/x", get(|State(_): State<AppState>| async { "x" }))
3082 .merge(metadata_router(Some(unreachable_validator())))
3083 .with_state(AppState);
3084 assert_eq!(
3085 get_path(&app, MCP_METADATA_PATH).await.status(),
3086 StatusCode::OK
3087 );
3088 }
3089
3090 async fn observed(resp: Response) -> (StatusCode, Vec<(String, Vec<u8>)>, Vec<u8>) {
3095 let status = resp.status();
3096 let headers = resp
3097 .headers()
3098 .iter()
3099 .map(|(k, v)| (k.to_string(), v.as_bytes().to_vec()))
3100 .collect();
3101 (status, headers, body_bytes(resp).await)
3102 }
3103
3104 fn json_reject(cx: RejectContext<'_>) -> Response {
3105 Response::new(Body::from(format!("refused {}", cx.status.as_u16())))
3106 }
3107
3108 fn optional_extractor_app(
3111 layer: AuthLayer,
3112 runs: Arc<std::sync::atomic::AtomicUsize>,
3113 ) -> Router {
3114 let handler = move |credential: Option<Credential>, token: Option<AuthorizedToken>| {
3115 let runs = Arc::clone(&runs);
3116 async move {
3117 runs.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
3118 match (credential, token) {
3119 (None, None) => "none".to_string(),
3120 (Some(Credential::StaticToken), None) => "static".to_string(),
3121 (Some(Credential::OAuth(c)), Some(t)) => {
3122 assert_eq!(c, t);
3123 format!("oauth {}", t.subject.as_deref().unwrap_or_default())
3124 }
3125 other => panic!("inconsistent extraction: {other:?}"),
3126 }
3127 }
3128 };
3129 Router::new()
3130 .route("/test", get(handler))
3131 .route_layer(layer)
3132 }
3133
3134 type Headers<'a> = Vec<(&'a str, &'a [u8])>;
3136
3137 async fn send_raw(app: &Router, headers: &[(&str, &[u8])]) -> Response {
3138 let mut req = Request::builder().uri("/test");
3139 for (name, value) in headers {
3140 req = req.header(*name, HeaderValue::from_bytes(value).unwrap());
3141 }
3142 app.clone()
3143 .oneshot(req.body(Body::empty()).unwrap())
3144 .await
3145 .unwrap()
3146 }
3147
3148 #[tokio::test]
3152 async fn an_oauth_extractor_refusing_a_static_token_names_its_kind() {
3153 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3154 let v = validator(&jwks.url);
3155 let layer = AuthLayer::builder()
3156 .static_token(STATIC)
3157 .oauth(Arc::clone(&v))
3158 .on_reject(|cx: RejectContext<'_>| {
3159 let label = match cx.rejection {
3160 TokenRejection::Invalid(invalid) => invalid.kind().as_str(),
3161 _ => "not invalid",
3162 };
3163 Response::new(Body::from(label))
3164 })
3165 .build()
3166 .unwrap();
3167 let app = Router::new()
3168 .route("/test", get(|_: AuthorizedToken| async { "ok" }))
3169 .route("/static", get(|_: StaticTokenMatch| async { "ok" }))
3170 .route_layer(layer);
3171 let resp = get_with_auth(&app, Some("Bearer secret")).await;
3172 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3173 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
3174 assert_eq!(body_bytes(resp).await, b"oauth_token_required");
3175 let resp = get_with_auth(&app, Some("Bearer wrong")).await;
3176 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3177 assert_eq!(body_bytes(resp).await, b"not_jwt");
3178 let resp = app
3180 .clone()
3181 .oneshot(
3182 Request::builder()
3183 .uri("/static")
3184 .header(
3185 "authorization",
3186 format!("Bearer {}", testing::valid_token()),
3187 )
3188 .body(Body::empty())
3189 .unwrap(),
3190 )
3191 .await
3192 .unwrap();
3193 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3194 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
3195 assert_eq!(body_bytes(resp).await, b"static_token_required");
3196 }
3197
3198 #[tokio::test]
3199 async fn the_extractors_read_a_valid_token_and_the_static_token() {
3200 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3201 let v = validator(&jwks.url);
3202 let layer = AuthLayer::builder()
3203 .static_token(STATIC)
3204 .oauth(Arc::clone(&v))
3205 .build()
3206 .unwrap();
3207 let app = Router::new()
3208 .route(
3209 "/credential",
3210 get(|credential: Credential| async move {
3211 match credential {
3212 Credential::StaticToken => "static".to_string(),
3213 Credential::OAuth(t) => format!("oauth {}", t.subject.unwrap_or_default()),
3214 }
3215 }),
3216 )
3217 .route(
3218 "/token",
3219 get(|token: AuthorizedToken| async move {
3220 format!(
3221 "{} {}",
3222 token.subject.as_deref().unwrap_or_default(),
3223 token.has_scope("mcp:read")
3224 )
3225 }),
3226 )
3227 .route_layer(layer);
3228 let get_at = |path: &'static str, header: String| {
3229 let app = app.clone();
3230 async move {
3231 app.oneshot(
3232 Request::builder()
3233 .uri(path)
3234 .header("authorization", header)
3235 .body(Body::empty())
3236 .unwrap(),
3237 )
3238 .await
3239 .unwrap()
3240 }
3241 };
3242
3243 let valid = format!("Bearer {}", testing::valid_token());
3244 let resp = get_at("/credential", valid.clone()).await;
3245 assert_eq!(resp.status(), StatusCode::OK);
3246 assert_eq!(body_bytes(resp).await, b"oauth user-1");
3247 let resp = get_at("/token", valid).await;
3248 assert_eq!(resp.status(), StatusCode::OK);
3249 assert_eq!(body_bytes(resp).await, b"user-1 true");
3250
3251 let resp = get_at("/credential", "Bearer secret".into()).await;
3252 assert_eq!(resp.status(), StatusCode::OK);
3253 assert_eq!(body_bytes(resp).await, b"static");
3254 let resp = get_at("/token", "Bearer secret".into()).await;
3257 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3258 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
3259 }
3260
3261 #[tokio::test]
3262 async fn an_extractor_outside_every_layer_fails_closed_with_500() {
3263 let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
3264 let count = |runs: &Arc<std::sync::atomic::AtomicUsize>| {
3265 runs.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
3266 };
3267 let (r1, r2, r3, r4) = (
3268 Arc::clone(&runs),
3269 Arc::clone(&runs),
3270 Arc::clone(&runs),
3271 Arc::clone(&runs),
3272 );
3273 let app = Router::new()
3274 .route(
3275 "/credential",
3276 get(move |_: Credential| async move { count(&r1) }),
3277 )
3278 .route(
3279 "/token",
3280 get(move |_: AuthorizedToken| async move { count(&r2) }),
3281 )
3282 .route(
3283 "/opt-credential",
3284 get(move |_: Option<Credential>| async move { count(&r3) }),
3285 )
3286 .route(
3287 "/opt-token",
3288 get(move |_: Option<AuthorizedToken>| async move { count(&r4) }),
3289 );
3290 for path in ["/credential", "/token", "/opt-credential", "/opt-token"] {
3291 for header in [None, Some("Bearer secret")] {
3292 let mut req = Request::builder().uri(path);
3293 if let Some(h) = header {
3294 req = req.header("authorization", h);
3295 }
3296 let resp = app
3297 .clone()
3298 .oneshot(req.body(Body::empty()).unwrap())
3299 .await
3300 .unwrap();
3301 assert_eq!(
3302 resp.status(),
3303 StatusCode::INTERNAL_SERVER_ERROR,
3304 "{path} {header:?}"
3305 );
3306 assert!(resp.headers().get(WWW_AUTHENTICATE).is_none());
3307 assert!(body_bytes(resp).await.is_empty(), "{path}");
3308 }
3309 }
3310 assert_eq!(runs.load(std::sync::atomic::Ordering::SeqCst), 0);
3311 }
3312
3313 #[test]
3314 fn optional_still_needs_a_credential_to_build() {
3315 assert_eq!(
3316 AuthLayer::builder().optional().build().unwrap_err(),
3317 AuthLayerError::NoCredential
3318 );
3319 assert_eq!(
3320 AuthLayer::builder()
3321 .static_token("")
3322 .optional()
3323 .build()
3324 .unwrap_err(),
3325 AuthLayerError::NoCredential
3326 );
3327 assert_eq!(
3328 AuthLayer::builder()
3329 .static_token(STATIC)
3330 .optional()
3331 .sources([])
3332 .build()
3333 .unwrap_err(),
3334 AuthLayerError::NoSources
3335 );
3336 let layer = AuthLayer::builder()
3337 .static_token("hunter2")
3338 .optional()
3339 .build()
3340 .unwrap();
3341 assert!(!layer.allows_unauthenticated());
3342 let rendered = format!("{layer:?}");
3343 assert!(!rendered.contains("hunter2") && rendered.contains("optional: true"));
3344 }
3345
3346 #[tokio::test]
3350 async fn an_optional_layer_passes_only_a_request_with_no_credential() {
3351 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3352 let v = validator(&jwks.url);
3353 let builder = || {
3354 AuthLayer::builder()
3355 .static_token(STATIC)
3356 .oauth(Arc::clone(&v))
3357 .sources([
3358 CredentialSource::authorization_bearer(),
3359 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
3360 ])
3361 .on_reject(json_reject)
3362 };
3363 let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
3364 let optional =
3365 optional_extractor_app(builder().optional().build().unwrap(), Arc::clone(&runs));
3366 let strict =
3367 optional_extractor_app(builder().build().unwrap(), Arc::new(Default::default()));
3368 let ran = || runs.load(std::sync::atomic::Ordering::SeqCst);
3369
3370 let blanks: &[&[(&str, &[u8])]] = &[
3372 &[],
3373 &[("authorization", b"")],
3374 &[("authorization", b"Bearer ")],
3375 &[("authorization", b"bearer ")],
3376 &[("authorization", b"Basic c2VjcmV0")],
3378 &[("x-api-key", b" ")],
3379 &[("authorization", b"Bearer "), ("x-api-key", b"")],
3380 &[("authorization", b"Bearer "), ("authorization", b" ")],
3381 ];
3382 for headers in blanks {
3383 let before = ran();
3384 let resp = send_raw(&optional, headers).await;
3385 assert_eq!(resp.status(), StatusCode::OK, "{headers:?}");
3386 assert!(resp.headers().get(WWW_AUTHENTICATE).is_none());
3387 assert_eq!(body_bytes(resp).await, b"none", "{headers:?}");
3388 assert_eq!(ran(), before + 1);
3389 let resp = send_raw(&strict, headers).await;
3391 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED, "{headers:?}");
3392 assert_eq!(www_authenticate(&resp), v.invalid_token_challenge());
3393 }
3394
3395 let invalid = format!("Bearer {}", testing::valid_token().replace('.', "x."));
3398 let expired = format!("Bearer {}", expired_token());
3399 let unscoped = format!("Bearer {}", unscoped_token());
3400 let refused: Vec<(Headers<'_>, StatusCode, String)> = vec![
3401 (
3402 vec![("authorization", b"Bearer not-a-jwt")],
3403 StatusCode::UNAUTHORIZED,
3404 v.invalid_token_challenge(),
3405 ),
3406 (
3407 vec![("authorization", invalid.as_bytes())],
3408 StatusCode::UNAUTHORIZED,
3409 v.invalid_token_challenge(),
3410 ),
3411 (
3412 vec![("authorization", expired.as_bytes())],
3413 StatusCode::UNAUTHORIZED,
3414 v.invalid_token_challenge(),
3415 ),
3416 (
3417 vec![("x-api-key", b"wrong-key")],
3418 StatusCode::UNAUTHORIZED,
3419 v.invalid_token_challenge(),
3420 ),
3421 (
3423 vec![("authorization", b"Bearer "), ("x-api-key", b"wrong-key")],
3424 StatusCode::UNAUTHORIZED,
3425 v.invalid_token_challenge(),
3426 ),
3427 (
3430 vec![
3431 ("authorization", b"Bearer "),
3432 ("authorization", b"Bearer junk"),
3433 ],
3434 StatusCode::UNAUTHORIZED,
3435 v.invalid_token_challenge(),
3436 ),
3437 (
3439 vec![("authorization", b"Bearer \xff")],
3440 StatusCode::UNAUTHORIZED,
3441 v.invalid_token_challenge(),
3442 ),
3443 (
3444 vec![("authorization", unscoped.as_bytes())],
3445 StatusCode::FORBIDDEN,
3446 v.insufficient_scope_challenge(),
3447 ),
3448 ];
3449 for (headers, status, challenge) in &refused {
3450 let before = ran();
3451 let resp = send_raw(&optional, headers).await;
3452 assert_eq!(resp.status(), *status, "{headers:?}");
3453 assert_eq!(&www_authenticate(&resp), challenge, "{headers:?}");
3454 let got = observed(resp).await;
3455 assert_eq!(got.2, format!("refused {}", status.as_u16()).as_bytes());
3456 assert_eq!(
3457 got,
3458 observed(send_raw(&strict, headers).await).await,
3459 "{headers:?}"
3460 );
3461 assert_eq!(ran(), before, "the handler ran for {headers:?}");
3462 }
3463
3464 let valid = format!("Bearer {}", testing::valid_token());
3466 let resp = send_raw(&optional, &[("authorization", valid.as_bytes())]).await;
3467 assert_eq!(resp.status(), StatusCode::OK);
3468 assert_eq!(body_bytes(resp).await, b"oauth user-1");
3469 let resp = send_raw(&optional, &[("x-api-key", STATIC.as_bytes())]).await;
3470 assert_eq!(resp.status(), StatusCode::OK);
3471 assert_eq!(body_bytes(resp).await, b"static");
3472 }
3473
3474 #[tokio::test]
3478 async fn a_required_extractor_behind_an_optional_layer_gets_the_layers_own_refusal() {
3479 let v = unreachable_validator();
3480 let builder = || {
3481 AuthLayer::builder()
3482 .static_token(STATIC)
3483 .oauth(Arc::clone(&v))
3484 .on_reject(json_reject)
3485 };
3486 let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
3487 let make = |layer: AuthLayer| {
3488 let (r1, r2) = (Arc::clone(&runs), Arc::clone(&runs));
3489 Router::new()
3490 .route(
3491 "/test",
3492 get(move |_: Credential| async move {
3493 r1.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
3494 }),
3495 )
3496 .route(
3497 "/token",
3498 get(move |_: AuthorizedToken| async move {
3499 r2.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
3500 }),
3501 )
3502 .route_layer(layer)
3503 };
3504 let optional = make(builder().optional().build().unwrap());
3505 let strict = make(builder().build().unwrap());
3506 for path in ["/test", "/token"] {
3507 let request = || Request::builder().uri(path).body(Body::empty()).unwrap();
3508 let got = observed(optional.clone().oneshot(request()).await.unwrap()).await;
3509 let want = observed(strict.clone().oneshot(request()).await.unwrap()).await;
3510 assert_eq!(got.0, StatusCode::UNAUTHORIZED, "{path}");
3511 assert_eq!(got, want, "{path}");
3512 assert!(
3513 got.1.iter().any(|(k, val)| k == "www-authenticate"
3514 && val == v.invalid_token_challenge().as_bytes()),
3515 "{path}"
3516 );
3517 }
3518 assert_eq!(runs.load(std::sync::atomic::Ordering::SeqCst), 0);
3519
3520 let optional = make(
3522 AuthLayer::builder()
3523 .static_token(STATIC)
3524 .optional()
3525 .build()
3526 .unwrap(),
3527 );
3528 let resp = optional
3529 .oneshot(Request::builder().uri("/test").body(Body::empty()).unwrap())
3530 .await
3531 .unwrap();
3532 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3533 assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
3534 }
3535
3536 #[tokio::test]
3537 async fn the_extractors_under_allow_unauthenticated() {
3538 let app = Router::new()
3539 .route(
3540 "/test",
3541 get(
3542 |c: Option<Credential>, t: Option<AuthorizedToken>| async move {
3543 assert!(c.is_none() && t.is_none());
3544 "none"
3545 },
3546 ),
3547 )
3548 .route("/required", get(|_: Credential| async { "unreachable" }))
3549 .route_layer(AuthLayer::allow_unauthenticated());
3550 let resp = get_with_auth(&app, Some("Bearer anything")).await;
3551 assert_eq!(resp.status(), StatusCode::OK);
3552 assert_eq!(body_bytes(resp).await, b"none");
3553 let resp = app
3554 .oneshot(
3555 Request::builder()
3556 .uri("/required")
3557 .body(Body::empty())
3558 .unwrap(),
3559 )
3560 .await
3561 .unwrap();
3562 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3563 assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
3564 }
3565
3566 #[tokio::test]
3569 async fn a_non_optional_layer_with_extractors_matches_extension_handlers() {
3570 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3571 let layer = AuthLayer::builder()
3572 .static_token(STATIC)
3573 .oauth(validator(&jwks.url))
3574 .build()
3575 .unwrap();
3576 let via_extension = Router::new()
3577 .route(
3578 "/test",
3579 get(|Extension(c): Extension<Credential>| async move { format!("{c:?}") }),
3580 )
3581 .route_layer(layer.clone());
3582 let via_extractor = Router::new()
3583 .route(
3584 "/test",
3585 get(|c: Credential| async move { format!("{c:?}") }),
3586 )
3587 .route_layer(layer);
3588 let valid = format!("Bearer {}", testing::valid_token());
3589 let unscoped = format!("Bearer {}", unscoped_token());
3590 let expired = format!("Bearer {}", expired_token());
3591 for header in [
3592 None,
3593 Some("Bearer "),
3594 Some("Bearer secret"),
3595 Some("Bearer wrong"),
3596 Some(valid.as_str()),
3597 Some(unscoped.as_str()),
3598 Some(expired.as_str()),
3599 ] {
3600 assert_eq!(
3601 observed(get_with_auth(&via_extractor, header).await).await,
3602 observed(get_with_auth(&via_extension, header).await).await,
3603 "{header:?}"
3604 );
3605 }
3606 }
3607
3608 fn claims_with(extra: serde_json::Value) -> serde_json::Value {
3609 let mut claims = serde_json::json!({
3610 "iss": testing::ISSUER, "aud": testing::AUDIENCE,
3611 "exp": testing::now() + 3600, "scope": "mcp:read mcp:write", "sub": "user-1",
3612 });
3613 for (k, v) in extra.as_object().unwrap() {
3614 claims[k] = v.clone();
3615 }
3616 claims
3617 }
3618
3619 #[test]
3620 fn names_a_token_only_for_dpop_and_tab_separated_bearer() {
3621 for value in [
3622 "DPoP x",
3623 "dpop x",
3624 "DPoP",
3625 "Bearer\tx",
3626 "bearer\t x",
3627 " Bearer x",
3628 "\tBEARER\tx",
3629 ] {
3630 assert!(names_a_token(value), "{value:?}");
3631 }
3632 for value in [
3633 "",
3634 "Bearer",
3635 "Bearer ",
3636 "Bearer\t",
3637 "Bearer \t ",
3638 "Basic x",
3639 "x",
3640 ] {
3641 assert!(!names_a_token(value), "{value:?}");
3642 }
3643 assert_eq!(bearer_credential("Bearer\tx"), "");
3645 assert_eq!(bearer_credential("DPoP x"), "");
3646 }
3647
3648 #[tokio::test]
3653 async fn an_optional_layer_refuses_every_presented_shape_like_the_strict_layer() {
3654 use base64::Engine;
3655 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
3656
3657 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3658 let v = validator(&jwks.url);
3659 let builder = || {
3660 AuthLayer::builder()
3661 .static_token(STATIC)
3662 .oauth(Arc::clone(&v))
3663 .sources([
3664 CredentialSource::authorization_bearer(),
3665 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
3666 ])
3667 .on_reject(json_reject)
3668 };
3669 let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
3670 let optional =
3671 optional_extractor_app(builder().optional().build().unwrap(), Arc::clone(&runs));
3672 let strict =
3673 optional_extractor_app(builder().build().unwrap(), Arc::new(Default::default()));
3674
3675 let forged = testing::mint(
3676 testing::KEY_B_PEM,
3677 testing::KID_A,
3678 &claims_with(serde_json::json!({})),
3679 );
3680 let wrong_aud = testing::mint(
3681 testing::KEY_A_PEM,
3682 testing::KID_A,
3683 &claims_with(serde_json::json!({ "aud": "some-other-client" })),
3684 );
3685 let cnf = testing::mint(
3686 testing::KEY_A_PEM,
3687 testing::KID_A,
3688 &claims_with(serde_json::json!({ "cnf": { "jkt": "abc" } })),
3689 );
3690 let valid = testing::valid_token();
3691 let crit = {
3692 let mut parts: Vec<String> = valid.split('.').map(str::to_string).collect();
3693 parts[0] =
3694 URL_SAFE_NO_PAD.encode(br#"{"alg":"RS256","kid":"test-key-a","crit":["exp"]}"#);
3695 parts.join(".")
3696 };
3697
3698 let cases: Vec<(&str, Vec<(&str, String)>)> = vec![
3699 (
3700 "forged",
3701 vec![("authorization", format!("Bearer {forged}"))],
3702 ),
3703 (
3704 "wrong aud",
3705 vec![("authorization", format!("Bearer {wrong_aud}"))],
3706 ),
3707 (
3708 "cnf bearer",
3709 vec![("authorization", format!("Bearer {cnf}"))],
3710 ),
3711 ("crit", vec![("authorization", format!("Bearer {crit}"))]),
3712 (
3713 "BEARER forged",
3714 vec![("authorization", format!("BEARER {forged}"))],
3715 ),
3716 (
3717 "bearer bad",
3718 vec![("authorization", "bearer not-a-jwt".to_string())],
3719 ),
3720 (
3721 "Bearer<TAB>forged",
3722 vec![("authorization", format!("Bearer\t{forged}"))],
3723 ),
3724 (
3725 "Bearer<TAB>valid",
3726 vec![("authorization", format!("Bearer\t{valid}"))],
3727 ),
3728 (
3729 " Bearer forged",
3730 vec![("authorization", format!(" Bearer {forged}"))],
3731 ),
3732 ("DPoP cnf", vec![("authorization", format!("DPoP {cnf}"))]),
3733 (
3734 "DPoP forged",
3735 vec![("authorization", format!("DPoP {forged}"))],
3736 ),
3737 (
3738 "Basic, then Bearer forged",
3739 vec![
3740 ("authorization", "Basic x".to_string()),
3741 ("authorization", format!("Bearer {forged}")),
3742 ],
3743 ),
3744 (
3745 "blank Bearer, then Bearer forged",
3746 vec![
3747 ("authorization", "Bearer ".to_string()),
3748 ("authorization", format!("Bearer {forged}")),
3749 ],
3750 ),
3751 ];
3752 for (name, headers) in &cases {
3753 let headers: Vec<(&str, &[u8])> =
3754 headers.iter().map(|(n, v)| (*n, v.as_bytes())).collect();
3755 let before = runs.load(std::sync::atomic::Ordering::SeqCst);
3756 let got = observed(send_raw(&optional, &headers).await).await;
3757 let want = observed(send_raw(&strict, &headers).await).await;
3758 assert_eq!(got.0, StatusCode::UNAUTHORIZED, "{name}");
3759 assert!(
3760 got.1.iter().any(|(k, val)| k == "www-authenticate"
3761 && val == v.invalid_token_challenge().as_bytes()),
3762 "{name}"
3763 );
3764 assert_eq!(got, want, "{name}");
3765 assert_eq!(
3766 runs.load(std::sync::atomic::Ordering::SeqCst),
3767 before,
3768 "the handler ran for {name}"
3769 );
3770 }
3771 }
3772
3773 #[tokio::test]
3776 async fn an_inner_optional_layer_does_not_leak_an_outer_layers_credential() {
3777 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3778 let outer = AuthLayer::builder()
3779 .oauth(validator(&jwks.url))
3780 .build()
3781 .unwrap();
3782 let inner = AuthLayer::builder()
3783 .static_token("inner-key")
3784 .sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
3785 .optional()
3786 .build()
3787 .unwrap();
3788 let runs = Arc::new(std::sync::atomic::AtomicUsize::new(0));
3789 let app = optional_extractor_app(inner, Arc::clone(&runs)).layer(outer);
3790 let valid = format!("Bearer {}", testing::valid_token());
3791
3792 let resp = send_raw(&app, &[("authorization", valid.as_bytes())]).await;
3793 assert_eq!(resp.status(), StatusCode::OK);
3794 assert_eq!(body_bytes(resp).await, b"none");
3795
3796 let resp = send_raw(
3798 &app,
3799 &[
3800 ("authorization", valid.as_bytes()),
3801 ("x-inner", b"inner-key"),
3802 ],
3803 )
3804 .await;
3805 assert_eq!(resp.status(), StatusCode::OK);
3806 assert_eq!(body_bytes(resp).await, b"static");
3807 assert_eq!(runs.load(std::sync::atomic::Ordering::SeqCst), 2);
3808 }
3809
3810 #[tokio::test]
3813 async fn strict_nested_layers_accumulate_extensions_as_documented() {
3814 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
3815 let outer = AuthLayer::builder()
3816 .oauth(validator(&jwks.url))
3817 .build()
3818 .unwrap();
3819 let inner = AuthLayer::builder()
3820 .static_token("inner-key")
3821 .sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
3822 .build()
3823 .unwrap();
3824 let app = Router::new()
3825 .route(
3826 "/test",
3827 get(|c: Credential, t: AuthorizedToken| async move {
3828 format!(
3829 "{} {}",
3830 matches!(c, Credential::StaticToken),
3831 t.subject.unwrap_or_default()
3832 )
3833 }),
3834 )
3835 .route_layer(inner)
3836 .layer(outer);
3837 let valid = format!("Bearer {}", testing::valid_token());
3838
3839 let resp = send_raw(
3840 &app,
3841 &[
3842 ("authorization", valid.as_bytes()),
3843 ("x-inner", b"inner-key"),
3844 ],
3845 )
3846 .await;
3847 assert_eq!(resp.status(), StatusCode::OK);
3848 assert_eq!(body_bytes(resp).await, b"true user-1");
3849
3850 let resp = send_raw(&app, &[("authorization", valid.as_bytes())]).await;
3852 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
3853 assert_eq!(resp.headers()[WWW_AUTHENTICATE], DEFAULT_STATIC_CHALLENGE);
3854 }
3855}
3856
3857#[cfg(test)]
3861mod shared_refusal_tests {
3862 use ::tower::{ServiceExt, service_fn};
3863
3864 use super::*;
3865 use crate::http_layer::HttpAuthLayer;
3866 use crate::testing;
3867 use crate::{Refusal, refusal, refusal_with_static_challenge};
3868
3869 const STATIC: &str = "secret";
3870 const CUSTOM: &str = "ApiKey realm=\"example\"";
3871
3872 fn validator(jwks_uri: &str) -> Arc<OAuthValidator> {
3873 Arc::new(OAuthValidator::new(&testing::resolved_config(jwks_uri)).unwrap())
3874 }
3875
3876 #[derive(Clone, Copy, Debug)]
3879 enum Static {
3880 Unset,
3881 Custom,
3882 Off,
3883 }
3884
3885 fn what_the_axum_layer_sends(
3886 oauth: Option<Arc<OAuthValidator>>,
3887 setting: Static,
3888 rejection: &TokenRejection,
3889 ) -> (u16, Vec<String>) {
3890 let mut builder = AuthLayer::builder()
3891 .static_token(STATIC)
3892 .optional_oauth(oauth)
3893 .on_reject(|_| {
3896 (StatusCode::IM_A_TEAPOT, [(WWW_AUTHENTICATE, "Callback x")]).into_response()
3897 });
3898 builder = match setting {
3899 Static::Unset => builder,
3900 Static::Custom => builder.static_challenge(Some(HeaderValue::from_static(CUSTOM))),
3901 Static::Off => builder.static_challenge(None),
3902 };
3903 let layer = builder.build().unwrap();
3904 let Mode::Enforce(enforce) = &*layer.inner else {
3905 unreachable!("an enforcing layer was built")
3906 };
3907 let (parts, ()) = Request::builder()
3908 .uri("/test")
3909 .body(())
3910 .unwrap()
3911 .into_parts();
3912 let response = enforce.reject(rejection, &parts);
3913 let challenges = response
3914 .headers()
3915 .get_all(WWW_AUTHENTICATE)
3916 .iter()
3917 .map(|v| v.to_str().unwrap().to_string())
3918 .collect();
3919 (response.status().as_u16(), challenges)
3920 }
3921
3922 #[test]
3923 fn refusal_gives_exactly_what_enforce_reject_gives() {
3924 let v = validator("http://127.0.0.1:1/jwks");
3925 let rejections = [
3926 TokenRejection::Missing,
3927 TokenRejection::Invalid("any reason".into()),
3928 TokenRejection::InsufficientScope,
3929 ];
3930 let mut rows = 0;
3931 for oauth in [None, Some(Arc::clone(&v))] {
3932 for setting in [Static::Unset, Static::Custom, Static::Off] {
3933 for rejection in &rejections {
3934 let static_str = match setting {
3935 Static::Unset => Some(DEFAULT_STATIC_CHALLENGE),
3936 Static::Custom => Some(CUSTOM),
3937 Static::Off => None,
3938 };
3939 let ours =
3940 refusal_with_static_challenge(rejection, oauth.as_deref(), static_str);
3941 if let Static::Unset = setting {
3942 assert_eq!(ours, refusal(rejection, oauth.as_deref()));
3943 }
3944 let (status, challenges) =
3945 what_the_axum_layer_sends(oauth.clone(), setting, rejection);
3946 let context = format!("oauth={} {setting:?} {rejection:?}", oauth.is_some());
3947 assert_eq!(ours.status, status, "{context}");
3948 let expected = match &ours {
3951 Refusal {
3952 www_authenticate: Some(c),
3953 ..
3954 } => vec![c.clone()],
3955 _ => vec!["Callback x".to_string()],
3956 };
3957 assert_eq!(challenges, expected, "{context}");
3958 rows += 1;
3959 }
3960 }
3961 }
3962 assert_eq!(rows, 18);
3963 }
3964
3965 #[test]
3966 fn the_status_and_challenge_are_what_rfc_6750_asks_for() {
3967 let v = validator("http://127.0.0.1:1/jwks");
3968 let r = refusal(&TokenRejection::Missing, Some(&v));
3969 assert_eq!(r.status, 401);
3970 assert_eq!(r.www_authenticate, Some(v.invalid_token_challenge()));
3971 let r = refusal(&TokenRejection::Invalid("x".into()), Some(&v));
3972 assert_eq!(r.status, 401);
3973 assert_eq!(r.www_authenticate, Some(v.invalid_token_challenge()));
3974 let r = refusal(&TokenRejection::InsufficientScope, Some(&v));
3975 assert_eq!(r.status, 403);
3976 assert_eq!(r.www_authenticate, Some(v.insufficient_scope_challenge()));
3977 assert_eq!(
3979 refusal_with_static_challenge(&TokenRejection::Missing, Some(&v), None),
3980 refusal(&TokenRejection::Missing, Some(&v))
3981 );
3982 }
3983
3984 type Seen = (u16, Vec<String>, String);
3986
3987 async fn through_axum(layer: AuthLayer, headers: &[(&str, &str)]) -> Seen {
3988 let app: Router = Router::new()
3989 .route(
3990 "/test",
3991 get(|credential: Option<Credential>| async move { format!("{credential:?}") }),
3992 )
3993 .route_layer(layer);
3994 let mut request = Request::builder().uri("/test");
3995 for (name, value) in headers {
3996 request = request.header(*name, *value);
3997 }
3998 let response = app
3999 .oneshot(request.body(Body::empty()).unwrap())
4000 .await
4001 .unwrap();
4002 let status = response.status().as_u16();
4003 let challenges = response
4004 .headers()
4005 .get_all(WWW_AUTHENTICATE)
4006 .iter()
4007 .map(|v| v.to_str().unwrap().to_string())
4008 .collect();
4009 let body = ::axum::body::to_bytes(response.into_body(), 64 * 1024)
4010 .await
4011 .unwrap();
4012 (
4013 status,
4014 challenges,
4015 String::from_utf8(body.to_vec()).unwrap(),
4016 )
4017 }
4018
4019 async fn through_tower(layer: HttpAuthLayer, headers: &[(&str, &str)]) -> Seen {
4020 let service = tower_layer::Layer::layer(
4021 &layer,
4022 service_fn(|request: http::Request<String>| async move {
4023 let credential = request.extensions().get::<Credential>().cloned();
4024 Ok::<_, std::convert::Infallible>(http::Response::new(format!("{credential:?}")))
4025 }),
4026 );
4027 let mut request = http::Request::builder().uri("/test");
4028 for (name, value) in headers {
4029 request = request.header(*name, *value);
4030 }
4031 let response = service
4032 .oneshot(request.body(String::new()).unwrap())
4033 .await
4034 .unwrap();
4035 let challenges = response
4036 .headers()
4037 .get_all(WWW_AUTHENTICATE)
4038 .iter()
4039 .map(|v| v.to_str().unwrap().to_string())
4040 .collect();
4041 (response.status().as_u16(), challenges, response.into_body())
4042 }
4043
4044 #[tokio::test]
4045 async fn the_axum_and_tower_layers_answer_every_request_identically() {
4046 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
4047 let v = validator(&jwks.url);
4048 let mint = |scope: &str, exp_offset: i64| {
4049 testing::mint(
4050 testing::KEY_A_PEM,
4051 testing::KID_A,
4052 &serde_json::json!({
4053 "iss": testing::ISSUER, "aud": testing::AUDIENCE, "sub": "user-1",
4054 "exp": testing::now() as i64 + exp_offset, "scope": scope,
4055 }),
4056 )
4057 };
4058 let valid = format!("Bearer {}", mint("mcp:read", 3600));
4059 let expired = format!("Bearer {}", mint("mcp:read", -3600));
4060 let unscoped = format!("Bearer {}", mint("openid", 3600));
4061 let requests: Vec<Vec<(&str, &str)>> = vec![
4062 vec![],
4063 vec![("authorization", valid.as_str())],
4064 vec![("authorization", expired.as_str())],
4065 vec![("authorization", unscoped.as_str())],
4066 vec![("authorization", "Bearer not-a-jwt")],
4067 vec![("authorization", "Bearer secret")],
4068 vec![("authorization", "bearer secret")],
4069 vec![("authorization", "Bearer ")],
4070 vec![("authorization", "Basic abc")],
4071 vec![("authorization", "DPoP abc")],
4072 vec![("x-api-key", "secret")],
4073 vec![("authorization", "Bearer wrong"), ("x-api-key", "secret")],
4074 ];
4075
4076 type Config = (
4078 Option<&'static str>,
4079 bool,
4080 Option<Option<&'static str>>,
4081 bool,
4082 bool,
4083 );
4084 let configs: [Config; 7] = [
4085 (None, true, None, false, false),
4086 (Some(STATIC), true, None, false, true),
4087 (Some(STATIC), false, None, false, false),
4088 (Some(STATIC), false, Some(None), false, false),
4089 (Some(STATIC), false, Some(Some(CUSTOM)), false, true),
4090 (Some(STATIC), true, None, true, false),
4091 (Some(STATIC), false, None, true, true),
4092 ];
4093 for (static_token, with_oauth, static_challenge, optional, api_key) in configs {
4094 let oauth = with_oauth.then(|| Arc::clone(&v));
4095 let sources = if api_key {
4096 vec![
4097 CredentialSource::authorization_bearer(),
4098 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
4099 ]
4100 } else {
4101 vec![CredentialSource::authorization_bearer()]
4102 };
4103 let challenge = static_challenge.map(|c| c.map(HeaderValue::from_static));
4104 let mut axum_builder = AuthLayer::builder()
4105 .optional_static_token(static_token.map(str::to_string))
4106 .optional_oauth(oauth.clone())
4107 .sources(sources.clone());
4108 let mut tower_builder = HttpAuthLayer::builder()
4109 .optional_static_token(static_token.map(str::to_string))
4110 .optional_oauth(oauth.clone())
4111 .sources(sources);
4112 if let Some(c) = challenge {
4113 axum_builder = axum_builder.static_challenge(c.clone());
4114 tower_builder = tower_builder.static_challenge(c);
4115 }
4116 if optional {
4117 axum_builder = axum_builder.optional();
4118 tower_builder = tower_builder.optional();
4119 }
4120 let axum_layer = axum_builder.build().unwrap();
4121 let tower_layer = tower_builder.build().unwrap();
4122 for headers in &requests {
4123 let a = through_axum(axum_layer.clone(), headers).await;
4124 let t = through_tower(tower_layer.clone(), headers).await;
4125 assert_eq!(
4126 a, t,
4127 "config {static_token:?} oauth={with_oauth} {static_challenge:?} \
4128 optional={optional} api_key={api_key}, request {headers:?}"
4129 );
4130 }
4131 }
4132 }
4133
4134 type SeenFull = (u16, Vec<String>, Option<String>, String);
4136
4137 fn seen_parts(headers: &HeaderMap, status: u16, body: String) -> SeenFull {
4138 let challenges = headers
4139 .get_all(WWW_AUTHENTICATE)
4140 .iter()
4141 .map(|v| v.to_str().unwrap().to_string())
4142 .collect();
4143 let content_type = headers
4144 .get(http::header::CONTENT_TYPE)
4145 .map(|v| v.to_str().unwrap().to_string());
4146 (status, challenges, content_type, body)
4147 }
4148
4149 fn describe(extensions: &http::Extensions) -> String {
4152 format!(
4153 "{:?} token={} {:?}",
4154 extensions.get::<Credential>(),
4155 extensions.get::<AuthorizedToken>().is_some(),
4156 extensions.get::<StaticTokenMatch>()
4157 )
4158 }
4159
4160 type RawHeaders = Vec<(&'static str, HeaderValue)>;
4163
4164 async fn axum_full(app: Router, headers: &RawHeaders) -> SeenFull {
4165 let mut request = Request::builder().uri("/test");
4166 for (name, value) in headers {
4167 request = request.header(*name, value.clone());
4168 }
4169 let response = app
4170 .oneshot(request.body(Body::empty()).unwrap())
4171 .await
4172 .unwrap();
4173 let (parts, body) = response.into_parts();
4174 let body = ::axum::body::to_bytes(body, 64 * 1024).await.unwrap();
4175 seen_parts(
4176 &parts.headers,
4177 parts.status.as_u16(),
4178 String::from_utf8(body.to_vec()).unwrap(),
4179 )
4180 }
4181
4182 async fn tower_full<S>(service: S, headers: &RawHeaders) -> SeenFull
4183 where
4184 S: tower_service::Service<
4185 http::Request<String>,
4186 Response = http::Response<String>,
4187 Error = std::convert::Infallible,
4188 >,
4189 {
4190 let mut request = http::Request::builder().uri("/test");
4191 for (name, value) in headers {
4192 request = request.header(*name, value.clone());
4193 }
4194 let response = service
4195 .oneshot(request.body(String::new()).unwrap())
4196 .await
4197 .unwrap();
4198 let (parts, body) = response.into_parts();
4199 seen_parts(&parts.headers, parts.status.as_u16(), body)
4200 }
4201
4202 fn axum_app(layer: AuthLayer) -> Router {
4203 Router::new()
4204 .route(
4205 "/test",
4206 get(|request: Request| async move { describe(request.extensions()) }),
4207 )
4208 .route_layer(layer)
4209 }
4210
4211 fn tower_handler(
4212 request: http::Request<String>,
4213 ) -> std::future::Ready<Result<http::Response<String>, std::convert::Infallible>> {
4214 let mut response = http::Response::new(describe(request.extensions()));
4217 response.headers_mut().insert(
4218 http::header::CONTENT_TYPE,
4219 HeaderValue::from_static("text/plain; charset=utf-8"),
4220 );
4221 std::future::ready(Ok(response))
4222 }
4223
4224 #[tokio::test]
4225 async fn the_layers_agree_on_callbacks_repeated_and_unreadable_headers() {
4226 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
4227 let v = validator(&jwks.url);
4228 let valid = HeaderValue::from_str(&format!("Bearer {}", testing::valid_token())).unwrap();
4229 let requests: Vec<RawHeaders> = vec![
4230 vec![],
4231 vec![("authorization", valid.clone())],
4232 vec![("authorization", HeaderValue::from_static("Bearer wrong"))],
4233 vec![
4236 ("authorization", valid.clone()),
4237 ("authorization", HeaderValue::from_static("Bearer wrong")),
4238 ],
4239 vec![
4240 ("authorization", HeaderValue::from_static("Bearer ")),
4241 ("authorization", HeaderValue::from_static("Bearer secret")),
4242 ],
4243 vec![
4244 ("x-api-key", HeaderValue::from_static("wrong")),
4245 ("x-api-key", HeaderValue::from_static("secret")),
4246 ],
4247 vec![(
4249 "authorization",
4250 HeaderValue::from_bytes(b"Bearer s\xe9cret").unwrap(),
4251 )],
4252 vec![("x-api-key", HeaderValue::from_bytes(b"\xff").unwrap())],
4253 vec![("x-api-key", HeaderValue::from_static("key-next"))],
4255 ];
4256 let next = || {
4257 crate::StaticTokens::new()
4258 .with(Some("next"), "key-next")
4259 .unwrap()
4260 };
4261 for with_oauth in [false, true] {
4262 for optional in [false, true] {
4263 for static_challenge in [None, Some(None)] {
4264 let sources = [
4265 CredentialSource::authorization_bearer(),
4266 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
4267 ];
4268 let oauth = with_oauth.then(|| Arc::clone(&v));
4269 let mut a = AuthLayer::builder()
4272 .static_token(STATIC)
4273 .static_tokens(next())
4274 .optional_oauth(oauth.clone())
4275 .sources(sources.clone())
4276 .on_reject(|cx| {
4277 (
4278 StatusCode::IM_A_TEAPOT,
4279 [
4280 (WWW_AUTHENTICATE, "Callback x"),
4281 (http::header::CONTENT_TYPE, "application/json"),
4282 ],
4283 format!("{{\"status\":{}}}", cx.status.as_u16()),
4284 )
4285 .into_response()
4286 });
4287 let mut t = HttpAuthLayer::builder()
4288 .static_token(STATIC)
4289 .static_tokens(next())
4290 .optional_oauth(oauth)
4291 .sources(sources)
4292 .on_reject(|cx: RejectContext<'_>| {
4293 http::Response::builder()
4294 .status(StatusCode::IM_A_TEAPOT)
4295 .header(WWW_AUTHENTICATE, "Callback x")
4296 .header(http::header::CONTENT_TYPE, "application/json")
4297 .body(format!("{{\"status\":{}}}", cx.status.as_u16()))
4298 .unwrap()
4299 });
4300 if let Some(c) = &static_challenge {
4301 a = a.static_challenge(c.clone());
4302 t = t.static_challenge(c.clone());
4303 }
4304 if optional {
4305 a = a.optional();
4306 t = t.optional();
4307 }
4308 let (a, t) = (a.build().unwrap(), t.build().unwrap());
4309 for headers in &requests {
4310 let service = tower_layer::Layer::layer(&t, service_fn(tower_handler));
4311 assert_eq!(
4312 axum_full(axum_app(a.clone()), headers).await,
4313 tower_full(service, headers).await,
4314 "oauth={with_oauth} optional={optional} \
4315 static_challenge={static_challenge:?} {headers:?}"
4316 );
4317 }
4318 }
4319 }
4320 }
4321 }
4322
4323 #[tokio::test]
4324 async fn the_layers_agree_that_optional_clears_an_outer_layers_credential() {
4325 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
4326 let v = validator(&jwks.url);
4327 let bearer = HeaderValue::from_str(&format!("Bearer {}", testing::valid_token())).unwrap();
4328 let outer_sources = [CredentialSource::Bearer(HeaderName::from_static("x-outer"))];
4331 let requests: Vec<RawHeaders> = vec![
4332 vec![("x-outer", bearer.clone())],
4333 vec![
4334 ("x-outer", bearer.clone()),
4335 ("authorization", HeaderValue::from_static("Bearer secret")),
4336 ],
4337 vec![
4338 ("x-outer", bearer.clone()),
4339 ("authorization", HeaderValue::from_static("Bearer wrong")),
4340 ],
4341 ];
4342 let axum_outer = AuthLayer::builder()
4343 .oauth(Arc::clone(&v))
4344 .sources(outer_sources.clone())
4345 .build()
4346 .unwrap();
4347 let axum_inner = AuthLayer::builder()
4348 .static_token(STATIC)
4349 .optional()
4350 .build()
4351 .unwrap();
4352 let tower_outer = HttpAuthLayer::builder()
4353 .oauth(Arc::clone(&v))
4354 .sources(outer_sources)
4355 .build()
4356 .unwrap();
4357 let tower_inner = HttpAuthLayer::builder()
4358 .static_token(STATIC)
4359 .optional()
4360 .build()
4361 .unwrap();
4362 let mut outcomes = Vec::new();
4363 for headers in &requests {
4364 let app = axum_app(axum_inner.clone()).layer(axum_outer.clone());
4365 let service = ::tower::ServiceBuilder::new()
4366 .layer(tower_outer.clone())
4367 .layer(tower_inner.clone())
4368 .service(service_fn(tower_handler));
4369 let a = axum_full(app, headers).await;
4370 assert_eq!(a, tower_full(service, headers).await, "{headers:?}");
4371 outcomes.push(a);
4372 }
4373 assert_eq!(outcomes[0].3, "None token=false None");
4375 assert_eq!(
4376 outcomes[1].3,
4377 "Some(StaticToken) token=false Some(StaticTokenMatch { label: None })"
4378 );
4379 assert_eq!(outcomes[2].0, 401);
4380 }
4381}
4382
4383#[cfg(test)]
4386mod static_tokens_tests {
4387 use ::tower::ServiceExt;
4388
4389 use super::*;
4390 use crate::testing;
4391 use http::HeaderName;
4392
4393 const STATIC: &str = "secret";
4394
4395 fn rotation() -> StaticTokens {
4396 StaticTokens::new()
4397 .with(Some("current"), "key-current")
4398 .and_then(|t| t.with(Some("next"), "key-next"))
4399 .unwrap()
4400 }
4401
4402 fn validator(jwks_uri: &str) -> Arc<OAuthValidator> {
4403 Arc::new(OAuthValidator::new(&testing::resolved_config(jwks_uri)).unwrap())
4404 }
4405
4406 async fn report(credential: Option<Credential>, matched: Option<StaticTokenMatch>) -> String {
4408 format!("{credential:?} {matched:?}")
4409 }
4410
4411 fn app(layer: AuthLayer) -> Router {
4412 Router::new()
4413 .route("/test", get(report))
4414 .route(
4415 "/required",
4416 get(|m: StaticTokenMatch| async move { format!("{:?}", m.label()) }),
4417 )
4418 .route_layer(layer)
4419 }
4420
4421 async fn send(app: &Router, path: &str, headers: &[(&str, &str)]) -> Response {
4422 let mut request = Request::builder().uri(path);
4423 for (name, value) in headers {
4424 request = request.header(*name, *value);
4425 }
4426 app.clone()
4427 .oneshot(request.body(Body::empty()).unwrap())
4428 .await
4429 .unwrap()
4430 }
4431
4432 async fn observed(resp: Response) -> (u16, Vec<(String, Vec<u8>)>, String) {
4434 let status = resp.status().as_u16();
4435 let headers = resp
4436 .headers()
4437 .iter()
4438 .map(|(k, v)| (k.to_string(), v.as_bytes().to_vec()))
4439 .collect();
4440 let body = ::axum::body::to_bytes(resp.into_body(), 64 * 1024)
4441 .await
4442 .unwrap();
4443 (status, headers, String::from_utf8(body.to_vec()).unwrap())
4444 }
4445
4446 fn json_reject(cx: RejectContext<'_>) -> Response {
4447 (
4448 cx.status,
4449 Json(serde_json::json!({ "status": cx.status.as_u16() })),
4450 )
4451 .into_response()
4452 }
4453
4454 #[tokio::test]
4455 async fn a_one_entry_set_answers_exactly_like_static_token() {
4456 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
4457 let v = validator(&jwks.url);
4458 let valid = format!("Bearer {}", testing::valid_token());
4459 let requests: Vec<Vec<(&str, &str)>> = vec![
4460 vec![],
4461 vec![("authorization", "Bearer secret")],
4462 vec![("authorization", "bearer secret")],
4463 vec![("authorization", "Bearer wrong")],
4464 vec![("authorization", "Bearer ")],
4465 vec![("authorization", "DPoP secret")],
4466 vec![("x-api-key", "secret")],
4467 vec![("authorization", "Bearer wrong"), ("x-api-key", "secret")],
4468 vec![("authorization", valid.as_str())],
4469 ];
4470 let configs = [
4472 (false, false, false, false),
4473 (false, false, true, false),
4474 (true, false, false, true),
4475 (false, true, false, true),
4476 (true, true, false, false),
4477 ];
4478 for (with_oauth, optional, no_challenge, callback) in configs {
4479 let build = |use_set: bool| {
4480 let mut b = AuthLayer::builder()
4481 .optional_oauth(with_oauth.then(|| Arc::clone(&v)))
4482 .sources([
4483 CredentialSource::authorization_bearer(),
4484 CredentialSource::Raw(HeaderName::from_static("x-api-key")),
4485 ]);
4486 b = if use_set {
4487 b.static_tokens(StaticTokens::single(STATIC).unwrap())
4488 } else {
4489 b.static_token(STATIC)
4490 };
4491 if optional {
4492 b = b.optional();
4493 }
4494 if no_challenge {
4495 b = b.static_challenge(None);
4496 }
4497 if callback {
4498 b = b.on_reject(json_reject);
4499 }
4500 app(b.build().unwrap())
4501 };
4502 let (old, new) = (build(false), build(true));
4503 for headers in &requests {
4504 for path in ["/test", "/required"] {
4505 let a = observed(send(&old, path, headers).await).await;
4506 let b = observed(send(&new, path, headers).await).await;
4507 assert_eq!(
4508 a,
4509 b,
4510 "config {:?} {path} {headers:.40?}",
4511 (with_oauth, optional, no_challenge, callback)
4512 );
4513 }
4514 }
4515 }
4516 }
4517
4518 #[tokio::test]
4519 async fn the_extractors_name_the_matching_key() {
4520 let app = app(AuthLayer::builder()
4521 .static_tokens(rotation())
4522 .build()
4523 .unwrap());
4524 for (secret, label) in [("key-current", "current"), ("key-next", "next")] {
4525 let bearer = format!("Bearer {secret}");
4526 let (status, _, body) =
4527 observed(send(&app, "/test", &[("authorization", &bearer)]).await).await;
4528 assert_eq!(status, 200);
4529 assert_eq!(
4530 body,
4531 format!("Some(StaticToken) Some(StaticTokenMatch {{ label: Some({label:?}) }})")
4532 );
4533 let (status, _, body) =
4534 observed(send(&app, "/required", &[("authorization", &bearer)]).await).await;
4535 assert_eq!((status, body), (200, format!("Some({label:?})")));
4536 }
4537 let (status, headers, _) =
4538 observed(send(&app, "/test", &[("authorization", "Bearer key-old")]).await).await;
4539 assert_eq!(status, 401);
4540 assert!(headers.contains(&(
4541 "www-authenticate".to_string(),
4542 DEFAULT_STATIC_CHALLENGE.as_bytes().to_vec()
4543 )));
4544 }
4545
4546 #[tokio::test]
4547 async fn the_extractor_refuses_an_oauth_request_and_a_route_outside_every_layer() {
4548 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
4549 let v = validator(&jwks.url);
4550 let layer = AuthLayer::builder()
4551 .oauth(Arc::clone(&v))
4552 .static_tokens(rotation())
4553 .build()
4554 .unwrap();
4555 let router = app(layer);
4556 let valid = format!("Bearer {}", testing::valid_token());
4557 let (status, _, body) =
4560 observed(send(&router, "/test", &[("authorization", &valid)]).await).await;
4561 assert_eq!(status, 200);
4562 assert!(
4563 body.starts_with("Some(OAuth(") && body.ends_with(" None"),
4564 "{body}"
4565 );
4566 let resp = send(&router, "/required", &[("authorization", &valid)]).await;
4567 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
4568 assert_eq!(
4569 resp.headers()[WWW_AUTHENTICATE],
4570 v.invalid_token_challenge().as_str()
4571 );
4572
4573 let oauth_only = app(AuthLayer::builder().oauth(Arc::clone(&v)).build().unwrap());
4575 let resp = send(&oauth_only, "/required", &[("authorization", &valid)]).await;
4576 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
4577
4578 let bare: Router = Router::new().route("/test", get(report)).route(
4580 "/required",
4581 get(|_: StaticTokenMatch| async { "unreachable" }),
4582 );
4583 for path in ["/test", "/required"] {
4584 let resp = send(&bare, path, &[("authorization", "Bearer key-current")]).await;
4585 assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR, "{path}");
4586 }
4587 }
4588
4589 #[tokio::test]
4590 async fn an_optional_layer_with_several_tokens() {
4591 let app = app(AuthLayer::builder()
4592 .static_tokens(rotation())
4593 .optional()
4594 .build()
4595 .unwrap());
4596 let (status, _, body) = observed(send(&app, "/test", &[]).await).await;
4597 assert_eq!((status, body.as_str()), (200, "None None"));
4598 let (status, _, _) = observed(send(&app, "/required", &[]).await).await;
4600 assert_eq!(status, 401);
4601 for (secret, label) in [("key-current", "current"), ("key-next", "next")] {
4602 let (status, _, body) = observed(
4603 send(
4604 &app,
4605 "/required",
4606 &[("authorization", &format!("Bearer {secret}"))],
4607 )
4608 .await,
4609 )
4610 .await;
4611 assert_eq!((status, body), (200, format!("Some({label:?})")));
4612 }
4613 let (status, _, _) =
4614 observed(send(&app, "/test", &[("authorization", "Bearer key-old")]).await).await;
4615 assert_eq!(status, 401);
4616 }
4617
4618 #[tokio::test]
4619 async fn nested_layers_keep_the_match_paired_with_the_credential() {
4620 let jwks = testing::spawn_jwks_server("200 OK", testing::jwks_body()).await;
4621 let v = validator(&jwks.url);
4622 let valid = format!("Bearer {}", testing::valid_token());
4623 let outer_static = AuthLayer::builder()
4624 .static_tokens(rotation())
4625 .sources([CredentialSource::Raw(HeaderName::from_static("x-outer"))])
4626 .build()
4627 .unwrap();
4628
4629 let inner_oauth = AuthLayer::builder().oauth(Arc::clone(&v)).build().unwrap();
4631 let router = app(inner_oauth).layer(outer_static.clone());
4632 let (status, _, body) = observed(
4633 send(
4634 &router,
4635 "/test",
4636 &[("x-outer", "key-next"), ("authorization", &valid)],
4637 )
4638 .await,
4639 )
4640 .await;
4641 assert_eq!(status, 200);
4642 assert!(
4643 body.starts_with("Some(OAuth(") && body.ends_with(" None"),
4644 "{body}"
4645 );
4646
4647 let inner_optional = AuthLayer::builder()
4649 .static_token("inner-key")
4650 .sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
4651 .optional()
4652 .build()
4653 .unwrap();
4654 let router = app(inner_optional).layer(outer_static.clone());
4655 let (_, _, body) = observed(send(&router, "/test", &[("x-outer", "key-next")]).await).await;
4656 assert_eq!(body, "None None");
4657
4658 let inner_static = AuthLayer::builder()
4660 .static_token("inner-key")
4661 .sources([CredentialSource::Raw(HeaderName::from_static("x-inner"))])
4662 .build()
4663 .unwrap();
4664 let router = app(inner_static).layer(outer_static);
4665 let (_, _, body) = observed(
4666 send(
4667 &router,
4668 "/test",
4669 &[("x-outer", "key-next"), ("x-inner", "inner-key")],
4670 )
4671 .await,
4672 )
4673 .await;
4674 assert_eq!(
4675 body,
4676 "Some(StaticToken) Some(StaticTokenMatch { label: None })"
4677 );
4678 }
4679
4680 #[test]
4681 fn build_with_decision_combines_the_set_as_documented() {
4682 let v = validator("http://127.0.0.1:1/jwks");
4683 assert!(
4684 AuthLayer::builder()
4685 .optional_oauth(Some(Arc::clone(&v)))
4686 .static_tokens(rotation())
4687 .build_with_decision(StaticTokenDecision::StaticAndOAuth("x".into()))
4688 .is_ok()
4689 );
4690 let dropped = AuthLayer::builder()
4691 .oauth(Arc::clone(&v))
4692 .static_tokens(rotation())
4693 .build_with_decision(StaticTokenDecision::StaticIgnored)
4694 .unwrap();
4695 assert!(
4696 format!("{dropped:?}").contains("static_tokens: None"),
4697 "{dropped:?}"
4698 );
4699 for (decision, oauth) in [
4700 (StaticTokenDecision::OAuthOnly, Some(Arc::clone(&v))),
4701 (StaticTokenDecision::Unauthenticated, None),
4702 ] {
4703 assert_eq!(
4704 AuthLayer::builder()
4705 .optional_oauth(oauth)
4706 .static_tokens(rotation())
4707 .build_with_decision(decision)
4708 .unwrap_err(),
4709 AuthLayerError::DecisionWithoutStaticToken
4710 );
4711 }
4712 assert!(
4714 AuthLayer::from_decision(StaticTokenDecision::Unauthenticated, None)
4715 .unwrap()
4716 .allows_unauthenticated()
4717 );
4718 assert_eq!(
4719 AuthLayer::builder()
4720 .static_tokens(StaticTokens::new())
4721 .build()
4722 .unwrap_err(),
4723 AuthLayerError::NoCredential
4724 );
4725 }
4726
4727 #[test]
4728 fn debug_never_prints_a_token_from_a_set() {
4729 let builder = AuthLayer::builder()
4730 .static_token("hunter2-single")
4731 .static_tokens(
4732 StaticTokens::new()
4733 .with(Some("next"), "hunter2-next")
4734 .unwrap(),
4735 );
4736 let rendered = format!("{builder:?}");
4737 assert!(
4738 !rendered.contains("hunter2") && rendered.contains("next"),
4739 "{rendered}"
4740 );
4741 let layer = builder.build().unwrap();
4742 let rendered = format!("{layer:?}");
4743 assert!(
4744 !rendered.contains("hunter2") && rendered.contains("<redacted>"),
4745 "{rendered}"
4746 );
4747 }
4748}