use std::{collections::BTreeMap, sync::Arc};
use anyhow::Context as _;
use zksync_concurrency::{ctx, ctx::channel, error::Wrap as _, io, scope, sync};
use crate::{frame, noise::bytes};
mod config;
mod handshake;
mod header;
mod reusable_stream;
#[cfg(test)]
mod tests;
mod transient_stream;
pub(crate) use config::*;
use handshake::Handshake;
use header::{FrameKind, Header, StreamId, StreamKind};
pub(crate) use reusable_stream::*;
pub(crate) use transient_stream::*;
pub(crate) struct Mux {
pub(crate) cfg: Arc<Config>,
pub(crate) accept: BTreeMap<CapabilityId, Arc<StreamQueue>>,
pub(crate) connect: BTreeMap<CapabilityId, Arc<StreamQueue>>,
}
fn saturating_sum(iter: impl Iterator<Item = u32>) -> u32 {
iter.fold(0, |x, v| x.saturating_add(v))
}
impl Mux {
fn handshake(&self) -> Handshake {
Handshake {
accept_max_streams: self
.accept
.iter()
.map(|(id, q)| (*id, q.max_streams))
.collect(),
connect_max_streams: self
.connect
.iter()
.map(|(id, q)| (*id, q.max_streams))
.collect(),
}
}
pub(crate) fn verify(&self) -> anyhow::Result<()> {
self.cfg.verify().context("cfg")?;
if saturating_sum(self.accept.values().map(|cfg| cfg.max_streams)) > MAX_STREAM_COUNT {
anyhow::bail!("sum of accept_inflight > {MAX_STREAM_COUNT}");
}
if saturating_sum(self.connect.values().map(|cfg| cfg.max_streams)) > MAX_STREAM_COUNT {
anyhow::bail!("sum of connect_inflight > {MAX_STREAM_COUNT}");
}
Ok(())
}
}
#[derive(Debug, thiserror::Error)]
#[allow(clippy::missing_docs_in_private_items)]
pub(crate) enum RunError {
#[error("config: {0:#}")]
Config(anyhow::Error),
#[error(transparent)]
Canceled(#[from] ctx::Canceled),
#[error("connection closed")]
Closed,
#[error("protocol: {0:#}")]
Protocol(anyhow::Error),
#[error(transparent)]
IO(#[from] io::Error),
}
impl From<RunError> for ctx::Error {
fn from(err: RunError) -> Self {
match err {
RunError::Canceled(err) => Self::Canceled(err),
err => Self::Internal(err.into()),
}
}
}
impl Mux {
async fn spawn_streams<'env>(
&self,
ctx: &'env ctx::Ctx,
scope: &scope::Scope<'env, RunError>,
stream_kind: StreamKind,
handshake: &Handshake,
write_send: &channel::Sender<WriteCommand>,
flush: &Arc<sync::Notify>,
) -> Vec<channel::UnboundedSender<Frame>> {
let mut streams = vec![];
let (queues, peer) = match stream_kind {
StreamKind::ACCEPT => (&self.accept, &handshake.connect_max_streams),
StreamKind::CONNECT => (&self.connect, &handshake.accept_max_streams),
_ => unreachable!("bad StreamKind"),
};
for (cap, queue) in queues {
let max_streams = std::cmp::min(queue.max_streams, *peer.get(cap).unwrap_or(&0));
for _ in 0..max_streams {
let (read_send, read_recv) = channel::unbounded();
let stream_id = StreamId::new(streams.len() as u16);
streams.push(read_send);
let stream = ReusableStream {
read: ReadReusableStream::new(read_recv),
write: WriteReusableStream::new(
self.cfg.clone(),
stream_id,
stream_kind,
write_send.clone(),
flush.clone(),
),
stream_queue: queue.clone(),
};
scope.spawn_bg(stream.run(ctx));
}
}
streams
}
async fn process_inbound_frames(
&self,
ctx: &ctx::Ctx,
mut read: impl io::AsyncRead + Send + Unpin,
accept_streams: Vec<channel::UnboundedSender<Frame>>,
connect_streams: Vec<channel::UnboundedSender<Frame>>,
) -> Result<(), RunError> {
let count_sem = Arc::new(sync::Semaphore::new(self.cfg.read_frame_count as usize));
let size_sem = Arc::new(sync::Semaphore::new(self.cfg.read_buffer_size as usize));
loop {
let mut header = [0u8, 2];
io::read_exact(ctx, &mut read, &mut header).await??;
let header = Header::from(header);
let streams = match header.stream_kind() {
StreamKind::ACCEPT => &connect_streams,
StreamKind::CONNECT => &accept_streams,
_ => unreachable!("bad StreamKind"),
};
let stream = streams
.get(header.stream_id().0 as usize)
.with_context(|| format!("bad stream id {:?}", header.stream_id()))
.map_err(RunError::Protocol)?;
match header.frame_kind() {
FrameKind::OPEN | FrameKind::CLOSE => {
let permit = Some(ReadPermit {
_count: sync::acquire_many_owned(ctx, count_sem.clone(), 1).await?,
_size: size_sem.clone().try_acquire_many_owned(0).unwrap(),
});
stream.send(Frame {
header,
data: None,
_permit: permit,
});
}
FrameKind::DATA => {
let mut length = [0u8, 2];
io::read_exact(ctx, &mut read, &mut length).await??;
let mut length = u16::from_le_bytes(length) as usize;
while length > 0 {
let size = std::cmp::min(length, self.cfg.read_frame_size as usize);
let permit = Some(ReadPermit {
_count: sync::acquire_many_owned(ctx, count_sem.clone(), 1).await?,
_size: sync::acquire_many_owned(ctx, size_sem.clone(), size as u32)
.await?,
});
let mut data = bytes::Buffer::new(size);
io::read_exact(ctx, &mut read, data.as_mut_capacity()).await??;
data.extend(size);
stream.send(Frame {
header,
data: Some(data),
_permit: permit,
});
length -= size;
}
}
_ => unreachable!("bad FrameKind"),
}
}
}
pub(crate) async fn run<S: io::AsyncRead + io::AsyncWrite + Send>(
self,
ctx: &ctx::Ctx,
transport: S,
) -> Result<(), RunError> {
self.verify().map_err(RunError::Config)?;
let (mut read, mut write) = io::split(transport);
let res = scope::run!(ctx, |ctx, s| async {
s.spawn(async {
let h = self.handshake();
frame::send_proto(ctx, &mut write, &h)
.await
.wrap("send_proto()")
});
frame::recv_proto(ctx, &mut read, handshake::MAX_FRAME)
.await
.wrap("recv_proto()")
})
.await;
let handshake = res.map_err(|err| match err {
ctx::Error::Canceled(err) => RunError::Canceled(err),
ctx::Error::Internal(err) => RunError::Protocol(err),
})?;
let (write_send, write_recv) = channel::bounded(1);
let flush = Arc::new(sync::Notify::new());
let res = scope::run!(ctx, |ctx, s| async {
let accept_streams = self
.spawn_streams(ctx, s, StreamKind::ACCEPT, &handshake, &write_send, &flush)
.await;
let connect_streams = self
.spawn_streams(ctx, s, StreamKind::CONNECT, &handshake, &write_send, &flush)
.await;
s.spawn_bg::<()>(async {
let mut write = write;
let mut write_recv = write_recv;
loop {
match write_recv.recv(ctx).await? {
WriteCommand::Flush => io::flush(ctx, &mut write).await??,
WriteCommand::Frame(frame) => {
io::write_all(ctx, &mut write, &frame.header.raw()).await??;
if let Some(data) = frame.data {
let length = (data.len() as u16).to_le_bytes();
io::write_all(ctx, &mut write, &length).await??;
io::write_all(ctx, &mut write, data.as_slice()).await??;
}
}
}
}
});
s.spawn_bg::<()>(async {
loop {
sync::notified(ctx, &flush).await?;
write_send.send(ctx, WriteCommand::Flush).await?;
}
});
self.process_inbound_frames(ctx, read, accept_streams, connect_streams)
.await
})
.await;
match res {
Ok(()) => unreachable!(),
Err(RunError::IO(err)) => match err.kind() {
io::ErrorKind::UnexpectedEof
| io::ErrorKind::ConnectionReset
| io::ErrorKind::BrokenPipe => Err(RunError::Closed),
_ => Err(RunError::IO(err)),
},
err => err,
}
}
}