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::mem::ManuallyDrop;
9
10#[cfg(feature = "std")]
11use std::sync::{Arc, atomic::AtomicU8};
12#[cfg(feature = "std")]
13use std::thread;
14
15#[cfg(feature = "std")]
16use moirai_utils::{CacheAligned, ResultCell};
17
18// ── ResultWaitPolicy sealed module ────────────────────────────────────────────
19
20#[cfg(feature = "std")]
21pub(super) mod result_wait {
22    pub(super) mod sealed {
23        pub trait Sealed {}
24    }
25
26    /// Compile-time wait policy for task result handoff.
27    ///
28    /// Implementors are zero-sized marker types. `TaskResultSlot` receives the
29    /// policy as a generic parameter, so the spin budget is const-folded and no
30    /// runtime policy value is stored in the handle or slot.
31    pub trait ResultWaitPolicy: sealed::Sealed {
32        /// Maximum number of spin-loop iterations before parking the thread.
33        const SPIN_ATTEMPTS: usize;
34    }
35
36    /// Zero-sized blocking wait policy: spins up to `MAX_SPIN_ATTEMPTS` then parks.
37    #[derive(Debug, Clone, Copy, Default)]
38    pub struct BlockingResultWait;
39
40    impl sealed::Sealed for BlockingResultWait {}
41
42    impl ResultWaitPolicy for BlockingResultWait {
43        const SPIN_ATTEMPTS: usize = super::super::MAX_SPIN_ATTEMPTS;
44    }
45}
46
47#[cfg(feature = "std")]
48pub use result_wait::{BlockingResultWait, ResultWaitPolicy};
49
50// ── TaskResultSlot (private) ──────────────────────────────────────────────────
51
52/// One-shot cell carrying the result of a single scheduled task.
53///
54/// # Cache-line layout
55///
56/// The `state` field is the synchronisation point between the *producer*
57/// (the worker that executes the task and writes the result) and the
58/// *consumer* (the thread that called `JoinHandle::wait`).  Keeping
59/// `state` on its own sector prevents false sharing: the producer
60/// invalidates only that sector when storing `RESULT_READY`, and the
61/// consumer accesses `result`/`waiter` beyond it.
62///
63/// [`CacheAligned`] supplies both halves of that layout — it aligns the slot
64/// (and therefore the cell's `state` word) to
65/// `moirai_utils::DESTRUCTIVE_INTERFERENCE_SIZE`, and its own size pushes
66/// `result`/`waiter` past that boundary. The separation is 128 bytes on
67/// x86-64/aarch64, where the adjacent-line prefetcher makes 64 too narrow;
68/// the per-target value lives in `moirai-utils`, not in a literal here.
69#[cfg(feature = "std")]
70struct TaskResultSlot<T> {
71    cell: ResultCell<Result<T, TaskError>, thread::Thread, CacheAligned<AtomicU8>>,
72}
73
74/// The blocking side of the completion path.
75///
76/// The protocol — the state machine, its cell-access invariants and its ordering
77/// argument — lives with [`ResultCell`] in `moirai-utils`, and the async handle
78/// runs it too. This type adds only what blocking needs: a parked
79/// [`thread::Thread`] as the waiter, the cache-aligned state word its layout note
80/// above describes, and [`wait`](Self::wait)'s spin-then-park loop, which needs
81/// `thread::park` and so does not belong in the shared cell.
82#[cfg(feature = "std")]
83impl<T> TaskResultSlot<T> {
84    fn new() -> Self {
85        Self {
86            cell: ResultCell::new(),
87        }
88    }
89
90    fn complete(&self, result: Result<T, TaskError>) {
91        self.cell.complete(result);
92    }
93
94    /// Block until the result is ready and take it.
95    ///
96    /// # Safety
97    ///
98    /// The caller is the slot's only consumer: no other thread waits on, polls, or
99    /// takes from this slot while this call runs.
100    unsafe fn wait<P>(&self) -> Result<T, TaskError>
101    where
102        P: ResultWaitPolicy,
103    {
104        if let Some(result) = self.try_take_ready() {
105            return result;
106        }
107
108        for _ in 0..P::SPIN_ATTEMPTS {
109            if let Some(result) = self.try_take_observed_ready() {
110                return result;
111            }
112            core::hint::spin_loop();
113        }
114
115        // SAFETY: `wait`'s caller is the slot's only consumer.
116        unsafe { self.register_waiter() };
117
118        loop {
119            if let Some(result) = self.try_take_observed_ready() {
120                return result;
121            }
122
123            thread::park();
124        }
125    }
126
127    fn is_completed(&self) -> bool {
128        self.cell.is_completed()
129    }
130
131    fn try_take_ready(&self) -> Option<Result<T, TaskError>> {
132        self.cell.try_take_ready()
133    }
134
135    fn try_take_observed_ready(&self) -> Option<Result<T, TaskError>> {
136        self.cell.try_take_observed_ready()
137    }
138
139    /// # Safety
140    ///
141    /// The caller is the slot's only consumer.
142    unsafe fn register_waiter(&self) {
143        // SAFETY: forwarded from this method's contract.
144        unsafe { self.cell.register(&thread::current()) };
145    }
146
147    #[cfg(feature = "result-diagnostics")]
148    fn has_registered_waiter(&self) -> bool {
149        self.cell.has_registered_waiter()
150    }
151}
152// ── Diagnostic helpers (feature = "result-diagnostics") ──────────────────────
153
154#[cfg(all(feature = "std", feature = "result-diagnostics"))]
155const DIAGNOSTIC_READY_VALUE: usize = 42;
156
157/// Diagnostic-only ready result-slot take path for benchmark attribution.
158#[cfg(all(feature = "std", feature = "result-diagnostics"))]
159#[doc(hidden)]
160pub fn diagnostic_result_slot_ready_take() -> usize {
161    let slot = TaskResultSlot::new();
162    slot.complete(Ok(DIAGNOSTIC_READY_VALUE));
163    match slot.try_take_ready() {
164        Some(Ok(value)) => value,
165        _ => 0,
166    }
167}
168
169/// Diagnostic-only pending spin miss path for benchmark attribution.
170#[cfg(all(feature = "std", feature = "result-diagnostics"))]
171#[doc(hidden)]
172pub fn diagnostic_result_slot_spin_miss() -> usize {
173    let slot = TaskResultSlot::<usize>::new();
174    let mut misses = 0usize;
175    for _ in 0..BlockingResultWait::SPIN_ATTEMPTS {
176        if slot.try_take_observed_ready().is_none() {
177            misses = misses.wrapping_add(1);
178        }
179        core::hint::spin_loop();
180    }
181    misses
182}
183
184/// Diagnostic-only waiter registration path for benchmark attribution.
185#[cfg(all(feature = "std", feature = "result-diagnostics"))]
186#[doc(hidden)]
187pub fn diagnostic_result_slot_register_waiter() -> usize {
188    let slot = TaskResultSlot::<usize>::new();
189    // SAFETY: `slot` is local, so this call is its only consumer.
190    unsafe { slot.register_waiter() };
191    usize::from(slot.has_registered_waiter())
192}
193
194/// Diagnostic-only waiting-result completion path for benchmark attribution.
195#[cfg(all(feature = "std", feature = "result-diagnostics"))]
196#[doc(hidden)]
197pub fn diagnostic_result_slot_complete_waiting() -> usize {
198    let slot = TaskResultSlot::new();
199    // SAFETY: `slot` is local, so this call is its only consumer.
200    unsafe { slot.register_waiter() };
201    slot.complete(Ok(DIAGNOSTIC_READY_VALUE));
202    match slot.try_take_ready() {
203        Some(Ok(value)) => value,
204        _ => 0,
205    }
206}
207
208// ── TaskHandle (std) ──────────────────────────────────────────────────────────
209
210/// A handle to a task that may be running on another thread.
211#[cfg(feature = "std")]
212#[allow(clippy::module_name_repetitions)]
213pub struct TaskHandle<T> {
214    id: TaskId,
215    result_slot: Option<Arc<TaskResultSlot<T>>>,
216}
217
218#[cfg(feature = "std")]
219impl<T> TaskHandle<T> {
220    /// Creates a new pending task handle and its completion sender.
221    #[must_use]
222    pub fn new_pending(id: TaskId) -> (Self, TaskResultSender<T>) {
223        let slot = Arc::new(TaskResultSlot::new());
224        (
225            Self {
226                id,
227                result_slot: Some(Arc::clone(&slot)),
228            },
229            TaskResultSender { slot: Some(slot) },
230        )
231    }
232
233    /// Creates a new task handle from an existing result.
234    #[must_use]
235    pub fn ready(id: TaskId, result: Result<T, TaskError>) -> Self {
236        let slot = Arc::new(TaskResultSlot::new());
237        slot.complete(result);
238        Self {
239            id,
240            result_slot: Some(slot),
241        }
242    }
243
244    /// Creates a new detached task handle (no result channel).
245    ///
246    /// # Arguments
247    /// * `id` - The unique identifier for this task
248    ///
249    /// # Returns
250    /// A new detached task handle instance
251    #[must_use]
252    pub fn new_detached(id: TaskId) -> Self {
253        Self {
254            id,
255            result_slot: None,
256        }
257    }
258
259    /// Returns the task ID.
260    ///
261    /// # Returns
262    /// The unique identifier for this task
263    #[must_use]
264    pub fn id(&self) -> TaskId {
265        self.id
266    }
267
268    /// Waits for the task to complete and returns the result.
269    ///
270    /// # Returns
271    /// - `Some(Ok(result))` if the task completed successfully
272    /// - `Some(Err(error))` if the task failed with an error
273    /// - `None` if the task was detached
274    #[must_use]
275    pub fn join(mut self) -> Option<Result<T, TaskError>> {
276        self.result_slot
277            .take()
278            // SAFETY: `join` consumes the only handle, and the completion sender
279            // never waits, so this call is the slot's only consumer.
280            .map(|slot| unsafe { slot.wait::<BlockingResultWait>() })
281    }
282
283    /// Checks if the task has finished execution.
284    ///
285    /// # Returns
286    /// `true` if the task has completed (successfully or with error), `false` if still running
287    #[must_use]
288    pub fn is_finished(&self) -> bool {
289        self.result_slot
290            .as_ref()
291            .is_some_and(|slot| slot.is_completed())
292    }
293}
294
295// ── TaskResultSender (std) ────────────────────────────────────────────────────
296
297/// Single-producer completion endpoint for a task result.
298#[cfg(feature = "std")]
299#[allow(clippy::module_name_repetitions)]
300pub struct TaskResultSender<T> {
301    slot: Option<Arc<TaskResultSlot<T>>>,
302}
303
304#[cfg(feature = "std")]
305impl<T> TaskResultSender<T> {
306    /// Complete the task result and wake any waiter.
307    pub fn send(self, result: Result<T, TaskError>) {
308        let mut sender = ManuallyDrop::new(self);
309        if let Some(slot) = sender.slot.take() {
310            slot.complete(result);
311        }
312    }
313}
314
315#[cfg(feature = "std")]
316impl<T> Drop for TaskResultSender<T> {
317    fn drop(&mut self) {
318        if let Some(slot) = self.slot.take() {
319            slot.complete(Err(TaskError::Cancelled));
320        }
321    }
322}
323
324// ── TaskHandle (no_std) ───────────────────────────────────────────────────────
325
326// For no_std environments, provide a simpler handle
327#[cfg(not(feature = "std"))]
328pub struct TaskHandle<T> {
329    id: TaskId,
330    _phantom: core::marker::PhantomData<T>,
331}
332
333#[cfg(not(feature = "std"))]
334impl<T> TaskHandle<T> {
335    /// Create a new task handle.
336    pub fn new(id: TaskId) -> Self {
337        Self {
338            id,
339            _phantom: core::marker::PhantomData,
340        }
341    }
342
343    /// Create a new detached task handle (alias for new in no_std).
344    pub fn new_detached(id: TaskId) -> Self {
345        Self::new(id)
346    }
347
348    /// Get the task ID.
349    pub fn id(&self) -> TaskId {
350        self.id
351    }
352}