1use actix_web::{
12 HttpRequest,
13 dev::ServiceRequest,
14 http::{
15 Method, header,
16 header::{HeaderName, InvalidHeaderName},
17 },
18};
19use rand::RngExt;
20use std::sync::OnceLock;
21use subtle::ConstantTimeEq;
22
23pub const CSRF_COOKIE: &str = "aster_csrf";
27pub const CSRF_HEADER: &str = "X-CSRF-Token";
31const DEFAULT_CSRF_HEADER_LOWER: &str = "x-csrf-token";
32
33const MAX_REQUEST_SCHEME_LEN: usize = 16;
34const MAX_REQUEST_HOST_LEN: usize = 512;
35const MAX_REFERER_AUTHORITY_LEN: usize = MAX_REQUEST_HOST_LEN + 16;
36const MAX_SOURCE_HEADER_LEN: usize = 2048;
37const MAX_SEC_FETCH_SITE_LEN: usize = 64;
38
39#[derive(Debug, Clone, Copy, PartialEq, Eq)]
41pub enum RequestSourceMode {
42 OptionalWhenPresent,
44 Required,
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq)]
50pub enum CsrfErrorKind {
51 TokenNameInvalid,
56 CookieMissing,
58 HeaderMissing,
60 TokenInvalid,
62 RequestSourceUntrusted,
64 RequestOriginUntrusted,
66 RequestRefererUntrusted,
68 RequestSourceMissing,
70 RequestSchemeInvalid,
72 RequestHostInvalid,
74 RequestOriginInvalid,
76 RequestRefererInvalid,
78 RequestHeaderValueInvalid,
80}
81
82#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
84#[error("{message}")]
85pub struct CsrfError {
86 kind: CsrfErrorKind,
87 message: String,
88}
89
90impl CsrfError {
91 fn new(kind: CsrfErrorKind, message: impl Into<String>) -> Self {
92 Self {
93 kind,
94 message: message.into(),
95 }
96 }
97
98 pub fn kind(&self) -> CsrfErrorKind {
100 self.kind
101 }
102
103 pub fn message(&self) -> &str {
105 &self.message
106 }
107}
108
109pub type Result<T> = std::result::Result<T, CsrfError>;
111
112#[derive(Debug, Clone, PartialEq, Eq)]
120pub struct CsrfTokenNames {
121 cookie_name: String,
122 header_name: HeaderName,
123}
124
125impl CsrfTokenNames {
126 pub fn new(cookie_name: impl Into<String>, header_name: impl AsRef<str>) -> Result<Self> {
132 let cookie_name = cookie_name.into();
133 validate_cookie_name(&cookie_name)?;
134 let header_name = parse_header_name(header_name.as_ref())?;
135 Ok(Self {
136 cookie_name,
137 header_name,
138 })
139 }
140
141 pub fn cookie_name(&self) -> &str {
143 &self.cookie_name
144 }
145
146 pub fn header_name(&self) -> &HeaderName {
148 &self.header_name
149 }
150
151 pub fn header_name_str(&self) -> &str {
156 self.header_name.as_str()
157 }
158}
159
160impl Default for CsrfTokenNames {
161 fn default() -> Self {
162 Self {
163 cookie_name: CSRF_COOKIE.to_string(),
164 header_name: HeaderName::from_static(DEFAULT_CSRF_HEADER_LOWER),
165 }
166 }
167}
168
169pub fn default_csrf_token_names() -> &'static CsrfTokenNames {
174 static DEFAULT_NAMES: OnceLock<CsrfTokenNames> = OnceLock::new();
175 DEFAULT_NAMES.get_or_init(CsrfTokenNames::default)
176}
177
178pub fn is_unsafe_method(method: &Method) -> bool {
180 !matches!(
181 *method,
182 Method::GET | Method::HEAD | Method::OPTIONS | Method::TRACE
183 )
184}
185
186pub fn build_csrf_token() -> String {
188 use base64::Engine;
189
190 let mut bytes = [0_u8; 32];
191 rand::rng().fill(&mut bytes);
192 base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
193}
194
195pub fn ensure_double_submit_token(req: &HttpRequest) -> Result<()> {
200 ensure_double_submit_token_with_names(req, default_csrf_token_names())
201}
202
203pub fn ensure_double_submit_token_with_names(
209 req: &HttpRequest,
210 names: &CsrfTokenNames,
211) -> Result<()> {
212 let cookie_token = req
213 .cookie(names.cookie_name())
214 .map(|cookie| cookie.value().to_string())
215 .ok_or_else(|| CsrfError::new(CsrfErrorKind::CookieMissing, "missing CSRF cookie"))?;
216 let header_token = req
217 .headers()
218 .get(names.header_name())
219 .and_then(|value| value.to_str().ok())
220 .map(str::trim)
221 .filter(|value| !value.is_empty())
222 .ok_or_else(|| {
223 CsrfError::new(
224 CsrfErrorKind::HeaderMissing,
225 format!("missing {} header", names.header_name_str()),
226 )
227 })?;
228
229 let tokens_match = header_token.len() == cookie_token.len()
233 && bool::from(header_token.as_bytes().ct_eq(cookie_token.as_bytes()));
234 if !tokens_match {
235 return Err(CsrfError::new(
236 CsrfErrorKind::TokenInvalid,
237 "invalid CSRF token",
238 ));
239 }
240
241 Ok(())
242}
243
244pub fn ensure_service_double_submit_token(req: &ServiceRequest) -> Result<()> {
249 ensure_double_submit_token(req.request())
250}
251
252pub fn ensure_service_double_submit_token_with_names(
255 req: &ServiceRequest,
256 names: &CsrfTokenNames,
257) -> Result<()> {
258 ensure_double_submit_token_with_names(req.request(), names)
259}
260
261pub fn ensure_request_source_allowed(
263 req: &HttpRequest,
264 public_site_origins: &[String],
265 mode: RequestSourceMode,
266) -> Result<()> {
267 let conn = req.connection_info();
268 let request_origin = request_origin(conn.scheme(), conn.host())?;
269 ensure_headers_allowed(
270 header_value(req, header::ORIGIN),
271 header_value(req, header::REFERER),
272 header_value(req, header::HeaderName::from_static("sec-fetch-site")),
273 &request_origin,
274 public_site_origins,
275 mode,
276 )
277}
278
279pub fn ensure_service_request_source_allowed(
281 req: &ServiceRequest,
282 public_site_origins: &[String],
283 mode: RequestSourceMode,
284) -> Result<()> {
285 let conn = req.connection_info();
286 let request_origin = request_origin(conn.scheme(), conn.host())?;
287 ensure_headers_allowed(
288 header_value(req.request(), header::ORIGIN),
289 header_value(req.request(), header::REFERER),
290 header_value(
291 req.request(),
292 header::HeaderName::from_static("sec-fetch-site"),
293 ),
294 &request_origin,
295 public_site_origins,
296 mode,
297 )
298}
299
300pub fn ensure_headers_allowed(
302 origin: Option<&str>,
303 referer: Option<&str>,
304 sec_fetch_site: Option<&str>,
305 request_origin: &str,
306 public_site_origins: &[String],
307 mode: RequestSourceMode,
308) -> Result<()> {
309 let fetch_site = source_header_value(
310 sec_fetch_site,
311 MAX_SEC_FETCH_SITE_LEN,
312 "Sec-Fetch-Site",
313 CsrfErrorKind::RequestHeaderValueInvalid,
314 )?
315 .map(|value| value.to_ascii_lowercase());
316
317 if let Some(fetch_site) = fetch_site.as_deref() {
318 match fetch_site {
319 "same-origin" | "same-site" => {}
320 "cross-site" | "none" => {
321 return Err(CsrfError::new(
322 CsrfErrorKind::RequestSourceUntrusted,
323 "untrusted request source for cookie-authenticated action",
324 ));
325 }
326 _ => {}
327 }
328 }
329 let same_site_fetch = fetch_site.as_deref() == Some("same-site");
330
331 if let Some(origin) = source_header_value(
332 origin,
333 MAX_SOURCE_HEADER_LEN,
334 "Origin",
335 CsrfErrorKind::RequestOriginInvalid,
336 )?
337 .map(|value| normalize_origin(value, CsrfErrorKind::RequestOriginInvalid))
338 .transpose()?
339 {
340 if origin_is_trusted(&origin, request_origin, public_site_origins) {
341 return Ok(());
342 }
343 return Err(CsrfError::new(
344 CsrfErrorKind::RequestOriginUntrusted,
345 "untrusted request origin for cookie-authenticated action",
346 ));
347 }
348
349 if let Some(referer) = trimmed_header_value(referer) {
350 let referer_origin = origin_from_url(referer)?;
351 if origin_is_trusted(&referer_origin, request_origin, public_site_origins) {
352 return Ok(());
353 }
354 return Err(CsrfError::new(
355 CsrfErrorKind::RequestRefererUntrusted,
356 "untrusted request referer for cookie-authenticated action",
357 ));
358 }
359
360 if same_site_fetch {
361 return Err(CsrfError::new(
362 CsrfErrorKind::RequestSourceUntrusted,
363 "missing trusted request source for same-site cookie-authenticated action",
364 ));
365 }
366
367 match mode {
368 RequestSourceMode::OptionalWhenPresent => Ok(()),
369 RequestSourceMode::Required => Err(CsrfError::new(
370 CsrfErrorKind::RequestSourceMissing,
371 "missing request source for cookie-authenticated action",
372 )),
373 }
374}
375
376fn header_value(req: &HttpRequest, name: header::HeaderName) -> Option<&str> {
377 req.headers()
378 .get(name)
379 .and_then(|value| value.to_str().ok())
380}
381
382fn validate_cookie_name(cookie_name: &str) -> Result<()> {
383 if cookie_name.is_empty() {
384 return Err(CsrfError::new(
385 CsrfErrorKind::TokenNameInvalid,
386 "CSRF cookie name cannot be empty",
387 ));
388 }
389 if cookie_name
390 .bytes()
391 .any(|byte| byte <= 0x20 || byte >= 0x7f || b"()<>@,;:\\\"/[]?={}".contains(&byte))
392 {
393 return Err(CsrfError::new(
394 CsrfErrorKind::TokenNameInvalid,
395 "CSRF cookie name contains invalid characters",
396 ));
397 }
398 Ok(())
399}
400
401fn parse_header_name(header_name: &str) -> Result<HeaderName> {
402 HeaderName::from_bytes(header_name.as_bytes()).map_err(header_name_error)
403}
404
405fn header_name_error(error: InvalidHeaderName) -> CsrfError {
406 CsrfError::new(
407 CsrfErrorKind::TokenNameInvalid,
408 format!("invalid CSRF header name: {error}"),
409 )
410}
411
412fn request_origin(scheme: &str, host: &str) -> Result<String> {
413 ensure_value_len(
414 scheme,
415 MAX_REQUEST_SCHEME_LEN,
416 "request scheme",
417 CsrfErrorKind::RequestSchemeInvalid,
418 )?;
419 ensure_value_len(
420 host,
421 MAX_REQUEST_HOST_LEN,
422 "request host",
423 CsrfErrorKind::RequestHostInvalid,
424 )?;
425 normalize_origin(
426 &format!("{scheme}://{host}"),
427 CsrfErrorKind::RequestHostInvalid,
428 )
429 .map_err(|_| CsrfError::new(CsrfErrorKind::RequestHostInvalid, "invalid request host"))
430}
431
432fn normalize_origin(origin: &str, kind: CsrfErrorKind) -> Result<String> {
433 aster_forge_utils::url::normalize_origin(origin, false)
434 .map_err(|_| CsrfError::new(kind, "invalid origin"))
435}
436
437fn origin_is_trusted(origin: &str, request_origin: &str, public_site_origins: &[String]) -> bool {
438 origin == request_origin || public_site_origins.iter().any(|allowed| allowed == origin)
439}
440
441fn source_header_value<'a>(
442 value: Option<&'a str>,
443 max_len: usize,
444 label: &str,
445 kind: CsrfErrorKind,
446) -> Result<Option<&'a str>> {
447 let Some(value) = trimmed_header_value(value) else {
448 return Ok(None);
449 };
450 ensure_value_len(value, max_len, label, kind)?;
451 Ok(Some(value))
452}
453
454fn trimmed_header_value(value: Option<&str>) -> Option<&str> {
455 value.map(str::trim).filter(|value| !value.is_empty())
456}
457
458fn ensure_value_len(value: &str, max_len: usize, label: &str, kind: CsrfErrorKind) -> Result<()> {
459 if value.len() > max_len {
460 return Err(CsrfError::new(
461 kind,
462 format!("{label} exceeds {max_len} bytes"),
463 ));
464 }
465 Ok(())
466}
467
468fn origin_from_url(url: &str) -> Result<String> {
469 let scheme_end = url.find("://").ok_or_else(|| {
470 CsrfError::new(
471 CsrfErrorKind::RequestSchemeInvalid,
472 "invalid Referer header",
473 )
474 })?;
475 let scheme = &url[..scheme_end];
476 ensure_value_len(
477 scheme,
478 MAX_REQUEST_SCHEME_LEN,
479 "Referer scheme",
480 CsrfErrorKind::RequestSchemeInvalid,
481 )?;
482
483 let authority_start = scheme_end + 3;
484 let authority_tail = &url[authority_start..];
485 let authority_end = authority_tail
486 .char_indices()
487 .find_map(|(idx, ch)| matches!(ch, '/' | '?' | '#').then_some(authority_start + idx))
488 .unwrap_or(url.len());
489 let authority = &url[authority_start..authority_end];
490 ensure_value_len(
491 authority,
492 MAX_REFERER_AUTHORITY_LEN,
493 "Referer authority",
494 CsrfErrorKind::RequestRefererInvalid,
495 )?;
496
497 normalize_origin(
498 &format!("{}://{}", scheme.to_ascii_lowercase(), authority),
499 CsrfErrorKind::RequestRefererInvalid,
500 )
501 .map_err(|_| {
502 CsrfError::new(
503 CsrfErrorKind::RequestRefererInvalid,
504 "invalid Referer header",
505 )
506 })
507}
508
509#[cfg(test)]
510mod tests {
511 use actix_web::cookie::Cookie;
512
513 use super::{
514 CSRF_COOKIE, CSRF_HEADER, CsrfErrorKind, CsrfTokenNames, RequestSourceMode,
515 build_csrf_token, ensure_double_submit_token, ensure_double_submit_token_with_names,
516 ensure_headers_allowed, ensure_request_source_allowed,
517 };
518
519 fn host_with_len(len: usize) -> String {
520 let suffix = ".example.com";
521 format!("{}{}", "a".repeat(len - suffix.len()), suffix)
522 }
523
524 #[test]
525 fn accepts_same_origin_and_public_site_origin() {
526 assert!(
527 ensure_headers_allowed(
528 Some("http://localhost"),
529 None,
530 Some("same-origin"),
531 "http://localhost",
532 &["https://forge.example.com".to_string()],
533 RequestSourceMode::Required,
534 )
535 .is_ok()
536 );
537
538 assert!(
539 ensure_headers_allowed(
540 Some("https://forge.example.com"),
541 None,
542 Some("same-origin"),
543 "http://127.0.0.1:3000",
544 &["https://forge.example.com".to_string()],
545 RequestSourceMode::Required,
546 )
547 .is_ok()
548 );
549 }
550
551 #[test]
552 fn same_site_fetch_metadata_requires_trusted_origin_or_referer() {
553 assert!(
554 ensure_headers_allowed(
555 Some("https://panel.example.com"),
556 None,
557 Some("same-site"),
558 "https://api.example.com",
559 &[
560 "https://api.example.com".to_string(),
561 "https://panel.example.com".to_string(),
562 ],
563 RequestSourceMode::OptionalWhenPresent,
564 )
565 .is_ok()
566 );
567
568 assert!(
569 ensure_headers_allowed(
570 None,
571 Some("https://panel.example.com/settings"),
572 Some("same-site"),
573 "https://api.example.com",
574 &[
575 "https://api.example.com".to_string(),
576 "https://panel.example.com".to_string(),
577 ],
578 RequestSourceMode::OptionalWhenPresent,
579 )
580 .is_ok()
581 );
582
583 let err = ensure_headers_allowed(
584 None,
585 None,
586 Some("same-site"),
587 "https://api.example.com",
588 &["https://api.example.com".to_string()],
589 RequestSourceMode::OptionalWhenPresent,
590 )
591 .unwrap_err();
592 assert_eq!(err.kind(), CsrfErrorKind::RequestSourceUntrusted);
593 assert!(err.message().contains("missing trusted request source"));
594 }
595
596 #[test]
597 fn rejects_untrusted_fetch_metadata_values() {
598 for fetch_site in ["cross-site", "none"] {
599 let err = ensure_headers_allowed(
600 None,
601 None,
602 Some(fetch_site),
603 "https://forge.example.com",
604 &[],
605 RequestSourceMode::OptionalWhenPresent,
606 )
607 .unwrap_err();
608 assert_eq!(err.kind(), CsrfErrorKind::RequestSourceUntrusted);
609 assert!(err.message().contains("untrusted request source"));
610 }
611 }
612
613 #[test]
614 fn rejects_untrusted_origin_and_missing_required_source() {
615 let err = ensure_headers_allowed(
616 Some("https://evil.example.com"),
617 None,
618 None,
619 "https://forge.example.com",
620 &[],
621 RequestSourceMode::OptionalWhenPresent,
622 )
623 .unwrap_err();
624 assert_eq!(err.kind(), CsrfErrorKind::RequestOriginUntrusted);
625
626 let err = ensure_headers_allowed(
627 None,
628 None,
629 None,
630 "https://forge.example.com",
631 &[],
632 RequestSourceMode::Required,
633 )
634 .unwrap_err();
635 assert_eq!(err.kind(), CsrfErrorKind::RequestSourceMissing);
636 }
637
638 #[test]
639 fn rejects_oversized_request_source_values_before_normalization() {
640 let max_host = host_with_len(512);
641 let req = actix_web::test::TestRequest::post()
642 .insert_header(("Host", max_host.as_str()))
643 .insert_header(("Origin", format!("http://{max_host}")))
644 .to_http_request();
645 assert!(ensure_request_source_allowed(&req, &[], RequestSourceMode::Required).is_ok());
646
647 let long_host = host_with_len(513);
648 let req = actix_web::test::TestRequest::post()
649 .insert_header(("Host", long_host))
650 .insert_header(("Origin", "https://forge.example.com"))
651 .to_http_request();
652 let err =
653 ensure_request_source_allowed(&req, &[], RequestSourceMode::Required).unwrap_err();
654 assert_eq!(err.kind(), CsrfErrorKind::RequestHostInvalid);
655
656 let req = actix_web::test::TestRequest::post()
657 .insert_header(("Host", "forge.example.com"))
658 .insert_header(("X-Forwarded-Proto", "x".repeat(17)))
659 .insert_header(("Origin", "https://forge.example.com"))
660 .to_http_request();
661 let err =
662 ensure_request_source_allowed(&req, &[], RequestSourceMode::Required).unwrap_err();
663 assert_eq!(err.kind(), CsrfErrorKind::RequestSchemeInvalid);
664
665 let max_origin = format!("https://{}", host_with_len(2040));
666 assert_eq!(max_origin.len(), 2048);
667 assert!(
668 ensure_headers_allowed(
669 Some(&max_origin),
670 None,
671 None,
672 "https://forge.example.com",
673 std::slice::from_ref(&max_origin),
674 RequestSourceMode::OptionalWhenPresent,
675 )
676 .is_ok()
677 );
678
679 let long_origin = format!("https://{}", host_with_len(2041));
680 assert_eq!(long_origin.len(), 2049);
681 let err = ensure_headers_allowed(
682 Some(&long_origin),
683 None,
684 None,
685 "https://forge.example.com",
686 &[],
687 RequestSourceMode::OptionalWhenPresent,
688 )
689 .unwrap_err();
690 assert_eq!(err.kind(), CsrfErrorKind::RequestOriginInvalid);
691
692 let max_referer_authority = host_with_len(528);
693 let max_referer_origin = format!("https://{max_referer_authority}");
694 let max_referer = format!("{max_referer_origin}/settings");
695 assert!(
696 ensure_headers_allowed(
697 None,
698 Some(&max_referer),
699 None,
700 "https://forge.example.com",
701 &[max_referer_origin],
702 RequestSourceMode::OptionalWhenPresent,
703 )
704 .is_ok()
705 );
706
707 let long_referer_authority = format!("https://{}.example.com/settings", "a".repeat(600));
708 let err = ensure_headers_allowed(
709 None,
710 Some(&long_referer_authority),
711 None,
712 "https://forge.example.com",
713 &[],
714 RequestSourceMode::OptionalWhenPresent,
715 )
716 .unwrap_err();
717 assert_eq!(err.kind(), CsrfErrorKind::RequestRefererInvalid);
718
719 let max_fetch_site = "x".repeat(64);
720 assert!(
721 ensure_headers_allowed(
722 None,
723 None,
724 Some(&max_fetch_site),
725 "https://forge.example.com",
726 &[],
727 RequestSourceMode::OptionalWhenPresent,
728 )
729 .is_ok()
730 );
731
732 let long_fetch_site = "x".repeat(65);
733 let err = ensure_headers_allowed(
734 None,
735 None,
736 Some(&long_fetch_site),
737 "https://forge.example.com",
738 &[],
739 RequestSourceMode::OptionalWhenPresent,
740 )
741 .unwrap_err();
742 assert_eq!(err.kind(), CsrfErrorKind::RequestHeaderValueInvalid);
743 }
744
745 #[test]
746 fn accepts_ipv6_request_host_origin_match() {
747 let req = actix_web::test::TestRequest::post()
748 .insert_header(("Host", "[2001:db8::1]:8443"))
749 .insert_header(("Origin", "http://[2001:db8::1]:8443"))
750 .to_http_request();
751
752 assert!(ensure_request_source_allowed(&req, &[], RequestSourceMode::Required).is_ok());
753 }
754
755 #[test]
756 fn referer_source_check_ignores_long_path_after_bounded_origin() {
757 let long_referer = format!("https://forge.example.com/settings/{}", "a".repeat(10_000));
758
759 assert!(
760 ensure_headers_allowed(
761 None,
762 Some(&long_referer),
763 Some("same-origin"),
764 "https://forge.example.com",
765 &[],
766 RequestSourceMode::Required,
767 )
768 .is_ok()
769 );
770 }
771
772 #[test]
773 fn invalid_referer_missing_scheme_reports_invalid_scheme() {
774 let err = ensure_headers_allowed(
775 None,
776 Some("forge.example.com/settings"),
777 None,
778 "https://forge.example.com",
779 &[],
780 RequestSourceMode::OptionalWhenPresent,
781 )
782 .unwrap_err();
783
784 assert_eq!(err.kind(), CsrfErrorKind::RequestSchemeInvalid);
785 }
786
787 #[test]
788 fn accepts_missing_optional_source() {
789 assert!(
790 ensure_headers_allowed(
791 None,
792 None,
793 None,
794 "https://forge.example.com",
795 &[],
796 RequestSourceMode::OptionalWhenPresent,
797 )
798 .is_ok()
799 );
800 }
801
802 #[test]
803 fn build_csrf_token_returns_url_safe_random_value() {
804 let token_a = build_csrf_token();
805 let token_b = build_csrf_token();
806
807 assert_ne!(token_a, token_b);
808 assert!(token_a.len() >= 32);
809 assert!(
810 token_a
811 .chars()
812 .all(|ch| ch.is_ascii_alphanumeric() || ch == '-' || ch == '_')
813 );
814 }
815
816 #[test]
817 fn csrf_token_check_requires_cookie_for_cookie_authenticated_writes() {
818 let req = actix_web::test::TestRequest::post()
819 .uri("/api/v1/auth/profile")
820 .to_http_request();
821
822 let err = ensure_double_submit_token(&req).unwrap_err();
823 assert_eq!(err.kind(), CsrfErrorKind::CookieMissing);
824 }
825
826 #[test]
827 fn csrf_token_check_requires_matching_cookie_and_header() {
828 let req = actix_web::test::TestRequest::patch()
829 .uri("/api/v1/auth/profile")
830 .insert_header(("Origin", "http://localhost"))
831 .cookie(Cookie::new(CSRF_COOKIE, "token-a"))
832 .insert_header((CSRF_HEADER, "token-a"))
833 .to_http_request();
834 assert!(ensure_double_submit_token(&req).is_ok());
835
836 let missing_header = actix_web::test::TestRequest::patch()
837 .uri("/api/v1/auth/profile")
838 .insert_header(("Origin", "http://localhost"))
839 .cookie(Cookie::new(CSRF_COOKIE, "token-a"))
840 .to_http_request();
841 let err = ensure_double_submit_token(&missing_header).unwrap_err();
842 assert_eq!(err.kind(), CsrfErrorKind::HeaderMissing);
843
844 let mismatch = actix_web::test::TestRequest::patch()
845 .uri("/api/v1/auth/profile")
846 .insert_header(("Origin", "http://localhost"))
847 .cookie(Cookie::new(CSRF_COOKIE, "token-a"))
848 .insert_header((CSRF_HEADER, "token-b"))
849 .to_http_request();
850 let err = ensure_double_submit_token(&mismatch).unwrap_err();
851 assert_eq!(err.kind(), CsrfErrorKind::TokenInvalid);
852 }
853
854 #[test]
855 fn csrf_token_check_rejects_tokens_of_different_lengths() {
856 let req = actix_web::test::TestRequest::patch()
857 .uri("/api/v1/auth/profile")
858 .cookie(Cookie::new(CSRF_COOKIE, "token-a"))
859 .insert_header((CSRF_HEADER, "token-a-with-a-longer-value"))
860 .to_http_request();
861 let err = ensure_double_submit_token(&req).unwrap_err();
862 assert_eq!(err.kind(), CsrfErrorKind::TokenInvalid);
863 }
864
865 #[test]
866 fn csrf_token_check_accepts_custom_cookie_and_header_names() {
867 let names = CsrfTokenNames::new("aster_yggdrasil_csrf", "X-Yggdrasil-CSRF-Token")
868 .expect("custom CSRF token names should be valid");
869 assert_eq!(names.cookie_name(), "aster_yggdrasil_csrf");
870 assert_eq!(names.header_name_str(), "x-yggdrasil-csrf-token");
871
872 let req = actix_web::test::TestRequest::patch()
873 .cookie(Cookie::new("aster_yggdrasil_csrf", "token-a"))
874 .insert_header(("X-Yggdrasil-CSRF-Token", "token-a"))
875 .to_http_request();
876 assert!(ensure_double_submit_token_with_names(&req, &names).is_ok());
877
878 let default_req = actix_web::test::TestRequest::patch()
879 .cookie(Cookie::new(CSRF_COOKIE, "token-a"))
880 .insert_header((CSRF_HEADER, "token-a"))
881 .to_http_request();
882 let err = ensure_double_submit_token_with_names(&default_req, &names).unwrap_err();
883 assert_eq!(err.kind(), CsrfErrorKind::CookieMissing);
884 }
885
886 #[test]
887 fn csrf_token_names_reject_invalid_cookie_and_header_names() {
888 let err = CsrfTokenNames::new("", "X-CSRF-Token").unwrap_err();
889 assert_eq!(err.kind(), CsrfErrorKind::TokenNameInvalid);
890
891 let err = CsrfTokenNames::new("aster csrf", "X-CSRF-Token").unwrap_err();
892 assert_eq!(err.kind(), CsrfErrorKind::TokenNameInvalid);
893
894 let err = CsrfTokenNames::new("aster_csrf", "bad header").unwrap_err();
895 assert_eq!(err.kind(), CsrfErrorKind::TokenNameInvalid);
896 }
897}