use std::cell::Cell;
use std::marker::PhantomPinned;
use std::pin::Pin;
use std::ptr::NonNull;
use o3::marker::ThreadBound;
use crate::task::{Context, Waker};
type WaiterPtr<T> = NonNull<Waiter<'static, T>>;
pub struct WaitQueue<T = ()> {
head: Cell<Option<WaiterPtr<T>>>,
tail: Cell<Option<WaiterPtr<T>>>,
len: Cell<usize>,
capacity: usize,
_pin: PhantomPinned,
_thread: ThreadBound,
}
pub struct Waiter<'d, T = ()> {
queue: Cell<Option<NonNull<WaitQueue<T>>>>,
previous: Cell<Option<WaiterPtr<T>>>,
next: Cell<Option<WaiterPtr<T>>>,
wake: Cell<Option<Waker<'d>>>,
assigned: Cell<Option<T>>,
_pin: PhantomPinned,
_thread: ThreadBound,
}
impl WaitQueue<()> {
pub const fn with_capacity(capacity: usize) -> Self {
Self::build(capacity)
}
}
impl<T> WaitQueue<T> {
const fn build(capacity: usize) -> Self {
Self {
head: Cell::new(None),
tail: Cell::new(None),
len: Cell::new(0),
capacity,
_pin: PhantomPinned,
_thread: ThreadBound::NEW,
}
}
pub const fn with_payload_capacity(capacity: usize) -> Self {
Self::build(capacity)
}
pub fn get_pinned(queues: Pin<&[Self]>, index: usize) -> Option<Pin<&Self>> {
let queue = queues.get_ref().get(index)?;
Some(unsafe { Pin::new_unchecked(queue) })
}
fn contains<'d>(self: Pin<&Self>, waiter: Pin<&Waiter<'d, T>>) -> bool {
waiter.queue.get() == Some(NonNull::from(self.get_ref()))
}
pub fn can_register<'d>(self: Pin<&Self>, waiter: Pin<&Waiter<'d, T>>) -> bool {
self.contains(waiter) || self.len.get() < self.capacity
}
#[must_use]
pub fn try_register<'d>(
self: Pin<&Self>,
waiter: Pin<&Waiter<'d, T>>,
context: Pin<&Context<'_, 'd>>,
) -> bool {
self.try_register_waker(waiter, unsafe { context.waker_unchecked() })
}
#[doc(hidden)]
pub fn try_register_waker<'d>(
self: Pin<&Self>,
waiter: Pin<&Waiter<'d, T>>,
waker: Waker<'d>,
) -> bool {
if self.contains(waiter) {
waiter.wake.set(Some(waker));
return true;
}
if self.len.get() == self.capacity {
return false;
}
waiter.unregister();
debug_assert!(waiter.previous.get().is_none());
debug_assert!(waiter.next.get().is_none());
let queue = NonNull::from(self.get_ref());
let node = NonNull::from(waiter.get_ref()).cast::<Waiter<'static, T>>();
let previous = self.tail.get();
waiter.queue.set(Some(queue));
waiter.previous.set(previous);
waiter.wake.set(Some(waker));
if let Some(previous) = previous {
unsafe { previous.as_ref() }.next.set(Some(node));
} else {
self.head.set(Some(node));
}
self.tail.set(Some(node));
self.len.set(self.len.get() + 1);
true
}
fn unlink<'d>(self: Pin<&Self>, waiter: NonNull<Waiter<'d, T>>) -> Option<Waker<'d>> {
let waiter = unsafe { waiter.as_ref() };
if waiter.queue.get() != Some(NonNull::from(self.get_ref())) {
return None;
}
let previous = waiter.previous.take();
let next = waiter.next.take();
if let Some(previous) = previous {
unsafe { previous.as_ref() }.next.set(next);
} else {
self.head.set(next);
}
if let Some(next) = next {
unsafe { next.as_ref() }.previous.set(previous);
} else {
self.tail.set(previous);
}
waiter.queue.set(None);
self.len.set(self.len.get() - 1);
waiter.wake.take()
}
fn pop_next(self: Pin<&Self>, assigned: Option<T>, wake: bool) -> Result<(), Option<T>> {
let Some(node) = self.head.get() else {
return Err(assigned);
};
let waiter = node.cast::<Waiter<'_, T>>();
let waker = self
.unlink(waiter)
.expect("dope-fiber: linked waiter missing its queue");
unsafe { waiter.as_ref() }.assigned.set(assigned);
if wake {
waker.wake();
}
Ok(())
}
pub fn wake(self: Pin<&Self>) {
while self.pop_next(None, true).is_ok() {}
}
pub fn wake_one(self: Pin<&Self>) {
let _ = self.pop_next(None, true);
}
pub fn assign_one(self: Pin<&Self>, value: T) -> Result<(), T> {
match self.pop_next(Some(value), true) {
Err(Some(value)) => Err(value),
_ => Ok(()),
}
}
pub fn len(&self) -> usize {
self.len.get()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl<T> Drop for WaitQueue<T> {
fn drop(&mut self) {
let queue = unsafe { Pin::new_unchecked(&*self) };
while queue.pop_next(None, false).is_ok() {}
}
}
impl<'d, T> Waiter<'d, T> {
pub const fn new() -> Self {
Self {
queue: Cell::new(None),
previous: Cell::new(None),
next: Cell::new(None),
wake: Cell::new(None),
assigned: Cell::new(None),
_pin: PhantomPinned,
_thread: ThreadBound::NEW,
}
}
pub fn unregister(self: Pin<&Self>) -> bool {
let Some(queue) = self.queue.get() else {
return false;
};
unsafe { Pin::new_unchecked(queue.as_ref()) }
.unlink(NonNull::from(self.get_ref()))
.is_some()
}
pub fn is_registered(&self) -> bool {
self.queue.get().is_some()
}
pub fn take_assigned(&self) -> Option<T> {
self.assigned.take()
}
}
impl<T> Default for Waiter<'_, T> {
fn default() -> Self {
Self::new()
}
}
impl<T> Drop for Waiter<'_, T> {
fn drop(&mut self) {
let Some(queue) = self.queue.get() else {
return;
};
let queue = unsafe { Pin::new_unchecked(queue.as_ref()) };
let _ = queue.unlink(NonNull::from(&*self));
}
}