1use std::{
7 fmt::Debug,
8 future::{ready, Ready},
9 rc::Rc,
10};
11
12#[cfg(feature = "csrf")]
13use actix_web::HttpMessage;
14use actix_web::{
15 dev::{forward_ready, Service, ServiceRequest, ServiceResponse, Transform},
16 web::ServiceConfig,
17 Route,
18};
19use anyhow::Result;
20use async_trait::async_trait;
21use futures::future::LocalBoxFuture;
22
23#[cfg(feature = "csrf")]
24pub fn build_router<F, Fut>(
27 router: Vec<Router>,
28 csrf: crate::csrf::Middleware<F>,
29) -> impl FnOnce(&mut ServiceConfig)
30where
31 F: Fn(actix_web::HttpRequest, String) -> Fut + 'static,
32 Fut: futures::Future<Output = Result<bool, actix_web::Error>>,
33{
34 move |cfg| {
35 for i in router {
36 if !i.path.is_empty() {
37 cfg.route(
38 &i.path,
39 i.route.wrap(csrf.clone()).wrap(RouterGuard {
40 checker: i.checker,
41 csrf: i.csrf,
42 }),
43 );
44 }
45 }
46 }
47}
48
49#[cfg(not(feature = "csrf"))]
50pub fn build_router(router: Vec<Router>) -> impl FnOnce(&mut ServiceConfig) {
52 |cfg| {
53 for i in router {
54 if !i.path.is_empty() {
55 cfg.route(&i.path, i.route.wrap(RouterGuard { checker: i.checker }));
56 }
57 }
58 }
59}
60
61#[async_trait(?Send)]
66pub trait Checker {
67 async fn check(&self, req: &mut ServiceRequest) -> Result<bool>;
68}
69
70#[cfg(feature = "csrf")]
71#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
76#[derive(Clone, Copy, enum_as_inner::EnumAsInner)]
77pub enum CSRFType {
78 Header,
79 Param,
80 ForceHeader,
81 ForceParam,
82 Disabled,
83}
84
85pub struct Router {
87 pub path: String,
89 pub route: Route,
91 pub checker: Option<Rc<dyn Checker>>,
93 #[cfg(feature = "csrf")]
94 pub csrf: CSRFType,
96}
97
98pub(crate) struct RouterGuard {
99 checker: Option<Rc<dyn Checker>>,
100 #[cfg(feature = "csrf")]
101 csrf: CSRFType,
102}
103
104impl<S, B> Transform<S, ServiceRequest> for RouterGuard
105where
106 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error> + 'static,
107 S::Future: 'static,
108 B: 'static + Debug,
109{
110 type Response = ServiceResponse<B>;
111 type Error = actix_web::Error;
112 type InitError = ();
113 type Transform = RouterGuardMiddleware<S>;
114 type Future = Ready<Result<Self::Transform, Self::InitError>>;
115
116 fn new_transform(&self, service: S) -> Self::Future {
117 ready(Ok(RouterGuardMiddleware {
118 service: Rc::new(service),
119 checker: self.checker.clone(),
120 #[cfg(feature = "csrf")]
121 csrf: self.csrf,
122 }))
123 }
124}
125
126pub(crate) struct RouterGuardMiddleware<S> {
127 service: Rc<S>,
128 checker: Option<Rc<dyn Checker>>,
129 #[cfg(feature = "csrf")]
130 csrf: CSRFType,
131}
132
133impl<S, B> Service<ServiceRequest> for RouterGuardMiddleware<S>
134where
135 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error> + 'static,
136 S::Future: 'static,
137 B: 'static + Debug,
138{
139 type Response = ServiceResponse<B>;
140 type Error = actix_web::Error;
141 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
142
143 forward_ready!(service);
144
145 fn call(&self, mut req: ServiceRequest) -> Self::Future {
146 let srv = self.service.clone();
147 let checker = self.checker.clone();
148 #[cfg(feature = "csrf")]
149 req.extensions_mut().insert(self.csrf);
150 Box::pin(async move {
151 if let Some(checker) = checker {
152 match checker.check(&mut req).await {
153 Ok(ok) => {
154 if ok {
155 srv.call(req).await
156 } else {
157 Err(actix_web::error::ErrorForbidden("Checker failed"))
158 }
159 }
160 Err(e) => Err(actix_web::error::ErrorInternalServerError(e)),
161 }
162 } else {
163 srv.call(req).await
164 }
165 })
166 }
167}