1use std::{future::Future, rc::Rc};
2
3use actix_web::{
4 dev::{forward_ready, Service, ServiceRequest, ServiceResponse, Transform},
5 HttpMessage, HttpRequest,
6};
7use futures::future::{ready, LocalBoxFuture, Ready};
8use qstring::QString;
9
10use crate::router::CSRFType;
11
12pub struct Middleware<F> {
13 cookie: Rc<String>,
14 header: Rc<String>,
15 checker: Rc<F>,
16}
17
18impl<F> Clone for Middleware<F> {
19 fn clone(&self) -> Self {
20 Self {
21 cookie: self.cookie.clone(),
22 header: self.header.clone(),
23 checker: self.checker.clone(),
24 }
25 }
26}
27
28impl<F, Fut> Middleware<F>
29where
30 F: Fn(HttpRequest, String) -> Fut,
31 Fut: Future<Output = Result<bool, actix_web::Error>>,
32{
33 pub fn new(cookie: String, header: String, checker: F) -> Self {
34 Self {
35 cookie: Rc::new(cookie),
36 header: Rc::new(header),
37 checker: Rc::new(checker),
38 }
39 }
40}
41
42impl<S, B, F, Fut> Transform<S, ServiceRequest> for Middleware<F>
43where
44 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error> + 'static,
45 S::Future: 'static,
46 B: 'static,
47 F: Fn(HttpRequest, String) -> Fut + 'static,
48 Fut: Future<Output = Result<bool, actix_web::Error>>,
49{
50 type Response = ServiceResponse<B>;
51 type Error = actix_web::Error;
52 type InitError = ();
53 type Transform = MiddlewareService<S, F>;
54 type Future = Ready<Result<Self::Transform, Self::InitError>>;
55
56 fn new_transform(&self, service: S) -> Self::Future {
57 ready(Ok(MiddlewareService {
58 service: Rc::new(service),
59 cookie: self.cookie.clone(),
60 header: self.header.clone(),
61 checker: self.checker.clone(),
62 }))
63 }
64}
65
66pub struct MiddlewareService<S, F> {
67 service: Rc<S>,
68 cookie: Rc<String>,
69 header: Rc<String>,
70 checker: Rc<F>,
71}
72
73impl<S, B, F, Fut> MiddlewareService<S, F>
74where
75 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error> + 'static,
76 S::Future: 'static,
77 B: 'static,
78 F: Fn(HttpRequest, String) -> Fut + 'static,
79 Fut: Future<Output = Result<bool, actix_web::Error>>,
80{
81 fn get_safe_header(req: &ServiceRequest, name: &str) -> Option<String> {
82 let mut ret: Vec<&str> = req
83 .headers()
84 .get_all(name)
85 .map(|x| x.to_str().unwrap_or_default())
86 .filter(|x| !x.is_empty())
87 .collect();
88 if ret.len() != 1 {
89 return None;
90 }
91 ret.pop().map(ToOwned::to_owned)
92 }
93
94 async fn check_csrf(
95 req: &ServiceRequest,
96 cookie: &str,
97 header: &str,
98 checker: Rc<F>,
99 allow_param: bool,
100 ) -> Result<bool, actix_web::Error> {
101 let Some(cookie) = req.cookie(cookie) else {
102 return Ok(false);
103 };
104 let mut csrf = Self::get_safe_header(req, header);
105 if csrf.is_none() && allow_param {
106 let qs = QString::from(req.query_string());
107 csrf = qs.get(header).map(ToOwned::to_owned);
108 }
109 let Some(csrf) = csrf else {
110 return Ok(false);
111 };
112 if csrf != cookie.value() {
113 return Ok(false);
114 }
115 checker(req.request().clone(), csrf).await
116 }
117}
118
119impl<S, B, F, Fut> Service<ServiceRequest> for MiddlewareService<S, F>
120where
121 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = actix_web::Error> + 'static,
122 S::Future: 'static,
123 B: 'static,
124 F: Fn(HttpRequest, String) -> Fut + 'static,
125 Fut: Future<Output = Result<bool, actix_web::Error>>,
126{
127 type Response = ServiceResponse<B>;
128 type Error = actix_web::Error;
129 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
130
131 forward_ready!(service);
132
133 fn call(&self, req: ServiceRequest) -> Self::Future {
134 let srv = self.service.clone();
135 let header = self.header.clone();
136 let cookie = self.cookie.clone();
137 let checker = self.checker.clone();
138 Box::pin(async move {
139 let csrf = req.extensions().get::<CSRFType>().unwrap().to_owned();
140 if csrf.is_force_header() || csrf.is_force_param() || !req.method().is_safe() {
141 let ret = match csrf {
142 CSRFType::Header => {
143 Self::check_csrf(&req, &cookie, &header, checker, false).await
144 }
145 CSRFType::Param => {
146 Self::check_csrf(&req, &cookie, &header, checker, true).await
147 }
148 CSRFType::ForceHeader => {
149 Self::check_csrf(&req, &cookie, &header, checker, false).await
150 }
151 CSRFType::ForceParam => {
152 Self::check_csrf(&req, &cookie, &header, checker, true).await
153 }
154 CSRFType::Disabled => Ok(true),
155 }?;
156 if !ret {
157 return Err(actix_web::error::ErrorBadRequest("CSRF check failed"));
158 }
159 }
160 srv.call(req).await
161 })
162 }
163}