1use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
34use std::sync::Arc;
35
36use serde_json::Value;
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40pub enum UsageUnit {
41 EmbedCalls,
44 FtsPasses,
46 VectorPasses,
48 GraphHops,
51 DbRoundTrips,
53 AnnJobsConsumed,
56 EventRows,
60}
61
62#[derive(Debug, Default)]
63struct UsageInner {
64 frozen: std::sync::OnceLock<Value>,
65 unmeasured: AtomicBool,
66 embed_calls: AtomicU64,
67 fts_passes: AtomicU64,
68 vector_passes: AtomicU64,
69 graph_hops: AtomicU64,
70 db_round_trips: AtomicU64,
71 ann_jobs_consumed: AtomicU64,
72 event_rows: AtomicU64,
73}
74
75#[derive(Debug, Clone, Default)]
77pub struct UsageContext {
78 inner: Arc<UsageInner>,
79}
80
81impl UsageContext {
82 pub fn new() -> Self {
83 Self::default()
84 }
85
86 pub fn mark_unmeasured(&self) {
88 self.inner.unmeasured.store(true, Ordering::Release);
89 }
90
91 pub fn shipping_snapshot(&self) -> Option<Value> {
94 if self.inner.unmeasured.load(Ordering::Acquire) {
95 None
96 } else {
97 Some(self.frozen_or_snapshot())
98 }
99 }
100
101 pub fn add(&self, unit: UsageUnit, n: u64) {
103 let cell = match unit {
104 UsageUnit::EmbedCalls => &self.inner.embed_calls,
105 UsageUnit::FtsPasses => &self.inner.fts_passes,
106 UsageUnit::VectorPasses => &self.inner.vector_passes,
107 UsageUnit::GraphHops => &self.inner.graph_hops,
108 UsageUnit::DbRoundTrips => &self.inner.db_round_trips,
109 UsageUnit::AnnJobsConsumed => &self.inner.ann_jobs_consumed,
110 UsageUnit::EventRows => &self.inner.event_rows,
111 };
112 let mut cur = cell.load(Ordering::Relaxed);
113 loop {
114 let next = cur.saturating_add(n);
115 match cell.compare_exchange_weak(cur, next, Ordering::Relaxed, Ordering::Relaxed) {
116 Ok(_) => break,
117 Err(observed) => cur = observed,
118 }
119 }
120 }
121
122 pub fn freeze(&self) -> Value {
129 self.inner.frozen.get_or_init(|| self.snapshot()).clone()
130 }
131
132 pub fn frozen_or_snapshot(&self) -> Value {
134 match self.inner.frozen.get() {
135 Some(v) => v.clone(),
136 None => self.snapshot(),
137 }
138 }
139
140 pub fn snapshot(&self) -> Value {
145 let mut map = serde_json::Map::new();
146 let mut put = |key: &str, cell: &AtomicU64| {
147 let v = cell.load(Ordering::Relaxed);
148 if v > 0 {
149 map.insert(key.to_string(), Value::from(v));
150 }
151 };
152 put("embed_calls", &self.inner.embed_calls);
153 put("fts_passes", &self.inner.fts_passes);
154 put("vector_passes", &self.inner.vector_passes);
155 put("graph_hops", &self.inner.graph_hops);
156 put("db_round_trips", &self.inner.db_round_trips);
157 put("ann_jobs_consumed", &self.inner.ann_jobs_consumed);
158 put("event_rows", &self.inner.event_rows);
159 Value::Object(map)
160 }
161}
162
163tokio::task_local! {
164 static CURRENT: UsageContext;
165}
166
167pub async fn scope<F: std::future::Future>(ctx: UsageContext, fut: F) -> F::Output {
169 CURRENT.scope(ctx, fut).await
170}
171
172pub fn current() -> Option<UsageContext> {
176 CURRENT.try_with(Clone::clone).ok()
177}
178
179pub fn count(unit: UsageUnit, n: u64) {
183 if n == 0 {
184 return;
185 }
186 if let Ok(ctx) = CURRENT.try_with(Clone::clone) {
187 ctx.add(unit, n);
188 }
189}
190
191pub fn account_event_write(outcome: Result<u64, &crate::StorageError>) {
196 match outcome {
197 Ok(committed_rows) => count(UsageUnit::EventRows, committed_rows),
198 Err(
199 crate::StorageError::WriterTaskRequestFailed {
200 request_state: crate::WriterTaskRequestState::SideEffectsUnknown,
201 ..
202 }
203 | crate::StorageError::WriterTaskTerminated {
204 request_state: crate::WriterTaskRequestState::SideEffectsUnknown,
205 },
206 ) => {
207 if let Some(context) = current() {
208 context.mark_unmeasured();
209 }
210 }
211 Err(_) => {}
212 }
213}
214
215#[cfg(test)]
216mod tests {
217 use super::*;
218
219 #[tokio::test]
220 async fn count_is_noop_without_scope_and_counts_inside() {
221 count(UsageUnit::EmbedCalls, 3);
222 let ctx = UsageContext::new();
223 scope(ctx.clone(), async {
224 count(UsageUnit::EmbedCalls, 2);
225 count(UsageUnit::FtsPasses, 1);
226 count(UsageUnit::GraphHops, 0);
227 })
228 .await;
229 let snap = ctx.snapshot();
230 assert_eq!(snap["embed_calls"], 2);
231 assert_eq!(snap["fts_passes"], 1);
232 assert!(
233 snap.get("graph_hops").is_none(),
234 "zero counters are omitted"
235 );
236 }
237
238 #[tokio::test]
239 async fn joined_spawned_child_counts_via_explicit_handle() {
240 let ctx = UsageContext::new();
241 scope(ctx.clone(), async {
242 let handle = current().expect("scope armed");
243 let child = tokio::spawn(scope(handle, async {
244 count(UsageUnit::VectorPasses, 2);
245 }));
246 child.await.expect("join child");
247 })
248 .await;
249 assert_eq!(ctx.snapshot()["vector_passes"], 2);
250 }
251
252 #[tokio::test]
253 async fn detached_spawn_without_handle_contributes_nothing() {
254 let ctx = UsageContext::new();
255 scope(ctx.clone(), async {
256 let orphan = tokio::spawn(async {
259 count(UsageUnit::EmbedCalls, 99);
260 });
261 orphan.await.expect("join orphan");
262 })
263 .await;
264 assert_eq!(
265 ctx.snapshot(),
266 serde_json::json!({}),
267 "task-locals do not cross tokio::spawn; only an explicit handle propagates"
268 );
269 }
270
271 #[test]
272 fn saturating_add_never_wraps() {
273 let ctx = UsageContext::new();
274 ctx.add(UsageUnit::EventRows, u64::MAX);
275 ctx.add(UsageUnit::EventRows, 5);
276 assert_eq!(ctx.snapshot()["event_rows"], u64::MAX);
277 }
278
279 #[test]
280 fn unmeasured_mark_wins_after_freeze_and_later_increments() {
281 let ctx = UsageContext::new();
282 ctx.add(UsageUnit::EventRows, 2);
283 let frozen = ctx.freeze();
284 assert_eq!(ctx.shipping_snapshot(), Some(frozen.clone()));
285
286 ctx.clone().mark_unmeasured();
287 assert_eq!(ctx.shipping_snapshot(), None);
288 ctx.add(UsageUnit::EventRows, 3);
289 ctx.add(UsageUnit::EmbedCalls, 1);
290 assert_eq!(ctx.shipping_snapshot(), None);
291 assert_eq!(ctx.freeze(), frozen);
292 assert_eq!(ctx.frozen_or_snapshot(), frozen);
293 assert_eq!(ctx.snapshot()["event_rows"], 5);
294 }
295
296 #[test]
297 fn unmeasured_mark_before_freeze_preserves_internal_readers() {
298 let ctx = UsageContext::new();
299 assert_eq!(ctx.shipping_snapshot(), Some(serde_json::json!({})));
300 ctx.mark_unmeasured();
301 ctx.mark_unmeasured();
302 ctx.add(UsageUnit::DbRoundTrips, 1);
303 assert_eq!(ctx.shipping_snapshot(), None);
304 assert_eq!(ctx.freeze(), serde_json::json!({"db_round_trips": 1}));
305 assert_eq!(ctx.shipping_snapshot(), None);
306 }
307}