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