use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use crate::config::ReadYourWrites;
struct Inner {
mode: ReadYourWrites,
incoming_pin: bool,
wrote: AtomicBool,
pin_traced: AtomicBool,
metrics: Option<crate::middleware::MetricsCollector>,
}
#[derive(Clone)]
pub struct RequestPin {
inner: Arc<Inner>,
}
impl RequestPin {
#[must_use]
pub fn new(mode: ReadYourWrites) -> Self {
Self {
inner: Arc::new(Inner {
mode,
incoming_pin: false,
wrote: AtomicBool::new(false),
pin_traced: AtomicBool::new(false),
metrics: None,
}),
}
}
#[must_use]
pub fn new_with_metrics(
mode: ReadYourWrites,
metrics: crate::middleware::MetricsCollector,
) -> Self {
Self {
inner: Arc::new(Inner {
mode,
incoming_pin: false,
wrote: AtomicBool::new(false),
pin_traced: AtomicBool::new(false),
metrics: Some(metrics),
}),
}
}
#[must_use]
pub fn with_session_cookie(
cookie: &str,
keys: &crate::security::config::ResolvedSigningKeys,
window_secs: u64,
) -> Self {
let incoming_pin = parse_session_cookie(cookie, keys, window_secs);
Self {
inner: Arc::new(Inner {
mode: ReadYourWrites::Session,
incoming_pin,
wrote: AtomicBool::new(false),
pin_traced: AtomicBool::new(false),
metrics: None,
}),
}
}
#[doc(hidden)]
#[must_use]
pub fn with_incoming_pin(mode: ReadYourWrites, incoming_pin: bool) -> Self {
Self {
inner: Arc::new(Inner {
mode,
incoming_pin,
wrote: AtomicBool::new(false),
pin_traced: AtomicBool::new(false),
metrics: None,
}),
}
}
#[must_use]
pub fn with_session_cookie_and_metrics(
cookie: &str,
keys: &crate::security::config::ResolvedSigningKeys,
window_secs: u64,
metrics: crate::middleware::MetricsCollector,
) -> Self {
let incoming_pin = parse_session_cookie(cookie, keys, window_secs);
Self {
inner: Arc::new(Inner {
mode: ReadYourWrites::Session,
incoming_pin,
wrote: AtomicBool::new(false),
pin_traced: AtomicBool::new(false),
metrics: Some(metrics),
}),
}
}
#[must_use]
pub fn incoming_pin(&self) -> bool {
self.inner.incoming_pin
}
#[must_use]
pub fn wrote(&self) -> bool {
self.inner.wrote.load(Ordering::Relaxed)
}
}
fn parse_session_cookie(
cookie: &str,
keys: &crate::security::config::ResolvedSigningKeys,
window_secs: u64,
) -> bool {
let Some((ts_str, sig)) = cookie.rsplit_once('.') else {
return false;
};
let Ok(ts) = ts_str.parse::<u64>() else {
return false;
};
if !keys.verify(ts_str.as_bytes(), sig) {
return false;
}
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
if ts > now {
ts - now < 5
} else {
now - ts < window_secs
}
}
#[must_use]
pub fn session_cookie_value(
pin: &RequestPin,
keys: &crate::security::config::ResolvedSigningKeys,
) -> Option<String> {
if !matches!(pin.inner.mode, ReadYourWrites::Session) {
return None;
}
if !pin.inner.wrote.load(Ordering::Relaxed) {
return None;
}
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let ts_str = now.to_string();
let sig = keys.sign(ts_str.as_bytes());
Some(format!("{ts_str}.{sig}"))
}
tokio::task_local! {
static PIN: RequestPin;
}
pub async fn scope<F: std::future::Future>(pin: RequestPin, fut: F) -> F::Output {
PIN.scope(pin, fut).await
}
pub fn mark_write() {
PIN.try_with(|pin| {
pin.inner.wrote.store(true, Ordering::Relaxed);
})
.ok();
}
#[inline]
#[must_use]
pub fn is_pinned() -> bool {
PIN.try_with(|pin| {
matches!(
pin.inner.mode,
ReadYourWrites::Request | ReadYourWrites::Session
) && (pin.inner.wrote.load(Ordering::Relaxed) || pin.inner.incoming_pin)
})
.unwrap_or(false)
}
pub fn note_pin_redirect() {
PIN.try_with(|pin| {
if let Some(ref metrics) = pin.inner.metrics {
metrics.record_read_your_writes_pin();
}
if !pin.inner.pin_traced.swap(true, Ordering::Relaxed) {
tracing::debug!(
target: "autumn::db",
ryw_pinned = true,
"read redirected to primary (read-your-own-writes pin active)"
);
}
})
.ok();
}
pub const RYW_COOKIE_NAME: &str = "autumn.ryw";
pub async fn middleware(
req: axum::http::Request<axum::body::Body>,
next: axum::middleware::Next,
mode: crate::config::ReadYourWrites,
window_secs: u64,
keys: Option<std::sync::Arc<crate::security::config::ResolvedSigningKeys>>,
metrics: crate::middleware::MetricsCollector,
) -> axum::http::Response<axum::body::Body> {
let pin = match mode {
crate::config::ReadYourWrites::Session => {
let cookie_val = extract_ryw_cookie_value(&req);
match (cookie_val, &keys) {
(Some(cv), Some(k)) => {
RequestPin::with_session_cookie_and_metrics(&cv, k, window_secs, metrics)
}
_ => RequestPin::new_with_metrics(mode, metrics),
}
}
crate::config::ReadYourWrites::Request => RequestPin::new_with_metrics(mode, metrics),
crate::config::ReadYourWrites::Off => unreachable!("RYW middleware installed in off mode"),
};
let pin_for_response = pin.clone();
let mut response = scope(pin, next.run(req)).await;
if mode == crate::config::ReadYourWrites::Session
&& let Some(k) = &keys
&& let Some(cv) = session_cookie_value(&pin_for_response, k)
{
let cookie_str = format!(
"{RYW_COOKIE_NAME}={cv}; Max-Age={window_secs}; HttpOnly; \
Secure; SameSite=Lax; Path=/"
);
if let Ok(hv) = axum::http::HeaderValue::from_str(&cookie_str) {
response
.headers_mut()
.append(axum::http::header::SET_COOKIE, hv);
}
}
response
}
fn extract_ryw_cookie_value(req: &axum::http::Request<axum::body::Body>) -> Option<String> {
crate::session::get_cookie(req.headers(), RYW_COOKIE_NAME)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn mark_write_no_op_outside_scope() {
mark_write(); assert!(!is_pinned());
}
#[tokio::test]
async fn is_pinned_false_outside_scope() {
assert!(!is_pinned());
}
#[tokio::test]
async fn is_pinned_request_mode_before_write() {
let pin = RequestPin::new(ReadYourWrites::Request);
scope(pin, async {
assert!(!is_pinned(), "no write yet");
})
.await;
}
#[tokio::test]
async fn is_pinned_request_mode_after_write() {
let pin = RequestPin::new(ReadYourWrites::Request);
scope(pin, async {
mark_write();
assert!(is_pinned(), "write marked");
})
.await;
}
#[tokio::test]
async fn is_pinned_off_mode_never_pins() {
let pin = RequestPin::new(ReadYourWrites::Off);
scope(pin, async {
mark_write();
assert!(!is_pinned(), "off mode must never pin");
})
.await;
}
#[tokio::test]
async fn incoming_pin_pins_without_write() {
let pin = RequestPin::with_incoming_pin(ReadYourWrites::Session, true);
scope(pin, async {
assert!(is_pinned(), "incoming_pin should activate the pin");
})
.await;
}
fn test_keys() -> crate::security::config::ResolvedSigningKeys {
crate::security::config::ResolvedSigningKeys::new(b"test-key-for-ryw-unit".to_vec(), vec![])
}
fn fresh_cookie(keys: &crate::security::config::ResolvedSigningKeys) -> String {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let ts = now.to_string();
let sig = keys.sign(ts.as_bytes());
format!("{ts}.{sig}")
}
#[test]
fn fresh_cookie_sets_incoming_pin() {
let keys = test_keys();
let cookie = fresh_cookie(&keys);
let pin = RequestPin::with_session_cookie(&cookie, &keys, 5);
assert!(
pin.incoming_pin(),
"fresh signed cookie must set incoming_pin"
);
}
#[test]
fn expired_cookie_does_not_set_incoming_pin() {
let keys = test_keys();
let ts = 1_000u64.to_string(); let sig = keys.sign(ts.as_bytes());
let cookie = format!("{ts}.{sig}");
let pin = RequestPin::with_session_cookie(&cookie, &keys, 5);
assert!(
!pin.incoming_pin(),
"expired cookie must NOT set incoming_pin"
);
}
#[test]
fn tampered_cookie_does_not_set_incoming_pin() {
let keys = test_keys();
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let cookie = format!("{now}.deadbeef");
let pin = RequestPin::with_session_cookie(&cookie, &keys, 5);
assert!(
!pin.incoming_pin(),
"cookie with invalid HMAC must NOT set incoming_pin"
);
}
#[test]
fn future_cookie_beyond_skew_tolerance_does_not_set_incoming_pin() {
let keys = test_keys();
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let ts = (now + 60).to_string();
let sig = keys.sign(ts.as_bytes());
let cookie = format!("{ts}.{sig}");
let pin = RequestPin::with_session_cookie(&cookie, &keys, 5);
assert!(
!pin.incoming_pin(),
"far-future cookie must NOT set incoming_pin"
);
}
#[test]
fn malformed_cookie_does_not_set_incoming_pin() {
let keys = test_keys();
for bad in &["", "notimestamp", "abc.def.ghi"] {
let pin = RequestPin::with_session_cookie(bad, &keys, 5);
assert!(
!pin.incoming_pin(),
"malformed cookie {bad:?} must NOT set incoming_pin"
);
}
}
#[test]
fn session_cookie_value_returns_none_for_request_mode() {
let keys = test_keys();
let pin = RequestPin::new(ReadYourWrites::Request);
assert!(session_cookie_value(&pin, &keys).is_none());
}
#[test]
fn session_cookie_value_returns_none_when_no_write() {
let keys = test_keys();
let pin = RequestPin::new(ReadYourWrites::Session);
assert!(session_cookie_value(&pin, &keys).is_none());
}
#[test]
fn session_cookie_value_returns_value_after_write() {
let keys = test_keys();
let pin = RequestPin::new(ReadYourWrites::Session);
pin.inner.wrote.store(true, Ordering::Relaxed);
let val = session_cookie_value(&pin, &keys);
assert!(
val.is_some(),
"session mode + wrote must produce a cookie value"
);
let val = val.unwrap();
let fresh_pin = RequestPin::with_session_cookie(&val, &keys, 5);
assert!(
fresh_pin.incoming_pin(),
"produced cookie must be parseable as a fresh incoming_pin"
);
}
#[tokio::test]
async fn note_pin_redirect_increments_metric_via_new_with_metrics() {
let metrics = crate::middleware::MetricsCollector::new();
let pin = RequestPin::new_with_metrics(ReadYourWrites::Request, metrics.clone());
scope(pin, async {
mark_write();
note_pin_redirect();
})
.await;
assert_eq!(
metrics.snapshot().read_your_writes_pins_total,
1,
"note_pin_redirect must increment the metric counter"
);
}
#[tokio::test]
async fn note_pin_redirect_trace_fires_only_once_per_request() {
let metrics = crate::middleware::MetricsCollector::new();
let pin = RequestPin::new_with_metrics(ReadYourWrites::Request, metrics.clone());
scope(pin, async {
mark_write();
note_pin_redirect();
note_pin_redirect();
note_pin_redirect();
})
.await;
assert_eq!(metrics.snapshot().read_your_writes_pins_total, 3);
}
#[test]
fn with_session_cookie_and_metrics_fresh_sets_incoming_pin() {
let keys = test_keys();
let cookie = fresh_cookie(&keys);
let metrics = crate::middleware::MetricsCollector::new();
let pin = RequestPin::with_session_cookie_and_metrics(&cookie, &keys, 5, metrics);
assert!(
pin.incoming_pin(),
"with_session_cookie_and_metrics: fresh cookie must set incoming_pin"
);
}
#[test]
fn with_session_cookie_and_metrics_expired_does_not_set_incoming_pin() {
let keys = test_keys();
let ts = 1_000u64.to_string();
let sig = keys.sign(ts.as_bytes());
let cookie = format!("{ts}.{sig}");
let metrics = crate::middleware::MetricsCollector::new();
let pin = RequestPin::with_session_cookie_and_metrics(&cookie, &keys, 5, metrics);
assert!(
!pin.incoming_pin(),
"with_session_cookie_and_metrics: expired cookie must NOT set incoming_pin"
);
}
#[tokio::test]
async fn middleware_request_mode_no_set_cookie() {
use axum::{Router, body::Body, http::Request, routing::get};
use tower::ServiceExt;
let metrics = crate::middleware::MetricsCollector::new();
let app =
Router::new()
.route("/", get(|| async { "ok" }))
.layer(axum::middleware::from_fn(move |req, next| {
middleware(req, next, ReadYourWrites::Request, 5, None, metrics.clone())
}));
let req = Request::builder().uri("/").body(Body::empty()).unwrap();
let resp = app.oneshot(req).await.unwrap();
assert!(
!resp.headers().contains_key("set-cookie"),
"request mode must not set a cookie"
);
}
#[tokio::test]
async fn middleware_session_mode_write_sets_cookie() {
use axum::{Router, body::Body, http::Request, routing::get};
use tower::ServiceExt;
let keys = std::sync::Arc::new(test_keys());
let metrics = crate::middleware::MetricsCollector::new();
let keys_clone = keys.clone();
let app = Router::new()
.route(
"/",
get(|| async {
mark_write();
"ok"
}),
)
.layer(axum::middleware::from_fn(move |req, next| {
middleware(
req,
next,
ReadYourWrites::Session,
5,
Some(keys_clone.clone()),
metrics.clone(),
)
}));
let req = Request::builder().uri("/").body(Body::empty()).unwrap();
let resp = app.oneshot(req).await.unwrap();
let set_cookie = resp.headers().get("set-cookie");
assert!(
set_cookie.is_some(),
"session mode + write must set the autumn.ryw cookie"
);
let cv = set_cookie.unwrap().to_str().unwrap();
assert!(
cv.starts_with("autumn.ryw="),
"Set-Cookie must be the autumn.ryw cookie, got: {cv}"
);
}
#[tokio::test]
async fn middleware_session_mode_no_write_no_set_cookie() {
use axum::{Router, body::Body, http::Request, routing::get};
use tower::ServiceExt;
let keys = std::sync::Arc::new(test_keys());
let metrics = crate::middleware::MetricsCollector::new();
let keys_clone = keys.clone();
let app =
Router::new()
.route("/", get(|| async { "ok" }))
.layer(axum::middleware::from_fn(move |req, next| {
middleware(
req,
next,
ReadYourWrites::Session,
5,
Some(keys_clone.clone()),
metrics.clone(),
)
}));
let req = Request::builder().uri("/").body(Body::empty()).unwrap();
let resp = app.oneshot(req).await.unwrap();
assert!(
!resp.headers().contains_key("set-cookie"),
"session mode without write must NOT set a cookie"
);
}
#[tokio::test]
async fn middleware_session_mode_incoming_cookie_pins_read() {
use axum::{Router, body::Body, http::Request, routing::get};
use tower::ServiceExt;
let keys = std::sync::Arc::new(test_keys());
let cookie = fresh_cookie(&keys);
let metrics = crate::middleware::MetricsCollector::new();
let keys_clone = keys.clone();
let saw_pin = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let saw_pin_clone = saw_pin.clone();
let app = Router::new()
.route(
"/",
get(move || {
let flag = saw_pin_clone.clone();
async move {
flag.store(is_pinned(), Ordering::Relaxed);
"ok"
}
}),
)
.layer(axum::middleware::from_fn(move |req, next| {
middleware(
req,
next,
ReadYourWrites::Session,
5,
Some(keys_clone.clone()),
metrics.clone(),
)
}));
let req = Request::builder()
.uri("/")
.header("cookie", format!("autumn.ryw={cookie}"))
.body(Body::empty())
.unwrap();
app.oneshot(req).await.unwrap();
assert!(
saw_pin.load(Ordering::Relaxed),
"a valid incoming autumn.ryw cookie must activate the pin inside the handler"
);
}
}