1#[cfg(not(feature = "std"))]
38use core::sync::atomic::{AtomicUsize, Ordering};
39
40#[cfg(all(feature = "threads", not(feature = "std")))]
41use crate::error::Error;
42use crate::error::Result;
43
44#[cfg(all(feature = "threads", not(feature = "std")))]
47pub const MAX_THREADS: usize = 64;
48
49#[cfg(not(feature = "std"))]
51const UNSET: usize = usize::MAX;
52
53#[cfg(feature = "std")]
58std::thread_local! {
59 static CURRENT: core::cell::Cell<Option<(usize, usize)>> = const { core::cell::Cell::new(None) };
60}
61
62#[cfg(feature = "std")]
64pub fn current() -> Option<(usize, usize)> {
65 CURRENT.with(|slot| slot.get())
66}
67
68#[cfg(feature = "std")]
70pub fn replace(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
71 CURRENT.with(|slot| slot.replace(next))
72}
73
74#[cfg(not(any(feature = "std", feature = "threads")))]
80pub fn current() -> Option<(usize, usize)> {
81 read_global()
82}
83
84#[cfg(not(any(feature = "std", feature = "threads")))]
86pub fn replace(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
87 replace_global(next)
88}
89
90#[cfg(all(feature = "threads", not(feature = "std")))]
95mod table {
96 use super::*;
97 use alloc::boxed::Box;
98 use once_cell::race::OnceBox;
99
100 pub type ThreadIdHook = Box<dyn Fn() -> usize + Send + Sync>;
102
103 static HOOK: OnceBox<ThreadIdHook> = OnceBox::new();
104
105 pub struct Slot {
107 pub token: AtomicUsize,
108 pub worker: AtomicUsize,
109 pub num_workers: AtomicUsize,
110 }
111
112 impl Slot {
113 const fn new() -> Slot {
114 Slot { token: AtomicUsize::new(UNSET), worker: AtomicUsize::new(UNSET), num_workers: AtomicUsize::new(1) }
115 }
116 }
117
118 #[allow(clippy::declare_interior_mutable_const)]
119 const EMPTY: Slot = Slot::new();
120 pub static SLOTS: [Slot; MAX_THREADS] = [EMPTY; MAX_THREADS];
121
122 pub fn set_hook(hook: impl Fn() -> usize + Send + Sync + 'static) -> Result<()> {
127 HOOK.set(Box::new(Box::new(hook))).map_err(|_| Error::value("the thread id hook has already been set"))
128 }
129
130 pub fn token() -> Option<usize> {
132 HOOK.get().map(|hook| hook())
133 }
134
135 fn find(token: usize) -> Option<&'static Slot> {
137 SLOTS.iter().find(|slot| slot.token.load(Ordering::Acquire) == token)
138 }
139
140 fn claim(token: usize) -> Option<&'static Slot> {
142 SLOTS.iter().find(|slot| slot.token.compare_exchange(UNSET, token, Ordering::AcqRel, Ordering::Acquire).is_ok())
143 }
144
145 pub fn current() -> Option<(usize, usize)> {
147 let Some(token) = token() else {
148 return read_global();
149 };
150 let slot = find(token)?;
151 match slot.worker.load(Ordering::Acquire) {
152 UNSET => None,
153 worker => Some((worker, slot.num_workers.load(Ordering::Acquire))),
154 }
155 }
156
157 pub fn replace(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
159 let Some(token) = token() else {
160 return replace_global(next);
163 };
164
165 let slot = match find(token).or_else(|| claim(token)) {
166 Some(slot) => slot,
167 None => {
168 log::warn!("more than {MAX_THREADS} worker threads; falling back to a shared identity");
171 return replace_global(next);
172 }
173 };
174
175 let previous = match slot.worker.load(Ordering::Acquire) {
176 UNSET => None,
177 worker => Some((worker, slot.num_workers.load(Ordering::Acquire))),
178 };
179 match next {
180 Some((worker, num_workers)) => {
181 slot.num_workers.store(num_workers.max(1), Ordering::Release);
182 slot.worker.store(worker, Ordering::Release);
183 }
184 None => {
185 slot.worker.store(UNSET, Ordering::Release);
186 slot.token.store(UNSET, Ordering::Release);
188 }
189 }
190 previous
191 }
192}
193
194#[cfg(all(feature = "threads", not(feature = "std")))]
195pub use table::{current, replace};
196
197#[cfg(not(feature = "std"))]
199static GLOBAL_FALLBACK: (AtomicUsize, AtomicUsize) = (AtomicUsize::new(UNSET), AtomicUsize::new(1));
200
201#[cfg(not(feature = "std"))]
202fn read_global() -> Option<(usize, usize)> {
203 match GLOBAL_FALLBACK.0.load(Ordering::Acquire) {
204 UNSET => None,
205 worker => Some((worker, GLOBAL_FALLBACK.1.load(Ordering::Acquire))),
206 }
207}
208
209#[cfg(not(feature = "std"))]
210fn replace_global(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
211 let previous = read_global();
212 let (worker, num_workers) = next.unwrap_or((UNSET, 1));
213 GLOBAL_FALLBACK.1.store(num_workers.max(1), Ordering::Release);
214 GLOBAL_FALLBACK.0.store(worker, Ordering::Release);
215 previous
216}
217
218#[cfg(any(feature = "std", not(feature = "threads")))]
224pub fn set_thread_id_hook(_hook: impl Fn() -> usize + Send + Sync + 'static) -> Result<()> {
225 Ok(())
226}
227
228#[cfg(all(feature = "threads", not(feature = "std")))]
232pub fn set_thread_id_hook(hook: impl Fn() -> usize + Send + Sync + 'static) -> Result<()> {
233 table::set_hook(hook)
234}
235
236pub fn with_worker<T>(worker: usize, num_workers: usize, body: impl FnOnce() -> T) -> T {
240 let previous = replace(Some((worker, num_workers)));
241 let result = body();
242 replace(previous);
243 result
244}
245
246pub fn set_worker(worker: usize, num_workers: usize) {
248 replace(Some((worker, num_workers)));
249}
250
251pub fn clear_worker() {
253 replace(None);
254}
255
256#[cfg(test)]
257extern crate std;
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262 #[allow(unused_imports)]
263 use std::vec::Vec;
264
265 #[test]
266 fn binds_and_restores_an_identity() {
267 assert_eq!(current(), None);
268 with_worker(3, 8, || {
269 assert_eq!(current(), Some((3, 8)));
270 with_worker(1, 2, || assert_eq!(current(), Some((1, 2))));
271 assert_eq!(current(), Some((3, 8)), "the outer binding is restored");
272 });
273 assert_eq!(current(), None);
274 }
275
276 #[test]
277 fn set_and_clear_persist_beyond_a_scope() {
278 set_worker(5, 6);
279 assert_eq!(current(), Some((5, 6)));
280 clear_worker();
281 assert_eq!(current(), None);
282 }
283
284 #[cfg(feature = "std")]
285 #[test]
286 fn identities_do_not_leak_between_threads() {
287 set_worker(1, 4);
288 let seen = std::thread::spawn(current).join().expect("thread");
289 assert_eq!(seen, None, "another thread has its own identity");
290 assert_eq!(current(), Some((1, 4)));
291 clear_worker();
292 }
293
294 #[cfg(feature = "std")]
295 #[test]
296 fn many_threads_keep_their_own_identity() {
297 let handles: Vec<_> =
298 std::vec::Vec::from_iter((0..16).map(|i| std::thread::spawn(move || with_worker(i, 16, current))));
299 for (i, handle) in handles.into_iter().enumerate() {
300 assert_eq!(handle.join().expect("thread"), Some((i, 16)));
301 }
302 }
303
304 #[cfg(all(feature = "threads", not(feature = "std")))]
309 #[test]
310 fn the_slot_table_keeps_threads_apart() {
311 std::thread_local! {
312 static TOKEN: core::cell::Cell<usize> = const { core::cell::Cell::new(0) };
313 }
314 set_thread_id_hook(|| TOKEN.with(|t| t.get())).expect("the hook is set once");
315 assert!(set_thread_id_hook(|| 0).is_err(), "the hook cannot be replaced");
316
317 let handles: Vec<_> = std::vec::Vec::from_iter((0..8).map(|i| {
318 std::thread::spawn(move || {
319 TOKEN.with(|t| t.set(i + 1));
321 let seen = with_worker(i, 8, current);
322 (seen, current())
323 })
324 }));
325
326 for (i, handle) in handles.into_iter().enumerate() {
327 let (inside, after) = handle.join().expect("thread");
328 assert_eq!(inside, Some((i, 8)), "thread {i} should see its own identity");
329 assert_eq!(after, None, "the binding is released when the scope ends");
330 }
331 }
332
333 #[cfg(all(not(feature = "threads"), not(feature = "std")))]
336 #[test]
337 fn without_threads_a_single_global_is_used() {
338 set_worker(2, 3);
339 assert_eq!(current(), Some((2, 3)));
340 assert_eq!(std::thread::spawn(current).join().expect("thread"), Some((2, 3)));
341 clear_worker();
342 }
343}