Skip to main content

shuttle_engine/
thread_support.rs

1use crate::runtime::execution::ExecutionState;
2use crate::runtime::thread;
3use std::marker::PhantomData;
4
5/// Cooperatively gives up a timeslice to the Shuttle scheduler.
6pub fn yield_now() {
7    let waker = ExecutionState::with(|state| state.current().waker());
8    waker.wake_by_ref();
9    ExecutionState::request_yield();
10    thread::switch();
11}
12
13/// The body of a spawned thread. Runs `f`, drops thread locals, publishes result.
14pub fn thread_fn<F, T>(
15    f: F,
16    switch_before_exit: bool,
17    result: std::sync::Arc<std::sync::Mutex<Option<std::thread::Result<T>>>>,
18) where
19    F: FnOnce() -> T,
20{
21    let ret = f();
22
23    if switch_before_exit && ExecutionState::with(|s| s.exit_current_truncates_execution()) {
24        thread::switch();
25    }
26
27    tracing::trace!("thread finished, dropping thread locals");
28    ExecutionState::drop_task_locals();
29    tracing::trace!("done dropping thread locals");
30
31    *result.lock().unwrap() = Some(Ok(ret));
32    ExecutionState::with(|state| {
33        if let Some(waiter) = state.current_mut().take_waiter() {
34            state.get_mut(waiter).unblock();
35        }
36    });
37}
38
39/// A key into Shuttle's thread-local storage.
40pub struct LocalKey<T: 'static> {
41    #[doc(hidden)]
42    pub init: fn() -> T,
43    #[doc(hidden)]
44    pub _p: PhantomData<T>,
45}
46
47// Safety: `LocalKey` implements thread-local storage; each thread sees its own value of the type T.
48unsafe impl<T> Send for LocalKey<T> {}
49unsafe impl<T> Sync for LocalKey<T> {}
50
51impl<T: 'static> std::fmt::Debug for LocalKey<T> {
52    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
53        f.debug_struct("LocalKey").finish_non_exhaustive()
54    }
55}
56
57impl<T: 'static> LocalKey<T> {
58    /// Acquires a reference to the value in this TLS key.
59    ///
60    /// This will lazily initialize the value if this thread has not referenced this key yet.
61    pub fn with<F, R>(&'static self, f: F) -> R
62    where
63        F: FnOnce(&T) -> R,
64    {
65        self.try_with(f).expect(
66            "cannot access a Thread Local Storage value \
67            during or after destruction",
68        )
69    }
70
71    /// Acquires a reference to the value in this TLS key.
72    ///
73    /// This will lazily initialize the value if this thread has not referenced this key yet. If the
74    /// key has been destroyed (which may happen if this is called in a destructor), this function
75    /// will return an AccessError.
76    pub fn try_with<F, R>(&'static self, f: F) -> std::result::Result<R, AccessError>
77    where
78        F: FnOnce(&T) -> R,
79    {
80        let value = self.get().unwrap_or_else(|| {
81            let value = (self.init)();
82
83            ExecutionState::with(move |state| {
84                state.current_mut().init_local(self, value);
85            });
86
87            self.get().unwrap()
88        })?;
89
90        Ok(f(value))
91    }
92
93    fn get(&'static self) -> Option<std::result::Result<&'static T, AccessError>> {
94        // Safety: see the usage below
95        unsafe fn extend_lt<'b, T>(t: &'_ T) -> &'b T {
96            std::mem::transmute(t)
97        }
98
99        ExecutionState::with(|state| {
100            if let Ok(value) = state.current().local(self)? {
101                // Safety: the `ExecutionState` outlives any thread, including the caller, and so
102                // it's safe to give the caller the lifetime it's asking for here.
103                Some(Ok(unsafe { extend_lt(value) }))
104            } else {
105                Some(Err(AccessError))
106            }
107        })
108    }
109}
110
111/// An error returned by [`LocalKey::try_with`]
112#[derive(Clone, Copy, PartialEq, Eq, Debug)]
113#[non_exhaustive]
114pub struct AccessError;
115
116impl std::fmt::Display for AccessError {
117    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
118        std::fmt::Display::fmt("already destroyed", f)
119    }
120}
121
122impl std::error::Error for AccessError {}