use crate::error::TaskError;
use super::id_and_context::TaskId;
#[cfg(feature = "std")]
use core::cell::UnsafeCell;
#[cfg(feature = "std")]
use core::mem::{ManuallyDrop, MaybeUninit};
#[cfg(feature = "std")]
use std::sync::{
atomic::{AtomicU8, Ordering},
Arc,
};
#[cfg(feature = "std")]
use std::thread;
#[cfg(feature = "std")]
const RESULT_PENDING: u8 = 0;
#[cfg(feature = "std")]
const RESULT_WRITING: u8 = 1;
#[cfg(feature = "std")]
const RESULT_READY: u8 = 2;
#[cfg(feature = "std")]
const RESULT_TAKEN: u8 = 3;
#[cfg(feature = "std")]
const RESULT_WAITING: u8 = 4;
#[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")]
#[repr(align(64))]
struct TaskResultSlot<T> {
state: AtomicU8,
_pad: [u8; 63],
result: UnsafeCell<MaybeUninit<Result<T, TaskError>>>,
waiter: UnsafeCell<MaybeUninit<thread::Thread>>,
}
#[cfg(feature = "std")]
unsafe impl<T: Send> Send for TaskResultSlot<T> {}
#[cfg(feature = "std")]
unsafe impl<T: Send> Sync for TaskResultSlot<T> {}
#[cfg(feature = "std")]
impl<T> TaskResultSlot<T> {
fn new() -> Self {
Self {
state: AtomicU8::new(RESULT_PENDING),
_pad: [0u8; 63],
result: UnsafeCell::new(MaybeUninit::uninit()),
waiter: UnsafeCell::new(MaybeUninit::uninit()),
}
}
fn complete(&self, result: Result<T, TaskError>) {
let Some(waiting) = self.begin_completion() else {
return;
};
unsafe {
(*self.result.get()).write(result);
}
self.state.store(RESULT_READY, Ordering::Release);
if waiting {
let thread = unsafe { (*self.waiter.get()).assume_init_read() };
thread.unpark();
}
}
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();
}
self.register_waiter();
loop {
if let Some(result) = self.try_take_observed_ready() {
return result;
}
thread::park();
}
}
fn is_completed(&self) -> bool {
self.state.load(Ordering::Acquire) == RESULT_READY
}
fn try_take_ready(&self) -> Option<Result<T, TaskError>> {
if self
.state
.compare_exchange(
RESULT_READY,
RESULT_TAKEN,
Ordering::Acquire,
Ordering::Relaxed,
)
.is_ok()
{
Some(unsafe { (*self.result.get()).assume_init_read() })
} else {
None
}
}
fn try_take_observed_ready(&self) -> Option<Result<T, TaskError>> {
if self.state.load(Ordering::Relaxed) == RESULT_READY {
self.try_take_ready()
} else {
None
}
}
fn register_waiter(&self) {
loop {
match self.state.load(Ordering::Acquire) {
RESULT_PENDING => {
unsafe {
(*self.waiter.get()).write(thread::current());
}
if self
.state
.compare_exchange(
RESULT_PENDING,
RESULT_WAITING,
Ordering::Release,
Ordering::Acquire,
)
.is_ok()
{
return;
}
unsafe {
(*self.waiter.get()).assume_init_drop();
}
}
RESULT_WRITING => core::hint::spin_loop(),
_ => return,
}
}
}
fn begin_completion(&self) -> Option<bool> {
match self.state.compare_exchange(
RESULT_PENDING,
RESULT_WRITING,
Ordering::Relaxed,
Ordering::Acquire,
) {
Ok(_) => Some(false),
Err(RESULT_WAITING) => {
if self
.state
.compare_exchange(
RESULT_WAITING,
RESULT_WRITING,
Ordering::Acquire,
Ordering::Acquire,
)
.is_ok()
{
Some(true)
} else {
None
}
}
Err(_) => None,
}
}
}
#[cfg(feature = "std")]
impl<T> Drop for TaskResultSlot<T> {
fn drop(&mut self) {
let state = *self.state.get_mut();
if state == RESULT_READY {
unsafe {
self.result.get_mut().assume_init_drop();
}
} else if state == RESULT_WAITING {
unsafe {
self.waiter.get_mut().assume_init_drop();
}
}
}
}
#[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();
slot.register_waiter();
usize::from(slot.state.load(Ordering::Acquire) == RESULT_WAITING)
}
#[cfg(all(feature = "std", feature = "result-diagnostics"))]
#[doc(hidden)]
pub fn diagnostic_result_slot_complete_waiting() -> usize {
let slot = TaskResultSlot::new();
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| 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
}
}