#![expect(
clippy::unwrap_used,
reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
)]
use super::recv::MpmcReceiver;
use super::send::MpmcSender;
use super::{MPMC_BLOCK_SPINS, MpmcState, block::backoff_step};
use crate::channel::CHANNEL_STORE_LOAD_ORDER;
use crate::channel::error::{Channel, ChannelError, Result};
use moirai_utils::queue::LockFreeQueue;
use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering, fence};
use std::sync::{Arc, Condvar, Mutex};
mod roles;
pub struct MpmcChannel<T> {
pub(super) state: (Mutex<MpmcState<T>>, Condvar, Condvar),
pub(super) bounded: Option<LockFreeQueue<T>>,
pub(super) closed: AtomicBool,
pub(super) sender_waiter_count: AtomicUsize,
pub(super) receiver_waiter_count: AtomicUsize,
}
const UNBOUNDED_PREALLOCATED_SLOTS: usize = 16;
impl<T> MpmcChannel<T> {
pub fn new(capacity: Option<usize>) -> Self {
let state = MpmcState {
queue: if capacity.is_some() {
VecDeque::new()
} else {
VecDeque::with_capacity(UNBOUNDED_PREALLOCATED_SLOTS)
},
capacity,
closed: false,
sender_count: 0,
receiver_count: 0,
};
let bounded = capacity.map(LockFreeQueue::with_capacity);
Self {
state: (Mutex::new(state), Condvar::new(), Condvar::new()),
bounded,
closed: AtomicBool::new(false),
sender_waiter_count: AtomicUsize::new(0),
receiver_waiter_count: AtomicUsize::new(0),
}
}
pub fn unbounded() -> Self {
Self::new(None)
}
pub fn bounded(capacity: usize) -> Self {
Self::new(Some(capacity))
}
pub fn channel(capacity: Option<usize>) -> (MpmcSender<T>, MpmcReceiver<T>) {
let channel = Arc::new(Self::new(capacity));
let (mutex, _, _) = &channel.state;
{
let mut state = mutex.lock().unwrap();
state.sender_count = 1;
state.receiver_count = 1;
}
(
MpmcSender {
channel: channel.clone(),
},
MpmcReceiver { channel },
)
}
fn send_bounded(&self, queue: &LockFreeQueue<T>, mut value: T) -> Result<()>
where
T: Send,
{
let mut spin_count = 0;
loop {
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
match queue.try_enqueue(value) {
Ok(()) => {
self.wake_receiver_after_push();
return Ok(());
}
Err(returned) => {
value = returned;
}
}
if spin_count < MPMC_BLOCK_SPINS {
backoff_step(&mut spin_count);
continue;
}
let (mutex, not_full, _) = &self.state;
let mut guard = mutex.lock().unwrap();
if self.closed.load(Ordering::Acquire) || guard.closed {
return Err(ChannelError::Closed);
}
self.sender_waiter_count
.fetch_add(1, CHANNEL_STORE_LOAD_ORDER);
fence(CHANNEL_STORE_LOAD_ORDER);
match queue.try_enqueue(value) {
Ok(()) => {
self.sender_waiter_count.fetch_sub(1, Ordering::Relaxed);
drop(guard);
if self.receiver_waiter_count.load(Ordering::Relaxed) > 0 {
let (_, _, not_empty) = &self.state;
not_empty.notify_one();
}
return Ok(());
}
Err(returned) => {
value = returned;
}
}
guard = not_full.wait(guard).unwrap();
self.sender_waiter_count.fetch_sub(1, Ordering::Relaxed);
}
}
fn wake_receiver_after_push(&self) {
fence(CHANNEL_STORE_LOAD_ORDER);
if self.receiver_waiter_count.load(CHANNEL_STORE_LOAD_ORDER) > 0 {
let (mutex, _, not_empty) = &self.state;
let _guard = mutex.lock().unwrap();
not_empty.notify_one();
}
}
fn wake_sender_after_pop(&self) {
fence(CHANNEL_STORE_LOAD_ORDER);
if self.sender_waiter_count.load(CHANNEL_STORE_LOAD_ORDER) > 0 {
let (mutex, not_full, _) = &self.state;
let _guard = mutex.lock().unwrap();
not_full.notify_one();
}
}
fn recv_bounded(&self, queue: &LockFreeQueue<T>) -> Result<T>
where
T: Send,
{
let mut spin_count = 0;
loop {
if let Some(value) = queue.try_dequeue() {
self.wake_sender_after_pop();
return Ok(value);
}
if self.closed.load(Ordering::Acquire) {
if queue.is_empty() {
return Err(ChannelError::Closed);
}
std::hint::spin_loop();
continue;
}
if spin_count < MPMC_BLOCK_SPINS {
backoff_step(&mut spin_count);
continue;
}
let (mutex, _, not_empty) = &self.state;
let mut guard = mutex.lock().unwrap();
self.receiver_waiter_count
.fetch_add(1, CHANNEL_STORE_LOAD_ORDER);
fence(CHANNEL_STORE_LOAD_ORDER);
if let Some(value) = queue.try_dequeue() {
self.receiver_waiter_count.fetch_sub(1, Ordering::Relaxed);
if self.sender_waiter_count.load(Ordering::Relaxed) > 0 {
let (_, not_full, _) = &self.state;
not_full.notify_one();
}
drop(guard);
return Ok(value);
}
if self.closed.load(Ordering::Acquire) || guard.closed {
self.receiver_waiter_count.fetch_sub(1, Ordering::Relaxed);
return Err(ChannelError::Closed);
}
guard = not_empty.wait(guard).unwrap();
self.receiver_waiter_count.fetch_sub(1, Ordering::Relaxed);
}
}
}
impl<T: Send> Channel<T> for MpmcChannel<T> {
fn send(&self, value: T) -> Result<()> {
if let Some(queue) = &self.bounded {
return self.send_bounded(queue, value);
}
self.send_unbounded(value)
}
fn try_send(&self, value: T) -> Result<()> {
if let Some(queue) = &self.bounded {
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
queue.try_enqueue(value).map_err(|_| ChannelError::Full)?;
self.wake_receiver_after_push();
return Ok(());
}
self.try_send_unbounded(value)
}
fn recv(&self) -> Result<T> {
if let Some(queue) = &self.bounded {
return self.recv_bounded(queue);
}
self.recv_unbounded()
}
fn try_recv(&self) -> Result<T> {
if let Some(queue) = &self.bounded {
if let Some(value) = queue.try_dequeue() {
self.wake_sender_after_pop();
return Ok(value);
}
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
return Err(ChannelError::Empty);
}
self.try_recv_unbounded()
}
fn is_empty(&self) -> bool {
if let Some(queue) = &self.bounded {
return queue.is_empty();
}
self.is_empty_unbounded()
}
fn is_full(&self) -> bool {
if let Some(queue) = &self.bounded {
return queue.is_full();
}
self.is_full_unbounded()
}
fn capacity(&self) -> Option<usize> {
if let Some(queue) = &self.bounded {
return Some(queue.capacity());
}
self.capacity_unbounded()
}
}