1use 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#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
15pub struct ThreadId {
16 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#[derive(Debug, Clone)]
28pub struct Thread {
29 name: Option<String>,
30 id: ThreadId,
31}
32
33impl Thread {
34 pub fn name(&self) -> Option<&str> {
36 self.name.as_deref()
37 }
38
39 pub fn id(&self) -> ThreadId {
41 self.id
42 }
43
44 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
54pub 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 #[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 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 ScopedJoinHandle {
121 handle: unsafe { spawn_named_unchecked(scope_closure, None, None, false, Location::caller()) },
122 finished,
123 _marker: PhantomData,
124 }
125 }
126}
127
128pub 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#[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 unsafe { spawn_named_unchecked(f, name, stack_size, true, caller) }
182}
183
184unsafe 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 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 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
222pub(crate) use shuttle_engine::thread_support::thread_fn;
224
225#[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 pub fn join(self) -> Result<T> {
238 self.handle.join()
239 }
240
241 pub fn thread(&self) -> &Thread {
243 self.handle.thread()
244 }
245
246 pub fn is_finished(&self) -> bool {
251 self.finished.load(Ordering::Relaxed)
252 }
253}
254
255#[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 pub fn join(self) -> Result<T> {
269 let is_finished = ExecutionState::with(|state| state.get(self.task_id).finished());
270 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 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 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 pub fn thread(&self) -> &Thread {
308 &self.thread
309 }
310}
311
312pub 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
323pub fn sleep(_dur: Duration) {
326 thread::switch();
327}
328
329pub 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
342pub fn park() {
344 let switch = ExecutionState::with(|s| s.current_mut().park());
345
346 if switch {
353 ExecutionState::request_yield();
354 thread::switch();
355 }
356}
357
358pub fn park_timeout(_dur: Duration) {
365 park();
366}
367
368#[derive(Debug, Default)]
370pub struct Builder {
371 name: Option<String>,
372 stack_size: Option<usize>,
373}
374
375impl Builder {
376 pub fn new() -> Self {
378 Self {
379 name: None,
380 stack_size: None,
381 }
382 }
383
384 pub fn name(mut self, name: String) -> Self {
386 self.name = Some(name);
387 self
388 }
389
390 pub fn stack_size(mut self, stack_size: usize) -> Self {
392 self.stack_size = Some(stack_size);
393 self
394 }
395
396 #[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};