use crate::channel::error::{CacheAligned, Channel, ChannelError, Result};
use std::cell::{Cell, UnsafeCell};
use std::mem::MaybeUninit;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
const SPSC_BLOCK_SPINS: usize = 6;
pub(crate) struct SpscChannel<T> {
buffer: Box<[UnsafeCell<MaybeUninit<T>>]>,
mask: usize,
head: CacheAligned<AtomicUsize>,
tail: CacheAligned<AtomicUsize>,
pub(super) closed: AtomicBool,
}
unsafe impl<T: Send> Send for SpscChannel<T> {}
unsafe impl<T: Send> Sync for SpscChannel<T> {}
impl<T> SpscChannel<T> {
pub fn new(capacity: usize) -> Self {
let capacity = capacity.next_power_of_two().max(2);
let buffer = (0..capacity)
.map(|_| UnsafeCell::new(MaybeUninit::uninit()))
.collect::<Vec<_>>()
.into_boxed_slice();
Self {
buffer,
mask: capacity - 1,
head: CacheAligned::new(AtomicUsize::new(0)),
tail: CacheAligned::new(AtomicUsize::new(0)),
closed: AtomicBool::new(false),
}
}
}
impl<T> SpscChannel<T> {
pub(super) fn indices(&self) -> (usize, usize) {
(
self.head.0.load(Ordering::Relaxed),
self.tail.0.load(Ordering::Relaxed),
)
}
}
impl<T> Drop for SpscChannel<T> {
fn drop(&mut self) {
let tail = self.tail.0.load(Ordering::Relaxed);
let head = self.head.0.load(Ordering::Relaxed);
let len = head.wrapping_sub(tail);
for i in 0..len {
unsafe {
let slot = &mut *self.buffer[(tail.wrapping_add(i)) & self.mask].get();
slot.assume_init_drop();
}
}
}
}
#[inline]
fn back_off(spin: &mut usize) {
if *spin < SPSC_BLOCK_SPINS {
for _ in 0..(1 << *spin) {
std::hint::spin_loop();
}
*spin += 1;
} else {
std::thread::yield_now();
}
}
pub(super) fn blocking<F, R>(mut attempt: F) -> Result<R>
where
F: FnMut() -> Result<R>,
{
let mut spin = 0;
loop {
match attempt() {
Err(ChannelError::Full | ChannelError::Empty) => back_off(&mut spin),
other => return other,
}
}
}
impl<T: Send> SpscChannel<T> {
#[inline]
pub(super) fn has_room(&self, head: usize, cached_tail: &Cell<usize>) -> bool {
if head.wrapping_sub(cached_tail.get()) < self.buffer.len() {
return true;
}
let tail = self.tail.0.load(Ordering::Acquire);
cached_tail.set(tail);
head.wrapping_sub(tail) < self.buffer.len()
}
#[inline]
pub(super) fn has_value(&self, tail: usize, cached_head: &Cell<usize>) -> bool {
if tail != cached_head.get() {
return true;
}
let head = self.head.0.load(Ordering::Acquire);
cached_head.set(head);
tail != head
}
pub(super) fn try_send_cached(&self, value: T, cached_tail: &Cell<usize>) -> Result<()> {
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
let head = self.head.0.load(Ordering::Relaxed);
if !self.has_room(head, cached_tail) {
return Err(ChannelError::Full);
}
unsafe {
let slot = &mut *self.buffer[head & self.mask].get();
slot.write(value);
}
self.head.0.store(head.wrapping_add(1), Ordering::Release);
Ok(())
}
pub(super) fn try_recv_cached(&self, cached_head: &Cell<usize>) -> Result<T> {
let tail = self.tail.0.load(Ordering::Relaxed);
if !self.has_value(tail, cached_head) {
if self.closed.load(Ordering::Acquire) {
let published = self.head.0.load(Ordering::Acquire);
cached_head.set(published);
if tail == published {
return Err(ChannelError::Closed);
}
} else {
return Err(ChannelError::Empty);
}
}
let value = unsafe {
let slot = &*self.buffer[tail & self.mask].get();
slot.assume_init_read()
};
self.tail.0.store(tail.wrapping_add(1), Ordering::Release);
Ok(value)
}
pub(super) fn send_cached(&self, value: T, cached_tail: &Cell<usize>) -> Result<()> {
let mut spin = 0;
loop {
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
let head = self.head.0.load(Ordering::Relaxed);
if self.has_room(head, cached_tail) {
unsafe {
let slot = &mut *self.buffer[head & self.mask].get();
slot.write(value);
}
self.head.0.store(head.wrapping_add(1), Ordering::Release);
return Ok(());
}
back_off(&mut spin);
}
}
}
impl<T: Send> Channel<T> for SpscChannel<T> {
fn send(&self, value: T) -> Result<()> {
let mut spin_count = 0;
loop {
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
let head = self.head.0.load(Ordering::Relaxed);
let tail = self.tail.0.load(Ordering::Acquire);
if head.wrapping_sub(tail) < self.buffer.len() {
unsafe {
let slot = &mut *self.buffer[head & self.mask].get();
slot.write(value);
}
self.head.0.store(head.wrapping_add(1), Ordering::Release);
return Ok(());
}
if spin_count < SPSC_BLOCK_SPINS {
for _ in 0..(1 << spin_count) {
std::hint::spin_loop();
}
spin_count += 1;
} else {
std::thread::yield_now();
}
}
}
fn try_send(&self, value: T) -> Result<()> {
if self.closed.load(Ordering::Acquire) {
return Err(ChannelError::Closed);
}
let head = self.head.0.load(Ordering::Relaxed);
let tail = self.tail.0.load(Ordering::Acquire);
if head.wrapping_sub(tail) >= self.buffer.len() {
return Err(ChannelError::Full);
}
unsafe {
let slot = &mut *self.buffer[head & self.mask].get();
slot.write(value);
}
self.head.0.store(head.wrapping_add(1), Ordering::Release);
Ok(())
}
fn recv(&self) -> Result<T> {
let mut spin_count = 0;
loop {
match self.try_recv() {
Ok(value) => return Ok(value),
Err(ChannelError::Empty) => {
if spin_count < SPSC_BLOCK_SPINS {
for _ in 0..(1 << spin_count) {
std::hint::spin_loop();
}
spin_count += 1;
} else {
std::thread::yield_now();
}
}
Err(e) => return Err(e), }
}
}
fn try_recv(&self) -> Result<T> {
let tail = self.tail.0.load(Ordering::Relaxed);
let head = self.head.0.load(Ordering::Acquire);
if tail == head {
if self.closed.load(Ordering::Acquire) {
let published_head = self.head.0.load(Ordering::Acquire);
if tail == published_head {
return Err(ChannelError::Closed);
}
} else {
return Err(ChannelError::Empty);
}
}
let value = unsafe {
let slot = &*self.buffer[tail & self.mask].get();
slot.assume_init_read()
};
self.tail.0.store(tail.wrapping_add(1), Ordering::Release);
Ok(value)
}
fn is_empty(&self) -> bool {
let tail = self.tail.0.load(Ordering::Relaxed);
let head = self.head.0.load(Ordering::Acquire);
tail == head
}
fn is_full(&self) -> bool {
let head = self.head.0.load(Ordering::Relaxed);
let tail = self.tail.0.load(Ordering::Acquire);
head.wrapping_sub(tail) >= self.buffer.len()
}
fn capacity(&self) -> Option<usize> {
Some(self.buffer.len())
}
}