1#![cfg_attr(test, allow(clippy::unwrap_used))]
2#![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 #[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 #[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 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 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 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 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 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 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 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 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 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 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 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 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
167pub const CSRF_TOKEN_KEY: &str = "salvo.csrf.token";
169
170#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
172#[non_exhaustive]
173pub enum CsrfRotationPolicy {
174 #[default]
176 PerSession,
177 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
188pub trait CsrfStore: Send + Sync + 'static {
190 type Error: StdError + Send + Sync + 'static;
192 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 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
210pub trait CsrfCipher: Send + Sync + 'static {
212 fn verify(&self, token: &str, proof: &str) -> bool;
214 fn generate(&self) -> (String, String);
216
217 fn random_bytes(&self, len: usize) -> Vec<u8> {
219 rand::rng().sample_iter(StandardUniform).take(len).collect()
220 }
221}
222
223pub trait CsrfDepotExt {
225 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
236pub 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 #[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 #[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 #[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 #[inline]
293 #[must_use]
294 pub fn rotation_policy(mut self, policy: CsrfRotationPolicy) -> Self {
295 self.rotation_policy = policy;
296 self
297 }
298
299 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}