use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use serde_json::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UsageUnit {
EmbedCalls,
FtsPasses,
VectorPasses,
GraphHops,
DbRoundTrips,
AnnJobsConsumed,
EventRows,
}
#[derive(Debug, Default)]
struct UsageInner {
frozen: std::sync::OnceLock<Value>,
embed_calls: AtomicU64,
fts_passes: AtomicU64,
vector_passes: AtomicU64,
graph_hops: AtomicU64,
db_round_trips: AtomicU64,
ann_jobs_consumed: AtomicU64,
event_rows: AtomicU64,
}
#[derive(Debug, Clone, Default)]
pub struct UsageContext {
inner: Arc<UsageInner>,
}
impl UsageContext {
pub fn new() -> Self {
Self::default()
}
pub fn add(&self, unit: UsageUnit, n: u64) {
let cell = match unit {
UsageUnit::EmbedCalls => &self.inner.embed_calls,
UsageUnit::FtsPasses => &self.inner.fts_passes,
UsageUnit::VectorPasses => &self.inner.vector_passes,
UsageUnit::GraphHops => &self.inner.graph_hops,
UsageUnit::DbRoundTrips => &self.inner.db_round_trips,
UsageUnit::AnnJobsConsumed => &self.inner.ann_jobs_consumed,
UsageUnit::EventRows => &self.inner.event_rows,
};
let mut cur = cell.load(Ordering::Relaxed);
loop {
let next = cur.saturating_add(n);
match cell.compare_exchange_weak(cur, next, Ordering::Relaxed, Ordering::Relaxed) {
Ok(_) => break,
Err(observed) => cur = observed,
}
}
}
pub fn freeze(&self) -> Value {
self.inner.frozen.get_or_init(|| self.snapshot()).clone()
}
pub fn frozen_or_snapshot(&self) -> Value {
match self.inner.frozen.get() {
Some(v) => v.clone(),
None => self.snapshot(),
}
}
pub fn snapshot(&self) -> Value {
let mut map = serde_json::Map::new();
let mut put = |key: &str, cell: &AtomicU64| {
let v = cell.load(Ordering::Relaxed);
if v > 0 {
map.insert(key.to_string(), Value::from(v));
}
};
put("embed_calls", &self.inner.embed_calls);
put("fts_passes", &self.inner.fts_passes);
put("vector_passes", &self.inner.vector_passes);
put("graph_hops", &self.inner.graph_hops);
put("db_round_trips", &self.inner.db_round_trips);
put("ann_jobs_consumed", &self.inner.ann_jobs_consumed);
put("event_rows", &self.inner.event_rows);
Value::Object(map)
}
}
tokio::task_local! {
static CURRENT: UsageContext;
}
pub async fn scope<F: std::future::Future>(ctx: UsageContext, fut: F) -> F::Output {
CURRENT.scope(ctx, fut).await
}
pub fn current() -> Option<UsageContext> {
CURRENT.try_with(Clone::clone).ok()
}
pub fn count(unit: UsageUnit, n: u64) {
if n == 0 {
return;
}
if let Ok(ctx) = CURRENT.try_with(Clone::clone) {
ctx.add(unit, n);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn count_is_noop_without_scope_and_counts_inside() {
count(UsageUnit::EmbedCalls, 3);
let ctx = UsageContext::new();
scope(ctx.clone(), async {
count(UsageUnit::EmbedCalls, 2);
count(UsageUnit::FtsPasses, 1);
count(UsageUnit::GraphHops, 0);
})
.await;
let snap = ctx.snapshot();
assert_eq!(snap["embed_calls"], 2);
assert_eq!(snap["fts_passes"], 1);
assert!(
snap.get("graph_hops").is_none(),
"zero counters are omitted"
);
}
#[tokio::test]
async fn joined_spawned_child_counts_via_explicit_handle() {
let ctx = UsageContext::new();
scope(ctx.clone(), async {
let handle = current().expect("scope armed");
let child = tokio::spawn(scope(handle, async {
count(UsageUnit::VectorPasses, 2);
}));
child.await.expect("join child");
})
.await;
assert_eq!(ctx.snapshot()["vector_passes"], 2);
}
#[tokio::test]
async fn detached_spawn_without_handle_contributes_nothing() {
let ctx = UsageContext::new();
scope(ctx.clone(), async {
let orphan = tokio::spawn(async {
count(UsageUnit::EmbedCalls, 99);
});
orphan.await.expect("join orphan");
})
.await;
assert_eq!(
ctx.snapshot(),
serde_json::json!({}),
"task-locals do not cross tokio::spawn; only an explicit handle propagates"
);
}
#[test]
fn saturating_add_never_wraps() {
let ctx = UsageContext::new();
ctx.add(UsageUnit::EventRows, u64::MAX);
ctx.add(UsageUnit::EventRows, 5);
assert_eq!(ctx.snapshot()["event_rows"], u64::MAX);
}
}