Skip to main content

salvo_csrf/
lib.rs

1#![cfg_attr(test, allow(clippy::unwrap_used))]
2//! CSRF middleware for Salvo web framework.
3//!
4//! CSRF middleware for Salvo that provides CSRF (Cross-Site Request Forgery) protection.
5//!
6//! CSRF token systems commonly use one of two rotation strategies:
7//!
8//! - [`CsrfRotationPolicy::PerSession`]: reuse the same token until it expires or disappears from
9//!   the configured store. This is the default because it works well with page refreshes, browser
10//!   back or forward navigation, and multiple tabs.
11//! - [`CsrfRotationPolicy::PerRequest`]: rotate the token after every accepted request. This
12//!   shortens the lifetime of each token, but clients must always submit the latest token from the
13//!   most recent response.
14//!
15//! Rotation policy is independent from storage. Tokens can be saved in cookies via
16//! [`CookieStore`](struct.CookieStore.html) or in session via
17//! [`SessionStore`](struct.SessionStore.html). [`SessionStore`](struct.SessionStore.html) need to
18//! work with `salvo-session` crate.
19//!
20//! Use [`Csrf::rotation_policy`] to opt into request-level rotation when needed.
21//!
22//! Read more: <https://salvo.rs>
23#![doc(html_favicon_url = "https://salvo.rs/favicon-32x32.png")]
24#![doc(html_logo_url = "https://salvo.rs/images/logo.svg")]
25#![cfg_attr(docsrs, feature(doc_cfg))]
26
27use std::error::Error as StdError;
28use std::fmt::{self, Debug, Formatter};
29
30mod finder;
31
32pub use finder::{CsrfTokenFinder, FormFinder, HeaderFinder, JsonFinder};
33use rand::RngExt;
34use rand::distr::StandardUniform;
35use salvo_core::handler::Skipper;
36use salvo_core::http::{Method, StatusCode};
37use salvo_core::{Depot, FlowCtrl, Handler, Request, Response, async_trait, cfg_feature};
38
39cfg_feature! {
40    #![feature = "cookie-store"]
41
42    mod cookie_store;
43    pub use cookie_store::CookieStore;
44
45    /// Helper function to create a `CookieStore`.
46    #[must_use]
47    pub fn cookie_store() -> CookieStore {
48        CookieStore::new()
49    }
50}
51cfg_feature! {
52    #![feature = "session-store"]
53
54    mod session_store;
55    pub use session_store::SessionStore;
56
57    /// Helper function to create a `SessionStore`.
58    #[must_use]
59    pub fn session_store() -> SessionStore {
60        SessionStore::new()
61    }
62}
63cfg_feature! {
64    #![feature = "bcrypt-cipher"]
65
66    mod bcrypt_cipher;
67    pub use bcrypt_cipher::BcryptCipher;
68
69    /// Helper function to create a `Csrf` use `BcryptCipher`.
70    pub fn bcrypt_csrf<S>(store: S, finder: impl CsrfTokenFinder) -> Csrf<BcryptCipher, S> where S: CsrfStore {
71        Csrf::new(BcryptCipher::new(), store, finder)
72    }
73}
74cfg_feature! {
75    #![all(feature = "bcrypt-cipher", feature = "cookie-store")]
76    /// Helper function to create a `Csrf` use `BcryptCipher` and `CookieStore`.
77    pub fn bcrypt_cookie_csrf(finder: impl CsrfTokenFinder) -> Csrf<BcryptCipher, CookieStore> {
78        Csrf::new(BcryptCipher::new(), CookieStore::new(), finder)
79    }
80}
81cfg_feature! {
82    #![all(feature = "bcrypt-cipher", feature = "session-store")]
83    /// Helper function to create a `Csrf` use `BcryptCipher` and `SessionStore`.
84    pub fn bcrypt_session_csrf(finder: impl CsrfTokenFinder) -> Csrf<BcryptCipher, SessionStore> {
85        Csrf::new(BcryptCipher::new(), SessionStore::new(), finder)
86    }
87}
88
89cfg_feature! {
90    #![feature = "hmac-cipher"]
91
92    mod hmac_cipher;
93    pub use hmac_cipher::HmacCipher;
94
95    /// Helper function to create a `Csrf` use `HmacCipher`.
96    pub fn hmac_csrf<S>(hmac_key: [u8; 32], store: S, finder: impl CsrfTokenFinder) -> Csrf<HmacCipher, S> where S: CsrfStore {
97        Csrf::new(HmacCipher::new(hmac_key), store, finder)
98    }
99}
100cfg_feature! {
101    #![all(feature = "hmac-cipher", feature = "cookie-store")]
102    /// Helper function to create a `Csrf` use `HmacCipher` and `CookieStore`.
103    pub fn hmac_cookie_csrf(hmac_key: [u8; 32], finder: impl CsrfTokenFinder) -> Csrf<HmacCipher, CookieStore> {
104        Csrf::new(HmacCipher::new(hmac_key), CookieStore::new(), finder)
105    }
106}
107cfg_feature! {
108    #![all(feature = "hmac-cipher", feature = "session-store")]
109    /// Helper function to create a `Csrf` use `HmacCipher` and `SessionStore`.
110    pub fn hmac_session_csrf(hmac_key: [u8; 32], finder: impl CsrfTokenFinder) -> Csrf<HmacCipher, SessionStore> {
111        Csrf::new(HmacCipher::new(hmac_key), SessionStore::new(), finder)
112    }
113}
114
115cfg_feature! {
116    #![feature = "aes-gcm-cipher"]
117
118    mod aes_gcm_cipher;
119    pub use aes_gcm_cipher::AesGcmCipher;
120
121    /// Helper function to create a `Csrf` use `AesGcmCipher`.
122    pub fn aes_gcm_csrf<S>(aead_key: [u8; 32], store: S, finder: impl CsrfTokenFinder) -> Csrf<AesGcmCipher, S> where S: CsrfStore {
123        Csrf::new(AesGcmCipher::new(aead_key), store, finder)
124    }
125}
126cfg_feature! {
127    #![all(feature = "aes-gcm-cipher", feature = "cookie-store")]
128    /// Helper function to create a `Csrf` use `AesGcmCipher` and `CookieStore`.
129    pub fn aes_gcm_cookie_csrf(aead_key: [u8; 32], finder: impl CsrfTokenFinder) -> Csrf<AesGcmCipher, CookieStore> {
130        Csrf::new(AesGcmCipher::new(aead_key), CookieStore::new(), finder)
131    }
132}
133cfg_feature! {
134    #![all(feature = "aes-gcm-cipher", feature = "session-store")]
135    /// Helper function to create a `Csrf` use `AesGcmCipher` and `SessionStore`.
136    pub fn aes_gcm_session_csrf(aead_key: [u8; 32], finder: impl CsrfTokenFinder) -> Csrf<AesGcmCipher, SessionStore> {
137        Csrf::new(AesGcmCipher::new(aead_key), SessionStore::new(), finder)
138    }
139}
140
141cfg_feature! {
142    #![feature = "ccp-cipher"]
143
144    mod ccp_cipher;
145    pub use ccp_cipher::CcpCipher;
146
147    /// Helper function to create a `Csrf` use `CcpCipher`.
148    pub fn ccp_csrf<S>(aead_key: [u8; 32], store: S, finder: impl CsrfTokenFinder) -> Csrf<CcpCipher, S> where S: CsrfStore {
149        Csrf::new(CcpCipher::new(aead_key), store, finder)
150    }
151}
152cfg_feature! {
153    #![all(feature = "ccp-cipher", feature = "cookie-store")]
154    /// Helper function to create a `Csrf` use `CcpCipher` and `CookieStore`.
155    pub fn ccp_cookie_csrf(aead_key: [u8; 32], finder: impl CsrfTokenFinder) -> Csrf<CcpCipher, CookieStore> {
156        Csrf::new(CcpCipher::new(aead_key), CookieStore::new(), finder)
157    }
158}
159cfg_feature! {
160    #![all(feature = "ccp-cipher", feature = "session-store")]
161    /// Helper function to create a `Csrf` use `CcpCipher` and `SessionStore`.
162    pub fn ccp_session_csrf(aead_key: [u8; 32], finder: impl CsrfTokenFinder) -> Csrf<CcpCipher, SessionStore> {
163        Csrf::new(CcpCipher::new(aead_key), SessionStore::new(), finder)
164    }
165}
166
167/// Key used to store the CSRF token in [`Depot`].
168pub const CSRF_TOKEN_KEY: &str = "salvo.csrf.token";
169
170/// Controls when a CSRF token is rotated.
171#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
172#[non_exhaustive]
173pub enum CsrfRotationPolicy {
174    /// Reuse the same token until it expires or disappears from the configured store.
175    #[default]
176    PerSession,
177    /// Generate a new token for every accepted request.
178    PerRequest,
179}
180
181fn default_skipper(req: &mut Request, _depot: &Depot) -> bool {
182    !matches!(
183        *req.method(),
184        Method::POST | Method::PATCH | Method::DELETE | Method::PUT
185    )
186}
187
188/// Storage backend for CSRF `(token, proof)` pairs.
189pub trait CsrfStore: Send + Sync + 'static {
190    /// Error type produced by store operations.
191    type Error: StdError + Send + Sync + 'static;
192    /// Load the previously saved `(token, proof)` pair from the store, if any.
193    fn load<C: CsrfCipher>(
194        &self,
195        req: &mut Request,
196        depot: &mut Depot,
197        cipher: &C,
198    ) -> impl Future<Output = Option<(String, String)>> + Send;
199    /// Save the `(token, proof)` pair to the store.
200    fn save(
201        &self,
202        req: &mut Request,
203        depot: &mut Depot,
204        res: &mut Response,
205        token: &str,
206        proof: &str,
207    ) -> impl Future<Output = Result<(), Self::Error>> + Send;
208}
209
210/// Generates and verifies CSRF token / proof pairs.
211pub trait CsrfCipher: Send + Sync + 'static {
212    /// Verify whether the given token matches the proof.
213    fn verify(&self, token: &str, proof: &str) -> bool;
214    /// Generate a new `(token, proof)` pair.
215    fn generate(&self) -> (String, String);
216
217    /// Generate `len` random bytes.
218    fn random_bytes(&self, len: usize) -> Vec<u8> {
219        rand::rng().sample_iter(StandardUniform).take(len).collect()
220    }
221}
222
223/// Extension for Depot.
224pub trait CsrfDepotExt {
225    /// Get csrf token reference from depot.
226    fn csrf_token(&self) -> Option<&str>;
227}
228
229impl CsrfDepotExt for Depot {
230    #[inline]
231    fn csrf_token(&self) -> Option<&str> {
232        self.get::<String>(CSRF_TOKEN_KEY).map(|v| &**v).ok()
233    }
234}
235
236/// Cross-Site Request Forgery (CSRF) protection middleware.
237pub struct Csrf<C, S> {
238    cipher: C,
239    store: S,
240    skipper: Box<dyn Skipper>,
241    finders: Vec<Box<dyn CsrfTokenFinder>>,
242    rotation_policy: CsrfRotationPolicy,
243}
244
245impl<C, S> Debug for Csrf<C, S>
246where
247    C: Debug,
248    S: Debug,
249{
250    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
251        f.debug_struct("Csrf")
252            .field("cipher", &self.cipher)
253            .field("store", &self.store)
254            .field("rotation_policy", &self.rotation_policy)
255            .finish()
256    }
257}
258
259impl<C: CsrfCipher, S: CsrfStore> Csrf<C, S> {
260    /// Create a new instance.
261    #[inline]
262    #[must_use]
263    pub fn new(cipher: C, store: S, finder: impl CsrfTokenFinder) -> Self {
264        Self {
265            cipher,
266            store,
267            skipper: Box::new(default_skipper),
268            finders: vec![Box::new(finder)],
269            rotation_policy: CsrfRotationPolicy::PerSession,
270        }
271    }
272
273    /// Add finder to find csrf token.
274    #[inline]
275    #[must_use]
276    pub fn add_finder(mut self, finder: impl CsrfTokenFinder) -> Self {
277        self.finders.push(Box::new(finder));
278        self
279    }
280
281    /// Sets the [`Skipper`] used to bypass CSRF validation for matching requests.
282    ///
283    /// This replaces the default skipper, which skips safe request methods.
284    #[inline]
285    #[must_use]
286    pub fn skipper(mut self, skipper: impl Skipper) -> Self {
287        self.skipper = Box::new(skipper);
288        self
289    }
290
291    /// Sets the token rotation policy. Defaults to [`CsrfRotationPolicy::PerSession`].
292    #[inline]
293    #[must_use]
294    pub fn rotation_policy(mut self, policy: CsrfRotationPolicy) -> Self {
295        self.rotation_policy = policy;
296        self
297    }
298
299    // /// Clear all finders.
300    // #[inline]
301    // pub fn clear_finders(mut self) -> Self {
302    //     self.finders = vec![];
303    //     self
304    // }
305
306    // /// Sets all finders.
307    // #[inline]
308    // pub fn with_finders(mut self, finders: Vec<Box<dyn CsrfTokenFinder>>) -> Self {
309    //     self.finders = finders;
310    //     self
311    // }
312
313    async fn find_token(&self, req: &mut Request) -> Option<String> {
314        for finder in self.finders.iter() {
315            if let Some(token) = finder.find_token(req).await {
316                return Some(token);
317            }
318        }
319        None
320    }
321
322    async fn issue_token(
323        &self,
324        req: &mut Request,
325        depot: &mut Depot,
326        res: &mut Response,
327    ) -> String {
328        let (token, proof) = self.cipher.generate();
329        if let Err(e) = self.store.save(req, depot, res, &token, &proof).await {
330            tracing::error!(error = ?e, "csrf token save failed");
331        }
332        tracing::debug!("new csrf token generated");
333        token
334    }
335}
336
337#[async_trait]
338impl<C: CsrfCipher, S: CsrfStore> Handler for Csrf<C, S> {
339    async fn handle(
340        &self,
341        req: &mut Request,
342        depot: &mut Depot,
343        res: &mut Response,
344        ctrl: &mut FlowCtrl,
345    ) {
346        let skipped = self.skipper.skipped(req, depot);
347        match self.store.load(req, depot, &self.cipher).await {
348            Some((current_token, proof)) => {
349                if !skipped {
350                    if let Some(token) = self.find_token(req).await {
351                        tracing::debug!("csrf token found in request");
352                        if !self.cipher.verify(&token, &proof) {
353                            tracing::debug!(
354                                "rejecting request due to invalid or expired csrf token"
355                            );
356                            res.status_code(StatusCode::FORBIDDEN);
357                            ctrl.skip_rest();
358                            return;
359                        } else {
360                            tracing::debug!("csrf token verification success");
361                        }
362                    } else {
363                        tracing::debug!("rejecting request due to missing csrf token");
364                        res.status_code(StatusCode::FORBIDDEN);
365                        ctrl.skip_rest();
366                        return;
367                    }
368                }
369
370                let token = if matches!(self.rotation_policy, CsrfRotationPolicy::PerRequest) {
371                    self.issue_token(req, depot, res).await
372                } else {
373                    current_token
374                };
375                depot.insert(CSRF_TOKEN_KEY, token);
376                ctrl.call_next(req, depot, res).await;
377            }
378            None => {
379                if !skipped {
380                    tracing::debug!("rejecting request due to missing csrf token",);
381                    res.status_code(StatusCode::FORBIDDEN);
382                    ctrl.skip_rest();
383                } else {
384                    let token = self.issue_token(req, depot, res).await;
385                    depot.insert(CSRF_TOKEN_KEY, token);
386                    ctrl.call_next(req, depot, res).await;
387                }
388            }
389        }
390    }
391}
392
393#[cfg(test)]
394mod tests {
395    use salvo_core::prelude::*;
396    use salvo_core::test::{ResponseExt, TestClient};
397
398    use super::*;
399
400    #[handler]
401    async fn get_index(depot: &mut Depot) -> String {
402        depot.csrf_token().unwrap().to_owned()
403    }
404    #[handler]
405    async fn post_index() -> &'static str {
406        "POST"
407    }
408    #[handler]
409    async fn post_token(depot: &mut Depot) -> String {
410        depot.csrf_token().unwrap().to_owned()
411    }
412
413    #[tokio::test]
414    async fn test_exposes_csrf_request_extensions() {
415        let csrf = Csrf::new(
416            BcryptCipher::new(),
417            CookieStore::new(),
418            HeaderFinder::new("x-csrf-token"),
419        );
420        let router = Router::new().hoop(csrf).get(get_index);
421        let res = TestClient::get("http://127.0.0.1:5801").send(router).await;
422        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
423    }
424
425    #[tokio::test]
426    async fn test_adds_csrf_cookie_sets_request_token() {
427        let csrf = Csrf::new(
428            BcryptCipher::new(),
429            CookieStore::new(),
430            HeaderFinder::new("x-csrf-token"),
431        );
432        let router = Router::new().hoop(csrf).get(get_index);
433
434        let mut res = TestClient::get("http://127.0.0.1:5801").send(router).await;
435
436        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
437        assert_ne!(res.take_string().await.unwrap(), "");
438        assert_ne!(res.cookie("salvo.csrf"), None);
439    }
440
441    #[cfg(feature = "session-store")]
442    #[tokio::test]
443    async fn test_session_store_without_session_middleware_does_not_panic() {
444        let csrf = Csrf::new(
445            BcryptCipher::new(),
446            SessionStore::new(),
447            HeaderFinder::new("x-csrf-token"),
448        );
449        let router = Router::new().hoop(csrf).get(get_index);
450
451        let mut res = TestClient::get("http://127.0.0.1:5801").send(router).await;
452
453        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
454        assert_ne!(res.take_string().await.unwrap(), "");
455    }
456
457    #[tokio::test]
458    async fn test_per_session_reuses_token_across_safe_requests() {
459        let csrf = Csrf::new(
460            BcryptCipher::new(),
461            CookieStore::new(),
462            HeaderFinder::new("x-csrf-token"),
463        );
464        let router = Router::new().hoop(csrf).get(get_index);
465        let service = Service::new(router);
466
467        let mut res = TestClient::get("http://127.0.0.1:5801")
468            .send(&service)
469            .await;
470        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
471        let token1 = res.take_string().await.unwrap();
472        let cookie = res.cookie("salvo.csrf").unwrap();
473        let cookie_header = cookie.to_string();
474
475        let mut res = TestClient::get("http://127.0.0.1:5801")
476            .add_header("cookie", cookie_header, true)
477            .send(&service)
478            .await;
479        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
480        let token2 = res.take_string().await.unwrap();
481
482        assert_eq!(token1, token2);
483    }
484
485    #[tokio::test]
486    async fn test_per_request_rotates_token_across_safe_requests() {
487        let csrf = Csrf::new(
488            BcryptCipher::new(),
489            CookieStore::new(),
490            HeaderFinder::new("x-csrf-token"),
491        )
492        .rotation_policy(CsrfRotationPolicy::PerRequest);
493        let router = Router::new().hoop(csrf).get(get_index);
494        let service = Service::new(router);
495
496        let mut res = TestClient::get("http://127.0.0.1:5801")
497            .send(&service)
498            .await;
499        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
500        let token1 = res.take_string().await.unwrap();
501        let cookie = res.cookie("salvo.csrf").unwrap();
502        let cookie_header = cookie.to_string();
503
504        let mut res = TestClient::get("http://127.0.0.1:5801")
505            .add_header("cookie", cookie_header, true)
506            .send(&service)
507            .await;
508        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
509        let token2 = res.take_string().await.unwrap();
510        let rotated_cookie = res.cookie("salvo.csrf").unwrap();
511        let rotated_token = rotated_cookie.value().split_once('.').unwrap().0;
512
513        assert_ne!(token1, token2);
514        assert_eq!(token2, rotated_token);
515    }
516
517    #[tokio::test]
518    async fn test_per_request_rotates_token_after_successful_unsafe_request() {
519        let csrf = Csrf::new(
520            BcryptCipher::new(),
521            CookieStore::new(),
522            HeaderFinder::new("x-csrf-token"),
523        )
524        .rotation_policy(CsrfRotationPolicy::PerRequest);
525        let router = Router::new().hoop(csrf).get(get_index).post(post_token);
526        let service = Service::new(router);
527
528        let mut res = TestClient::get("http://127.0.0.1:5801")
529            .send(&service)
530            .await;
531        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
532        let token1 = res.take_string().await.unwrap();
533        let cookie1 = res.cookie("salvo.csrf").unwrap();
534        let cookie1_header = cookie1.to_string();
535
536        let mut res = TestClient::post("http://127.0.0.1:5801")
537            .add_header("x-csrf-token", token1.clone(), true)
538            .add_header("cookie", cookie1_header, true)
539            .send(&service)
540            .await;
541        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
542        let token2 = res.take_string().await.unwrap();
543        let cookie2 = res.cookie("salvo.csrf").unwrap();
544        let cookie2_header = cookie2.to_string();
545        let rotated_token = cookie2.value().split_once('.').unwrap().0;
546
547        assert_ne!(token1, token2);
548        assert_eq!(token2, rotated_token);
549
550        let res = TestClient::post("http://127.0.0.1:5801")
551            .add_header("x-csrf-token", token1, true)
552            .add_header("cookie", cookie2_header.clone(), true)
553            .send(&service)
554            .await;
555        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
556
557        let mut res = TestClient::post("http://127.0.0.1:5801")
558            .add_header("x-csrf-token", token2.clone(), true)
559            .add_header("cookie", cookie2_header, true)
560            .send(&service)
561            .await;
562        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
563        let token3 = res.take_string().await.unwrap();
564
565        assert_ne!(token2, token3);
566    }
567
568    #[tokio::test]
569    async fn test_validates_token_in_header() {
570        let csrf = Csrf::new(
571            BcryptCipher::new(),
572            CookieStore::new(),
573            HeaderFinder::new("x-csrf-token"),
574        );
575        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
576        let service = Service::new(router);
577
578        let mut res = TestClient::get("http://127.0.0.1:5801")
579            .send(&service)
580            .await;
581        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
582
583        let csrf_token = res.take_string().await.unwrap();
584        let cookie = res.cookie("salvo.csrf").unwrap();
585
586        let res = TestClient::post("http://127.0.0.1:5801")
587            .send(&service)
588            .await;
589        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
590
591        let mut res = TestClient::post("http://127.0.0.1:5801")
592            .add_header("x-csrf-token", csrf_token, true)
593            .add_header("cookie", cookie.to_string(), true)
594            .send(&service)
595            .await;
596        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
597        assert_eq!(res.take_string().await.unwrap(), "POST");
598    }
599
600    #[tokio::test]
601    async fn test_custom_skipper_bypasses_csrf_validation() {
602        let csrf = Csrf::new(
603            BcryptCipher::new(),
604            CookieStore::new(),
605            HeaderFinder::new("x-csrf-token"),
606        )
607        .skipper(|req: &mut Request, _depot: &Depot| *req.method() == Method::POST);
608        let router = Router::new().hoop(csrf).post(post_index);
609
610        let mut res = TestClient::post("http://127.0.0.1:5801").send(router).await;
611
612        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
613        assert_eq!(res.take_string().await.unwrap(), "POST");
614    }
615
616    #[tokio::test]
617    async fn test_validates_token_in_custom_header() {
618        let csrf = Csrf::new(
619            BcryptCipher::new(),
620            CookieStore::new(),
621            HeaderFinder::new("x-mycsrf-header"),
622        );
623        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
624        let service = Service::new(router);
625
626        let mut res = TestClient::get("http://127.0.0.1:5801")
627            .send(&service)
628            .await;
629        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
630
631        let csrf_token = res.take_string().await.unwrap();
632        let cookie = res.cookie("salvo.csrf").unwrap();
633
634        let res = TestClient::post("http://127.0.0.1:5801")
635            .send(&service)
636            .await;
637        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
638
639        let mut res = TestClient::post("http://127.0.0.1:5801")
640            .add_header("x-mycsrf-header", csrf_token, true)
641            .add_header("cookie", cookie.to_string(), true)
642            .send(&service)
643            .await;
644        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
645        assert_eq!(res.take_string().await.unwrap(), "POST");
646    }
647
648    #[tokio::test]
649    async fn test_validates_token_in_query() {
650        let csrf = Csrf::new(
651            BcryptCipher::new(),
652            CookieStore::new(),
653            HeaderFinder::new("csrf-token"),
654        );
655        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
656        let service = Service::new(router);
657
658        let mut res = TestClient::get("http://127.0.0.1:5801")
659            .send(&service)
660            .await;
661        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
662
663        let csrf_token = res.take_string().await.unwrap();
664        let cookie = res.cookie("salvo.csrf").unwrap();
665
666        let res = TestClient::post("http://127.0.0.1:5801")
667            .send(&service)
668            .await;
669        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
670
671        let mut res = TestClient::post("http://127.0.0.1:5801?a=1&b=2")
672            .add_header("csrf-token", csrf_token, true)
673            .add_header("cookie", cookie.to_string(), true)
674            .send(&service)
675            .await;
676        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
677        assert_eq!(res.take_string().await.unwrap(), "POST");
678    }
679    #[cfg(feature = "hmac-cipher")]
680    #[tokio::test]
681    async fn test_validates_token_in_alternate_query() {
682        let csrf = Csrf::new(
683            HmacCipher::new(*b"01234567012345670123456701234567"),
684            CookieStore::new(),
685            HeaderFinder::new("my-csrf-token"),
686        );
687        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
688        let service = Service::new(router);
689
690        let mut res = TestClient::get("http://127.0.0.1:5801")
691            .send(&service)
692            .await;
693        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
694
695        let csrf_token = res.take_string().await.unwrap();
696        let cookie = res.cookie("salvo.csrf").unwrap();
697
698        let res = TestClient::post("http://127.0.0.1:5801")
699            .send(&service)
700            .await;
701        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
702
703        let mut res = TestClient::post("http://127.0.0.1:5801?a=1&b=2")
704            .add_header("my-csrf-token", csrf_token, true)
705            .add_header("cookie", cookie.to_string(), true)
706            .send(&service)
707            .await;
708        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
709        assert_eq!(res.take_string().await.unwrap(), "POST");
710    }
711
712    #[cfg(feature = "hmac-cipher")]
713    #[tokio::test]
714    async fn test_validates_token_in_form() {
715        let csrf = Csrf::new(
716            HmacCipher::new(*b"01234567012345670123456701234567"),
717            CookieStore::new(),
718            FormFinder::new("csrf-token"),
719        );
720        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
721        let service = Service::new(router);
722
723        let mut res = TestClient::get("http://127.0.0.1:5801")
724            .send(&service)
725            .await;
726        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
727
728        let csrf_token = res.take_string().await.unwrap();
729        let cookie = res.cookie("salvo.csrf").unwrap();
730
731        let res = TestClient::post("http://127.0.0.1:5801")
732            .send(&service)
733            .await;
734        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
735
736        let mut res = TestClient::post("http://127.0.0.1:5801")
737            .add_header("cookie", cookie.to_string(), true)
738            .form(&[("a", "1"), ("csrf-token", &*csrf_token), ("b", "2")])
739            .send(&service)
740            .await;
741        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
742        assert_eq!(res.take_string().await.unwrap(), "POST");
743    }
744    #[tokio::test]
745    async fn test_validates_token_in_alternate_form() {
746        let csrf = Csrf::new(
747            BcryptCipher::new(),
748            CookieStore::new(),
749            FormFinder::new("my-csrf-token"),
750        );
751        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
752        let service = Service::new(router);
753
754        let mut res = TestClient::get("http://127.0.0.1:5801")
755            .send(&service)
756            .await;
757        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
758
759        let csrf_token = res.take_string().await.unwrap();
760        let cookie = res.cookie("salvo.csrf").unwrap();
761
762        let res = TestClient::post("http://127.0.0.1:5801")
763            .send(&service)
764            .await;
765        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
766        let mut res = TestClient::post("http://127.0.0.1:5801")
767            .add_header("cookie", cookie.to_string(), true)
768            .form(&[("a", "1"), ("my-csrf-token", &*csrf_token), ("b", "2")])
769            .send(&service)
770            .await;
771        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
772        assert_eq!(res.take_string().await.unwrap(), "POST");
773    }
774
775    #[tokio::test]
776    async fn test_rejects_short_token() {
777        let csrf = Csrf::new(
778            BcryptCipher::new(),
779            CookieStore::new(),
780            HeaderFinder::new("x-csrf-token"),
781        );
782        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
783        let service = Service::new(router);
784
785        let res = TestClient::get("http://127.0.0.1:5801")
786            .send(&service)
787            .await;
788        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
789
790        let cookie = res.cookie("salvo.csrf").unwrap();
791
792        let res = TestClient::post("http://127.0.0.1:5801")
793            .send(&service)
794            .await;
795        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
796
797        let res = TestClient::post("http://127.0.0.1:5801")
798            .add_header("x-csrf-token", "aGVsbG8=", true)
799            .add_header(
800                "cookie",
801                cookie.to_string().split_once('.').unwrap().0,
802                true,
803            )
804            .send(&service)
805            .await;
806        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
807    }
808
809    #[tokio::test]
810    async fn test_rejects_invalid_base64_token() {
811        let csrf = Csrf::new(
812            BcryptCipher::new(),
813            CookieStore::new(),
814            HeaderFinder::new("x-csrf-token"),
815        );
816        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
817        let service = Service::new(router);
818
819        let res = TestClient::get("http://127.0.0.1:5801")
820            .send(&service)
821            .await;
822        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
823
824        let cookie = res.cookie("salvo.csrf").unwrap();
825
826        let res = TestClient::post("http://127.0.0.1:5801")
827            .send(&service)
828            .await;
829        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
830
831        let res = TestClient::post("http://127.0.0.1:5801")
832            .add_header("x-csrf-token", "aGVsbG8", true)
833            .add_header(
834                "cookie",
835                cookie.to_string().split_once('.').unwrap().0,
836                true,
837            )
838            .send(&service)
839            .await;
840        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
841    }
842
843    #[tokio::test]
844    async fn test_rejects_mismatched_token() {
845        let csrf = Csrf::new(
846            BcryptCipher::new(),
847            CookieStore::new(),
848            HeaderFinder::new("x-csrf-token"),
849        );
850        let router = Router::new().hoop(csrf).get(get_index).post(post_index);
851        let service = Service::new(router);
852
853        let mut res = TestClient::get("http://127.0.0.1:5801")
854            .send(&service)
855            .await;
856        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
857        let csrf_token = res.take_string().await.unwrap();
858
859        let res = TestClient::get("http://127.0.0.1:5801")
860            .send(&service)
861            .await;
862        assert_eq!(res.status_code.unwrap(), StatusCode::OK);
863        let cookie = res.cookie("salvo.csrf").unwrap();
864
865        let res = TestClient::post("http://127.0.0.1:5801")
866            .send(&service)
867            .await;
868        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
869
870        let res = TestClient::post("http://127.0.0.1:5801")
871            .add_header("x-csrf-token", csrf_token, true)
872            .add_header(
873                "cookie",
874                cookie.to_string().split_once('.').unwrap().0,
875                true,
876            )
877            .send(&service)
878            .await;
879        assert_eq!(res.status_code.unwrap(), StatusCode::FORBIDDEN);
880    }
881}