Skip to main content

moirai_core/task/
handle.rs

1use crate::error::TaskError;
2
3use super::id_and_context::TaskId;
4
5// ── std-only block ────────────────────────────────────────────────────────────
6
7#[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// State constants for TaskResultSlot
21#[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// ── ResultWaitPolicy sealed module ────────────────────────────────────────────
33
34#[cfg(feature = "std")]
35pub(super) mod result_wait {
36    pub(super) mod sealed {
37        pub trait Sealed {}
38    }
39
40    /// Compile-time wait policy for task result handoff.
41    ///
42    /// Implementors are zero-sized marker types. `TaskResultSlot` receives the
43    /// policy as a generic parameter, so the spin budget is const-folded and no
44    /// runtime policy value is stored in the handle or slot.
45    pub trait ResultWaitPolicy: sealed::Sealed {
46        /// Maximum number of spin-loop iterations before parking the thread.
47        const SPIN_ATTEMPTS: usize;
48    }
49
50    /// Zero-sized blocking wait policy: spins up to `MAX_SPIN_ATTEMPTS` then parks.
51    #[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// ── TaskResultSlot (private) ──────────────────────────────────────────────────
65
66/// One-shot cell carrying the result of a single scheduled task.
67///
68/// # Cache-line layout
69///
70/// The `state` field is the synchronisation point between the *producer*
71/// (the worker that executes the task and writes the result) and the
72/// *consumer* (the thread that called `JoinHandle::wait`).  Keeping
73/// `state` on its own 64-byte cache line prevents false sharing:
74/// the producer invalidates only its line when storing `RESULT_READY`,
75/// and the consumer accesses `result`/`waiter` on a separate line.
76///
77/// `#[repr(align(64))]` on `TaskResultSlot` itself aligns the start of
78/// the struct to a cache line; `_pad` pushes `result` and `waiter` past
79/// the first 64 bytes so they land on a second line.
80#[cfg(feature = "std")]
81#[repr(align(64))]
82struct TaskResultSlot<T> {
83    /// Synchronisation state — written by the producer, read by consumer.
84    /// Placed first so it occupies the beginning of the first cache line.
85    state: AtomicU8,
86    /// Padding to push `result` and `waiter` onto a separate cache line,
87    /// eliminating producer-consumer false sharing on the `state` field.
88    _pad: [u8; 63],
89    result: UnsafeCell<MaybeUninit<Result<T, TaskError>>>,
90    waiter: UnsafeCell<MaybeUninit<thread::Thread>>,
91}
92
93// Safety: the slot is a single-producer/single-consumer one-shot cell.
94// `complete` wins the PENDING/WAITING -> WRITING transition before writing.
95// WAITING is entered only after the waiter thread is stored in the waiter cell,
96// and `wait` takes the value only after an acquire READY -> TAKEN transition.
97#[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        // Safety: the WRITING state is reachable only through
120        // `begin_completion`, so no other thread can read, write, or drop the
121        // result cell until READY publishes.
122        unsafe {
123            (*self.result.get()).write(result);
124        }
125
126        self.state.store(RESULT_READY, Ordering::Release);
127
128        if waiting {
129            // Safety: WAITING is reachable only after `register_waiter` writes
130            // the thread handle and publishes it with a release CAS.
131            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            // Safety: READY is published only after `complete` initializes the
178            // cell. The READY -> TAKEN transition is unique, so this read moves
179            // the result exactly once.
180            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                    // Safety: there is only one consumer. If the publish CAS
199                    // fails, this local thread handle is dropped before retry.
200                    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                    // Safety: the CAS failed, so no producer can observe this
218                    // waiter cell as initialized through the WAITING state.
219                    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            // Safety: READY means the cell is initialized and no consuming join
264            // took it because `drop` has exclusive access to the slot.
265            unsafe {
266                self.result.get_mut().assume_init_drop();
267            }
268        } else if state == RESULT_WAITING {
269            // Safety: WAITING means the waiter thread handle is initialized and
270            // no producer unparked it because `drop` has exclusive access.
271            unsafe {
272                self.waiter.get_mut().assume_init_drop();
273            }
274        }
275    }
276}
277
278// ── Diagnostic helpers (feature = "result-diagnostics") ──────────────────────
279
280#[cfg(all(feature = "std", feature = "result-diagnostics"))]
281const DIAGNOSTIC_READY_VALUE: usize = 42;
282
283/// Diagnostic-only ready result-slot take path for benchmark attribution.
284#[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/// Diagnostic-only pending spin miss path for benchmark attribution.
296#[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/// Diagnostic-only waiter registration path for benchmark attribution.
311#[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/// Diagnostic-only waiting-result completion path for benchmark attribution.
320#[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// ── TaskHandle (std) ──────────────────────────────────────────────────────────
333
334/// A handle to a task that may be running on another thread.
335#[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    /// Creates a new pending task handle and its completion sender.
345    #[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    /// Creates a new task handle from an existing result.
358    #[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    /// Creates a new detached task handle (no result channel).
369    ///
370    /// # Arguments
371    /// * `id` - The unique identifier for this task
372    ///
373    /// # Returns
374    /// A new detached task handle instance
375    #[must_use]
376    pub fn new_detached(id: TaskId) -> Self {
377        Self {
378            id,
379            result_slot: None,
380        }
381    }
382
383    /// Returns the task ID.
384    ///
385    /// # Returns
386    /// The unique identifier for this task
387    #[must_use]
388    pub fn id(&self) -> TaskId {
389        self.id
390    }
391
392    /// Waits for the task to complete and returns the result.
393    ///
394    /// # Returns
395    /// - `Some(Ok(result))` if the task completed successfully
396    /// - `Some(Err(error))` if the task failed with an error
397    /// - `None` if the task was detached
398    #[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    /// Checks if the task has finished execution.
406    ///
407    /// # Returns
408    /// `true` if the task has completed (successfully or with error), `false` if still running
409    #[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// ── TaskResultSender (std) ────────────────────────────────────────────────────
418
419/// Single-producer completion endpoint for a task result.
420#[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    /// Complete the task result and wake any waiter.
429    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// ── TaskHandle (no_std) ───────────────────────────────────────────────────────
447
448// For no_std environments, provide a simpler handle
449#[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    /// Create a new task handle.
458    pub fn new(id: TaskId) -> Self {
459        Self {
460            id,
461            _phantom: core::marker::PhantomData,
462        }
463    }
464
465    /// Create a new detached task handle (alias for new in no_std).
466    pub fn new_detached(id: TaskId) -> Self {
467        Self::new(id)
468    }
469
470    /// Get the task ID.
471    pub fn id(&self) -> TaskId {
472        self.id
473    }
474}