Skip to main content

rskit_auth/
bearer.rs

1//! Tower middleware for header-only bearer-token authentication.
2//!
3//! The layer rejects missing credentials unless [`BearerAuthLayer::accept_missing`] is selected explicitly.
4//! Rejections include a neutral `WWW-Authenticate: Bearer` challenge;
5//! this crate does not hard-code an application realm.
6
7use std::{future::Future, marker::PhantomData, pin::Pin, sync::Arc};
8
9use http::{Request, Response, StatusCode, header};
10use rskit_security::BEARER_AUTH_SCHEME;
11use tower::{Layer, Service};
12
13use crate::{AuthClaims, AuthOutcome, MissingCredentialPolicy, traits::TokenValidator};
14
15/// Tower layer that validates `Authorization: Bearer <token>` headers.
16///
17/// The layer stores successful claims in request extensions as [`AuthClaims<C>`](crate::AuthClaims)
18/// and [`AuthOutcome<C>`](crate::AuthOutcome). Invalid
19/// or rejected requests receive `401 Unauthorized` with `WWW-Authenticate: Bearer`.
20#[derive(Clone)]
21pub struct BearerAuthLayer<V, C> {
22    validator: Arc<V>,
23    missing_policy: MissingCredentialPolicy,
24    _claims: PhantomData<fn() -> C>,
25}
26
27impl<V: 'static, C> BearerAuthLayer<V, C> {
28    /// Create a new bearer-auth layer that rejects missing credentials.
29    #[must_use]
30    pub fn new(validator: V) -> Self {
31        Self {
32            validator: Arc::new(validator),
33            missing_policy: MissingCredentialPolicy::RejectMissing,
34            _claims: PhantomData,
35        }
36    }
37
38    /// Explicitly accept requests with no credentials.
39    #[must_use]
40    pub const fn accept_missing(mut self) -> Self {
41        self.missing_policy = MissingCredentialPolicy::AcceptMissing;
42        self
43    }
44}
45
46impl<S, V, C> Layer<S> for BearerAuthLayer<V, C>
47where
48    V: 'static,
49{
50    type Service = BearerAuthService<S, V, C>;
51
52    fn layer(&self, inner: S) -> Self::Service {
53        BearerAuthService {
54            inner,
55            validator: Arc::clone(&self.validator),
56            missing_policy: self.missing_policy,
57            _claims: PhantomData,
58        }
59    }
60}
61
62/// Tower service produced by [`BearerAuthLayer`].
63#[derive(Clone)]
64pub struct BearerAuthService<S, V, C> {
65    inner: S,
66    validator: Arc<V>,
67    missing_policy: MissingCredentialPolicy,
68    _claims: PhantomData<fn() -> C>,
69}
70
71impl<S, V, C, ReqBody, ResBody> Service<Request<ReqBody>> for BearerAuthService<S, V, C>
72where
73    S: Service<Request<ReqBody>, Response = Response<ResBody>> + Clone + Send + 'static,
74    S::Future: Send + 'static,
75    V: TokenValidator<C> + 'static,
76    C: Clone + Send + Sync + 'static,
77    ReqBody: Send + 'static,
78    ResBody: Default + Send + 'static,
79{
80    type Response = S::Response;
81    type Error = S::Error;
82    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
83
84    fn poll_ready(
85        &mut self,
86        cx: &mut std::task::Context<'_>,
87    ) -> std::task::Poll<Result<(), Self::Error>> {
88        self.inner.poll_ready(cx)
89    }
90
91    fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
92        let clone = self.inner.clone();
93        let mut inner = std::mem::replace(&mut self.inner, clone);
94        let validator = Arc::clone(&self.validator);
95        let missing_policy = self.missing_policy;
96
97        Box::pin(async move {
98            match extract_bearer_token(&req) {
99                CredentialExtraction::Missing => {
100                    if missing_policy == MissingCredentialPolicy::AcceptMissing {
101                        let mut req = req;
102                        req.extensions_mut().remove::<AuthClaims<C>>();
103                        req.extensions_mut().insert(AuthOutcome::<C>::Missing);
104                        inner.call(req).await
105                    } else {
106                        Ok(unauthorized_bearer_response())
107                    }
108                }
109                CredentialExtraction::Invalid => Ok(unauthorized_bearer_response()),
110                CredentialExtraction::Present(token) => match validator.validate(token).await {
111                    Ok(claims) => {
112                        let mut req = req;
113                        req.extensions_mut().insert(AuthClaims(claims.clone()));
114                        req.extensions_mut()
115                            .insert(AuthOutcome::Authenticated(claims));
116                        inner.call(req).await
117                    }
118                    Err(_) => Ok(unauthorized_bearer_response()),
119                },
120            }
121        })
122    }
123}
124
125enum CredentialExtraction<'a> {
126    Missing,
127    Invalid,
128    Present(&'a str),
129}
130
131fn extract_bearer_token<B>(req: &Request<B>) -> CredentialExtraction<'_> {
132    let mut values = req.headers().get_all(header::AUTHORIZATION).iter();
133    let Some(value) = values.next() else {
134        return CredentialExtraction::Missing;
135    };
136    if values.next().is_some() {
137        return CredentialExtraction::Invalid;
138    }
139    let Ok(value) = value.to_str() else {
140        return CredentialExtraction::Invalid;
141    };
142    let Some((scheme, token)) = value.split_once(' ') else {
143        return CredentialExtraction::Invalid;
144    };
145    if !scheme.eq_ignore_ascii_case(BEARER_AUTH_SCHEME) {
146        return CredentialExtraction::Invalid;
147    }
148    let token = token.trim_start_matches(' ');
149    if token.is_empty() || token.chars().any(char::is_whitespace) {
150        return CredentialExtraction::Invalid;
151    }
152    CredentialExtraction::Present(token)
153}
154
155fn unauthorized_bearer_response<ResBody: Default>() -> Response<ResBody> {
156    let mut response = Response::new(ResBody::default());
157    *response.status_mut() = StatusCode::UNAUTHORIZED;
158    response.headers_mut().insert(
159        header::WWW_AUTHENTICATE,
160        http::HeaderValue::from_static(BEARER_AUTH_SCHEME),
161    );
162    response
163}
164
165#[cfg(test)]
166mod tests {
167    use std::{convert::Infallible, future::Ready};
168
169    use async_trait::async_trait;
170    use http::{Request, Response, StatusCode};
171    use rskit_errors::{AppError, AppResult};
172    use rskit_security::BEARER_AUTH_SCHEME;
173    use serde::{Deserialize, Serialize};
174    use tower::{Layer, Service, ServiceExt};
175
176    use super::{BearerAuthLayer, CredentialExtraction, extract_bearer_token};
177    use crate::{AuthClaims, AuthOutcome, TokenValidator};
178
179    #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
180    struct Claims {
181        sub: String,
182    }
183
184    struct Validator;
185
186    #[async_trait]
187    impl TokenValidator<Claims> for Validator {
188        async fn validate(&self, token: &str) -> AppResult<Claims> {
189            if token == "good.token" {
190                Ok(Claims {
191                    sub: "user-1".into(),
192                })
193            } else {
194                Err(AppError::invalid_token())
195            }
196        }
197    }
198
199    #[derive(Clone)]
200    struct ExtensionCheckingService;
201
202    impl Service<Request<()>> for ExtensionCheckingService {
203        type Response = Response<()>;
204        type Error = Infallible;
205        type Future = Ready<Result<Self::Response, Self::Error>>;
206
207        fn poll_ready(
208            &mut self,
209            _cx: &mut std::task::Context<'_>,
210        ) -> std::task::Poll<Result<(), Self::Error>> {
211            std::task::Poll::Ready(Ok(()))
212        }
213
214        fn call(&mut self, request: Request<()>) -> Self::Future {
215            let has_claims = request.extensions().get::<AuthClaims<Claims>>().is_some();
216            let has_outcome = request.extensions().get::<AuthOutcome<Claims>>().is_some();
217            let status = if has_claims && has_outcome {
218                StatusCode::OK
219            } else if matches!(
220                request.extensions().get::<AuthOutcome<Claims>>(),
221                Some(AuthOutcome::Missing)
222            ) {
223                StatusCode::NO_CONTENT
224            } else {
225                StatusCode::IM_A_TEAPOT
226            };
227            std::future::ready(Ok(Response::builder().status(status).body(()).unwrap()))
228        }
229    }
230
231    #[test]
232    fn bearer_extraction_requires_single_authorization_header() {
233        for value in [
234            "Bearer abc.def.ghi",
235            "bearer abc.def.ghi",
236            "Bearer  abc.def.ghi",
237        ] {
238            let request = http::Request::builder()
239                .header(http::header::AUTHORIZATION, value)
240                .body(())
241                .unwrap();
242
243            assert!(matches!(
244                extract_bearer_token(&request),
245                CredentialExtraction::Present("abc.def.ghi")
246            ));
247        }
248
249        let request = http::Request::builder()
250            .header(http::header::AUTHORIZATION, "Bearer one")
251            .header(http::header::AUTHORIZATION, "Bearer two")
252            .body(())
253            .unwrap();
254
255        assert!(matches!(
256            extract_bearer_token(&request),
257            CredentialExtraction::Invalid
258        ));
259    }
260
261    #[test]
262    fn bearer_extraction_rejects_missing_or_malformed_values() {
263        let missing = http::Request::builder().body(()).unwrap();
264        assert!(matches!(
265            extract_bearer_token(&missing),
266            CredentialExtraction::Missing
267        ));
268
269        for value in [
270            "Basic token",
271            "Bearer ",
272            "Bearer token with-space",
273            "Bearer\ttoken",
274            " Bearer token",
275        ] {
276            let request = http::Request::builder()
277                .header(http::header::AUTHORIZATION, value)
278                .body(())
279                .unwrap();
280            assert!(matches!(
281                extract_bearer_token(&request),
282                CredentialExtraction::Invalid
283            ));
284        }
285    }
286
287    #[tokio::test]
288    async fn bearer_layer_rejects_missing_by_default_and_accepts_valid_tokens() {
289        let mut service =
290            BearerAuthLayer::<_, Claims>::new(Validator).layer(ExtensionCheckingService);
291
292        let missing = service
293            .ready()
294            .await
295            .unwrap()
296            .call(Request::builder().body(()).unwrap())
297            .await
298            .unwrap();
299        assert_eq!(missing.status(), StatusCode::UNAUTHORIZED);
300        assert_eq!(
301            missing
302                .headers()
303                .get(http::header::WWW_AUTHENTICATE)
304                .expect("missing authenticate challenge")
305                .to_str()
306                .expect("challenge should be visible ASCII"),
307            BEARER_AUTH_SCHEME
308        );
309
310        let valid = service
311            .ready()
312            .await
313            .unwrap()
314            .call(
315                Request::builder()
316                    .header(http::header::AUTHORIZATION, "Bearer good.token")
317                    .body(())
318                    .unwrap(),
319            )
320            .await
321            .unwrap();
322        assert_eq!(valid.status(), StatusCode::OK);
323    }
324
325    #[tokio::test]
326    async fn bearer_layer_accept_missing_is_explicit_and_invalid_still_fails() {
327        let mut service = BearerAuthLayer::<_, Claims>::new(Validator)
328            .accept_missing()
329            .layer(ExtensionCheckingService);
330
331        let missing = service
332            .call(Request::builder().body(()).unwrap())
333            .await
334            .unwrap();
335        assert_eq!(missing.status(), StatusCode::NO_CONTENT);
336
337        let invalid = service
338            .call(
339                Request::builder()
340                    .header(http::header::AUTHORIZATION, "Bearer bad.token")
341                    .body(())
342                    .unwrap(),
343            )
344            .await
345            .unwrap();
346        assert_eq!(invalid.status(), StatusCode::UNAUTHORIZED);
347    }
348
349    #[tokio::test]
350    async fn bearer_layer_accept_missing_clears_stale_claims() {
351        let mut service = BearerAuthLayer::<_, Claims>::new(Validator)
352            .accept_missing()
353            .layer(ExtensionCheckingService);
354        let mut request = Request::builder().body(()).unwrap();
355        request.extensions_mut().insert(AuthClaims(Claims {
356            sub: "stale-user".into(),
357        }));
358
359        let response = service.call(request).await.unwrap();
360
361        assert_eq!(response.status(), StatusCode::NO_CONTENT);
362    }
363}