1use std::{
4 future::Future,
5 marker::PhantomData,
6 pin::Pin,
7 rc::Rc,
8 sync::Arc,
9 task::{Context, Poll},
10};
11
12use actix_web::{
13 body::{EitherBody, MessageBody},
14 dev::{Service, ServiceRequest, ServiceResponse, Transform},
15 Error, FromRequest,
16};
17use futures_core::ready;
18use futures_util::future::{self, LocalBoxFuture, TryFutureExt as _};
19
20use crate::extractors::{basic, bearer};
21
22#[derive(Debug, Clone)]
34pub struct HttpAuthentication<T, F>
35where
36 T: FromRequest,
37{
38 process_fn: Arc<F>,
39 _extractor: PhantomData<T>,
40}
41
42impl<T, F, O> HttpAuthentication<T, F>
43where
44 T: FromRequest,
45 F: Fn(ServiceRequest, T) -> O,
46 O: Future<Output = Result<ServiceRequest, (Error, ServiceRequest)>>,
47{
48 pub fn with_fn(process_fn: F) -> HttpAuthentication<T, F> {
101 HttpAuthentication {
102 process_fn: Arc::new(process_fn),
103 _extractor: PhantomData,
104 }
105 }
106}
107
108impl<F, O> HttpAuthentication<basic::BasicAuth, F>
109where
110 F: Fn(ServiceRequest, basic::BasicAuth) -> O,
111 O: Future<Output = Result<ServiceRequest, (Error, ServiceRequest)>>,
112{
113 pub fn basic(process_fn: F) -> Self {
133 Self::with_fn(process_fn)
134 }
135}
136
137impl<F, O> HttpAuthentication<bearer::BearerAuth, F>
138where
139 F: Fn(ServiceRequest, bearer::BearerAuth) -> O,
140 O: Future<Output = Result<ServiceRequest, (Error, ServiceRequest)>>,
141{
142 pub fn bearer(process_fn: F) -> Self {
170 Self::with_fn(process_fn)
171 }
172}
173
174impl<S, B, T, F, O> Transform<S, ServiceRequest> for HttpAuthentication<T, F>
175where
176 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
177 S::Future: 'static,
178 F: Fn(ServiceRequest, T) -> O + 'static,
179 O: Future<Output = Result<ServiceRequest, (Error, ServiceRequest)>> + 'static,
180 T: FromRequest + 'static,
181 B: MessageBody + 'static,
182{
183 type Response = ServiceResponse<EitherBody<B>>;
184 type Error = Error;
185 type Transform = AuthenticationMiddleware<S, F, T>;
186 type InitError = ();
187 type Future = future::Ready<Result<Self::Transform, Self::InitError>>;
188
189 fn new_transform(&self, service: S) -> Self::Future {
190 future::ok(AuthenticationMiddleware {
191 service: Rc::new(service),
192 process_fn: self.process_fn.clone(),
193 _extractor: PhantomData,
194 })
195 }
196}
197
198#[doc(hidden)]
199pub struct AuthenticationMiddleware<S, F, T>
200where
201 T: FromRequest,
202{
203 service: Rc<S>,
204 process_fn: Arc<F>,
205 _extractor: PhantomData<T>,
206}
207
208impl<S, B, F, T, O> Service<ServiceRequest> for AuthenticationMiddleware<S, F, T>
209where
210 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
211 S::Future: 'static,
212 F: Fn(ServiceRequest, T) -> O + 'static,
213 O: Future<Output = Result<ServiceRequest, (Error, ServiceRequest)>> + 'static,
214 T: FromRequest + 'static,
215 B: MessageBody + 'static,
216{
217 type Response = ServiceResponse<EitherBody<B>>;
218 type Error = S::Error;
219 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
220
221 actix_web::dev::forward_ready!(service);
222
223 fn call(&self, req: ServiceRequest) -> Self::Future {
224 let process_fn = Arc::clone(&self.process_fn);
225 let service = Rc::clone(&self.service);
226
227 Box::pin(async move {
228 let (req, credentials) = match Extract::<T>::new(req).await {
229 Ok(req) => req,
230 Err((err, req)) => {
231 return Ok(req.error_response(err).map_into_right_body());
232 }
233 };
234
235 let req = match process_fn(req, credentials).await {
236 Ok(req) => req,
237 Err((err, req)) => {
238 return Ok(req.error_response(err).map_into_right_body());
239 }
240 };
241
242 service.call(req).await.map(|res| res.map_into_left_body())
243 })
244 }
245}
246
247struct Extract<T> {
248 req: Option<ServiceRequest>,
249 fut: Option<LocalBoxFuture<'static, Result<T, Error>>>,
250 _extractor: PhantomData<fn() -> T>,
251}
252
253impl<T> Extract<T> {
254 pub fn new(req: ServiceRequest) -> Self {
255 Extract {
256 req: Some(req),
257 fut: None,
258 _extractor: PhantomData,
259 }
260 }
261}
262
263impl<T> Future for Extract<T>
264where
265 T: FromRequest,
266 T::Future: 'static,
267 T::Error: 'static,
268{
269 type Output = Result<(ServiceRequest, T), (Error, ServiceRequest)>;
270
271 fn poll(mut self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Self::Output> {
272 if self.fut.is_none() {
273 let req = self.req.as_mut().expect("Extract future was polled twice!");
274 let fut = req.extract::<T>().map_err(Into::into);
275 self.fut = Some(Box::pin(fut));
276 }
277
278 let fut = self
279 .fut
280 .as_mut()
281 .expect("Extraction future should be initialized at this point");
282
283 let credentials = ready!(fut.as_mut().poll(ctx)).map_err(|err| {
284 (
285 err,
286 self.req.take().expect("Extract future was polled twice!"),
288 )
289 })?;
290
291 let req = self.req.take().expect("Extract future was polled twice!");
292 Poll::Ready(Ok((req, credentials)))
293 }
294}
295
296#[cfg(test)]
297mod tests {
298 use actix_service::into_service;
299 use actix_web::{
300 error::{self, ErrorForbidden},
301 http::StatusCode,
302 test::TestRequest,
303 web, App, HttpResponse,
304 };
305
306 use super::*;
307 use crate::extractors::{basic::BasicAuth, bearer::BearerAuth};
308
309 #[actix_web::test]
311 async fn test_middleware_panic() {
312 let middleware = AuthenticationMiddleware {
313 service: Rc::new(into_service(|_: ServiceRequest| async move {
314 actix_web::rt::time::sleep(std::time::Duration::from_secs(1)).await;
315 Err::<ServiceResponse, _>(error::ErrorBadRequest("error"))
316 })),
317 process_fn: Arc::new(|req, _: BearerAuth| async { Ok(req) }),
318 _extractor: PhantomData,
319 };
320
321 let req = TestRequest::get()
322 .append_header(("Authorization", "Bearer 1"))
323 .to_srv_request();
324
325 let f = middleware.call(req).await;
326
327 let _res = futures_util::future::lazy(|cx| middleware.poll_ready(cx)).await;
328
329 assert!(f.is_err());
330 }
331
332 #[actix_web::test]
334 async fn test_middleware_panic_several_orders() {
335 let middleware = AuthenticationMiddleware {
336 service: Rc::new(into_service(|_: ServiceRequest| async move {
337 actix_web::rt::time::sleep(std::time::Duration::from_secs(1)).await;
338 Err::<ServiceResponse, _>(error::ErrorBadRequest("error"))
339 })),
340 process_fn: Arc::new(|req, _: BearerAuth| async { Ok(req) }),
341 _extractor: PhantomData,
342 };
343
344 let req = TestRequest::get()
345 .append_header(("Authorization", "Bearer 1"))
346 .to_srv_request();
347
348 let f1 = middleware.call(req).await;
349
350 let req = TestRequest::get()
351 .append_header(("Authorization", "Bearer 1"))
352 .to_srv_request();
353
354 let f2 = middleware.call(req).await;
355
356 let req = TestRequest::get()
357 .append_header(("Authorization", "Bearer 1"))
358 .to_srv_request();
359
360 let f3 = middleware.call(req).await;
361
362 let _res = futures_util::future::lazy(|cx| middleware.poll_ready(cx)).await;
363
364 assert!(f1.is_err());
365 assert!(f2.is_err());
366 assert!(f3.is_err());
367 }
368
369 #[actix_web::test]
370 async fn test_middleware_opt_extractor() {
371 let middleware = AuthenticationMiddleware {
372 service: Rc::new(into_service(|req: ServiceRequest| async move {
373 Ok::<ServiceResponse, _>(req.into_response(HttpResponse::Ok().finish()))
374 })),
375 process_fn: Arc::new(|req, auth: Option<BearerAuth>| {
376 assert!(auth.is_none());
377 async { Ok(req) }
378 }),
379 _extractor: PhantomData,
380 };
381
382 let req = TestRequest::get()
383 .append_header(("Authorization996", "Bearer 1"))
384 .to_srv_request();
385
386 let f = middleware.call(req).await;
387
388 let _res = futures_util::future::lazy(|cx| middleware.poll_ready(cx)).await;
389
390 assert!(f.is_ok());
391 }
392
393 #[actix_web::test]
394 async fn test_middleware_res_extractor() {
395 let middleware = AuthenticationMiddleware {
396 service: Rc::new(into_service(|req: ServiceRequest| async move {
397 Ok::<ServiceResponse, _>(req.into_response(HttpResponse::Ok().finish()))
398 })),
399 process_fn: Arc::new(
400 |req, auth: Result<BearerAuth, <BearerAuth as FromRequest>::Error>| {
401 assert!(auth.is_err());
402 async { Ok(req) }
403 },
404 ),
405 _extractor: PhantomData,
406 };
407
408 let req = TestRequest::get()
409 .append_header(("Authorization", "BearerLOL"))
410 .to_srv_request();
411
412 let f = middleware.call(req).await;
413
414 let _res = futures_util::future::lazy(|cx| middleware.poll_ready(cx)).await;
415
416 assert!(f.is_ok());
417 }
418
419 #[actix_web::test]
420 async fn test_middleware_works_with_app() {
421 async fn validator(
422 req: ServiceRequest,
423 _credentials: BasicAuth,
424 ) -> Result<ServiceRequest, (actix_web::Error, ServiceRequest)> {
425 Err((ErrorForbidden("You are not welcome!"), req))
426 }
427 let middleware = HttpAuthentication::basic(validator);
428
429 let srv = actix_web::test::init_service(
430 App::new()
431 .wrap(middleware)
432 .route("/", web::get().to(HttpResponse::Ok)),
433 )
434 .await;
435
436 let req = actix_web::test::TestRequest::with_uri("/")
437 .append_header(("Authorization", "Basic DoNotCare"))
438 .to_request();
439
440 let resp = srv.call(req).await.unwrap();
441 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
442 }
443
444 #[actix_web::test]
445 async fn test_middleware_works_with_scope() {
446 async fn validator(
447 req: ServiceRequest,
448 _credentials: BasicAuth,
449 ) -> Result<ServiceRequest, (actix_web::Error, ServiceRequest)> {
450 Err((ErrorForbidden("You are not welcome!"), req))
451 }
452 let middleware = actix_web::middleware::Compat::new(HttpAuthentication::basic(validator));
453
454 let srv = actix_web::test::init_service(
455 App::new().service(
456 web::scope("/")
457 .wrap(middleware)
458 .route("/", web::get().to(HttpResponse::Ok)),
459 ),
460 )
461 .await;
462
463 let req = actix_web::test::TestRequest::with_uri("/")
464 .append_header(("Authorization", "Basic DontCare"))
465 .to_request();
466
467 let resp = srv.call(req).await.unwrap();
468 assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
469 }
470}