1use 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#[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 #[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 #[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#[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}