Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -39,3 +39,4 @@ All notable changes to this project will be documented in this file.

* Remove the `slab` dependency in favor of a focused internal waiter arena.
* Describe disconnected channel states consistently in channel error messages.
* Replace the legacy standard-library MPSC backend with an owned queue core and remove the receiver types' manual `Sync` implementations.
58 changes: 24 additions & 34 deletions asyncband/src/mpsc/bounded.rs
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,8 @@ use super::RecvError;
use super::SendError;
use super::TryRecvError;
use super::TrySendError;
use super::queue::BoundedQueue;
use super::queue::PushError;
use crate::internal::atomic_waker::AtomicWaker;
use crate::internal::semaphore::Acquire;
use crate::internal::semaphore::Semaphore;
Expand All @@ -50,23 +52,20 @@ use crate::internal::semaphore::Semaphore;
pub fn bounded<T>(buffer: usize) -> (BoundedSender<T>, BoundedReceiver<T>) {
assert!(buffer > 0, "mpsc bounded channel requires buffer > 0");
let state = Arc::new(BoundedState {
queue: BoundedQueue::new(buffer),
senders: AtomicUsize::new(1),
tx_permits: Semaphore::new(0),
rx_waker: AtomicWaker::new(),
});
let (sender, receiver) = std::sync::mpsc::sync_channel(buffer);
let sender = BoundedSender {
state: state.clone(),
sender: Some(sender),
};
let receiver = BoundedReceiver {
state: state.clone(),
receiver: Some(receiver),
};
let receiver = BoundedReceiver { state };
(sender, receiver)
}

struct BoundedState {
struct BoundedState<T> {
queue: BoundedQueue<T>,
senders: AtomicUsize,
tx_permits: Semaphore,
rx_waker: AtomicWaker,
Expand All @@ -76,16 +75,14 @@ struct BoundedState {
///
/// Instances are created by the [`bounded`] function.
pub struct BoundedSender<T> {
state: Arc<BoundedState>,
sender: Option<std::sync::mpsc::SyncSender<T>>,
state: Arc<BoundedState<T>>,
}

impl<T> Clone for BoundedSender<T> {
fn clone(&self) -> Self {
self.state.senders.fetch_add(1, Ordering::Release);
BoundedSender {
state: self.state.clone(),
sender: self.sender.clone(),
}
}
}
Expand All @@ -98,9 +95,6 @@ impl<T> fmt::Debug for BoundedSender<T> {

impl<T> Drop for BoundedSender<T> {
fn drop(&mut self) {
// Dropping the final underlying sender disconnects the channel.
drop(self.sender.take());

match self.state.senders.fetch_sub(1, Ordering::AcqRel) {
1 => {
// Wake the receiver so it can observe the channel's disconnected state.
Expand Down Expand Up @@ -197,18 +191,14 @@ impl<T> BoundedSender<T> {
/// # }
/// ```
pub fn try_send(&self, value: T) -> Result<(), TrySendError<T>> {
// INVARIANT: A shared borrow of the endpoint cannot overlap its destructor.
let sender = self.sender.as_ref().unwrap();
match sender.try_send(value) {
match self.state.queue.try_push(value) {
Ok(()) => {
self.state.rx_waker.wake();

Ok(())
}
Err(std::sync::mpsc::TrySendError::Full(value)) => Err(TrySendError::Full(value)),
Err(std::sync::mpsc::TrySendError::Disconnected(value)) => {
Err(TrySendError::Disconnected(value))
}
Err(PushError::Full(value)) => Err(TrySendError::Full(value)),
Err(PushError::Disconnected(value)) => Err(TrySendError::Disconnected(value)),
}
}
}
Expand All @@ -217,14 +207,9 @@ impl<T> BoundedSender<T> {
///
/// Instances are created by the [`bounded`] function.
pub struct BoundedReceiver<T> {
state: Arc<BoundedState>,
receiver: Option<std::sync::mpsc::Receiver<T>>,
state: Arc<BoundedState<T>>,
}

/// The only `!Sync` field `receiver` is protected by `&mut self` in `recv` and `try_recv`.
/// That is, `BoundedReceiver` can only be accessed by one thread at a time.
unsafe impl<T: Send> Sync for BoundedReceiver<T> {}

impl<T> fmt::Debug for BoundedReceiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BoundedReceiver").finish_non_exhaustive()
Expand All @@ -233,7 +218,7 @@ impl<T> fmt::Debug for BoundedReceiver<T> {

impl<T> Drop for BoundedReceiver<T> {
fn drop(&mut self) {
drop(self.receiver.take());
self.state.queue.disconnect_receiver();
self.state.tx_permits.notify_all();
}
}
Expand Down Expand Up @@ -274,15 +259,20 @@ impl<T> BoundedReceiver<T> {
/// # }
/// ```
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
// INVARIANT: A mutable borrow of the endpoint cannot overlap its destructor.
let receiver = self.receiver.as_ref().unwrap();
match receiver.try_recv() {
Ok(v) => {
if let Some(value) = self.state.queue.pop() {
self.state.tx_permits.release_if_nonempty(1);
Ok(value)
} else if self.state.senders.load(Ordering::Acquire) == 0 {
// The final sender can enqueue between the first empty observation and decrementing
// the sender count, so check the queue again before reporting disconnection.
if let Some(value) = self.state.queue.pop() {
self.state.tx_permits.release_if_nonempty(1);
Ok(v)
Ok(value)
} else {
Err(TryRecvError::Disconnected)
}
Err(std::sync::mpsc::TryRecvError::Disconnected) => Err(TryRecvError::Disconnected),
Err(std::sync::mpsc::TryRecvError::Empty) => Err(TryRecvError::Empty),
} else {
Err(TryRecvError::Empty)
}
}

Expand Down
1 change: 1 addition & 0 deletions asyncband/src/mpsc/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@

mod bounded;
mod error;
mod queue;
mod unbounded;

pub use self::bounded::BoundedReceiver;
Expand Down
Loading