use bytes::Bytes;
use futures::{
Future, FutureExt, future, pin_mut,
sink::{Sink, SinkExt},
stream::{Stream, StreamExt},
};
use std::{
collections::{BTreeSet, HashMap, HashSet, VecDeque},
convert::TryFrom,
error::Error,
fmt,
marker::PhantomData,
mem::{self, size_of},
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::Duration,
};
use tokio::{
sync::{mpsc, mpsc::Permit, oneshot},
try_join,
};
use tokio_util::sync::ReusableBoxFuture;
use super::{
AnyStorage, Cfg, ChMuxError, PROTOCOL_VERSION, PROTOCOL_VERSION_PORT_ID, PortReq,
client::{Client, ConnectRequest, ConnectResponse},
credit::{
CreditPool, CreditProvider, CreditUser, GlobalCreditMonitor, MixedCreditUser, PortCreditMonitor,
UsedGlobalCredit, credit_send_pair, port_credit_monitor,
},
listener::{Listener, RemoteConnectMsg, Request},
msg::{DataCredits, ExchangedCfg, GlobalCredits, MultiplexMsg},
port_allocator::{PortAllocator, PortNumber},
receiver::{PortReceiveMsg, ReceivedData, ReceivedPortRequests, Receiver},
sender::Sender,
sizer::{BufferSizer, DummySizer, GlobalCreditsReport},
};
use crate::exec::time::{sleep, timeout};
fn protocol_err<SinkError, StreamError>(msg: impl AsRef<str>) -> super::ChMuxError<SinkError, StreamError> {
super::ChMuxError::Protocol(msg.as_ref().to_string())
}
#[derive(Debug)]
enum PortState {
Connecting {
response_tx: oneshot::Sender<ConnectResponse>,
},
Connected {
remote_port: u32,
sender_credit_provider: CreditProvider,
receiver_tx_data: Option<mpsc::UnboundedSender<PortReceiveMsg>>,
receiver_credit_monitor: PortCreditMonitor,
receiver_closed: bool,
receiver_dropped: bool,
sender_dropped: bool,
remote_receiver_closed: Arc<AtomicBool>,
remote_receiver_closed_notify: Arc<std::sync::Mutex<Option<Vec<oneshot::Sender<()>>>>>,
remote_receiver_dropped: bool,
report_processed: VecDeque<oneshot::Sender<()>>,
},
}
#[derive(custom_debug::Debug)]
pub(crate) enum PortEvt {
Accepted {
local_port: PortNumber,
remote_port: u32,
port_tx: oneshot::Sender<(Sender, Receiver)>,
},
Rejected {
remote_port: u32,
no_ports: bool,
},
SendData {
remote_port: u32,
#[debug(with = crate::util::dbg_bytes)]
data: Bytes,
first: bool,
last: bool,
credits: DataCredits,
},
SendPorts {
remote_port: u32,
first: bool,
last: bool,
wait: bool,
ports: Vec<(PortReq, oneshot::Sender<ConnectResponse>)>,
},
Flush {
flushed_tx: oneshot::Sender<()>,
},
RequestReceivedReport {
local_port: u32,
processed_tx: oneshot::Sender<()>,
},
ReceivedReport {
remote_port: u32,
},
ChangeGlobalCreditUsage {
remote_port: u32,
allow: bool,
},
ReturnCredits {
remote_port: u32,
credits: u32,
},
SenderDropped {
local_port: u32,
},
ReceiverClosed {
local_port: u32,
},
ReceiverDropped {
local_port: u32,
},
}
#[derive(Debug)]
enum GlobalEvt {
ConnectReq(ConnectRequest),
AllClientsDropped,
ListenerDropped,
Port(PortEvt),
ReturnGlobalCredits(GlobalCredits),
GlobalCreditsReport(GlobalCreditsReport),
SendGoodbye,
}
#[derive(custom_debug::Debug)]
struct TransportMsg {
msg: MultiplexMsg,
#[debug(with = crate::util::dbg_option_bytes)]
data: Option<Bytes>,
}
impl TransportMsg {
fn new(msg: MultiplexMsg) -> Self {
assert!(!matches!(&msg, &MultiplexMsg::Data { .. }), "MultiplexMsg::Data with missing data");
Self { msg, data: None }
}
fn with_data(msg: MultiplexMsg, data: Bytes) -> Self {
assert!(matches!(&msg, &MultiplexMsg::Data { .. }), "MultiplexMsg with unexpected data");
Self { msg, data: Some(data) }
}
}
#[derive(Debug)]
enum SendReq {
Feed(TransportMsg),
Flush(Option<oneshot::Sender<()>>),
}
#[must_use = "You must call run() on the ChMux object for the connection to work."]
pub struct ChMux<TransportSink, TransportStream> {
local_cfg: Cfg,
remote_cfg: ExchangedCfg,
remote_protocol_version: u8,
connect_rx: Option<mpsc::UnboundedReceiver<ConnectRequest>>,
listen_tx: Option<(mpsc::Sender<RemoteConnectMsg>, mpsc::Sender<RemoteConnectMsg>)>,
port_allocator: PortAllocator,
ports: HashMap<PortNumber, PortState>,
outstanding_remote_port_requests: HashSet<u32>,
channel_tx: mpsc::Sender<PortEvt>,
channel_rx: Option<mpsc::Receiver<PortEvt>>,
high_priority_channel_tx: mpsc::Sender<PortEvt>,
high_priority_channel_rx: Option<mpsc::Receiver<PortEvt>>,
terminate_rx: Option<mpsc::UnboundedReceiver<()>>,
all_clients_dropped: bool,
remote_client_dropped: bool,
remote_listener_dropped: Arc<AtomicBool>,
goodbye_sent: bool,
goodbye_received: bool,
transport_sink: Option<TransportSink>,
transport_stream: Option<TransportStream>,
storage: AnyStorage,
send_credit_provider: Option<CreditProvider>,
send_credit_user: Arc<Option<CreditUser>>,
send_credit_report: Option<GlobalCreditsReport>,
receive_buffer_sizer: Box<dyn BufferSizer>,
receive_credit_monitor: Arc<GlobalCreditMonitor>,
remote_credits_report: GlobalCreditsReport,
outstanding_inhibit_global_credit_usage_ports: BTreeSet<u32>,
}
impl<TransportSink, TransportStream> fmt::Debug for ChMux<TransportSink, TransportStream> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("ChMux")
.field("local_cfg", &self.local_cfg)
.field("remote_cfg", &self.remote_cfg)
.field("local_protocol_version", &PROTOCOL_VERSION)
.field("remote_protocol_version", &self.remote_protocol_version)
.finish()
}
}
impl<TransportSink, TransportSinkError, TransportStream, TransportStreamError>
ChMux<TransportSink, TransportStream>
where
TransportSink: Sink<Bytes, Error = TransportSinkError> + Send + Unpin,
TransportSinkError: Error + Send + Sync + 'static,
TransportStream: Stream<Item = Result<Bytes, TransportStreamError>> + Send + Unpin,
TransportStreamError: Error + Send + Sync + 'static,
{
#[tracing::instrument(level = "trace", skip_all, fields(cfg))]
pub async fn new(
mut cfg: Cfg, mut transport_sink: TransportSink, mut transport_stream: TransportStream,
) -> Result<(Self, Client, Listener), ChMuxError<TransportSinkError, TransportStreamError>> {
cfg.check();
let mut receive_buffer_sizer = mem::replace(&mut cfg.shared_receive_buffer, DummySizer::new());
let receive_credit_monitor = GlobalCreditMonitor::new(receive_buffer_sizer.initial());
let initial_global_credits = receive_credit_monitor.total();
let remote_credits_report = GlobalCreditsReport::initial(initial_global_credits);
let exchanged_cfg = ExchangedCfg::new(&cfg, initial_global_credits);
let fut = Self::exchange_hello(exchanged_cfg, &mut transport_sink, &mut transport_stream);
let (remote_protocol_version, remote_cfg) = match cfg.connection_timeout {
Some(dur) => timeout(dur, fut).await.map_err(|_| ChMuxError::Timeout)??,
None => fut.await?,
};
let (send_credit_provider, send_credit_user) =
remote_cfg.global_credits.map(|initial| credit_send_pair(CreditPool::Global, initial)).unzip();
let (channel_tx, channel_rx) = mpsc::channel(cfg.shared_send_queue);
let (high_priority_channel_tx, high_priority_channel_rx) = mpsc::channel(cfg.shared_send_queue);
let (listen_wait_tx, listen_wait_rx) = mpsc::channel(usize::from(cfg.connect_queue) + 1);
let (listen_no_wait_tx, listen_no_wait_rx) = mpsc::channel(usize::from(cfg.connect_queue) + 1);
let (connect_tx, connect_rx) = mpsc::unbounded_channel();
let (terminate_tx, terminate_rx) = mpsc::unbounded_channel();
let port_allocator = PortAllocator::new(cfg.max_ports);
let remote_listener_dropped = Arc::new(AtomicBool::new(false));
let multiplexer = ChMux {
remote_protocol_version,
local_cfg: cfg,
remote_cfg: remote_cfg.clone(),
connect_rx: Some(connect_rx),
listen_tx: Some((listen_wait_tx, listen_no_wait_tx)),
port_allocator: port_allocator.clone(),
ports: HashMap::new(),
outstanding_remote_port_requests: HashSet::new(),
channel_tx,
channel_rx: Some(channel_rx),
high_priority_channel_tx,
high_priority_channel_rx: Some(high_priority_channel_rx),
terminate_rx: Some(terminate_rx),
remote_client_dropped: false,
remote_listener_dropped: remote_listener_dropped.clone(),
all_clients_dropped: false,
goodbye_sent: false,
goodbye_received: false,
transport_sink: Some(transport_sink),
transport_stream: Some(transport_stream),
storage: AnyStorage::new(),
send_credit_provider,
send_credit_user: Arc::new(send_credit_user),
send_credit_report: None,
receive_buffer_sizer,
receive_credit_monitor,
remote_credits_report,
outstanding_inhibit_global_credit_usage_ports: BTreeSet::new(),
};
let client = Client::new(
connect_tx,
remote_cfg.connect_queue,
port_allocator.clone(),
remote_listener_dropped,
terminate_tx.clone(),
);
let listener = Listener::new(listen_wait_rx, listen_no_wait_rx, port_allocator, terminate_tx);
Ok((multiplexer, client, listener))
}
#[tracing::instrument(level = "trace", skip_all, fields(?msg))]
async fn feed_msg(
msg: TransportMsg, sink: &mut TransportSink,
) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
sink.feed(msg.msg.to_vec().into()).await.map_err(ChMuxError::SinkError)?;
if let Some(data) = msg.data {
sink.feed(data).await.map_err(ChMuxError::SinkError)?;
}
Ok(())
}
#[tracing::instrument(level = "trace", skip_all)]
async fn flush(sink: &mut TransportSink) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
sink.flush().await.map_err(ChMuxError::SinkError)
}
#[tracing::instrument(level = "trace", skip_all, fields(msg, data))]
async fn recv_msg(
stream: &mut TransportStream,
) -> Result<TransportMsg, ChMuxError<TransportSinkError, TransportStreamError>> {
let msg_data = match stream.next().await {
Some(Ok(msg_data)) => msg_data,
Some(Err(err)) => return Err(ChMuxError::StreamError(err)),
None => return Err(ChMuxError::StreamClosed),
};
let msg = MultiplexMsg::from_slice(&msg_data)?;
let data = if let MultiplexMsg::Data { .. } = &msg {
match stream.next().await {
Some(Ok(data)) => Some(data),
Some(Err(err)) => return Err(ChMuxError::StreamError(err)),
None => return Err(ChMuxError::StreamClosed),
}
} else {
None
};
tracing::Span::current().record("msg", tracing::field::debug(&msg));
if let Some(data) = &data {
tracing::Span::current().record("data", tracing::field::debug(&data));
}
Ok(TransportMsg { msg, data })
}
#[tracing::instrument(level = "trace", skip_all)]
async fn exchange_hello(
exchanged_cfg: ExchangedCfg, sink: &mut TransportSink, stream: &mut TransportStream,
) -> Result<(u8, ExchangedCfg), ChMuxError<TransportSinkError, TransportStreamError>> {
let send_task = async {
Self::feed_msg(TransportMsg::new(MultiplexMsg::Reset), sink).await?;
Self::flush(sink).await?;
Self::feed_msg(
TransportMsg::new(MultiplexMsg::Hello { version: PROTOCOL_VERSION, cfg: exchanged_cfg }),
sink,
)
.await?;
Self::flush(sink).await?;
Ok(())
};
let recv_task = async {
loop {
match Self::recv_msg(stream).await {
Ok(TransportMsg { msg: MultiplexMsg::Hello { version, cfg }, .. }) => {
break Ok((version, cfg));
}
Ok(_) => (),
Err(ChMuxError::Protocol(_)) => (),
Err(err) => return Err(err),
}
}
};
Ok(try_join!(send_task, recv_task)?.1)
}
fn should_terminate(&self) -> bool {
let mut terminate = true;
terminate &= self.ports.is_empty();
terminate &= self.all_clients_dropped || self.remote_listener_dropped.load(Ordering::Relaxed);
terminate &= self.listen_tx.is_none() || self.remote_client_dropped;
terminate &= self.outstanding_remote_port_requests.is_empty();
terminate |= self.goodbye_sent;
terminate |= self.goodbye_received;
if terminate {
tracing::trace!("should terminate");
}
terminate
}
#[tracing::instrument(level = "trace", skip(self))]
fn create_port(&mut self, local_port: PortNumber, remote_port: u32) -> (Sender, Receiver) {
let local_port_num = *local_port;
let sender_tx = self.channel_tx.clone();
let (sender_credit_provider, sender_credit_user) =
credit_send_pair(CreditPool::Port, self.remote_cfg.port_receive_buffer);
let receiver_tx = self.channel_tx.clone();
let high_priority_receiver_tx = self.high_priority_channel_tx.clone();
let (receiver_tx_data, receiver_rx_data) = mpsc::unbounded_channel();
let (receiver_credit_monitor, receiver_credit_returner) =
port_credit_monitor(self.local_cfg.port_receive_buffer, self.local_cfg.port_receive_throttle);
let hangup_notify = Arc::new(std::sync::Mutex::new(Some(Vec::new())));
let hangup_recved = Arc::new(AtomicBool::new(false));
if let Some(PortState::Connected { remote_port, .. }) = self.ports.insert(
local_port,
PortState::Connected {
remote_port,
sender_credit_provider,
receiver_tx_data: Some(receiver_tx_data),
receiver_credit_monitor,
remote_receiver_closed_notify: hangup_notify.clone(),
remote_receiver_closed: hangup_recved.clone(),
receiver_closed: false,
receiver_dropped: false,
sender_dropped: false,
remote_receiver_dropped: false,
report_processed: VecDeque::new(),
},
) {
panic!(
"create_port called for local port {local_port_num} already connected to remote port {remote_port}"
);
}
let sender = Sender::new(
local_port_num,
remote_port,
self.remote_cfg.chunk_size as usize,
self.local_cfg.max_data_size,
sender_tx,
MixedCreditUser::new(sender_credit_user, self.send_credit_user.clone()),
Arc::downgrade(&hangup_recved),
Arc::downgrade(&hangup_notify),
self.port_allocator.clone(),
self.storage.clone(),
self.remote_cfg.received_report,
);
let receiver = Receiver::new(
local_port_num,
remote_port,
self.local_cfg.max_data_size,
self.local_cfg.max_received_ports,
receiver_tx,
high_priority_receiver_tx,
receiver_rx_data,
receiver_credit_returner,
self.port_allocator.clone(),
self.storage.clone(),
);
(sender, receiver)
}
fn maybe_free_port(&mut self, local_port: u32) {
let mut free = true;
if let Some(PortState::Connected {
receiver_tx_data,
receiver_dropped,
sender_dropped,
remote_receiver_dropped,
..
}) = self.ports.get(&local_port)
{
free &= *sender_dropped;
free &= *receiver_dropped;
free &= receiver_tx_data.is_none();
free &= *remote_receiver_dropped;
} else {
panic!("maybe_free_port called for port {} not in connected state.", local_port);
}
if free {
tracing::trace!(local_port, "freed port");
self.ports.remove(&local_port);
}
}
async fn send_task(
mut sink: &mut TransportSink, ping_interval: Option<Duration>, flush_interval: Option<Duration>,
mut rx: mpsc::Receiver<SendReq>, mut high_priority_rx: mpsc::Receiver<SendReq>,
) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
async fn sleep_opt(duration: Option<Duration>) {
match duration {
Some(duration) => sleep(duration).await,
None => future::pending().await,
}
}
let mut next_ping = ReusableBoxFuture::new(sleep_opt(ping_interval));
let mut next_flush = ReusableBoxFuture::new(sleep_opt(None));
let mut need_flush = false;
loop {
SinkReady::new(&mut sink).await.map_err(ChMuxError::SinkError)?;
let req_opt = tokio::select! {
biased;
() = &mut next_flush => {
next_flush.set(sleep_opt(None));
Some(SendReq::Flush(None))
}
msg_opt = high_priority_rx.recv() => msg_opt,
msg_opt = rx.recv() => msg_opt,
() = &mut next_ping => {
Self::feed_msg(TransportMsg::new(MultiplexMsg::Ping), sink).await?;
next_ping.set(sleep_opt(ping_interval));
need_flush = true;
continue;
}
() = future::ready(()), if need_flush => Some(SendReq::Flush(None)),
};
match req_opt {
Some(SendReq::Feed(goodbye_msg @ TransportMsg { msg: MultiplexMsg::Goodbye, .. })) => {
Self::feed_msg(goodbye_msg, sink).await?;
break;
}
Some(SendReq::Feed(msg)) => {
Self::feed_msg(msg, sink).await?;
next_ping.set(sleep_opt(ping_interval));
if !need_flush {
need_flush = true;
next_flush.set(sleep_opt(flush_interval));
}
}
Some(SendReq::Flush(opt_flushed_tx)) => {
if opt_flushed_tx.as_ref().is_none_or(|flushed_tx| !flushed_tx.is_closed()) && need_flush {
Self::flush(sink).await?;
need_flush = false;
next_flush.set(sleep_opt(None));
}
if let Some(flushed_tx) = opt_flushed_tx {
let _ = flushed_tx.send(());
}
}
None => break,
}
}
let _ = Self::flush(sink).await;
Ok(())
}
async fn recv_task(
stream: &mut TransportStream, connection_timeout: Option<Duration>, tx: mpsc::Sender<TransportMsg>,
) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
async fn get_connection_timeout(connection_timeout: Option<Duration>) {
match connection_timeout {
Some(timeout) => sleep(timeout).await,
None => future::pending().await,
}
}
let mut next_timeout = ReusableBoxFuture::new(get_connection_timeout(connection_timeout));
while let Ok(tx_permit) = tx.reserve().await {
tokio::select! {
biased;
msg = Self::recv_msg(stream) => {
let msg = msg?;
let is_goodbye = matches!(&msg, TransportMsg {msg: MultiplexMsg::Goodbye, ..});
tx_permit.send(msg);
if is_goodbye {
break;
}
next_timeout.set(get_connection_timeout(connection_timeout));
},
() = &mut next_timeout => return Err(ChMuxError::Timeout),
}
}
Ok(())
}
#[tracing::instrument(name = "remoc::chmux", level = "debug", skip_all, ret)]
pub async fn run(mut self) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
let mut transport_sink = self.transport_sink.take().unwrap();
let mut transport_stream = self.transport_stream.take().unwrap();
let (send_tx, send_rx) = mpsc::channel(self.local_cfg.transport_send_queue);
let (high_priority_send_tx, high_priority_send_rx) = mpsc::channel(8);
let send_task = Self::send_task(
&mut transport_sink,
self.remote_cfg.connection_timeout.map(|d| d / 2),
self.local_cfg.flush_interval,
send_rx,
high_priority_send_rx,
)
.fuse();
pin_mut!(send_task);
let (recv_tx, mut recv_rx) = mpsc::channel(self.local_cfg.transport_receive_queue);
let recv_task = Self::recv_task(&mut transport_stream, self.local_cfg.connection_timeout, recv_tx).fuse();
pin_mut!(recv_task);
let mut channel_rx = self.channel_rx.take().unwrap();
let mut high_priority_channel_rx = self.high_priority_channel_rx.take().unwrap();
let mut connect_rx = self.connect_rx.take().unwrap();
let mut terminate_rx = self.terminate_rx.take().unwrap();
let mut send_task_ended = false;
while !(self.goodbye_sent && self.goodbye_received && send_task_ended) {
let should_terminate = self.should_terminate();
let next_high_priority_event = async {
let Ok(permit) = high_priority_send_tx.reserve().await else { return None };
loop {
while let Some(port) = self.outstanding_inhibit_global_credit_usage_ports.pop_first() {
if let Some(PortState::Connected { remote_port, receiver_credit_monitor, .. }) =
self.ports.get(&port)
&& receiver_credit_monitor.inhibiting_global_credit_usage(true)
{
return Some((
permit,
GlobalEvt::Port(PortEvt::ChangeGlobalCreditUsage {
remote_port: *remote_port,
allow: false,
}),
));
}
}
if let Some(report) = self.send_credit_report.take() {
return Some((permit, GlobalEvt::GlobalCreditsReport(report)));
}
if let Some(credits) = self
.receive_credit_monitor
.return_to_remote(&mut *self.receive_buffer_sizer, &self.remote_credits_report)
{
return Some((permit, GlobalEvt::ReturnGlobalCredits(credits)));
}
tokio::select! {
biased;
Some(msg) = high_priority_channel_rx.recv() => {
return Some((permit, GlobalEvt::Port(msg)));
}
() = self.receive_credit_monitor.wait_for_returnable() => continue,
}
}
};
let next_event = async {
let Ok(permit) = send_tx.reserve().await else { return None };
let server_dropped = async {
match &self.listen_tx {
Some((listen_wait_tx, _)) => listen_wait_tx.closed().await,
None => future::pending().await,
}
};
tokio::select! {
biased;
() = server_dropped => {
Some((permit, GlobalEvt::ListenerDropped))
}
connect_req_opt = connect_rx.recv(), if !self.all_clients_dropped => {
match connect_req_opt {
Some(connect_req) => Some((permit, GlobalEvt::ConnectReq(connect_req))),
None => Some((permit, GlobalEvt::AllClientsDropped)),
}
}
Some(msg) = channel_rx.recv() => {
Some((permit, GlobalEvt::Port(msg)))
}
Some(()) = terminate_rx.recv(), if !self.goodbye_sent => {
Some((permit, GlobalEvt::SendGoodbye))
}
() = future::ready(()), if should_terminate && !self.goodbye_sent => {
Some((permit, GlobalEvt::SendGoodbye))
}
}
};
tokio::select! {
biased;
res = &mut send_task => {
match res {
Ok(()) => send_task_ended = true,
Err(err) => return Err(err),
}
}
Err(err) = &mut recv_task => return Err(err),
Some((permit, event)) = next_high_priority_event => self.handle_event(permit, event).await?,
Some(msg) = recv_rx.recv() => self.handle_received_msg(msg).await?,
Some((permit, event)) = next_event => self.handle_event(permit, event).await?,
}
}
Ok(())
}
#[tracing::instrument(level = "trace", skip_all, fields(event=?event))]
async fn handle_event(
&mut self, permit: Permit<'_, SendReq>, event: GlobalEvt,
) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
let send_msg = |permit: Permit<'_, SendReq>, msg: MultiplexMsg| {
tracing::trace!(op="send", msg=?msg);
permit.send(SendReq::Feed(TransportMsg::new(msg)))
};
match event {
GlobalEvt::ConnectReq(ConnectRequest { local_port, id, sent_tx: _sent_tx, response_tx, wait }) => {
if !self.remote_listener_dropped.load(Ordering::Relaxed) {
let local_port_num = *local_port;
if self.ports.insert(local_port, PortState::Connecting { response_tx }).is_some() {
panic!("ConnectRequest for already used local port {local_port_num}");
}
let id = (self.remote_protocol_version >= PROTOCOL_VERSION_PORT_ID).then_some(id);
send_msg(permit, MultiplexMsg::OpenPort { client_port: local_port_num, wait, id });
} else {
let _ = response_tx.send(ConnectResponse::Rejected { no_ports: false });
}
}
GlobalEvt::Port(PortEvt::Accepted { local_port, remote_port, port_tx }) => {
if !self.outstanding_remote_port_requests.remove(&remote_port) {
panic!("Accepted non-outstanding remote port {remote_port} request");
}
let local_port_num = *local_port;
send_msg(
permit,
MultiplexMsg::PortOpened { client_port: remote_port, server_port: local_port_num },
);
let (sender, receiver) = self.create_port(local_port, remote_port);
let _ = port_tx.send((sender, receiver));
}
GlobalEvt::Port(PortEvt::Rejected { remote_port, no_ports }) => {
if !self.outstanding_remote_port_requests.remove(&remote_port) {
panic!("Rejected non-outstanding remote port {remote_port} request");
}
send_msg(permit, MultiplexMsg::Rejected { client_port: remote_port, no_ports });
}
GlobalEvt::Port(PortEvt::SendData { remote_port, data, first, last, credits }) => {
let msg = MultiplexMsg::Data { port: remote_port, first, last, credits };
tracing::trace!(op = "send", msg =? msg);
permit.send(SendReq::Feed(TransportMsg::with_data(msg, data)));
}
GlobalEvt::Port(PortEvt::SendPorts { remote_port, ports, first, last, wait }) => {
let mut port_nums = Vec::new();
let mut ids = (self.remote_protocol_version >= PROTOCOL_VERSION_PORT_ID).then_some(Vec::new());
for (PortReq { port, id }, response_tx) in ports {
let port_num = *port;
if self.ports.insert(port, PortState::Connecting { response_tx }).is_some() {
panic!("SendPorts with already used local port {port_num}");
}
port_nums.push(port_num);
if let Some(ids) = &mut ids {
ids.push(id);
}
}
send_msg(
permit,
MultiplexMsg::PortData { port: remote_port, first, last, wait, ports: port_nums, ids },
);
}
GlobalEvt::Port(PortEvt::Flush { flushed_tx }) => {
tracing::trace!(op = "flush");
permit.send(SendReq::Flush(Some(flushed_tx)));
}
GlobalEvt::Port(PortEvt::RequestReceivedReport { local_port, processed_tx }) => {
let Some(PortState::Connected { remote_port, report_processed, .. }) =
self.ports.get_mut(&local_port)
else {
panic!("RequestReportProcessed for {local_port} in invalid state")
};
report_processed.push_back(processed_tx);
send_msg(permit, MultiplexMsg::RequestReceivedReport { port: *remote_port });
}
GlobalEvt::Port(PortEvt::ReceivedReport { remote_port }) => {
send_msg(permit, MultiplexMsg::ReceivedReport { port: remote_port });
}
GlobalEvt::Port(PortEvt::ReturnCredits { remote_port, credits }) => {
send_msg(permit, MultiplexMsg::PortCredits { port: remote_port, credits });
}
GlobalEvt::Port(PortEvt::ChangeGlobalCreditUsage { remote_port, allow }) => {
let msg = if allow {
MultiplexMsg::AllowGlobalCreditUsageByPort { port: remote_port }
} else {
MultiplexMsg::InhibitGlobalCreditUsageByPort { port: remote_port }
};
send_msg(permit, msg);
}
GlobalEvt::Port(PortEvt::SenderDropped { local_port }) => {
let Some(PortState::Connected { remote_port, sender_dropped, .. }) =
self.ports.get_mut(&local_port)
else {
panic!("PortEvt SenderDropped for port {local_port} in invalid state");
};
if *sender_dropped {
panic!("PortEvt SenderDropped more than once for port {}", local_port);
}
*sender_dropped = true;
send_msg(permit, MultiplexMsg::SendFinish { port: *remote_port });
self.maybe_free_port(local_port);
}
GlobalEvt::Port(PortEvt::ReceiverClosed { local_port }) => {
let Some(PortState::Connected { remote_port, receiver_closed, receiver_dropped, .. }) =
self.ports.get_mut(&local_port)
else {
panic!("PortEvt ReceiverClosed for non-connected port {local_port}");
};
if *receiver_closed || *receiver_dropped {
panic!("PortEvt ReceiverClosed or ReceiverDropped more than once for port {local_port}");
}
*receiver_closed = true;
send_msg(permit, MultiplexMsg::ReceiveClose { port: *remote_port });
}
GlobalEvt::Port(PortEvt::ReceiverDropped { local_port }) => {
let Some(PortState::Connected { remote_port, receiver_dropped, .. }) =
self.ports.get_mut(&local_port)
else {
panic!("PortEvt ReceiverDropped for port {local_port} in invalid state.");
};
if *receiver_dropped {
panic!("PortEvt ReceiverDropped more than once for port {local_port}");
}
*receiver_dropped = true;
send_msg(permit, MultiplexMsg::ReceiveFinish { port: *remote_port });
self.maybe_free_port(local_port);
}
GlobalEvt::AllClientsDropped => {
self.all_clients_dropped = true;
send_msg(permit, MultiplexMsg::ClientFinish);
}
GlobalEvt::ListenerDropped => {
self.listen_tx = None;
send_msg(permit, MultiplexMsg::ListenerFinish);
}
GlobalEvt::ReturnGlobalCredits(credits) => {
send_msg(permit, MultiplexMsg::GlobalCredits(credits));
}
GlobalEvt::GlobalCreditsReport(report) => {
send_msg(permit, MultiplexMsg::GlobalCreditsReport(report));
}
GlobalEvt::SendGoodbye => {
self.goodbye_sent = true;
send_msg(permit, MultiplexMsg::Goodbye);
}
}
Ok(())
}
#[tracing::instrument(level = "trace", skip_all, fields(msg=?received_msg))]
async fn handle_received_msg(
&mut self, received_msg: TransportMsg,
) -> Result<(), ChMuxError<TransportSinkError, TransportStreamError>> {
let TransportMsg { msg, data } = received_msg;
match msg {
MultiplexMsg::Reset => {
return Err(ChMuxError::Reset);
}
MultiplexMsg::Hello { .. } => {
return Err(protocol_err(
"received Hello message for already established multiplexer connection",
));
}
MultiplexMsg::Ping => (),
MultiplexMsg::OpenPort { client_port, wait, id } => {
if !self.outstanding_remote_port_requests.insert(client_port) {
return Err(protocol_err(format!(
"remote endpoint sent OpenPort request for same remote port {client_port} twice"
)));
}
let req = RemoteConnectMsg::Request(Request::new(
client_port,
id.unwrap_or(client_port),
wait,
self.port_allocator.clone(),
self.high_priority_channel_tx.clone(),
));
if let Some((listen_wait_tx, listen_no_wait_tx)) = &self.listen_tx {
let res = if wait { listen_wait_tx.try_send(req) } else { listen_no_wait_tx.try_send(req) };
if let Err(mpsc::error::TrySendError::Full(_)) = res {
return Err(protocol_err("remote endpoint sent too many OpenPort requests"));
}
}
}
MultiplexMsg::PortOpened { client_port, server_port } => {
let Some((local_port, PortState::Connecting { response_tx })) =
self.ports.remove_entry(&client_port)
else {
return Err(protocol_err(format!(
"received PortOpened message for port {client_port} not in connecting state"
)));
};
let (sender, receiver) = self.create_port(local_port, server_port);
let _ = response_tx.send(ConnectResponse::Accepted(sender, receiver));
}
MultiplexMsg::Rejected { client_port, no_ports } => {
let Some(PortState::Connecting { response_tx }) = self.ports.remove(&client_port) else {
return Err(protocol_err(format!(
"received Rejected message for port {client_port} not in connecting state"
)));
};
let _ = response_tx.send(ConnectResponse::Rejected { no_ports });
}
MultiplexMsg::Data { port, first, last, credits } => {
let Some(PortState::Connected {
receiver_tx_data: Some(receiver_tx_data),
receiver_credit_monitor,
..
}) = self.ports.get_mut(&port)
else {
return Err(protocol_err(format!(
"received data for non-connected or finished local port {port}"
)));
};
let data = data.unwrap();
let Ok(size) = u32::try_from(data.len()) else {
return Err(protocol_err(format!("received data exceeds maximum size on port {port}")));
};
if size > self.local_cfg.chunk_size {
return Err(protocol_err(format!(
"received data exceeds maximum chunk size {} on port {port}",
self.local_cfg.chunk_size
)));
};
let total = size.max(1);
let port_credit;
let global_credit;
match credits {
DataCredits::PortOnly => {
global_credit = UsedGlobalCredit::default();
port_credit = receiver_credit_monitor.use_credits(total, 0)?;
}
DataCredits::GlobalOnly => {
global_credit = self.receive_credit_monitor.use_credits(total)?;
port_credit = receiver_credit_monitor.use_credits(0, total)?;
}
DataCredits::GlobalAndPort(global) => {
global_credit = self.receive_credit_monitor.use_credits(global)?;
port_credit =
receiver_credit_monitor.use_credits(total.saturating_sub(global), global)?;
}
}
self.remote_credits_report.consume(global_credit.credits());
if receiver_credit_monitor.inhibiting_global_credit_usage(false) {
self.outstanding_inhibit_global_credit_usage_ports.insert(port);
}
let _ = receiver_tx_data.send(PortReceiveMsg::Data(ReceivedData {
buf: data,
first,
last,
port_credit,
global_credit,
}));
}
MultiplexMsg::PortData { port, first, last, wait, ports, ids } => {
let Some(PortState::Connected {
receiver_tx_data: Some(receiver_tx_data),
receiver_credit_monitor,
..
}) = self.ports.get_mut(&port)
else {
return Err(protocol_err(format!(
"received port data for non-connected or finished local port {port}",
)));
};
for port in &ports {
if !self.outstanding_remote_port_requests.insert(*port) {
return Err(protocol_err(format!(
"remote endpoint sent PortData request for same remote port {port} twice"
)));
}
}
let used_credit =
match ports.len().checked_mul(size_of::<u32>()).and_then(|v| u32::try_from(v).ok()) {
Some(size) if size <= self.local_cfg.chunk_size => {
receiver_credit_monitor.use_credits(size, 0)?
}
_ => {
return Err(protocol_err(format!(
"received ports exceeds maximum chunk size on port {port}",
)));
}
};
let port_allocator = self.port_allocator.clone();
let channel_tx = self.channel_tx.clone();
let ids = ids.unwrap_or_else(|| ports.clone());
let requests = ports
.into_iter()
.zip(ids)
.map(|(remote_port, id)| {
Request::new(remote_port, id, wait, port_allocator.clone(), channel_tx.clone())
})
.collect();
let _ = receiver_tx_data.send(PortReceiveMsg::PortRequests(ReceivedPortRequests {
requests,
first,
last,
credit: used_credit,
}));
}
MultiplexMsg::RequestReceivedReport { port } => {
let Some(PortState::Connected { receiver_tx_data: Some(receiver_tx_data), .. }) =
self.ports.get(&port)
else {
return Err(protocol_err(format!(
"received processed report request for non-connected or finished local port {port}",
)));
};
let _ = receiver_tx_data.send(PortReceiveMsg::RequestReceivedReport);
}
MultiplexMsg::ReceivedReport { port } => {
if let Some(PortState::Connected { report_processed, .. }) = self.ports.get_mut(&port)
&& let Some(processed_tx) = report_processed.pop_front()
{
let _ = processed_tx.send(());
}
}
MultiplexMsg::PortCredits { port, credits } => {
if let Some(PortState::Connected { sender_credit_provider, .. }) = self.ports.get_mut(&port) {
sender_credit_provider.provide(credits)?;
}
}
MultiplexMsg::InhibitGlobalCreditUsageByPort { port } => {
if let Some(PortState::Connected { sender_credit_provider, .. }) = self.ports.get_mut(&port) {
tracing::trace!(%port, "global credit usage inhibited");
sender_credit_provider.set_use_global_credits(false);
}
}
MultiplexMsg::AllowGlobalCreditUsageByPort { port } => {
if let Some(PortState::Connected { sender_credit_provider, .. }) = self.ports.get_mut(&port) {
tracing::trace!(%port, "global credit usage allowed");
sender_credit_provider.set_use_global_credits(true);
}
}
MultiplexMsg::SendFinish { port } => {
let Some(PortState::Connected { receiver_tx_data, .. }) = self.ports.get_mut(&port) else {
return Err(protocol_err(format!(
"received SendFinish message for local port {port} not in connected state",
)));
};
let Some(receiver_tx_data) = receiver_tx_data.take() else {
return Err(protocol_err(format!(
"received SendFinish message for local port {port} more than once",
)));
};
let _ = receiver_tx_data.send(PortReceiveMsg::Finished);
self.maybe_free_port(port);
}
MultiplexMsg::ReceiveClose { port } => {
let Some(PortState::Connected {
sender_credit_provider,
remote_receiver_closed_notify,
remote_receiver_closed,
..
}) = self.ports.get_mut(&port)
else {
return Err(protocol_err(format!(
"received ReceiveClose message for port {port} not in connected state",
)));
};
if remote_receiver_closed.load(Ordering::Relaxed) {
return Err(protocol_err(format!(
"received more than one ReceiveClose message for port {port}",
)));
}
sender_credit_provider.close(true);
remote_receiver_closed.store(true, Ordering::Relaxed);
let notifies = remote_receiver_closed_notify.lock().unwrap().take().unwrap();
for tx in notifies {
let _ = tx.send(());
}
self.maybe_free_port(port);
}
MultiplexMsg::ReceiveFinish { port } => {
let Some(PortState::Connected {
sender_credit_provider,
remote_receiver_closed_notify,
remote_receiver_closed,
remote_receiver_dropped,
report_processed,
..
}) = self.ports.get_mut(&port)
else {
return Err(protocol_err(format!(
"received ReceiveFinish message for port {port} not in connected state",
)));
};
if !remote_receiver_closed.load(Ordering::Relaxed) {
sender_credit_provider.close(false);
remote_receiver_closed.store(true, Ordering::Relaxed);
let notifies = remote_receiver_closed_notify.lock().unwrap().take().unwrap();
for tx in notifies {
let _ = tx.send(());
}
}
for processed_tx in report_processed.drain(..) {
let _ = processed_tx.send(());
}
*remote_receiver_dropped = true;
self.maybe_free_port(port);
}
MultiplexMsg::GlobalCredits(GlobalCredits { credits, seq }) => {
let Some(send_credit_provider) = &mut self.send_credit_provider else {
return Err(protocol_err(
"received GlobalCredits message from peer without global credits support",
));
};
send_credit_provider.provide(credits)?;
let mut new_report = send_credit_provider.take_status(seq);
if let Some(prev_report) = &self.send_credit_report {
new_report.min = new_report.min.min(prev_report.min);
}
self.send_credit_report = Some(new_report);
}
MultiplexMsg::GlobalCreditsReport(report) => {
self.remote_credits_report = report;
}
MultiplexMsg::ClientFinish => {
if let Some((listen_wait_tx, listen_no_wait_tx)) = &self.listen_tx {
let mut failed = false;
if let Err(mpsc::error::TrySendError::Full(_)) =
listen_wait_tx.try_send(RemoteConnectMsg::ClientDropped)
{
failed = true;
}
if let Err(mpsc::error::TrySendError::Full(_)) =
listen_no_wait_tx.try_send(RemoteConnectMsg::ClientDropped)
{
failed = true;
}
if failed {
return Err(protocol_err(
"remote endpoint sent too many OpenPort or ClientFinish requests",
));
}
}
self.remote_client_dropped = true;
}
MultiplexMsg::ListenerFinish => {
self.remote_listener_dropped.store(true, Ordering::Relaxed);
}
MultiplexMsg::Goodbye => {
self.goodbye_received = true;
}
}
Ok(())
}
}
impl<TransportSink, TransportStream> Drop for ChMux<TransportSink, TransportStream> {
fn drop(&mut self) {
}
}
struct SinkReady<S, Item> {
sink: S,
_item: PhantomData<Item>,
}
impl<S, Item> SinkReady<S, Item> {
fn new(sink: S) -> Self {
Self { sink, _item: PhantomData }
}
}
impl<S, Item> Future for SinkReady<S, Item>
where
S: Sink<Item> + Unpin,
Item: Unpin,
{
type Output = Result<(), S::Error>;
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
Pin::into_inner(self).sink.poll_ready_unpin(cx)
}
}