use crate::sink::AsyncSink;
#[cfg(feature = "trace_log")]
use crate::tokio_task_id;
use crate::{shared::*, trace_log, MTx, Tx};
use std::cell::Cell;
use std::fmt;
use std::future::Future;
use std::marker::PhantomData;
use std::mem::{needs_drop, MaybeUninit};
use std::ops::Deref;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
pub struct AsyncTx<T> {
pub(crate) shared: Arc<ChannelShared<T>>,
_phan: PhantomData<Cell<()>>,
}
impl<T> fmt::Debug for AsyncTx<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "AsyncTx")
}
}
impl<T> fmt::Display for AsyncTx<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "AsyncTx")
}
}
unsafe impl<T: Send> Send for AsyncTx<T> {}
impl<T> Drop for AsyncTx<T> {
fn drop(&mut self) {
self.shared.close_tx();
}
}
impl<T> From<Tx<T>> for AsyncTx<T> {
fn from(value: Tx<T>) -> Self {
value.add_tx();
Self::new(value.shared.clone())
}
}
impl<T> AsyncTx<T> {
#[inline]
pub(crate) fn new(shared: Arc<ChannelShared<T>>) -> Self {
Self { shared, _phan: Default::default() }
}
#[inline]
pub fn into_sink(self) -> AsyncSink<T> {
AsyncSink::new(self)
}
#[inline]
pub fn into_blocking(self) -> Tx<T> {
self.into()
}
}
impl<T: Unpin + Send + 'static> AsyncTx<T> {
#[inline(always)]
pub fn send<'a>(&'a self, item: T) -> SendFuture<'a, T> {
return SendFuture { tx: &self, item: MaybeUninit::new(item), waker: None };
}
#[inline]
pub fn try_send(&self, item: T) -> Result<(), TrySendError<T>> {
if self.shared.is_disconnected() {
return Err(TrySendError::Disconnected(item));
}
let _item = MaybeUninit::new(item);
if self.shared.inner.try_send(&_item) {
self.shared.on_send();
return Ok(());
} else {
return unsafe { Err(TrySendError::Full(_item.assume_init())) };
}
}
#[cfg(any(feature = "tokio", feature = "async_std"))]
#[cfg_attr(docsrs, doc(cfg(any(feature = "tokio", feature = "async_std"))))]
#[inline]
pub fn send_timeout<'a>(
&'a self, item: T, duration: std::time::Duration,
) -> SendTimeoutFuture<'a, T, ()> {
let sleep = {
#[cfg(feature = "tokio")]
{
tokio::time::sleep(duration)
}
#[cfg(feature = "async_std")]
{
async_std::task::sleep(duration)
}
};
self.send_with_timer(item, sleep)
}
#[inline]
pub fn send_with_timer<'a, F, R>(&'a self, item: T, fut: F) -> SendTimeoutFuture<'a, T, R>
where
F: Future<Output = R> + 'static,
{
SendTimeoutFuture {
tx: &self,
item: MaybeUninit::new(item),
waker: None,
sleep: Box::pin(fut),
}
}
#[inline(always)]
pub(crate) fn poll_send<'a>(
&self, ctx: &'a mut Context, item: &MaybeUninit<T>, o_waker: &'a mut Option<SendWaker<T>>,
sink: bool,
) -> Poll<Result<(), ()>> {
let shared = &self.shared;
if shared.is_disconnected() {
trace_log!("tx{:?}: closed {:?}", tokio_task_id!(), o_waker);
return Poll::Ready(Err(()));
}
let mut state;
loop {
if shared.inner.try_send(item) {
shared.on_send();
if let Some(_waker) = o_waker.take() {
trace_log!("tx{:?}: send {:?}", tokio_task_id!(), _waker);
} else {
trace_log!("tx{:?}: send", tokio_task_id!());
}
return Poll::Ready(Ok(()));
}
if let Some(waker) = o_waker.as_ref() {
match waker.try_change_state(WakerState::Woken, WakerState::Init) {
Ok(_) => {
if !waker.will_wake(ctx) {
let _ = o_waker.take();
}
}
Err(state) => {
if state < WakerState::Woken as u8 {
if waker.will_wake(ctx) {
trace_log!("tx{:?}: will_wake {:?}", tokio_task_id!(), waker);
return Poll::Pending;
} else {
self.senders.cancel_waker(waker);
trace_log!("tx{:?}: drop waker {:?}", tokio_task_id!(), waker);
let _ = o_waker.take();
}
} else if state == WakerState::Closed as u8 {
return Poll::Ready(Err(()));
}
}
}
} else {
if let Some(mut backoff) = shared.get_async_backoff() {
loop {
backoff.spin();
if shared.inner.try_send(item) {
shared.on_send();
trace_log!("tx{:?}: send", tokio_task_id!());
return Poll::Ready(Ok(()));
}
if backoff.is_completed() {
break;
}
}
}
}
(state, *o_waker) = if let Some(waker) = o_waker.take() {
shared.sender_reg_and_try(item, waker, sink)
} else {
let waker = SendWaker::<T>::new_async(ctx, std::ptr::null_mut());
shared.sender_reg_and_try(item, waker, sink)
};
trace_log!("tx{:?}: sender_reg_and_try {:?} {}", tokio_task_id!(), o_waker, state);
if state < WakerState::Woken as u8 {
return Poll::Pending;
} else if state > WakerState::Woken as u8 {
if state == WakerState::Done as u8 {
trace_log!("tx{:?}: send {:?} done", o_waker, tokio_task_id!());
let _ = o_waker.take();
return Poll::Ready(Ok(()));
} else {
debug_assert_eq!(state, WakerState::Closed as u8);
trace_log!("tx{:?}: closed {:?}", o_waker, tokio_task_id!());
let _ = o_waker.take();
return Poll::Ready(Err(()));
}
}
debug_assert_eq!(state, WakerState::Woken as u8);
continue;
}
}
}
#[must_use]
pub struct SendFuture<'a, T: Unpin> {
tx: &'a AsyncTx<T>,
item: MaybeUninit<T>,
waker: Option<SendWaker<T>>,
}
unsafe impl<T: Unpin + Send> Send for SendFuture<'_, T> {}
impl<T: Unpin> Drop for SendFuture<'_, T> {
fn drop(&mut self) {
if let Some(waker) = self.waker.take() {
if self.tx.shared.abandon_send_waker(waker) {
if needs_drop::<T>() {
unsafe { self.item.assume_init_drop() };
}
}
}
}
}
impl<T: Unpin + Send + 'static> Future for SendFuture<'_, T> {
type Output = Result<(), SendError<T>>;
fn poll(self: Pin<&mut Self>, ctx: &mut Context) -> Poll<Self::Output> {
let mut _self = self.get_mut();
match _self.tx.poll_send(ctx, &_self.item, &mut _self.waker, false) {
Poll::Ready(Ok(())) => {
debug_assert!(_self.waker.is_none());
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(())) => {
let _ = _self.waker.take();
return Poll::Ready(Err(SendError(unsafe { _self.item.assume_init_read() })));
}
Poll::Pending => return Poll::Pending,
}
}
}
#[must_use]
pub struct SendTimeoutFuture<'a, T: Unpin, R> {
tx: &'a AsyncTx<T>,
sleep: Pin<Box<dyn Future<Output = R>>>,
item: MaybeUninit<T>,
waker: Option<SendWaker<T>>,
}
unsafe impl<T: Unpin + Send, R> Send for SendTimeoutFuture<'_, T, R> {}
impl<T: Unpin, R> Drop for SendTimeoutFuture<'_, T, R> {
fn drop(&mut self) {
if let Some(waker) = self.waker.take() {
if self.tx.shared.abandon_send_waker(waker) {
if needs_drop::<T>() {
unsafe { self.item.assume_init_drop() };
}
}
}
}
}
impl<T: Unpin + Send + 'static, R> Future for SendTimeoutFuture<'_, T, R> {
type Output = Result<(), SendTimeoutError<T>>;
fn poll(self: Pin<&mut Self>, ctx: &mut Context) -> Poll<Self::Output> {
let mut _self = self.get_mut();
match _self.tx.poll_send(ctx, &_self.item, &mut _self.waker, false) {
Poll::Ready(Ok(())) => {
debug_assert!(_self.waker.is_none());
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(())) => {
let _ = _self.waker.take();
return Poll::Ready(Err(SendTimeoutError::Disconnected(unsafe {
_self.item.assume_init_read()
})));
}
Poll::Pending => {
if let Poll::Ready(_) = _self.sleep.as_mut().poll(ctx) {
if let Some(waker) = _self.waker.take() {
if _self.tx.shared.abandon_send_waker(waker) {
return Poll::Ready(Err(SendTimeoutError::Timeout(unsafe {
_self.item.assume_init_read()
})));
} else {
return Poll::Ready(Ok(()));
}
} else {
unreachable!();
}
}
return Poll::Pending;
}
}
}
}
pub trait AsyncTxTrait<T: Unpin + Send + 'static>:
Send + 'static + fmt::Debug + fmt::Display + AsRef<ChannelShared<T>> + Into<AsyncSink<T>>
{
fn try_send(&self, item: T) -> Result<(), TrySendError<T>>;
#[inline(always)]
fn len(&self) -> usize {
self.as_ref().len()
}
#[inline(always)]
fn capacity(&self) -> Option<usize> {
self.as_ref().capacity()
}
#[inline(always)]
fn is_empty(&self) -> bool {
self.as_ref().is_empty()
}
#[inline(always)]
fn is_full(&self) -> bool {
self.as_ref().is_full()
}
#[inline(always)]
fn is_disconnected(&self) -> bool {
self.as_ref().is_disconnected()
}
fn clone_to_vec(self, count: usize) -> Vec<Self>
where
Self: Sized;
fn send<'a>(&'a self, item: T) -> SendFuture<'a, T>;
#[cfg(any(feature = "tokio", feature = "async_std"))]
#[cfg_attr(docsrs, doc(cfg(any(feature = "tokio", feature = "async_std"))))]
fn send_timeout<'a>(
&'a self, item: T, duration: std::time::Duration,
) -> SendTimeoutFuture<'a, T, ()>;
fn send_with_timer<'a, F, R>(&'a self, item: T, fut: F) -> SendTimeoutFuture<'a, T, R>
where
F: Future<Output = R> + 'static;
}
impl<T: Unpin + Send + 'static> AsyncTxTrait<T> for AsyncTx<T> {
#[inline(always)]
fn clone_to_vec(self, count: usize) -> Vec<Self> {
assert_eq!(count, 1);
vec![self]
}
#[inline(always)]
fn try_send(&self, item: T) -> Result<(), TrySendError<T>> {
AsyncTx::try_send(self, item)
}
#[inline(always)]
fn send<'a>(&'a self, item: T) -> SendFuture<'a, T> {
AsyncTx::send(self, item)
}
#[cfg(any(feature = "tokio", feature = "async_std"))]
#[cfg_attr(docsrs, doc(cfg(any(feature = "tokio", feature = "async_std"))))]
#[inline(always)]
fn send_timeout<'a>(
&'a self, item: T, duration: std::time::Duration,
) -> SendTimeoutFuture<'a, T, ()> {
AsyncTx::send_timeout(self, item, duration)
}
#[inline(always)]
fn send_with_timer<'a, F, R>(&'a self, item: T, fut: F) -> SendTimeoutFuture<'a, T, R>
where
F: Future<Output = R> + 'static,
{
AsyncTx::send_with_timer(self, item, fut)
}
}
pub struct MAsyncTx<T>(pub(crate) AsyncTx<T>);
impl<T> fmt::Debug for MAsyncTx<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "MAsyncTx")
}
}
impl<T> fmt::Display for MAsyncTx<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "MAsyncTx")
}
}
unsafe impl<T: Send> Sync for MAsyncTx<T> {}
impl<T: Unpin> Clone for MAsyncTx<T> {
#[inline]
fn clone(&self) -> Self {
let inner = &self.0;
inner.shared.add_tx();
Self(AsyncTx::new(inner.shared.clone()))
}
}
impl<T> From<MAsyncTx<T>> for AsyncTx<T> {
fn from(tx: MAsyncTx<T>) -> Self {
tx.0
}
}
impl<T> MAsyncTx<T> {
#[inline]
pub(crate) fn new(shared: Arc<ChannelShared<T>>) -> Self {
Self(AsyncTx::new(shared))
}
#[inline]
pub fn into_sink(self) -> AsyncSink<T> {
AsyncSink::new(self.0)
}
#[inline]
pub fn into_blocking(self) -> MTx<T> {
self.into()
}
}
impl<T> Deref for MAsyncTx<T> {
type Target = AsyncTx<T>;
#[inline(always)]
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<T> From<MTx<T>> for MAsyncTx<T> {
fn from(value: MTx<T>) -> Self {
value.add_tx();
Self::new(value.shared.clone())
}
}
impl<T: Unpin + Send + 'static> AsyncTxTrait<T> for MAsyncTx<T> {
#[inline(always)]
fn clone_to_vec(self, count: usize) -> Vec<Self> {
let mut v = Vec::with_capacity(count);
for _ in 0..count - 1 {
v.push(self.clone());
}
v.push(self);
v
}
#[inline(always)]
fn try_send(&self, item: T) -> Result<(), TrySendError<T>> {
self.0.try_send(item)
}
#[inline(always)]
fn send<'a>(&'a self, item: T) -> SendFuture<'a, T> {
self.0.send(item)
}
#[cfg(any(feature = "tokio", feature = "async_std"))]
#[cfg_attr(docsrs, doc(cfg(any(feature = "tokio", feature = "async_std"))))]
#[inline(always)]
fn send_timeout<'a>(
&'a self, item: T, duration: std::time::Duration,
) -> SendTimeoutFuture<'a, T, ()> {
self.0.send_timeout(item, duration)
}
#[inline(always)]
fn send_with_timer<'a, F, R>(&'a self, item: T, fut: F) -> SendTimeoutFuture<'a, T, R>
where
F: Future<Output = R> + 'static,
{
self.0.send_with_timer(item, fut)
}
}
impl<T> Deref for AsyncTx<T> {
type Target = ChannelShared<T>;
#[inline(always)]
fn deref(&self) -> &ChannelShared<T> {
&self.shared
}
}
impl<T> AsRef<ChannelShared<T>> for AsyncTx<T> {
#[inline(always)]
fn as_ref(&self) -> &ChannelShared<T> {
&self.shared
}
}
impl<T> AsRef<ChannelShared<T>> for MAsyncTx<T> {
#[inline(always)]
fn as_ref(&self) -> &ChannelShared<T> {
&self.0.shared
}
}