use std::fmt::Write as _;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll};
use std::time::Instant;
use axum::http::{HeaderName, HeaderValue, Request, Response};
use pin_project_lite::pin_project;
use tower::{Layer, Service};
#[cfg(feature = "db")]
use crate::db::{REQUEST_DB_TIMINGS, RequestDbTimings};
#[cfg(feature = "db")]
type DbTimings = Arc<RequestDbTimings>;
#[cfg(not(feature = "db"))]
type DbTimings = ();
#[cfg(feature = "db")]
type ScopedInner<F> = tokio::task::futures::TaskLocalFuture<Arc<RequestDbTimings>, F>;
#[cfg(not(feature = "db"))]
type ScopedInner<F> = F;
#[cfg(feature = "db")]
fn scope_inner<F: Future>(fut: F) -> (DbTimings, ScopedInner<F>) {
let timings = Arc::new(RequestDbTimings::default());
let scope = REQUEST_DB_TIMINGS.scope(Arc::clone(&timings), fut);
(timings, scope)
}
#[cfg(not(feature = "db"))]
const fn scope_inner<F>(fut: F) -> (DbTimings, ScopedInner<F>) {
((), fut)
}
#[cfg(feature = "db")]
fn read_db_snapshot(timings: &DbTimings) -> (Option<f64>, usize) {
let query_count = timings.query_count.load(Ordering::Relaxed);
let db_micros = timings.total_us.load(Ordering::Relaxed);
let db_ms = if query_count == 0 {
None
} else {
#[allow(clippy::cast_precision_loss)]
Some((db_micros as f64) / 1000.0)
};
(db_ms, query_count)
}
#[cfg(not(feature = "db"))]
#[allow(clippy::trivially_copy_pass_by_ref)]
const fn read_db_snapshot(_timings: &DbTimings) -> (Option<f64>, usize) {
(None, 0)
}
static SERVER_TIMING: HeaderName = HeaderName::from_static("server-timing");
#[derive(Clone, Debug, Default)]
pub struct ServerTimingEmitted(Arc<AtomicBool>);
impl ServerTimingEmitted {
pub(crate) fn mark(&self) {
self.0.store(true, Ordering::Release);
}
fn is_marked(&self) -> bool {
self.0.load(Ordering::Acquire)
}
}
#[derive(Clone, Debug)]
pub struct ServerTimingLayer {
enabled: bool,
fallback: bool,
}
impl ServerTimingLayer {
#[must_use]
pub const fn new(enabled: bool) -> Self {
Self {
enabled,
fallback: false,
}
}
#[must_use]
pub const fn fallback(enabled: bool) -> Self {
Self {
enabled,
fallback: true,
}
}
}
impl<S> Layer<S> for ServerTimingLayer {
type Service = ServerTimingService<S>;
fn layer(&self, inner: S) -> Self::Service {
ServerTimingService {
inner,
enabled: self.enabled,
fallback: self.fallback,
}
}
}
#[derive(Clone, Debug)]
pub struct ServerTimingService<S> {
inner: S,
enabled: bool,
fallback: bool,
}
impl<S, ReqBody, ResBody> Service<Request<ReqBody>> for ServerTimingService<S>
where
S: Service<Request<ReqBody>, Response = Response<ResBody>>,
{
type Response = S::Response;
type Error = S::Error;
type Future = ServerTimingFuture<S::Future>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
if !self.enabled {
return ServerTimingFuture::Disabled {
inner: self.inner.call(req),
};
}
let mut req = req;
let sentinel = if self.fallback {
let sentinel = ServerTimingEmitted::default();
req.extensions_mut().insert(sentinel.clone());
Some(sentinel)
} else {
req.extensions().get::<ServerTimingEmitted>().cloned()
};
let (timings, inner) = scope_inner(self.inner.call(req));
ServerTimingFuture::Enabled {
start: Instant::now(),
timings,
inner,
fallback: self.fallback,
sentinel,
}
}
}
pin_project! {
#[project = ServerTimingProj]
pub enum ServerTimingFuture<F: Future> {
Disabled {
#[pin] inner: F,
},
Enabled {
start: Instant,
timings: DbTimings,
#[pin] inner: ScopedInner<F>,
fallback: bool,
sentinel: Option<ServerTimingEmitted>,
},
}
}
impl<F, ResBody, E> Future for ServerTimingFuture<F>
where
F: Future<Output = Result<Response<ResBody>, E>>,
{
type Output = Result<Response<ResBody>, E>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
match self.project() {
ServerTimingProj::Disabled { inner } => inner.poll(cx),
ServerTimingProj::Enabled {
start,
timings,
inner,
fallback,
sentinel,
} => match inner.poll(cx) {
Poll::Ready(Ok(mut response)) => {
let already_emitted = *fallback
&& sentinel
.as_ref()
.is_some_and(ServerTimingEmitted::is_marked);
if !already_emitted {
let total_ms = start.elapsed().as_secs_f64() * 1000.0;
let (db_ms, query_count) = read_db_snapshot(timings);
let streaming = is_streaming_response(&response);
let value = build_header_value(total_ms, db_ms, query_count, streaming);
if let Ok(hv) = HeaderValue::from_str(&value) {
response.headers_mut().append(SERVER_TIMING.clone(), hv);
}
if !*fallback && let Some(sentinel) = sentinel.as_ref() {
sentinel.mark();
}
}
Poll::Ready(Ok(response))
}
other => other,
},
}
}
}
fn is_streaming_response<B>(response: &Response<B>) -> bool {
response
.headers()
.get(http::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.is_some_and(|ct| {
ct.trim_start()
.to_ascii_lowercase()
.starts_with("text/event-stream")
})
}
pub fn build_header_value(
total_ms: f64,
db_ms: Option<f64>,
query_count: usize,
streaming: bool,
) -> String {
let mut out = format!("total;dur={total_ms:.3}");
if streaming {
return out;
}
if let Some(db) = db_ms
&& query_count > 0
{
let noun = if query_count == 1 { "query" } else { "queries" };
let _ = write!(out, ", db;dur={db:.3};desc=\"{query_count} {noun}\"");
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn build_header_value_total_only_when_no_queries() {
let v = build_header_value(12.345, None, 0, false);
assert_eq!(v, "total;dur=12.345");
}
#[test]
fn build_header_value_db_uses_singular_query_word() {
let v = build_header_value(20.0, Some(3.5), 1, false);
assert_eq!(v, "total;dur=20.000, db;dur=3.500;desc=\"1 query\"");
}
#[test]
fn build_header_value_db_uses_plural_queries_word() {
let v = build_header_value(50.0, Some(21.25), 14, false);
assert_eq!(v, "total;dur=50.000, db;dur=21.250;desc=\"14 queries\"");
}
#[test]
fn build_header_value_streaming_omits_db_even_when_present() {
let v = build_header_value(9.0, Some(4.0), 2, true);
assert_eq!(v, "total;dur=9.000");
}
#[test]
fn build_header_value_zero_query_count_omits_db_even_with_ms() {
let v = build_header_value(1.0, Some(0.0), 0, false);
assert_eq!(v, "total;dur=1.000");
}
#[cfg(feature = "db")]
#[tokio::test]
async fn scope_inner_isolates_nested_scopes() {
let outer = Arc::new(RequestDbTimings::default());
REQUEST_DB_TIMINGS
.scope(Arc::clone(&outer), async {
let (inner, _scope) = scope_inner(std::future::ready(()));
assert!(
!Arc::ptr_eq(&inner, &outer),
"nested scope_inner must isolate with a fresh accumulator, \
not share the outer one (which would double-emit the db metric)"
);
})
.await;
}
#[cfg(feature = "db")]
#[tokio::test]
async fn scope_inner_creates_fresh_accumulator_off_scope() {
let (fresh, _scope) = scope_inner(std::future::ready(()));
assert_eq!(
fresh.query_count.load(Ordering::Relaxed),
0,
"a fresh off-scope accumulator starts with no recorded queries"
);
}
#[cfg(feature = "db")]
#[tokio::test]
async fn nested_scopes_do_not_duplicate_capture_and_stay_isolated() {
use std::sync::Mutex;
use std::time::Duration;
let capture = Arc::new(Mutex::new(Vec::new()));
crate::db::REQUEST_QUERY_CAPTURE
.scope(Arc::clone(&capture), async {
let (outer_timings, outer_scope) = scope_inner(async {
let (inner_timings, inner_scope) = scope_inner(async {
crate::db::record_request_db_query(
Duration::from_micros(500),
Some("SELECT * FROM books WHERE id = $1"),
);
});
inner_scope.await;
inner_timings
});
let inner_timings = outer_scope.await;
assert!(
!Arc::ptr_eq(&outer_timings, &inner_timings),
"nested timing scopes must be distinct Arcs (no shared counts)"
);
})
.await;
let captured = capture
.lock()
.map(|v| v.clone())
.expect("capture mutex poisoned");
assert_eq!(
captured.len(),
1,
"the query must be captured exactly once — no duplication across \
nested Server-Timing scopes"
);
assert_eq!(captured[0].sql, "SELECT * FROM books WHERE id = $1");
}
#[cfg(feature = "db")]
#[tokio::test]
async fn db_accumulator_feeds_header_value() {
use std::time::Duration;
let timings = Arc::new(RequestDbTimings::default());
REQUEST_DB_TIMINGS
.scope(Arc::clone(&timings), async {
crate::db::record_request_db_query(Duration::from_micros(1_500), None);
crate::db::record_request_db_query(Duration::from_micros(2_500), None);
})
.await;
let query_count = timings.query_count.load(Ordering::Relaxed);
let db_micros = timings.total_us.load(Ordering::Relaxed);
assert_eq!(
query_count, 2,
"both instrumented queries should be counted"
);
assert_eq!(
db_micros, 4_000,
"microsecond totals should accumulate (1500 + 2500)"
);
#[allow(clippy::cast_precision_loss)]
let db_ms = (db_micros as f64) / 1000.0;
let value = build_header_value(50.0, Some(db_ms), query_count, false);
assert_eq!(
value, "total;dur=50.000, db;dur=4.000;desc=\"2 queries\"",
"accumulated db time and query count should surface in the header"
);
}
}