use std::any::type_name;
use std::fmt;
use std::mem::ManuallyDrop;
use std::ptr::NonNull;
use std::sync::atomic::Ordering;
use std::sync::atomic::fence;
use crate::oneshot::Channel;
use crate::oneshot::DISCONNECTED;
use crate::oneshot::EMPTY;
use crate::oneshot::MESSAGE;
use crate::oneshot::RECEIVING;
#[cfg(doc)]
use crate::oneshot::Receiver;
use crate::oneshot::deallocate_empty_channel;
use crate::oneshot::drop_message_and_deallocate_channel;
pub struct Sender<T> {
channel_ptr: NonNull<Channel<T>>,
}
impl<T> fmt::Debug for Sender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Sender").finish_non_exhaustive()
}
}
unsafe impl<T: Send> Send for Sender<T> {}
unsafe impl<T: Sync> Sync for Sender<T> {}
impl<T> Sender<T> {
pub fn send(self, message: T) -> Result<(), SendError<T>> {
let sender = ManuallyDrop::new(self);
let channel_ptr = sender.channel_ptr;
let channel = unsafe { channel_ptr.as_ref() };
unsafe { channel.write_message(message) };
match channel.state.fetch_add(1, Ordering::Release) {
EMPTY => Ok(()),
RECEIVING => {
let (waker, receiver_owns_allocation) =
unsafe { channel.finish_sender_awakening(MESSAGE) };
if receiver_owns_allocation {
waker.wake();
} else {
unsafe { drop_message_and_deallocate_channel(channel_ptr) };
}
Ok(())
}
DISCONNECTED => {
fence(Ordering::Acquire);
Err(SendError { channel_ptr })
}
state => unreachable!("unexpected channel state: {}", state),
}
}
pub fn is_closed(&self) -> bool {
let channel = unsafe { self.channel_ptr.as_ref() };
matches!(channel.state.load(Ordering::Relaxed), DISCONNECTED)
}
pub(super) fn new(channel_ptr: NonNull<Channel<T>>) -> Self {
Self { channel_ptr }
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
let channel = unsafe { self.channel_ptr.as_ref() };
match channel.state.fetch_xor(0b001, Ordering::Release) {
EMPTY => {}
RECEIVING => {
let (waker, receiver_owns_allocation) =
unsafe { channel.finish_sender_awakening(DISCONNECTED) };
if receiver_owns_allocation {
waker.wake();
} else {
unsafe { deallocate_empty_channel(self.channel_ptr) };
}
}
DISCONNECTED => {
fence(Ordering::Acquire);
unsafe { deallocate_empty_channel(self.channel_ptr) };
}
state => unreachable!("unexpected channel state: {}", state),
}
}
}
pub struct SendError<T> {
channel_ptr: NonNull<Channel<T>>,
}
unsafe impl<T: Send> Send for SendError<T> {}
unsafe impl<T: Sync> Sync for SendError<T> {}
impl<T> SendError<T> {
pub fn as_inner(&self) -> &T {
unsafe { self.channel_ptr.as_ref().message() }
}
pub fn into_inner(self) -> T {
let error = ManuallyDrop::new(self);
let channel_ptr = error.channel_ptr;
let channel: &Channel<T> = unsafe { channel_ptr.as_ref() };
let message = unsafe { channel.take_message() };
unsafe { deallocate_empty_channel(channel_ptr) };
message
}
}
impl<T> Drop for SendError<T> {
fn drop(&mut self) {
unsafe { drop_message_and_deallocate_channel(self.channel_ptr) };
}
}
impl<T> fmt::Display for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("sending on a closed channel")
}
}
impl<T> fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SendError<{}>(..)", type_name::<T>())
}
}
impl<T> std::error::Error for SendError<T> {}
#[cfg(test)]
impl<T> Sender<T> {
pub(super) fn channel_ptr(&self) -> NonNull<Channel<T>> {
self.channel_ptr
}
}