use crate::error::TaskError;
use super::id_and_context::TaskId;
#[cfg(feature = "std")]
use core::mem::ManuallyDrop;
#[cfg(feature = "std")]
use std::sync::{Arc, atomic::AtomicU8};
#[cfg(feature = "std")]
use std::thread;
#[cfg(feature = "std")]
use moirai_utils::{CacheAligned, ResultCell};
#[cfg(feature = "std")]
pub(super) mod result_wait {
pub(super) mod sealed {
pub trait Sealed {}
}
pub trait ResultWaitPolicy: sealed::Sealed {
const SPIN_ATTEMPTS: usize;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct BlockingResultWait;
impl sealed::Sealed for BlockingResultWait {}
impl ResultWaitPolicy for BlockingResultWait {
const SPIN_ATTEMPTS: usize = super::super::MAX_SPIN_ATTEMPTS;
}
}
#[cfg(feature = "std")]
pub use result_wait::{BlockingResultWait, ResultWaitPolicy};
#[cfg(feature = "std")]
struct TaskResultSlot<T> {
cell: ResultCell<Result<T, TaskError>, thread::Thread, CacheAligned<AtomicU8>>,
}
#[cfg(feature = "std")]
impl<T> TaskResultSlot<T> {
fn new() -> Self {
Self {
cell: ResultCell::new(),
}
}
fn complete(&self, result: Result<T, TaskError>) {
self.cell.complete(result);
}
unsafe fn wait<P>(&self) -> Result<T, TaskError>
where
P: ResultWaitPolicy,
{
if let Some(result) = self.try_take_ready() {
return result;
}
for _ in 0..P::SPIN_ATTEMPTS {
if let Some(result) = self.try_take_observed_ready() {
return result;
}
core::hint::spin_loop();
}
unsafe { self.register_waiter() };
loop {
if let Some(result) = self.try_take_observed_ready() {
return result;
}
thread::park();
}
}
fn is_completed(&self) -> bool {
self.cell.is_completed()
}
fn try_take_ready(&self) -> Option<Result<T, TaskError>> {
self.cell.try_take_ready()
}
fn try_take_observed_ready(&self) -> Option<Result<T, TaskError>> {
self.cell.try_take_observed_ready()
}
unsafe fn register_waiter(&self) {
unsafe { self.cell.register(&thread::current()) };
}
#[cfg(feature = "result-diagnostics")]
fn has_registered_waiter(&self) -> bool {
self.cell.has_registered_waiter()
}
}
#[cfg(all(feature = "std", feature = "result-diagnostics"))]
const DIAGNOSTIC_READY_VALUE: usize = 42;
#[cfg(all(feature = "std", feature = "result-diagnostics"))]
#[doc(hidden)]
pub fn diagnostic_result_slot_ready_take() -> usize {
let slot = TaskResultSlot::new();
slot.complete(Ok(DIAGNOSTIC_READY_VALUE));
match slot.try_take_ready() {
Some(Ok(value)) => value,
_ => 0,
}
}
#[cfg(all(feature = "std", feature = "result-diagnostics"))]
#[doc(hidden)]
pub fn diagnostic_result_slot_spin_miss() -> usize {
let slot = TaskResultSlot::<usize>::new();
let mut misses = 0usize;
for _ in 0..BlockingResultWait::SPIN_ATTEMPTS {
if slot.try_take_observed_ready().is_none() {
misses = misses.wrapping_add(1);
}
core::hint::spin_loop();
}
misses
}
#[cfg(all(feature = "std", feature = "result-diagnostics"))]
#[doc(hidden)]
pub fn diagnostic_result_slot_register_waiter() -> usize {
let slot = TaskResultSlot::<usize>::new();
unsafe { slot.register_waiter() };
usize::from(slot.has_registered_waiter())
}
#[cfg(all(feature = "std", feature = "result-diagnostics"))]
#[doc(hidden)]
pub fn diagnostic_result_slot_complete_waiting() -> usize {
let slot = TaskResultSlot::new();
unsafe { slot.register_waiter() };
slot.complete(Ok(DIAGNOSTIC_READY_VALUE));
match slot.try_take_ready() {
Some(Ok(value)) => value,
_ => 0,
}
}
#[cfg(feature = "std")]
#[allow(clippy::module_name_repetitions)]
pub struct TaskHandle<T> {
id: TaskId,
result_slot: Option<Arc<TaskResultSlot<T>>>,
}
#[cfg(feature = "std")]
impl<T> TaskHandle<T> {
#[must_use]
pub fn new_pending(id: TaskId) -> (Self, TaskResultSender<T>) {
let slot = Arc::new(TaskResultSlot::new());
(
Self {
id,
result_slot: Some(Arc::clone(&slot)),
},
TaskResultSender { slot: Some(slot) },
)
}
#[must_use]
pub fn ready(id: TaskId, result: Result<T, TaskError>) -> Self {
let slot = Arc::new(TaskResultSlot::new());
slot.complete(result);
Self {
id,
result_slot: Some(slot),
}
}
#[must_use]
pub fn new_detached(id: TaskId) -> Self {
Self {
id,
result_slot: None,
}
}
#[must_use]
pub fn id(&self) -> TaskId {
self.id
}
#[must_use]
pub fn join(mut self) -> Option<Result<T, TaskError>> {
self.result_slot
.take()
.map(|slot| unsafe { slot.wait::<BlockingResultWait>() })
}
#[must_use]
pub fn is_finished(&self) -> bool {
self.result_slot
.as_ref()
.is_some_and(|slot| slot.is_completed())
}
}
#[cfg(feature = "std")]
#[allow(clippy::module_name_repetitions)]
pub struct TaskResultSender<T> {
slot: Option<Arc<TaskResultSlot<T>>>,
}
#[cfg(feature = "std")]
impl<T> TaskResultSender<T> {
pub fn send(self, result: Result<T, TaskError>) {
let mut sender = ManuallyDrop::new(self);
if let Some(slot) = sender.slot.take() {
slot.complete(result);
}
}
}
#[cfg(feature = "std")]
impl<T> Drop for TaskResultSender<T> {
fn drop(&mut self) {
if let Some(slot) = self.slot.take() {
slot.complete(Err(TaskError::Cancelled));
}
}
}
#[cfg(not(feature = "std"))]
pub struct TaskHandle<T> {
id: TaskId,
_phantom: core::marker::PhantomData<T>,
}
#[cfg(not(feature = "std"))]
impl<T> TaskHandle<T> {
pub fn new(id: TaskId) -> Self {
Self {
id,
_phantom: core::marker::PhantomData,
}
}
pub fn new_detached(id: TaskId) -> Self {
Self::new(id)
}
pub fn id(&self) -> TaskId {
self.id
}
}