moirai_core/task/
handle.rs1use crate::error::TaskError;
2
3use super::id_and_context::TaskId;
4
5#[cfg(feature = "std")]
8use core::cell::UnsafeCell;
9#[cfg(feature = "std")]
10use core::mem::{ManuallyDrop, MaybeUninit};
11
12#[cfg(feature = "std")]
13use std::sync::{
14 atomic::{AtomicU8, Ordering},
15 Arc,
16};
17#[cfg(feature = "std")]
18use std::thread;
19
20#[cfg(feature = "std")]
22const RESULT_PENDING: u8 = 0;
23#[cfg(feature = "std")]
24const RESULT_WRITING: u8 = 1;
25#[cfg(feature = "std")]
26const RESULT_READY: u8 = 2;
27#[cfg(feature = "std")]
28const RESULT_TAKEN: u8 = 3;
29#[cfg(feature = "std")]
30const RESULT_WAITING: u8 = 4;
31
32#[cfg(feature = "std")]
35pub(super) mod result_wait {
36 pub(super) mod sealed {
37 pub trait Sealed {}
38 }
39
40 pub trait ResultWaitPolicy: sealed::Sealed {
46 const SPIN_ATTEMPTS: usize;
48 }
49
50 #[derive(Debug, Clone, Copy, Default)]
52 pub struct BlockingResultWait;
53
54 impl sealed::Sealed for BlockingResultWait {}
55
56 impl ResultWaitPolicy for BlockingResultWait {
57 const SPIN_ATTEMPTS: usize = super::super::MAX_SPIN_ATTEMPTS;
58 }
59}
60
61#[cfg(feature = "std")]
62pub use result_wait::{BlockingResultWait, ResultWaitPolicy};
63
64#[cfg(feature = "std")]
81#[repr(align(64))]
82struct TaskResultSlot<T> {
83 state: AtomicU8,
86 _pad: [u8; 63],
89 result: UnsafeCell<MaybeUninit<Result<T, TaskError>>>,
90 waiter: UnsafeCell<MaybeUninit<thread::Thread>>,
91}
92
93#[cfg(feature = "std")]
98unsafe impl<T: Send> Send for TaskResultSlot<T> {}
99
100#[cfg(feature = "std")]
101unsafe impl<T: Send> Sync for TaskResultSlot<T> {}
102
103#[cfg(feature = "std")]
104impl<T> TaskResultSlot<T> {
105 fn new() -> Self {
106 Self {
107 state: AtomicU8::new(RESULT_PENDING),
108 _pad: [0u8; 63],
109 result: UnsafeCell::new(MaybeUninit::uninit()),
110 waiter: UnsafeCell::new(MaybeUninit::uninit()),
111 }
112 }
113
114 fn complete(&self, result: Result<T, TaskError>) {
115 let Some(waiting) = self.begin_completion() else {
116 return;
117 };
118
119 unsafe {
123 (*self.result.get()).write(result);
124 }
125
126 self.state.store(RESULT_READY, Ordering::Release);
127
128 if waiting {
129 let thread = unsafe { (*self.waiter.get()).assume_init_read() };
132 thread.unpark();
133 }
134 }
135
136 fn wait<P>(&self) -> Result<T, TaskError>
137 where
138 P: ResultWaitPolicy,
139 {
140 if let Some(result) = self.try_take_ready() {
141 return result;
142 }
143
144 for _ in 0..P::SPIN_ATTEMPTS {
145 if let Some(result) = self.try_take_observed_ready() {
146 return result;
147 }
148 core::hint::spin_loop();
149 }
150
151 self.register_waiter();
152
153 loop {
154 if let Some(result) = self.try_take_observed_ready() {
155 return result;
156 }
157
158 thread::park();
159 }
160 }
161
162 fn is_completed(&self) -> bool {
163 self.state.load(Ordering::Acquire) == RESULT_READY
164 }
165
166 fn try_take_ready(&self) -> Option<Result<T, TaskError>> {
167 if self
168 .state
169 .compare_exchange(
170 RESULT_READY,
171 RESULT_TAKEN,
172 Ordering::Acquire,
173 Ordering::Relaxed,
174 )
175 .is_ok()
176 {
177 Some(unsafe { (*self.result.get()).assume_init_read() })
181 } else {
182 None
183 }
184 }
185
186 fn try_take_observed_ready(&self) -> Option<Result<T, TaskError>> {
187 if self.state.load(Ordering::Relaxed) == RESULT_READY {
188 self.try_take_ready()
189 } else {
190 None
191 }
192 }
193
194 fn register_waiter(&self) {
195 loop {
196 match self.state.load(Ordering::Acquire) {
197 RESULT_PENDING => {
198 unsafe {
201 (*self.waiter.get()).write(thread::current());
202 }
203
204 if self
205 .state
206 .compare_exchange(
207 RESULT_PENDING,
208 RESULT_WAITING,
209 Ordering::Release,
210 Ordering::Acquire,
211 )
212 .is_ok()
213 {
214 return;
215 }
216
217 unsafe {
220 (*self.waiter.get()).assume_init_drop();
221 }
222 }
223 RESULT_WRITING => core::hint::spin_loop(),
224 _ => return,
225 }
226 }
227 }
228
229 fn begin_completion(&self) -> Option<bool> {
230 match self.state.compare_exchange(
231 RESULT_PENDING,
232 RESULT_WRITING,
233 Ordering::Relaxed,
234 Ordering::Acquire,
235 ) {
236 Ok(_) => Some(false),
237 Err(RESULT_WAITING) => {
238 if self
239 .state
240 .compare_exchange(
241 RESULT_WAITING,
242 RESULT_WRITING,
243 Ordering::Acquire,
244 Ordering::Acquire,
245 )
246 .is_ok()
247 {
248 Some(true)
249 } else {
250 None
251 }
252 }
253 Err(_) => None,
254 }
255 }
256}
257
258#[cfg(feature = "std")]
259impl<T> Drop for TaskResultSlot<T> {
260 fn drop(&mut self) {
261 let state = *self.state.get_mut();
262 if state == RESULT_READY {
263 unsafe {
266 self.result.get_mut().assume_init_drop();
267 }
268 } else if state == RESULT_WAITING {
269 unsafe {
272 self.waiter.get_mut().assume_init_drop();
273 }
274 }
275 }
276}
277
278#[cfg(all(feature = "std", feature = "result-diagnostics"))]
281const DIAGNOSTIC_READY_VALUE: usize = 42;
282
283#[cfg(all(feature = "std", feature = "result-diagnostics"))]
285#[doc(hidden)]
286pub fn diagnostic_result_slot_ready_take() -> usize {
287 let slot = TaskResultSlot::new();
288 slot.complete(Ok(DIAGNOSTIC_READY_VALUE));
289 match slot.try_take_ready() {
290 Some(Ok(value)) => value,
291 _ => 0,
292 }
293}
294
295#[cfg(all(feature = "std", feature = "result-diagnostics"))]
297#[doc(hidden)]
298pub fn diagnostic_result_slot_spin_miss() -> usize {
299 let slot = TaskResultSlot::<usize>::new();
300 let mut misses = 0usize;
301 for _ in 0..BlockingResultWait::SPIN_ATTEMPTS {
302 if slot.try_take_observed_ready().is_none() {
303 misses = misses.wrapping_add(1);
304 }
305 core::hint::spin_loop();
306 }
307 misses
308}
309
310#[cfg(all(feature = "std", feature = "result-diagnostics"))]
312#[doc(hidden)]
313pub fn diagnostic_result_slot_register_waiter() -> usize {
314 let slot = TaskResultSlot::<usize>::new();
315 slot.register_waiter();
316 usize::from(slot.state.load(Ordering::Acquire) == RESULT_WAITING)
317}
318
319#[cfg(all(feature = "std", feature = "result-diagnostics"))]
321#[doc(hidden)]
322pub fn diagnostic_result_slot_complete_waiting() -> usize {
323 let slot = TaskResultSlot::new();
324 slot.register_waiter();
325 slot.complete(Ok(DIAGNOSTIC_READY_VALUE));
326 match slot.try_take_ready() {
327 Some(Ok(value)) => value,
328 _ => 0,
329 }
330}
331
332#[cfg(feature = "std")]
336#[allow(clippy::module_name_repetitions)]
337pub struct TaskHandle<T> {
338 id: TaskId,
339 result_slot: Option<Arc<TaskResultSlot<T>>>,
340}
341
342#[cfg(feature = "std")]
343impl<T> TaskHandle<T> {
344 #[must_use]
346 pub fn new_pending(id: TaskId) -> (Self, TaskResultSender<T>) {
347 let slot = Arc::new(TaskResultSlot::new());
348 (
349 Self {
350 id,
351 result_slot: Some(Arc::clone(&slot)),
352 },
353 TaskResultSender { slot: Some(slot) },
354 )
355 }
356
357 #[must_use]
359 pub fn ready(id: TaskId, result: Result<T, TaskError>) -> Self {
360 let slot = Arc::new(TaskResultSlot::new());
361 slot.complete(result);
362 Self {
363 id,
364 result_slot: Some(slot),
365 }
366 }
367
368 #[must_use]
376 pub fn new_detached(id: TaskId) -> Self {
377 Self {
378 id,
379 result_slot: None,
380 }
381 }
382
383 #[must_use]
388 pub fn id(&self) -> TaskId {
389 self.id
390 }
391
392 #[must_use]
399 pub fn join(mut self) -> Option<Result<T, TaskError>> {
400 self.result_slot
401 .take()
402 .map(|slot| slot.wait::<BlockingResultWait>())
403 }
404
405 #[must_use]
410 pub fn is_finished(&self) -> bool {
411 self.result_slot
412 .as_ref()
413 .is_some_and(|slot| slot.is_completed())
414 }
415}
416
417#[cfg(feature = "std")]
421#[allow(clippy::module_name_repetitions)]
422pub struct TaskResultSender<T> {
423 slot: Option<Arc<TaskResultSlot<T>>>,
424}
425
426#[cfg(feature = "std")]
427impl<T> TaskResultSender<T> {
428 pub fn send(self, result: Result<T, TaskError>) {
430 let mut sender = ManuallyDrop::new(self);
431 if let Some(slot) = sender.slot.take() {
432 slot.complete(result);
433 }
434 }
435}
436
437#[cfg(feature = "std")]
438impl<T> Drop for TaskResultSender<T> {
439 fn drop(&mut self) {
440 if let Some(slot) = self.slot.take() {
441 slot.complete(Err(TaskError::Cancelled));
442 }
443 }
444}
445
446#[cfg(not(feature = "std"))]
450pub struct TaskHandle<T> {
451 id: TaskId,
452 _phantom: core::marker::PhantomData<T>,
453}
454
455#[cfg(not(feature = "std"))]
456impl<T> TaskHandle<T> {
457 pub fn new(id: TaskId) -> Self {
459 Self {
460 id,
461 _phantom: core::marker::PhantomData,
462 }
463 }
464
465 pub fn new_detached(id: TaskId) -> Self {
467 Self::new(id)
468 }
469
470 pub fn id(&self) -> TaskId {
472 self.id
473 }
474}