use bytes::Buf;
use futures::{FutureExt, future::BoxFuture};
use std::{
fmt,
future::Future,
mem,
pin::Pin,
task::{Context, Poll, ready},
};
use super::{ClosedReason, RemoteSendError, Sending, base};
use crate::{
RemoteSend, chmux,
codec::{self, AnySend, ErasedDeserializer, ErasedSerializer},
exec,
rch::{BACKCHANNEL_MSG_CLOSE, BACKCHANNEL_MSG_ERROR},
};
mod distributor;
mod receiver;
mod sender;
pub use distributor::{DistributedReceiverHandle, Distributor};
pub use receiver::{Receiver, RecvError, TryRecvError};
pub use sender::{Permit, SendError, Sender, SenderSink, TrySendError};
pub fn channel<T, Codec>(local_buffer: usize) -> (Sender<T, Codec>, Receiver<T, Codec>)
where
T: RemoteSend,
{
assert!(local_buffer > 0, "local_buffer must not be zero");
let (tx, rx) = tokio::sync::mpsc::channel(local_buffer);
let (closed_tx, closed_rx) = tokio::sync::watch::channel(None);
let (remote_send_err_tx, remote_send_err_rx) = tokio::sync::watch::channel(None);
let sender = Sender::new(tx, closed_rx, remote_send_err_rx);
let receiver = Receiver::new(rx, closed_tx, false, remote_send_err_tx, None);
(sender, receiver)
}
pub fn forward<T, Codec>(mut local_rx: tokio::sync::mpsc::Receiver<T>) -> (Forwarding, Receiver<T, Codec>)
where
T: RemoteSend,
Codec: codec::Codec,
{
let (tx, rx) = channel(1);
let hnd = exec::spawn(async move {
loop {
let permit = match tx.reserve().await {
Ok(permit) => permit,
Err(err) if err.is_closed() => break,
Err(err) => return Err(err),
};
match local_rx.recv().await {
Some(v) => {
permit.send(v);
}
None => break,
}
}
Ok(())
});
(Forwarding(hnd), rx)
}
pub struct Forwarding(exec::task::JoinHandle<Result<(), SendError<()>>>);
impl fmt::Debug for Forwarding {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("Forwarding").finish()
}
}
impl Future for Forwarding {
type Output = Result<(), SendError<()>>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
match ready!(self.0.poll_unpin(cx)) {
Ok(res) => Poll::Ready(res),
Err(_) => Poll::Ready(Err(SendError::Closed(()))),
}
}
}
impl Forwarding {
pub fn stop(self) {
self.0.abort();
}
}
pub trait MpscExt<T, Codec, const BUFFER: usize, const MAX_ITEM_SIZE: usize> {
fn with_buffer<const NEW_BUFFER: usize>(
self,
) -> (Sender<T, Codec, NEW_BUFFER>, Receiver<T, Codec, NEW_BUFFER, MAX_ITEM_SIZE>);
fn with_max_item_size<const NEW_MAX_ITEM_SIZE: usize>(
self,
) -> (Sender<T, Codec, BUFFER>, Receiver<T, Codec, BUFFER, NEW_MAX_ITEM_SIZE>);
}
impl<T, Codec, const BUFFER: usize, const MAX_ITEM_SIZE: usize> MpscExt<T, Codec, BUFFER, MAX_ITEM_SIZE>
for (Sender<T, Codec, BUFFER>, Receiver<T, Codec, BUFFER, MAX_ITEM_SIZE>)
where
T: Send + 'static,
{
fn with_buffer<const NEW_BUFFER: usize>(
self,
) -> (Sender<T, Codec, NEW_BUFFER>, Receiver<T, Codec, NEW_BUFFER, MAX_ITEM_SIZE>) {
let (tx, rx) = self;
let tx = tx.set_buffer();
let rx = rx.set_buffer();
(tx, rx)
}
fn with_max_item_size<const NEW_MAX_ITEM_SIZE: usize>(
self,
) -> (Sender<T, Codec, BUFFER>, Receiver<T, Codec, BUFFER, NEW_MAX_ITEM_SIZE>) {
let (mut tx, rx) = self;
tx.set_max_item_size(NEW_MAX_ITEM_SIZE);
let rx = rx.set_max_item_size();
(tx, rx)
}
}
pub(crate) struct SendReq<T> {
pub value: Result<T, RecvError>,
pub result_tx: Option<tokio::sync::oneshot::Sender<Result<(), base::SendError<T>>>>,
}
impl<T> SendReq<T> {
fn new(value: Result<T, RecvError>) -> Self {
Self { value, result_tx: None }
}
fn ack(self) -> Result<T, RecvError> {
let Self { value, result_tx } = self;
if let Some(result_tx) = result_tx {
let _ = result_tx.send(Ok(()));
}
value
}
}
pub(crate) trait ErasedSendReq {
fn take_value(&mut self) -> AnySend;
fn result_ok(&mut self);
fn result_err(&mut self, err: base::SendError<AnySend>) -> Result<(), base::SendError<AnySend>>;
}
impl<T> ErasedSendReq for SendReq<T>
where
T: Send + 'static,
{
fn take_value(&mut self) -> AnySend {
let value = mem::replace(&mut self.value, Err(RecvError::RemoteConnect(chmux::ConnectError::Rejected)));
Box::new(value)
}
fn result_ok(&mut self) {
if let Some(result_tx) = self.result_tx.take() {
let _ = result_tx.send(Ok(()));
}
}
fn result_err(&mut self, err: base::SendError<AnySend>) -> Result<(), base::SendError<AnySend>> {
let item: Result<T, RecvError> = *err.item.downcast().expect("type mismatch in SendReq");
let Ok(item) = item else { return Ok(()) };
let err = base::SendError { kind: err.kind, item };
let err = match self.result_tx.take() {
Some(result_tx) => match result_tx.send(Err(err)) {
Ok(()) => return Ok(()),
Err(res) => res.expect_err("sent item was error"),
},
None => err,
};
Err(base::SendError { kind: err.kind, item: Box::new(err.item) as AnySend })
}
}
pub(crate) fn send_req<T>(value: Result<T, RecvError>) -> (SendReq<T>, Sending<T>) {
let (result_tx, result_rx) = tokio::sync::oneshot::channel();
let this = SendReq { value, result_tx: Some(result_tx) };
let sent = Sending(result_rx);
(this, sent)
}
trait ErasedMpscRx {
fn recv_erased(&'_ mut self) -> BoxFuture<'_, Option<Box<dyn ErasedSendReq + Send>>>;
}
impl<T> ErasedMpscRx for tokio::sync::mpsc::Receiver<SendReq<T>>
where
T: Send + 'static,
{
fn recv_erased(&'_ mut self) -> BoxFuture<'_, Option<Box<dyn ErasedSendReq + Send>>> {
async { self.recv().await.map(|send_req| Box::new(send_req) as Box<dyn ErasedSendReq + Send>) }.boxed()
}
}
async fn send_impl(
erased_serializer: ErasedSerializer, mut rx: Box<dyn ErasedMpscRx + Send>, raw_tx: chmux::Sender,
mut raw_rx: chmux::Receiver, remote_send_err_tx: tokio::sync::watch::Sender<Option<RemoteSendError>>,
closed_tx: tokio::sync::watch::Sender<Option<ClosedReason>>, max_item_size: usize,
) {
let mut remote_tx = base::ErasedSender::new(erased_serializer, raw_tx);
remote_tx.set_max_item_size(max_item_size);
loop {
tokio::select! {
biased;
backchannel_msg = raw_rx.recv() => {
match backchannel_msg {
Ok(Some(mut msg)) if msg.remaining() >= 1 => {
match msg.get_u8() {
BACKCHANNEL_MSG_CLOSE => {
let _ = remote_send_err_tx.send(Some(RemoteSendError::Closed));
let _ = closed_tx.send(Some(ClosedReason::Closed));
break;
}
BACKCHANNEL_MSG_ERROR => {
let _ = remote_send_err_tx.send(Some(RemoteSendError::Forward));
let _ = closed_tx.send(Some(ClosedReason::Failed));
break;
}
_ => (),
}
},
Ok(Some(_)) => (),
Ok(None) => {
let _ = remote_send_err_tx.send(Some(RemoteSendError::Send(
base::SendErrorKind::Send(chmux::SendError::Closed { gracefully: false })
)));
let _ = closed_tx.send(Some(ClosedReason::Dropped));
break;
}
_ => {
let _ = remote_send_err_tx.send(Some(RemoteSendError::Send(
base::SendErrorKind::Send(chmux::SendError::ChMux)
)));
let _ = closed_tx.send(Some(ClosedReason::Failed));
break;
},
}
}
send_req_opt = rx.recv_erased() => {
let Some(mut send_req) = send_req_opt else { break };
match remote_tx.send_erased(send_req.take_value()).await {
Ok(()) => send_req.result_ok(),
Err(err) => {
let _ = remote_send_err_tx.send(Some(RemoteSendError::Send(err.kind.clone())));
let _ = closed_tx.send(Some(ClosedReason::Failed));
if let Err(err) = send_req.result_err(err) && err.is_item_specific() {
tracing::warn!(%err, "sending over remote channel failed");
}
}
}
}
}
}
}
trait ErasedMpscTx {
fn send(&'_ self, value: AnySend) -> BoxFuture<'_, Result<(), ()>>;
fn send_err(&'_ self, err: RecvError) -> BoxFuture<'_, Result<(), ()>>;
}
impl<T> ErasedMpscTx for tokio::sync::mpsc::Sender<SendReq<T>>
where
T: Send + 'static,
{
fn send(&'_ self, value: AnySend) -> BoxFuture<'_, Result<(), ()>> {
let value: Result<T, RecvError> = *value.downcast().expect("type mismatch in mpsc receiver");
async { self.send(SendReq::new(value)).await.map_err(|_| ()) }.boxed()
}
fn send_err(&'_ self, err: RecvError) -> BoxFuture<'_, Result<(), ()>> {
async { self.send(SendReq::new(Err(err))).await.map_err(|_| ()) }.boxed()
}
}
async fn recv_impl(
erased_deserializer: ErasedDeserializer, tx: &(dyn ErasedMpscTx + Send + Sync), mut raw_tx: chmux::Sender,
raw_rx: chmux::Receiver, mut remote_send_err_rx: tokio::sync::watch::Receiver<Option<RemoteSendError>>,
mut closed_rx: tokio::sync::watch::Receiver<Option<ClosedReason>>, max_item_size: usize,
) {
let mut remote_rx = base::ErasedReceiver::new(erased_deserializer, raw_rx);
remote_rx.set_max_item_size(max_item_size);
loop {
tokio::select! {
biased;
res = closed_rx.changed() => {
match res {
Ok(()) => {
let reason = closed_rx.borrow().clone();
match reason {
Some(ClosedReason::Closed) => {
let _ = raw_tx.send(vec![BACKCHANNEL_MSG_CLOSE].into()).await;
}
Some(ClosedReason::Dropped) => break,
Some(ClosedReason::Failed) => {
let _ = raw_tx.send(vec![BACKCHANNEL_MSG_ERROR].into()).await;
}
None => (),
}
},
Err(_) => break,
}
}
Ok(()) = remote_send_err_rx.changed() => {
if remote_send_err_rx.borrow().as_ref().is_some() {
let _ = raw_tx.send(vec![BACKCHANNEL_MSG_ERROR].into()).await;
}
}
res = remote_rx.recv_erased() => {
match res {
Ok(Some(value)) => {
if tx.send(value).await.is_err() {
break
}
}
Ok(None) => break,
Err(err) => {
let is_final_err = err.is_final();
if tx.send_err(RecvError::RemoteReceive(err)).await.is_err() || is_final_err {
break
}
}
}
}
}
}
}