use core::cell::UnsafeCell;
use core::future::Future;
use core::mem::MaybeUninit;
use core::pin::Pin;
use core::task::{Context, Poll};
use crate::waker;
const NO_WAITER: u32 = 0xFFFF_FFFF;
pub struct Channel<T, const N: usize> {
buffer: UnsafeCell<[MaybeUninit<T>; N]>,
head: crate::sync::atomic::AtomicUsize,
tail: crate::sync::atomic::AtomicUsize,
taken: crate::sync::atomic::AtomicBool,
recv_waiter: crate::sync::atomic::AtomicU32,
send_waiter: crate::sync::atomic::AtomicU32,
}
impl<T, const N: usize> Channel<T, N> {
#[cfg(not(loom))]
pub const fn new() -> Self {
Self {
buffer: UnsafeCell::new(unsafe {
MaybeUninit::uninit().assume_init()
}),
head: crate::sync::atomic::AtomicUsize::new(0),
tail: crate::sync::atomic::AtomicUsize::new(0),
taken: crate::sync::atomic::AtomicBool::new(false),
recv_waiter: crate::sync::atomic::AtomicU32::new(NO_WAITER),
send_waiter: crate::sync::atomic::AtomicU32::new(NO_WAITER),
}
}
#[cfg(loom)]
pub fn new() -> Self {
Self {
buffer: UnsafeCell::new(unsafe {
MaybeUninit::uninit().assume_init()
}),
head: crate::sync::atomic::AtomicUsize::new(0),
tail: crate::sync::atomic::AtomicUsize::new(0),
taken: crate::sync::atomic::AtomicBool::new(false),
recv_waiter: crate::sync::atomic::AtomicU32::new(NO_WAITER),
send_waiter: crate::sync::atomic::AtomicU32::new(NO_WAITER),
}
}
pub fn split(&'static self) -> Option<(Sender<'static, T, N>, Receiver<'static, T, N>)> {
if self.taken.swap(true, crate::sync::atomic::Ordering::AcqRel) {
return None;
}
Some((Sender { chan: self }, Receiver { chan: self }))
}
}
impl<T, const N: usize> Default for Channel<T, N> {
fn default() -> Self {
Self::new()
}
}
fn wake_waiter(slot: &crate::sync::atomic::AtomicU32) {
let w = slot.swap(NO_WAITER, crate::sync::atomic::Ordering::AcqRel);
if w != NO_WAITER {
waker::wake_task(crate::task::TaskId::from_u16(w as u16));
}
}
fn register_waiter(slot: &crate::sync::atomic::AtomicU32, id: crate::task::TaskId) {
slot.store(id.as_u16() as u32, crate::sync::atomic::Ordering::Release);
}
pub struct Sender<'a, T, const N: usize> {
chan: &'a Channel<T, N>,
}
pub struct Receiver<'a, T, const N: usize> {
chan: &'a Channel<T, N>,
}
impl<'a, T, const N: usize> Sender<'a, T, N> {
pub fn try_send(&self, val: T) -> Result<(), T> {
let chan = self.chan;
let head = chan.head.load(crate::sync::atomic::Ordering::Acquire);
let tail = chan.tail.load(crate::sync::atomic::Ordering::Relaxed);
let next_tail = (tail + 1) % N;
if next_tail == head {
return Err(val);
}
unsafe {
(*chan.buffer.get())[tail].write(val);
}
chan.tail
.store(next_tail, crate::sync::atomic::Ordering::Release);
wake_waiter(&chan.recv_waiter);
Ok(())
}
pub fn send(&self, val: T) -> SendFut<'_, 'a, T, N> {
SendFut {
tx: self,
val: Some(val),
registered: false,
}
}
}
impl<'a, T, const N: usize> Receiver<'a, T, N> {
pub fn try_recv(&self) -> Option<T> {
let chan = self.chan;
let head = chan.head.load(crate::sync::atomic::Ordering::Relaxed);
let tail = chan.tail.load(crate::sync::atomic::Ordering::Acquire);
if head == tail {
return None;
}
let val = unsafe { (*chan.buffer.get())[head].assume_init_read() };
chan.head
.store((head + 1) % N, crate::sync::atomic::Ordering::Release);
wake_waiter(&chan.send_waiter);
Some(val)
}
pub fn recv(&self) -> Recv<'_, 'a, T, N> {
Recv {
rx: self,
registered: false,
}
}
}
pub struct SendFut<'b, 'a, T, const N: usize> {
tx: &'b Sender<'a, T, N>,
val: Option<T>,
registered: bool,
}
impl<'b, 'a, T, const N: usize> SendFut<'b, 'a, T, N> {
fn clear_registration(&self) {
if self.registered {
self.tx
.chan
.send_waiter
.store(NO_WAITER, crate::sync::atomic::Ordering::Release);
}
}
}
impl<'b, 'a, T, const N: usize> Drop for SendFut<'b, 'a, T, N> {
fn drop(&mut self) {
self.clear_registration();
}
}
impl<'b, 'a, T, const N: usize> Future for SendFut<'b, 'a, T, N> {
type Output = ();
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<()> {
let this = unsafe { self.get_unchecked_mut() };
let val = this
.val
.take()
.expect("Send future polled after completion");
match this.tx.try_send(val) {
Ok(()) => {
this.clear_registration();
Poll::Ready(())
}
Err(v) => {
let id = crate::executor::current_task()
.expect("Sender::send().await polled outside of a task context");
register_waiter(&this.tx.chan.send_waiter, id);
this.registered = true;
match this.tx.try_send(v) {
Ok(()) => {
this.clear_registration();
Poll::Ready(())
}
Err(v) => {
this.val = Some(v);
Poll::Pending
}
}
}
}
}
}
pub struct Recv<'b, 'a, T, const N: usize> {
rx: &'b Receiver<'a, T, N>,
registered: bool,
}
impl<'b, 'a, T, const N: usize> Recv<'b, 'a, T, N> {
fn clear_registration(&self) {
if self.registered {
self.rx
.chan
.recv_waiter
.store(NO_WAITER, crate::sync::atomic::Ordering::Release);
}
}
}
impl<'b, 'a, T, const N: usize> Drop for Recv<'b, 'a, T, N> {
fn drop(&mut self) {
self.clear_registration();
}
}
impl<'b, 'a, T, const N: usize> Future for Recv<'b, 'a, T, N> {
type Output = T;
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<T> {
let this = unsafe { self.get_unchecked_mut() };
if let Some(v) = this.rx.try_recv() {
this.clear_registration();
return Poll::Ready(v);
}
let id = crate::executor::current_task()
.expect("Receiver::recv().await polled outside of a task context");
register_waiter(&this.rx.chan.recv_waiter, id);
this.registered = true;
if let Some(v) = this.rx.try_recv() {
this.clear_registration();
return Poll::Ready(v);
}
Poll::Pending
}
}
unsafe impl<T: Send, const N: usize> Sync for Channel<T, N> {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn send_recv_roundtrip() {
crate::kernel_test! {
static CHAN: Channel<u32, 8> = Channel::new();
let (tx, rx) = CHAN.split().expect("split once");
assert!(tx.try_send(42).is_ok());
assert_eq!(rx.try_recv(), Some(42));
assert_eq!(rx.try_recv(), None);
}
}
#[test]
fn full_channel_blocks_send() {
crate::kernel_test! {
static CHAN: Channel<u32, 3> = Channel::new(); let (tx, rx) = CHAN.split().expect("split once");
assert!(tx.try_send(1).is_ok());
assert!(tx.try_send(2).is_ok());
assert_eq!(tx.try_send(3), Err(3));
assert_eq!(rx.try_recv(), Some(1));
assert!(tx.try_send(3).is_ok());
assert_eq!(rx.try_recv(), Some(2));
assert_eq!(rx.try_recv(), Some(3));
}
}
#[test]
fn empty_channel_try_recv_none() {
crate::kernel_test! {
static CHAN: Channel<u32, 4> = Channel::new();
let (_tx, rx) = CHAN.split().expect("split once");
assert_eq!(rx.try_recv(), None);
}
}
#[test]
fn recv_future_ready_when_data_present() {
crate::kernel_test! {
static CHAN: Channel<u32, 4> = Channel::new();
let (tx, rx) = CHAN.split().expect("split once");
tx.try_send(7).unwrap();
let waker = crate::waker::task_waker(crate::task::TaskId::new(0, 0));
let mut cx = Context::from_waker(&waker);
let mut fut = rx.recv();
let pinned = unsafe { Pin::new_unchecked(&mut fut) };
assert_eq!(pinned.poll(&mut cx), Poll::Ready(7));
}
}
#[test]
#[should_panic(expected = "outside of a task context")]
fn recv_future_panics_without_task_context_when_empty() {
crate::kernel_test! {
static CHAN: Channel<u32, 4> = Channel::new();
let (_tx, rx) = CHAN.split().expect("split once");
let waker = crate::waker::task_waker(crate::task::TaskId::new(0, 0));
let mut cx = Context::from_waker(&waker);
let mut fut = rx.recv();
let pinned = unsafe { Pin::new_unchecked(&mut fut) };
let _ = pinned.poll(&mut cx);
}
}
}