Skip to main content

provide_telemetry/
context.rs

1// SPDX-FileCopyrightText: Copyright (C) 2026 provide.io llc
2// SPDX-License-Identifier: Apache-2.0
3// SPDX-Comment: Part of provide-telemetry.
4//
5
6use std::cell::RefCell;
7use std::collections::{BTreeMap, HashMap};
8use std::sync::atomic::{AtomicU64, Ordering};
9use std::sync::{Mutex, OnceLock};
10
11use serde_json::Value;
12
13#[derive(Clone, Debug, Default, PartialEq)]
14pub struct ContextSnapshot {
15    pub fields: BTreeMap<String, Value>,
16    pub session_id: Option<String>,
17    pub trace_id: Option<String>,
18    pub span_id: Option<String>,
19}
20
21#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Hash)]
22enum ContextScopeKey {
23    Task(tokio::task::Id),
24    #[default]
25    Thread,
26}
27
28#[derive(Default)]
29pub struct ContextGuard {
30    key: ContextScopeKey,
31    previous: ContextSnapshot,
32    epoch: u64,
33}
34
35thread_local! {
36    static THREAD_CONTEXT: RefCell<ContextSnapshot> = RefCell::new(ContextSnapshot::default());
37}
38
39static TASK_CONTEXTS: OnceLock<Mutex<HashMap<tokio::task::Id, ContextSnapshot>>> = OnceLock::new();
40static CONTEXT_EPOCH: AtomicU64 = AtomicU64::new(0);
41
42#[cfg_attr(test, mutants::skip)] // Equivalent mutants only swap in Mutex::default().
43fn empty_task_contexts_mutex() -> Mutex<HashMap<tokio::task::Id, ContextSnapshot>> {
44    Mutex::new(HashMap::new())
45}
46
47#[cfg_attr(test, mutants::skip)] // cargo-mutants synthesizes tokio::task::Id::default(), which cannot compile.
48fn task_contexts() -> &'static Mutex<HashMap<tokio::task::Id, ContextSnapshot>> {
49    TASK_CONTEXTS.get_or_init(empty_task_contexts_mutex)
50}
51
52fn current_scope_key() -> ContextScopeKey {
53    tokio::task::try_id()
54        .map(ContextScopeKey::Task)
55        .unwrap_or(ContextScopeKey::Thread)
56}
57
58fn current_snapshot() -> ContextSnapshot {
59    match current_scope_key() {
60        ContextScopeKey::Task(task_id) => crate::_lock::lock(task_contexts())
61            .get(&task_id)
62            .cloned()
63            .unwrap_or_default(),
64        ContextScopeKey::Thread => THREAD_CONTEXT.with(|ctx| ctx.borrow().clone()),
65    }
66}
67
68fn set_snapshot_for_key(key: ContextScopeKey, snapshot: ContextSnapshot) {
69    match key {
70        ContextScopeKey::Task(task_id) => {
71            let mut map = crate::_lock::lock(task_contexts());
72            if snapshot == ContextSnapshot::default() {
73                // Remove the entry when restoring to default to prevent unbounded growth
74                map.remove(&task_id);
75            } else {
76                map.insert(task_id, snapshot);
77            }
78        }
79        ContextScopeKey::Thread => {
80            THREAD_CONTEXT.with(|ctx| {
81                *ctx.borrow_mut() = snapshot;
82            });
83        }
84    }
85}
86
87fn replace_snapshot(next: ContextSnapshot) -> ContextGuard {
88    let key = current_scope_key();
89    let previous = current_snapshot();
90    let epoch = CONTEXT_EPOCH.load(Ordering::SeqCst);
91    set_snapshot_for_key(key, next);
92    ContextGuard {
93        key,
94        previous,
95        epoch,
96    }
97}
98
99pub fn get_context() -> BTreeMap<String, Value> {
100    current_snapshot().fields
101}
102
103pub fn bind_context<I, K>(fields: I) -> ContextGuard
104where
105    I: IntoIterator<Item = (K, Value)>,
106    K: Into<String>,
107{
108    let mut next = current_snapshot();
109    for (key, value) in fields {
110        next.fields.insert(key.into(), value);
111    }
112    replace_snapshot(next)
113}
114
115pub fn unbind_context(keys: &[&str]) -> ContextGuard {
116    let mut next = current_snapshot();
117    for key in keys {
118        next.fields.remove(*key);
119    }
120    replace_snapshot(next)
121}
122
123pub fn clear_context() -> ContextGuard {
124    let mut next = current_snapshot();
125    next.fields.clear();
126    replace_snapshot(next)
127}
128
129pub fn bind_session_context(session_id: impl Into<String>) -> ContextGuard {
130    let mut next = current_snapshot();
131    let session_id = session_id.into();
132    next.session_id = Some(session_id.clone());
133    next.fields
134        .insert("session_id".to_string(), Value::String(session_id));
135    replace_snapshot(next)
136}
137
138pub fn get_session_id() -> Option<String> {
139    current_snapshot().session_id
140}
141
142pub fn clear_session_context() -> ContextGuard {
143    let mut next = current_snapshot();
144    next.session_id = None;
145    next.fields.remove("session_id");
146    replace_snapshot(next)
147}
148
149pub(crate) fn set_trace_context_internal(
150    trace_id: Option<String>,
151    span_id: Option<String>,
152) -> ContextGuard {
153    let mut next = current_snapshot();
154    next.trace_id = trace_id;
155    next.span_id = span_id;
156    replace_snapshot(next)
157}
158
159pub(crate) fn trace_snapshot() -> ContextSnapshot {
160    current_snapshot()
161}
162
163pub(crate) fn reset_context_for_tests() {
164    CONTEXT_EPOCH.fetch_add(1, Ordering::SeqCst);
165    THREAD_CONTEXT.with(|ctx| {
166        *ctx.borrow_mut() = ContextSnapshot::default();
167    });
168    crate::_lock::lock(task_contexts()).clear();
169}
170
171pub(crate) fn reset_trace_context_for_tests() {
172    CONTEXT_EPOCH.fetch_add(1, Ordering::SeqCst);
173    THREAD_CONTEXT.with(|ctx| {
174        let mut snapshot = ctx.borrow_mut();
175        snapshot.trace_id = None;
176        snapshot.span_id = None;
177    });
178    let mut tasks = crate::_lock::lock(task_contexts());
179    for snapshot in tasks.values_mut() {
180        snapshot.trace_id = None;
181        snapshot.span_id = None;
182    }
183}
184
185impl Drop for ContextGuard {
186    fn drop(&mut self) {
187        if CONTEXT_EPOCH.load(Ordering::SeqCst) == self.epoch {
188            set_snapshot_for_key(self.key, self.previous.clone());
189        }
190    }
191}
192
193#[cfg(test)]
194mod tests {
195    use super::*;
196
197    use serde_json::json;
198
199    use crate::testing::acquire_test_state_lock;
200
201    #[test]
202    fn context_test_bind_context_roundtrip_restores_previous_fields() {
203        let _guard = acquire_test_state_lock();
204        reset_context_for_tests();
205
206        let outer = bind_context([
207            ("request_id".to_string(), json!("req-1")),
208            ("tenant_id".to_string(), json!("tenant-1")),
209        ]);
210        assert_eq!(get_context().get("request_id"), Some(&json!("req-1")));
211        assert_eq!(get_context().get("tenant_id"), Some(&json!("tenant-1")));
212
213        {
214            let cleared = clear_context();
215            assert!(get_context().is_empty());
216            drop(cleared);
217        }
218
219        assert_eq!(get_context().get("request_id"), Some(&json!("req-1")));
220        assert_eq!(get_context().get("tenant_id"), Some(&json!("tenant-1")));
221        drop(outer);
222        assert!(get_context().is_empty());
223    }
224
225    #[test]
226    fn context_test_unbind_context_removes_selected_fields_and_restores_them() {
227        let _guard = acquire_test_state_lock();
228        reset_context_for_tests();
229
230        let outer = bind_context([
231            ("request_id".to_string(), json!("req-1")),
232            ("tenant_id".to_string(), json!("tenant-1")),
233        ]);
234
235        {
236            let unbound = unbind_context(&["tenant_id"]);
237            assert_eq!(get_context().get("request_id"), Some(&json!("req-1")));
238            assert!(!get_context().contains_key("tenant_id"));
239            drop(unbound);
240        }
241
242        assert_eq!(get_context().get("request_id"), Some(&json!("req-1")));
243        assert_eq!(get_context().get("tenant_id"), Some(&json!("tenant-1")));
244        drop(outer);
245        assert!(get_context().is_empty());
246    }
247
248    #[test]
249    fn context_test_a_session_context_roundtrip_restores_session_id() {
250        let _guard = acquire_test_state_lock();
251        reset_context_for_tests();
252
253        let outer = bind_session_context("session-123");
254        assert_eq!(get_session_id(), Some("session-123".to_string()));
255        assert_eq!(get_context().get("session_id"), Some(&json!("session-123")));
256
257        {
258            let cleared = clear_session_context();
259            assert_eq!(get_session_id(), None);
260            assert!(!get_context().contains_key("session_id"));
261            drop(cleared);
262        }
263
264        assert_eq!(get_session_id(), Some("session-123".to_string()));
265        assert_eq!(get_context().get("session_id"), Some(&json!("session-123")));
266        drop(outer);
267        assert_eq!(get_session_id(), None);
268        assert!(!get_context().contains_key("session_id"));
269    }
270
271    #[test]
272    fn context_test_current_scope_key_tracks_tasks_without_leaking_to_threads() {
273        let _guard = acquire_test_state_lock();
274        reset_context_for_tests();
275
276        assert_eq!(current_scope_key(), ContextScopeKey::Thread);
277
278        let runtime = tokio::runtime::Builder::new_current_thread()
279            .enable_all()
280            .build()
281            .expect("runtime");
282        runtime.block_on(async {
283            tokio::spawn(async {
284                let task_id = tokio::task::id();
285                assert_eq!(current_scope_key(), ContextScopeKey::Task(task_id));
286
287                let _bound = bind_context([("task_id".to_string(), json!(task_id.to_string()))]);
288                assert_eq!(
289                    get_context().get("task_id"),
290                    Some(&json!(task_id.to_string()))
291                );
292            })
293            .await
294            .expect("task should complete");
295        });
296
297        assert!(!get_context().contains_key("task_id"));
298    }
299
300    #[test]
301    fn context_test_set_trace_context_roundtrip_restores_previous_snapshot() {
302        let _guard = acquire_test_state_lock();
303        reset_context_for_tests();
304
305        let outer = set_trace_context_internal(
306            Some("outer-trace".to_string()),
307            Some("outer-span".to_string()),
308        );
309        let snapshot = trace_snapshot();
310        assert_eq!(snapshot.trace_id.as_deref(), Some("outer-trace"));
311        assert_eq!(snapshot.span_id.as_deref(), Some("outer-span"));
312
313        {
314            let inner = set_trace_context_internal(
315                Some("inner-trace".to_string()),
316                Some("inner-span".to_string()),
317            );
318            let snapshot = trace_snapshot();
319            assert_eq!(snapshot.trace_id.as_deref(), Some("inner-trace"));
320            assert_eq!(snapshot.span_id.as_deref(), Some("inner-span"));
321            drop(inner);
322        }
323
324        let snapshot = trace_snapshot();
325        assert_eq!(snapshot.trace_id.as_deref(), Some("outer-trace"));
326        assert_eq!(snapshot.span_id.as_deref(), Some("outer-span"));
327        drop(outer);
328
329        let snapshot = trace_snapshot();
330        assert_eq!(snapshot.trace_id, None);
331        assert_eq!(snapshot.span_id, None);
332    }
333
334    #[test]
335    fn context_test_reset_trace_context_clears_task_snapshots_too() {
336        let _guard = acquire_test_state_lock();
337        reset_context_for_tests();
338
339        let runtime = tokio::runtime::Builder::new_current_thread()
340            .enable_all()
341            .build()
342            .expect("runtime");
343        runtime.block_on(async {
344            tokio::spawn(async {
345                let task_id = tokio::task::id();
346                crate::_lock::lock(task_contexts()).insert(
347                    task_id,
348                    ContextSnapshot {
349                        trace_id: Some("trace".to_string()),
350                        span_id: Some("span".to_string()),
351                        ..ContextSnapshot::default()
352                    },
353                );
354
355                reset_trace_context_for_tests();
356
357                let tasks = crate::_lock::lock(task_contexts());
358                let snapshot = tasks
359                    .get(&task_id)
360                    .expect("task snapshot should still exist");
361                assert_eq!(snapshot.trace_id, None);
362                assert_eq!(snapshot.span_id, None);
363            })
364            .await
365            .expect("task should complete");
366        });
367    }
368}