Skip to main content

webdataset_core/
workers.rs

1//! Which worker the calling thread is, and how that is discovered.
2//!
3//! A sharded pipeline needs to know two things: how many workers are reading
4//! it, and which one this thread is. Everything else — splitting shards across
5//! workers, deriving per-worker seeds — follows from that.
6//!
7//! How the answer is stored depends on what the target offers:
8//!
9//! | build | storage |
10//! |---|---|
11//! | `std` | a thread-local, set for the duration of [`with_worker`] |
12//! | `no_std` + `threads` | a lock-free slot table, keyed by a host-supplied thread token |
13//! | `no_std` without `threads` | a single global, which is all a single-threaded target needs |
14//!
15//! The middle row is the interesting one. Without the standard library there is
16//! no portable way to ask "which thread am I on?", so the host has to say. Call
17//! [`set_thread_id_hook`] once with something that identifies the current
18//! execution context — a task id, a core id, a worker index — and per-worker
19//! state starts working:
20//!
21//! ```
22//! # #[cfg(feature = "threads")] {
23//! use webdataset_core::workers::{set_thread_id_hook, with_worker};
24//! use webdataset_core::worker_info;
25//!
26//! // On a real RTOS this would return the current task's id.
27//! set_thread_id_hook(|| 0).ok();
28//!
29//! let info = with_worker(2, 4, worker_info);
30//! assert_eq!((info.worker, info.num_workers), (2, 4));
31//! # }
32//! ```
33//!
34//! Without a hook the table is bypassed and the single global is used, so a
35//! single-threaded `no_std` program needs no setup at all.
36
37#[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/// How many execution contexts can hold a worker identity at once in a
45/// `no_std` build. Beyond this the global fallback is used.
46#[cfg(all(feature = "threads", not(feature = "std")))]
47pub const MAX_THREADS: usize = 64;
48
49/// The sentinel for "no identity bound".
50#[cfg(not(feature = "std"))]
51const UNSET: usize = usize::MAX;
52
53// ---------------------------------------------------------------------------
54// std: a thread-local, which is exactly what this needs and costs nothing.
55// ---------------------------------------------------------------------------
56
57#[cfg(feature = "std")]
58std::thread_local! {
59    static CURRENT: core::cell::Cell<Option<(usize, usize)>> = const { core::cell::Cell::new(None) };
60}
61
62/// The worker identity bound to this thread, if any.
63#[cfg(feature = "std")]
64pub fn current() -> Option<(usize, usize)> {
65    CURRENT.with(|slot| slot.get())
66}
67
68/// Bind an identity to this thread, returning the one it replaced.
69#[cfg(feature = "std")]
70pub fn replace(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
71    CURRENT.with(|slot| slot.replace(next))
72}
73
74// ---------------------------------------------------------------------------
75// no_std without threads: one global is enough, and it is still atomic.
76// ---------------------------------------------------------------------------
77
78/// The worker identity bound to this program, if any.
79#[cfg(not(any(feature = "std", feature = "threads")))]
80pub fn current() -> Option<(usize, usize)> {
81    read_global()
82}
83
84/// Bind an identity, returning the one it replaced.
85#[cfg(not(any(feature = "std", feature = "threads")))]
86pub fn replace(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
87    replace_global(next)
88}
89
90// ---------------------------------------------------------------------------
91// no_std with threads: a slot table keyed by a host-supplied thread token.
92// ---------------------------------------------------------------------------
93
94#[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    /// Tells the library which execution context is running.
101    pub type ThreadIdHook = Box<dyn Fn() -> usize + Send + Sync>;
102
103    static HOOK: OnceBox<ThreadIdHook> = OnceBox::new();
104
105    /// One bound identity: a token and the worker it maps to.
106    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    /// Install the hook that identifies the current execution context.
123    ///
124    /// Can only be done once; a second call reports an error rather than
125    /// silently changing how every existing binding is interpreted.
126    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    /// The current context's token, if a hook is installed.
131    pub fn token() -> Option<usize> {
132        HOOK.get().map(|hook| hook())
133    }
134
135    /// Find the slot holding `token`.
136    fn find(token: usize) -> Option<&'static Slot> {
137        SLOTS.iter().find(|slot| slot.token.load(Ordering::Acquire) == token)
138    }
139
140    /// Claim a free slot for `token`.
141    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    /// The identity bound to the current context.
146    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    /// Bind `next` to the current context, returning what it replaced.
158    pub fn replace(next: Option<(usize, usize)>) -> Option<(usize, usize)> {
159        let Some(token) = token() else {
160            // No hook: there is no way to tell contexts apart, so the single
161            // global is the best available answer.
162            return replace_global(next);
163        };
164
165        let slot = match find(token).or_else(|| claim(token)) {
166            Some(slot) => slot,
167            None => {
168                // More live contexts than slots. Falling back keeps the
169                // pipeline correct for one of them rather than wrong for all.
170                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                // Release the slot so a later context can reuse it.
187                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/// The global fallback, used when no per-context storage is available.
198#[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/// Install the hook that tells the library which execution context is running.
219///
220/// Only meaningful in a `no_std` build with the `threads` feature; with the
221/// standard library a thread-local is used instead and this is a no-op that
222/// reports success. Can only be set once.
223#[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/// Install the hook that tells the library which execution context is running.
229///
230/// See the [module documentation](self) for what the token should be.
231#[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
236/// Run `body` with this context presenting as worker `worker` of `num_workers`.
237///
238/// The identity is restored afterwards, so nesting works.
239pub 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
246/// Bind a worker identity to this context until it is replaced.
247pub fn set_worker(worker: usize, num_workers: usize) {
248    replace(Some((worker, num_workers)));
249}
250
251/// Forget this context's worker identity.
252pub 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    /// The `no_std` slot table, driven by real threads.
305    ///
306    /// The hook stands in for whatever a real host would provide — an RTOS task
307    /// id, a core number, a worker index.
308    #[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                // A real host would derive this; the test just assigns one.
320                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    /// Without a hook there is nothing to tell contexts apart, so the single
334    /// global is used and shared — correct for a single-threaded program.
335    #[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}