use std::cell::RefCell;
use std::future::Future;
#[derive(Default)]
struct PageStaging {
cache_control: Option<String>,
csrf: String,
csp_nonce: String,
status: Option<u16>,
redirect: Option<String>,
title: Option<String>,
description: Option<String>,
robots: Option<String>,
canonical: Option<String>,
json_ld: Option<String>,
dir: Option<String>,
theme: Option<String>,
}
tokio::task_local! {
static PAGE_STAGING: RefCell<PageStaging>;
}
thread_local! {
static FALLBACK_STAGING: RefCell<PageStaging> = RefCell::new(PageStaging::default());
}
pub async fn scope_page_staging<F: Future>(fut: F) -> F::Output {
PAGE_STAGING
.scope(RefCell::new(PageStaging::default()), fut)
.await
}
fn with_staging<R>(f: impl FnOnce(&mut PageStaging) -> R) -> R {
let mut f = Some(f);
match PAGE_STAGING.try_with(|cell| (f.take().expect("staging fn"))(&mut cell.borrow_mut())) {
Ok(out) => out,
Err(_) => {
FALLBACK_STAGING.with(|cell| (f.take().expect("staging fn"))(&mut cell.borrow_mut()))
}
}
}
pub fn clear_request_staging() {
with_staging(|s| *s = PageStaging::default());
}
pub fn stage_response_status(status: u16) {
with_staging(|s| {
s.status = Some(status);
if status != 200 && s.robots.is_none() {
s.robots = Some("noindex".into());
}
});
}
pub fn take_response_status() -> Option<u16> {
with_staging(|s| s.status.take())
}
pub fn stage_response_redirect(location: impl Into<String>) {
let location = location.into();
with_staging(|s| s.redirect = Some(location));
}
pub fn take_response_redirect() -> Option<String> {
with_staging(|s| s.redirect.take())
}
pub fn stage_response_cache_control(value: impl Into<String>) {
let value = value.into();
with_staging(|s| s.cache_control = Some(value));
}
pub fn sanitize_cache_for_session(
cache: Option<String>,
sets_session_cookie: bool,
) -> Option<String> {
if !sets_session_cookie {
return cache;
}
if cache.as_deref().is_some_and(|c| {
let lower = c.to_ascii_lowercase();
lower.contains("public") || lower.contains("max-age")
}) {
tracing::warn!("overriding Cache-Control to private, no-store — page sets CSRF cookie");
}
Some("private, no-store".to_string())
}
pub fn take_response_cache_control() -> Option<String> {
with_staging(|s| s.cache_control.take())
}
pub fn set_page_title(title: impl Into<String>) {
let title = title.into();
with_staging(|s| s.title = Some(title));
}
pub fn page_title_override() -> Option<String> {
with_staging(|s| s.title.clone())
}
pub fn set_page_description(description: impl Into<String>) {
let description = description.into();
with_staging(|s| s.description = Some(description));
}
pub fn page_description_override() -> Option<String> {
with_staging(|s| s.description.clone())
}
pub fn set_page_robots(robots: impl Into<String>) {
let robots = robots.into();
with_staging(|s| s.robots = Some(robots));
}
pub fn page_robots_override() -> Option<String> {
with_staging(|s| s.robots.clone())
}
pub fn set_page_canonical(canonical: impl Into<String>) {
let canonical = canonical.into();
with_staging(|s| s.canonical = Some(canonical));
}
pub fn page_canonical_override() -> Option<String> {
with_staging(|s| s.canonical.clone())
}
pub fn set_page_json_ld(json_ld: impl Into<String>) {
let json_ld = json_ld.into();
with_staging(|s| s.json_ld = Some(json_ld));
}
pub fn page_json_ld_override() -> Option<String> {
with_staging(|s| s.json_ld.clone())
}
pub fn set_page_dir(dir: impl Into<String>) {
let dir = dir.into();
with_staging(|s| s.dir = Some(dir));
}
pub fn page_dir_override() -> Option<String> {
with_staging(|s| s.dir.clone())
}
pub fn set_page_theme(theme: impl Into<String>) {
let theme = theme.into();
with_staging(|s| s.theme = Some(theme));
}
pub fn page_theme_override() -> Option<String> {
with_staging(|s| s.theme.clone())
}
pub fn stage_page_csrf(token: impl Into<String>) {
let token = token.into();
with_staging(|s| s.csrf = token);
}
pub fn page_csrf() -> String {
with_staging(|s| s.csrf.clone())
}
pub fn stage_page_csp_nonce(nonce: impl Into<String>) {
let nonce = nonce.into();
with_staging(|s| s.csp_nonce = nonce);
}
pub fn page_csp_nonce() -> String {
with_staging(|s| s.csp_nonce.clone())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sanitize_cache_for_session_overrides_public() {
let out = sanitize_cache_for_session(Some("public, max-age=3600".into()), true);
assert_eq!(out.as_deref(), Some("private, no-store"));
}
#[test]
fn redirect_staging_round_trips_and_is_consumed_once() {
clear_request_staging();
assert_eq!(take_response_redirect(), None);
stage_response_redirect("/login");
assert_eq!(take_response_redirect(), Some("/login".to_string()));
assert_eq!(take_response_redirect(), None);
}
#[test]
fn sanitize_cache_for_session_keeps_when_no_cookie() {
let cache = "public, max-age=60".to_string();
let out = sanitize_cache_for_session(Some(cache.clone()), false);
assert_eq!(out, Some(cache));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn page_staging_isolated_per_scoped_task() {
let mut handles = Vec::new();
for i in 0..32u32 {
handles.push(tokio::spawn(async move {
scope_page_staging(async {
let token = format!("csrf-{i:04}");
let nonce = format!("nonce-{i:04}");
stage_page_csrf(token.clone());
stage_page_csp_nonce(nonce.clone());
tokio::task::yield_now().await;
assert_eq!(page_csrf(), token);
assert_eq!(page_csp_nonce(), nonce);
})
.await
}));
}
for h in handles {
h.await.unwrap();
}
}
}