use std::collections::VecDeque;
use std::fmt;
use std::future::Future;
use std::mem;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use crate::internal::arena::Arena;
use crate::internal::arena::SlotId;
use crate::internal::mutex::Mutex;
use crate::internal::wake_all;
use crate::internal::wakerset::WakerSet;
use crate::internal::wakerset::WakerToken;
#[cfg(test)]
mod tests;
pub fn unbounded<T: Clone>() -> (UnboundedSender<T>, UnboundedReceiver<T>) {
let mut receivers = Arena::new();
let key = receivers.insert(0);
let shared = Arc::new(Shared {
inner: Mutex::new(Inner {
buffer: VecDeque::new(),
head: 0,
head_receivers: 1,
tail: 0,
receivers,
peak_len: 0,
waiters: WakerSet::new(),
}),
senders: AtomicUsize::new(1),
});
let sender = UnboundedSender {
shared: shared.clone(),
};
let receiver = UnboundedReceiver { shared, key };
(sender, receiver)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RecvError {
Disconnected,
}
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
RecvError::Disconnected => write!(f, "receiving on a disconnected channel"),
}
}
}
impl std::error::Error for RecvError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TryRecvError {
Empty,
Disconnected,
}
impl fmt::Display for TryRecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TryRecvError::Empty => write!(f, "receiving on an empty channel"),
TryRecvError::Disconnected => write!(f, "receiving on a disconnected channel"),
}
}
}
impl std::error::Error for TryRecvError {}
const MIN_RETAINED_CAPACITY: usize = 64;
struct Inner<T> {
buffer: VecDeque<Arc<T>>,
head: u64,
head_receivers: usize,
tail: u64,
receivers: Arena<u64>,
peak_len: usize,
waiters: WakerSet,
}
struct Reclaimed<T> {
first: Option<Arc<T>>,
rest: Vec<Arc<T>>,
}
impl<T> Reclaimed<T> {
fn empty() -> Self {
Self {
first: None,
rest: vec![],
}
}
fn first(&self) -> Option<&Arc<T>> {
self.first.as_ref()
}
fn is_empty(&self) -> bool {
self.first.is_none()
}
fn drop_messages(self) {
let Self { first, rest } = self;
drop((first, rest));
}
}
impl<T> Inner<T> {
fn insert_receiver(&mut self, head: u64) -> SlotId {
if head == self.head {
self.head_receivers += 1;
}
self.receivers.insert(head)
}
fn remove_receiver(&mut self, key: SlotId) -> Reclaimed<T> {
let head = self.receivers.remove(key);
if head == self.head {
self.release_head_receiver()
} else {
Reclaimed::empty()
}
}
fn release_head_receiver(&mut self) -> Reclaimed<T> {
self.head_receivers -= 1;
if self.head_receivers == 0 {
self.reclaim_consumed()
} else {
Reclaimed::empty()
}
}
fn receive(&mut self, key: SlotId) -> Option<(Arc<T>, Reclaimed<T>)> {
let head = {
let cursor = self
.receivers
.get_mut(key)
.expect("active broadcast receiver must be registered");
if *cursor >= self.tail {
return None;
}
let head = *cursor;
*cursor += 1;
head
};
debug_assert!(head >= self.head);
let offset = (head - self.head) as usize;
let msg = self.buffer[offset].clone();
let reclaimed = if head == self.head {
self.release_head_receiver()
} else {
Reclaimed::empty()
};
debug_assert!(
reclaimed
.first()
.is_none_or(|first| Arc::ptr_eq(first, &msg))
);
Some((msg, reclaimed))
}
fn reclaim_consumed(&mut self) -> Reclaimed<T> {
let mut next_head = self.tail;
let mut head_receivers = 0;
for head in self.receivers.values() {
if *head < next_head {
next_head = *head;
head_receivers = 1;
} else if *head == next_head {
head_receivers += 1;
}
}
debug_assert!(next_head >= self.head);
let consumed = usize::try_from(next_head - self.head)
.expect("retained broadcast message count exceeds usize");
let first = if consumed == 0 {
None
} else {
self.buffer.pop_front()
};
let rest = self.buffer.drain(..consumed.saturating_sub(1)).collect();
let reclaimed = Reclaimed { first, rest };
self.head = next_head;
self.head_receivers = head_receivers;
self.shrink_buffer();
reclaimed
}
fn shrink_buffer(&mut self) {
if !self.buffer.is_empty() {
return;
}
let peak = mem::take(&mut self.peak_len);
let capacity = self.buffer.capacity();
if capacity > MIN_RETAINED_CAPACITY && peak <= capacity / 4 {
self.buffer.shrink_to(MIN_RETAINED_CAPACITY.max(peak * 2));
}
}
}
struct Shared<T> {
inner: Mutex<Inner<T>>,
senders: AtomicUsize,
}
pub struct UnboundedSender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for UnboundedSender<T> {
fn clone(&self) -> Self {
self.shared.senders.fetch_add(1, Ordering::Relaxed);
Self {
shared: self.shared.clone(),
}
}
}
impl<T> fmt::Debug for UnboundedSender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UnboundedSender").finish_non_exhaustive()
}
}
impl<T> Drop for UnboundedSender<T> {
fn drop(&mut self) {
match self.shared.senders.fetch_sub(1, Ordering::AcqRel) {
1 => {
let wakers = {
let mut inner = self.shared.inner.lock();
inner.waiters.take_all()
};
wake_all(wakers);
}
_ => {
}
}
}
}
impl<T> UnboundedSender<T> {
pub fn send(&self, msg: T) {
let msg = Arc::new(msg);
let wakers = {
let mut inner = self.shared.inner.lock();
inner.tail = inner
.tail
.checked_add(1)
.expect("broadcast channel version counter overflowed");
if inner.receivers.is_empty() {
debug_assert!(inner.buffer.is_empty());
debug_assert_eq!(inner.head_receivers, 0);
inner.head = inner.tail;
} else {
inner.buffer.push_back(msg);
inner.peak_len = inner.peak_len.max(inner.buffer.len());
}
inner.waiters.drain()
};
wake_all(wakers);
}
pub fn retained_message_count(&self) -> usize {
self.shared.inner.lock().buffer.len()
}
#[must_use = "the receiver is dropped immediately if it is not retained"]
pub fn subscribe(&self) -> UnboundedReceiver<T> {
let mut inner = self.shared.inner.lock();
let head = inner.tail;
let key = inner.insert_receiver(head);
let shared = self.shared.clone();
UnboundedReceiver { shared, key }
}
}
pub struct UnboundedReceiver<T> {
shared: Arc<Shared<T>>,
key: SlotId,
}
impl<T> fmt::Debug for UnboundedReceiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UnboundedReceiver").finish_non_exhaustive()
}
}
impl<T> Drop for UnboundedReceiver<T> {
fn drop(&mut self) {
let reclaimed = {
let mut inner = self.shared.inner.lock();
inner.remove_receiver(self.key)
};
drop(reclaimed);
}
}
impl<T: Clone> UnboundedReceiver<T> {
pub async fn recv(&mut self) -> Result<T, RecvError> {
Recv {
receiver: self,
token: None,
}
.await
}
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
let (msg, reclaimed) = self.try_recv_shared()?;
Ok(take_msg(msg, reclaimed))
}
}
fn take_msg<T: Clone>(msg: Arc<T>, reclaimed: Reclaimed<T>) -> T {
let sole_owner = !reclaimed.is_empty();
reclaimed.drop_messages();
if !sole_owner {
return (*msg).clone();
}
Arc::try_unwrap(msg).unwrap_or_else(|msg| (*msg).clone())
}
impl<T> UnboundedReceiver<T> {
fn try_recv_shared(&mut self) -> Result<(Arc<T>, Reclaimed<T>), TryRecvError> {
let mut inner = self.shared.inner.lock();
if let Some(received) = inner.receive(self.key) {
return Ok(received);
}
if self.shared.senders.load(Ordering::Acquire) == 0 {
Err(TryRecvError::Disconnected)
} else {
Err(TryRecvError::Empty)
}
}
#[must_use = "the receiver is dropped immediately if it is not retained"]
pub fn resubscribe(&self) -> Self {
let mut inner = self.shared.inner.lock();
let head = inner.tail;
let key = inner.insert_receiver(head);
let shared = self.shared.clone();
Self { shared, key }
}
pub fn unread_message_count(&self) -> usize {
let inner = self.shared.inner.lock();
let head = *inner
.receivers
.get(self.key)
.expect("active broadcast receiver must be registered");
usize::try_from(inner.tail - head).expect("unread broadcast message count exceeds usize")
}
}
struct Recv<'a, T> {
receiver: &'a mut UnboundedReceiver<T>,
token: Option<WakerToken>,
}
impl<T> Drop for Recv<'_, T> {
fn drop(&mut self) {
if self.token.is_none() {
return;
}
let mut inner = self.receiver.shared.inner.lock();
let cursor = *inner
.receivers
.get(self.receiver.key)
.expect("active broadcast receiver must be registered");
if cursor != inner.tail || self.receiver.shared.senders.load(Ordering::Acquire) == 0 {
self.token = None;
return;
}
let waker = inner.waiters.unregister(&mut self.token);
drop(inner);
drop(waker);
}
}
impl<T: Clone> Future for Recv<'_, T> {
type Output = Result<T, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { receiver, token } = self.get_mut();
let received = {
let mut inner = receiver.shared.inner.lock();
match inner.receive(receiver.key) {
Some(received) => received,
None => {
if receiver.shared.senders.load(Ordering::Acquire) == 0 {
*token = None;
return Poll::Ready(Err(RecvError::Disconnected));
}
let retired_waker = inner.waiters.register(token, cx.waker());
drop(inner);
drop(retired_waker);
return Poll::Pending;
}
}
};
let (msg, reclaimed) = received;
*token = None;
Poll::Ready(Ok(take_msg(msg, reclaimed)))
}
}