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 ) => {
208 if let Some(context) = current() {
209 context.mark_unmeasured();
210 }
211 }
212 Err(_) => {}
213 }
214}
215
216#[cfg(test)]
217mod tests {
218 use super::*;
219
220 #[tokio::test]
221 async fn count_is_noop_without_scope_and_counts_inside() {
222 count(UsageUnit::EmbedCalls, 3);
223 let ctx = UsageContext::new();
224 scope(ctx.clone(), async {
225 count(UsageUnit::EmbedCalls, 2);
226 count(UsageUnit::FtsPasses, 1);
227 count(UsageUnit::GraphHops, 0);
228 })
229 .await;
230 let snap = ctx.snapshot();
231 assert_eq!(snap["embed_calls"], 2);
232 assert_eq!(snap["fts_passes"], 1);
233 assert!(
234 snap.get("graph_hops").is_none(),
235 "zero counters are omitted"
236 );
237 }
238
239 #[tokio::test]
240 async fn joined_spawned_child_counts_via_explicit_handle() {
241 let ctx = UsageContext::new();
242 scope(ctx.clone(), async {
243 let handle = current().expect("scope armed");
244 let child = tokio::spawn(scope(handle, async {
245 count(UsageUnit::VectorPasses, 2);
246 }));
247 child.await.expect("join child");
248 })
249 .await;
250 assert_eq!(ctx.snapshot()["vector_passes"], 2);
251 }
252
253 #[tokio::test]
254 async fn detached_spawn_without_handle_contributes_nothing() {
255 let ctx = UsageContext::new();
256 scope(ctx.clone(), async {
257 let orphan = tokio::spawn(async {
260 count(UsageUnit::EmbedCalls, 99);
261 });
262 orphan.await.expect("join orphan");
263 })
264 .await;
265 assert_eq!(
266 ctx.snapshot(),
267 serde_json::json!({}),
268 "task-locals do not cross tokio::spawn; only an explicit handle propagates"
269 );
270 }
271
272 #[test]
273 fn saturating_add_never_wraps() {
274 let ctx = UsageContext::new();
275 ctx.add(UsageUnit::EventRows, u64::MAX);
276 ctx.add(UsageUnit::EventRows, 5);
277 assert_eq!(ctx.snapshot()["event_rows"], u64::MAX);
278 }
279
280 #[test]
281 fn unmeasured_mark_wins_after_freeze_and_later_increments() {
282 let ctx = UsageContext::new();
283 ctx.add(UsageUnit::EventRows, 2);
284 let frozen = ctx.freeze();
285 assert_eq!(ctx.shipping_snapshot(), Some(frozen.clone()));
286
287 ctx.clone().mark_unmeasured();
288 assert_eq!(ctx.shipping_snapshot(), None);
289 ctx.add(UsageUnit::EventRows, 3);
290 ctx.add(UsageUnit::EmbedCalls, 1);
291 assert_eq!(ctx.shipping_snapshot(), None);
292 assert_eq!(ctx.freeze(), frozen);
293 assert_eq!(ctx.frozen_or_snapshot(), frozen);
294 assert_eq!(ctx.snapshot()["event_rows"], 5);
295 }
296
297 #[test]
298 fn unmeasured_mark_before_freeze_preserves_internal_readers() {
299 let ctx = UsageContext::new();
300 assert_eq!(ctx.shipping_snapshot(), Some(serde_json::json!({})));
301 ctx.mark_unmeasured();
302 ctx.mark_unmeasured();
303 ctx.add(UsageUnit::DbRoundTrips, 1);
304 assert_eq!(ctx.shipping_snapshot(), None);
305 assert_eq!(ctx.freeze(), serde_json::json!({"db_round_trips": 1}));
306 assert_eq!(ctx.shipping_snapshot(), None);
307 }
308}