Skip to main content

aster_forge_actix_middleware/
csrf.rs

1//! CSRF helpers for Actix Web services.
2//!
3//! This module implements product-neutral CSRF mechanics: URL-safe token generation, double-submit
4//! cookie/header checks, and request source validation using `Origin`, `Referer`, and
5//! `Sec-Fetch-Site`. Callers map [`CsrfErrorKind`] into their own product error codes.
6//!
7//! The default cookie and header names are compatibility defaults, not a requirement. Services
8//! that share a browser origin should pass [`CsrfTokenNames`] into the `*_with_names` helpers so
9//! each product can use names that will not collide with other products on the same domain.
10
11use 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
23/// Default CSRF cookie name used by compatibility helpers.
24///
25/// Prefer [`CsrfTokenNames`] when a product can share a browser origin with another Aster service.
26pub const CSRF_COOKIE: &str = "aster_csrf";
27/// Default CSRF request header name used by compatibility helpers.
28///
29/// Prefer [`CsrfTokenNames`] when a product can share a browser origin with another Aster service.
30pub 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/// Whether source headers are required or only validated when present.
40#[derive(Debug, Clone, Copy, PartialEq, Eq)]
41pub enum RequestSourceMode {
42    /// Accept requests without source headers, but validate them when present.
43    OptionalWhenPresent,
44    /// Require a trusted `Origin` or `Referer` header for unsafe cookie-authenticated actions.
45    Required,
46}
47
48/// Product-neutral CSRF failure category.
49#[derive(Debug, Clone, Copy, PartialEq, Eq)]
50pub enum CsrfErrorKind {
51    /// A configured CSRF cookie or header name was invalid.
52    ///
53    /// This can only be returned while constructing [`CsrfTokenNames`], not while validating a
54    /// normal request.
55    TokenNameInvalid,
56    /// The CSRF cookie was missing.
57    CookieMissing,
58    /// The CSRF header was missing.
59    HeaderMissing,
60    /// The CSRF cookie and header did not match.
61    TokenInvalid,
62    /// `Sec-Fetch-Site` reported an untrusted source.
63    RequestSourceUntrusted,
64    /// `Origin` was present but not trusted.
65    RequestOriginUntrusted,
66    /// `Referer` was present but not trusted.
67    RequestRefererUntrusted,
68    /// Required source headers were missing.
69    RequestSourceMissing,
70    /// Request scheme was malformed or too long.
71    RequestSchemeInvalid,
72    /// Request host was malformed or too long.
73    RequestHostInvalid,
74    /// Origin header was malformed or too long.
75    RequestOriginInvalid,
76    /// Referer header was malformed or too long.
77    RequestRefererInvalid,
78    /// Generic source header validation failure.
79    RequestHeaderValueInvalid,
80}
81
82/// Error returned by CSRF helper functions.
83#[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    /// Returns the product-neutral failure category.
99    pub fn kind(&self) -> CsrfErrorKind {
100        self.kind
101    }
102
103    /// Returns the diagnostic message.
104    pub fn message(&self) -> &str {
105        &self.message
106    }
107}
108
109/// Result type returned by CSRF helper functions.
110pub type Result<T> = std::result::Result<T, CsrfError>;
111
112/// Cookie and header names used by the double-submit token check.
113///
114/// Services that share a browser origin should configure service-specific names during startup to
115/// avoid cookie/header collisions. Store this value in the product's startup state, app data, or a
116/// process-wide `OnceLock`; do not switch names while a process is serving traffic because active
117/// browser sessions would still hold the previous cookie name and frontend code may still send the
118/// previous header.
119#[derive(Debug, Clone, PartialEq, Eq)]
120pub struct CsrfTokenNames {
121    cookie_name: String,
122    header_name: HeaderName,
123}
124
125impl CsrfTokenNames {
126    /// Builds CSRF token names after validating the cookie and header names.
127    ///
128    /// Cookie names are validated against the conservative RFC 6265 token character set. Header
129    /// names are parsed through Actix's HTTP header type and are stored in canonical lower-case
130    /// form, which makes comparisons and CORS allow-list generation stable.
131    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    /// Returns the configured CSRF cookie name.
142    pub fn cookie_name(&self) -> &str {
143        &self.cookie_name
144    }
145
146    /// Returns the configured CSRF request header name.
147    pub fn header_name(&self) -> &HeaderName {
148        &self.header_name
149    }
150
151    /// Returns the configured CSRF request header name as a lower-case string.
152    ///
153    /// This is useful when building `Access-Control-Allow-Headers` values for browser preflight
154    /// responses.
155    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
169/// Returns the shared default CSRF token names.
170///
171/// This is intended for compatibility helpers and tests. Product integrations that support
172/// service-specific names should construct and store their own [`CsrfTokenNames`] instead.
173pub 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
178/// Returns whether `method` can mutate state and should be protected by CSRF checks.
179pub fn is_unsafe_method(method: &Method) -> bool {
180    !matches!(
181        *method,
182        Method::GET | Method::HEAD | Method::OPTIONS | Method::TRACE
183    )
184}
185
186/// Builds a URL-safe random CSRF token.
187pub 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
195/// Ensures an Actix request contains matching CSRF cookie and header values.
196///
197/// This uses [`default_csrf_token_names`]. Prefer [`ensure_double_submit_token_with_names`] in
198/// products that can run beside another Aster service on the same browser origin.
199pub fn ensure_double_submit_token(req: &HttpRequest) -> Result<()> {
200    ensure_double_submit_token_with_names(req, default_csrf_token_names())
201}
202
203/// Ensures an Actix request contains matching CSRF cookie and header values using custom names.
204///
205/// The helper only performs the double-submit comparison. Product middleware should decide when to
206/// call it, usually for unsafe methods authenticated by cookies. Pair it with request-source
207/// validation to reject cross-site writes before checking the token value.
208pub 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    // The token length is not secret (issued tokens are fixed-length random
230    // values), so the length pre-check leaks nothing; the byte comparison runs
231    // in constant time to avoid a timing side channel on the token value.
232    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
244/// Ensures an Actix service request contains matching CSRF cookie and header values.
245///
246/// This uses [`default_csrf_token_names`]. Prefer [`ensure_service_double_submit_token_with_names`]
247/// in products that configure service-specific token names.
248pub fn ensure_service_double_submit_token(req: &ServiceRequest) -> Result<()> {
249    ensure_double_submit_token(req.request())
250}
251
252/// Ensures an Actix service request contains matching CSRF cookie and header values using custom
253/// names.
254pub 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
261/// Validates source headers for an Actix request.
262pub 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
279/// Validates source headers for an Actix service request.
280pub 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
300/// Validates raw source header values against the request and public-site origins.
301pub 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}