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 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 ScopedJoinHandle {
113 handle: unsafe { spawn_named_unchecked(scope_closure, None, None, false, Location::caller()) },
114 finished,
115 _marker: PhantomData,
116 }
117 }
118}
119
120pub 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#[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 unsafe { spawn_named_unchecked(f, name, stack_size, true, caller) }
174}
175
176unsafe 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 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 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
214pub(crate) use shuttle_engine::thread_support::thread_fn;
216
217#[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 pub fn join(self) -> Result<T> {
230 self.handle.join()
231 }
232
233 pub fn thread(&self) -> &Thread {
235 self.handle.thread()
236 }
237
238 pub fn is_finished(&self) -> bool {
243 self.finished.load(Ordering::Relaxed)
244 }
245}
246
247#[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 pub fn join(self) -> Result<T> {
261 let is_finished = ExecutionState::with(|state| state.get(self.task_id).finished());
262 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 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 pub fn thread(&self) -> &Thread {
294 &self.thread
295 }
296}
297
298pub 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
309pub fn sleep(_dur: Duration) {
312 thread::switch();
313}
314
315pub 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
328pub fn park() {
330 let switch = ExecutionState::with(|s| s.current_mut().park());
331
332 if switch {
339 ExecutionState::request_yield();
340 thread::switch();
341 }
342}
343
344pub fn park_timeout(_dur: Duration) {
351 park();
352}
353
354#[derive(Debug, Default)]
356pub struct Builder {
357 name: Option<String>,
358 stack_size: Option<usize>,
359}
360
361impl Builder {
362 pub fn new() -> Self {
364 Self {
365 name: None,
366 stack_size: None,
367 }
368 }
369
370 pub fn name(mut self, name: String) -> Self {
372 self.name = Some(name);
373 self
374 }
375
376 pub fn stack_size(mut self, stack_size: usize) -> Self {
378 self.stack_size = Some(stack_size);
379 self
380 }
381
382 #[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};