Skip to main content

shuttle_std/
thread.rs

1//! Shuttle's implementation of [`std::thread`].
2
3use shuttle_engine::runtime::execution::ExecutionState;
4use shuttle_engine::runtime::task::TaskId;
5use shuttle_engine::runtime::thread;
6use std::marker::PhantomData;
7use std::panic::Location;
8use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
9use std::time::Duration;
10
11pub use std::thread::{panicking, Result};
12
13/// A unique identifier for a running thread
14#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
15pub struct ThreadId {
16    // TODO Should we add an execution id here, like Loom does?
17    task_id: TaskId,
18}
19
20impl From<ThreadId> for usize {
21    fn from(id: ThreadId) -> usize {
22        id.task_id.into()
23    }
24}
25
26/// A handle to a thread.
27#[derive(Debug, Clone)]
28pub struct Thread {
29    name: Option<String>,
30    id: ThreadId,
31}
32
33impl Thread {
34    /// Gets the thread's name.
35    pub fn name(&self) -> Option<&str> {
36        self.name.as_deref()
37    }
38
39    /// Gets the thread's unique identifier
40    pub fn id(&self) -> ThreadId {
41        self.id
42    }
43
44    /// Atomically makes the handle's token available if it is not already.
45    pub fn unpark(&self) {
46        thread::switch();
47
48        ExecutionState::with(|s| {
49            s.get_mut(self.id.task_id).unpark();
50        });
51    }
52}
53
54/// A scope to spawn scoped threads in.
55///
56/// See [`scope`] for details.
57pub struct Scope<'scope, 'env: 'scope> {
58    num_running_threads: AtomicUsize,
59    main_task: TaskId,
60    scope: PhantomData<&'scope mut &'scope ()>,
61    env: PhantomData<&'env mut &'env ()>,
62}
63
64impl std::fmt::Debug for Scope<'_, '_> {
65    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66        f.debug_struct("Scope")
67            .field("num_running_threads", &self.num_running_threads.load(Ordering::Relaxed))
68            .field("main_thread", &self.main_task)
69            .finish_non_exhaustive()
70    }
71}
72
73impl<'scope> Scope<'scope, '_> {
74    /// Spawns a new thread within a scope, returning a [`ScopedJoinHandle`] for it.
75    ///
76    /// Unlike non-scoped threads, threads spawned with this function may
77    /// borrow non-`'static` data from the outside the scope. See [`scope`] for
78    /// details.
79    #[track_caller]
80    pub fn spawn<F, T>(&'scope self, f: F) -> ScopedJoinHandle<'scope, T>
81    where
82        F: FnOnce() -> T + Send + 'scope,
83        T: Send + 'scope,
84    {
85        self.num_running_threads.fetch_add(1, Ordering::Relaxed);
86
87        let finished = std::sync::Arc::new(AtomicBool::new(false));
88        let scope_closure = {
89            let finished = finished.clone();
90            move || {
91                let ret = f();
92
93                if ExecutionState::with(|s| s.exit_current_truncates_execution()) {
94                    thread::switch();
95                }
96
97                finished.store(true, Ordering::Relaxed);
98
99                if self.num_running_threads.fetch_sub(1, Ordering::Relaxed) == 1 {
100                    ExecutionState::with(|s| s.get_mut(self.main_task).unblock());
101                }
102
103                ret
104            }
105        };
106
107        // Note: Scoped threads wrap their inner function in some additional logic (above) to update the `finished` variable on termination.
108        // This logic is expected to run atomically with scoped thread termination, so there *cannot* be a switch after the atomic store.
109        // To avoid violating this invariant, we pass `switch_before_exit = false` (below). Instead, we provide our own context switch on exit
110        // (above) in the `scope_closure` *before* setting `finished` to be `true`.
111        // SAFETY: main task is blocked until all scoped closures complete so all captured references remain valid
112        ScopedJoinHandle {
113            handle: unsafe { spawn_named_unchecked(scope_closure, None, None, false, Location::caller()) },
114            finished,
115            _marker: PhantomData,
116        }
117    }
118}
119
120/// Creates a scope for spawning scoped threads.
121///
122/// The function passed to `scope` will be provided a [`Scope`] object,
123/// through which scoped threads can be [spawned][`Scope::spawn`].
124pub fn scope<'env, F, T>(f: F) -> T
125where
126    F: for<'scope> FnOnce(&'scope Scope<'scope, 'env>) -> T,
127{
128    let scope = Scope {
129        num_running_threads: AtomicUsize::new(0),
130        main_task: ExecutionState::with(|s| s.current().id()),
131        env: PhantomData,
132        scope: PhantomData,
133    };
134
135    let ret = f(&scope);
136
137    if scope.num_running_threads.load(Ordering::Relaxed) != 0 {
138        tracing::info!("thread blocked, waiting for completion of scoped threads");
139        ExecutionState::with(|s| s.current_mut().block(false));
140        thread::switch();
141    }
142
143    ret
144}
145
146/// Spawn a new thread, returning a JoinHandle for it.
147///
148/// The join handle can be used (via the `join` method) to block until the child thread has
149/// finished.
150#[track_caller]
151pub fn spawn<F, T>(f: F) -> JoinHandle<T>
152where
153    F: FnOnce() -> T,
154    F: Send + 'static,
155    T: Send + 'static,
156{
157    spawn_named(f, None, None, Location::caller())
158}
159
160fn spawn_named<F, T>(
161    f: F,
162    name: Option<String>,
163    stack_size: Option<usize>,
164    caller: &'static Location<'static>,
165) -> JoinHandle<T>
166where
167    F: FnOnce() -> T,
168    F: Send + 'static,
169    T: Send + 'static,
170{
171    // SAFETY: F is static so all captured references must be `static and therefore
172    // will outlive the spawned continuation
173    unsafe { spawn_named_unchecked(f, name, stack_size, true, caller) }
174}
175
176/// Must ensure all captured references in f are valid for at least as long as the spawned continuation will run
177unsafe fn spawn_named_unchecked<F, T>(
178    f: F,
179    name: Option<String>,
180    stack_size: Option<usize>,
181    switch_before_exit: bool,
182    caller: &'static Location<'static>,
183) -> JoinHandle<T>
184where
185    F: FnOnce() -> T,
186    T: Send,
187{
188    // TODO Check if it's worth avoiding the call to `ExecutionState::config()` if we're going
189    // TODO to use an existing continuation from the pool.
190    let stack_size = stack_size.unwrap_or_else(|| ExecutionState::with(|s| s.config.stack_size));
191    let result = std::sync::Arc::new(std::sync::Mutex::new(None));
192    let task_id = {
193        let result = std::sync::Arc::clone(&result);
194
195        // Allocate `thread_fn` on the heap and assume a `'static` bound.
196        let f: Box<dyn FnOnce()> = Box::new(move || thread_fn(f, switch_before_exit, result));
197        let f: Box<dyn FnOnce() + 'static> = unsafe { std::mem::transmute(f) };
198
199        ExecutionState::spawn_thread(f, stack_size, name.clone(), None, caller)
200    };
201
202    let thread = Thread {
203        id: ThreadId { task_id },
204        name,
205    };
206
207    JoinHandle {
208        task_id,
209        thread,
210        result,
211    }
212}
213
214/// Body of a Shuttle thread, that runs the given closure, handles thread-local destructors, and
215pub(crate) use shuttle_engine::thread_support::thread_fn;
216
217/// An owned permission to join on a scoped thread (block on its termination).
218///
219/// See [`Scope::spawn`] for details.
220#[derive(Debug)]
221pub struct ScopedJoinHandle<'scope, T> {
222    handle: JoinHandle<T>,
223    finished: std::sync::Arc<AtomicBool>,
224    _marker: PhantomData<&'scope T>,
225}
226
227impl<T> ScopedJoinHandle<'_, T> {
228    /// Waits for the associated thread to finish.
229    pub fn join(self) -> Result<T> {
230        self.handle.join()
231    }
232
233    /// Extracts a handle to the underlying thread.
234    pub fn thread(&self) -> &Thread {
235        self.handle.thread()
236    }
237
238    /// Checks if the associated thread has finished running its main function.
239    ///
240    /// This might return `true` for a brief moment after the thread's main
241    /// function has returned, but before the thread itself has stopped running.
242    pub fn is_finished(&self) -> bool {
243        self.finished.load(Ordering::Relaxed)
244    }
245}
246
247/// An owned permission to join on a thread (block on its termination).
248#[derive(Debug)]
249pub struct JoinHandle<T> {
250    task_id: TaskId,
251    thread: Thread,
252    result: std::sync::Arc<std::sync::Mutex<Option<Result<T>>>>,
253}
254
255unsafe impl<T> Send for JoinHandle<T> {}
256unsafe impl<T> Sync for JoinHandle<T> {}
257
258impl<T> JoinHandle<T> {
259    /// Waits for the associated thread to finish.
260    pub fn join(self) -> Result<T> {
261        let is_finished = ExecutionState::with(|state| state.get(self.task_id).finished());
262        // If the joinee task is finished then the joiner will not block
263        if is_finished {
264            thread::switch();
265        }
266
267        let should_block = ExecutionState::with(|state| {
268            let me = state.current().id();
269            let target = state.get_mut(self.task_id);
270            if target.set_waiter(me) {
271                state.current_mut().block(false);
272                true
273            } else {
274                false
275            }
276        });
277
278        if should_block {
279            thread::switch();
280        }
281
282        // Waiting thread inherits the clock of the finished thread
283        ExecutionState::with(|state| {
284            let target = state.get_mut(self.task_id);
285            let clock = target.clock.clone();
286            state.update_clock(&clock);
287        });
288
289        self.result.lock().unwrap().take().expect("target should have finished")
290    }
291
292    /// Extracts a handle to the underlying thread.
293    pub fn thread(&self) -> &Thread {
294        &self.thread
295    }
296}
297
298/// Cooperatively gives up a timeslice to the Shuttle scheduler.
299///
300/// Some Shuttle schedulers use this as a hint to deprioritize the current thread in order for other
301/// threads to make progress (e.g., in a spin loop).
302pub fn yield_now() {
303    let waker = ExecutionState::with(|state| state.current().waker());
304    waker.wake_by_ref();
305    ExecutionState::request_yield();
306    thread::switch();
307}
308
309/// Puts the current thread to sleep for at least the specified amount of time.
310// Note that Shuttle does not model time, so this behaves just like a context switch.
311pub fn sleep(_dur: Duration) {
312    thread::switch();
313}
314
315/// Get a handle to the thread that invokes it
316pub fn current() -> Thread {
317    let (task_id, name) = ExecutionState::with(|s| {
318        let me = s.current();
319        (me.id(), me.name())
320    });
321
322    Thread {
323        id: ThreadId { task_id },
324        name,
325    }
326}
327
328/// Blocks unless or until the current thread's token is made available (may wake spuriously).
329pub fn park() {
330    let switch = ExecutionState::with(|s| s.current_mut().park());
331
332    // We only need to context switch if the park token was unavailable. If it was available, then
333    // any execution reachable by context switching here would also be reachable by having not
334    // chosen this thread at the last context switch, because the park state of a thread is only
335    // observable by the thread itself. We also mark it as an explicit yield request by the task,
336    // since otherwise some schedulers might prefer to to reschedule the current task, which in this
337    // context would result in spurious wakeups triggering nearly every time.
338    if switch {
339        ExecutionState::request_yield();
340        thread::switch();
341    }
342}
343
344/// Blocks unless or until the current thread's token is made available or the specified duration
345/// has been reached (may wake spuriously).
346///
347/// Note that Shuttle does not model time, so this behaves identically to `park`. In particular,
348/// Shuttle does not assume that the timeout will ever fire, so if all threads are blocked in a call
349/// to `park_timeout` it will be treated as a deadlock.
350pub fn park_timeout(_dur: Duration) {
351    park();
352}
353
354/// Thread factory, which can be used in order to configure the properties of a new thread.
355#[derive(Debug, Default)]
356pub struct Builder {
357    name: Option<String>,
358    stack_size: Option<usize>,
359}
360
361impl Builder {
362    /// Generates the base configuration for spawning a thread, from which configuration methods can be chained.
363    pub fn new() -> Self {
364        Self {
365            name: None,
366            stack_size: None,
367        }
368    }
369
370    /// Names the thread-to-be. Currently the name is used for identification only in panic messages.
371    pub fn name(mut self, name: String) -> Self {
372        self.name = Some(name);
373        self
374    }
375
376    /// Sets the size of the stack (in bytes) for the new thread.
377    pub fn stack_size(mut self, stack_size: usize) -> Self {
378        self.stack_size = Some(stack_size);
379        self
380    }
381
382    /// Spawns a new thread by taking ownership of the Builder, and returns an `io::Result` to its `JoinHandle`.
383    #[track_caller]
384    pub fn spawn<F, T>(self, f: F) -> std::io::Result<JoinHandle<T>>
385    where
386        F: FnOnce() -> T,
387        F: Send + 'static,
388        T: Send + 'static,
389    {
390        Ok(spawn_named(f, self.name, self.stack_size, Location::caller()))
391    }
392}
393
394pub use shuttle_engine::thread_support::{AccessError, LocalKey};