1use 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)] fn empty_task_contexts_mutex() -> Mutex<HashMap<tokio::task::Id, ContextSnapshot>> {
44 Mutex::new(HashMap::new())
45}
46
47#[cfg_attr(test, mutants::skip)] fn 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 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}