use std::sync::Arc;
use zksync_concurrency::{ctx, ctx::channel, limiter, oneshot, scope, sync};
use super::{
Config, FrameKind, Header, ReadStream, RunError, Stream, StreamId, StreamKind, WriteStream,
};
use crate::noise::bytes;
#[derive(Debug)]
pub(super) struct ReadPermit {
pub(super) _count: sync::OwnedSemaphorePermit,
pub(super) _size: sync::OwnedSemaphorePermit,
}
#[derive(Debug)]
pub(super) struct Frame {
pub(super) header: Header,
pub(super) data: Option<bytes::Buffer>,
pub(super) _permit: Option<ReadPermit>,
}
#[derive(Debug)]
pub(super) enum WriteCommand {
Frame(Frame),
Flush,
}
type Reservation = oneshot::Sender<Stream>;
pub(crate) struct ReservedStream(oneshot::Sender<Reservation>);
impl ReservedStream {
pub(crate) async fn open(
self,
ctx: &ctx::Ctx,
) -> ctx::OrCanceled<Result<Stream, sync::Disconnected>> {
let (send, recv) = oneshot::channel();
if self.0.send(send).is_err() {
return Ok(Err(sync::Disconnected));
}
recv.recv_or_disconnected(ctx).await
}
}
pub(crate) struct StreamQueue {
pub(super) max_streams: u32,
limiter: limiter::Limiter,
send: channel::Sender<ReservedStream>,
recv: sync::Mutex<channel::Receiver<ReservedStream>>,
}
impl StreamQueue {
pub(crate) fn new(ctx: &ctx::Ctx, max_streams: u32, rate: limiter::Rate) -> Arc<Self> {
let (send, recv) = channel::bounded(1);
Arc::new(Self {
max_streams,
limiter: limiter::Limiter::new(ctx, rate),
send,
recv: sync::Mutex::new(recv),
})
}
pub(crate) async fn reserve(&self, ctx: &ctx::Ctx) -> ctx::OrCanceled<ReservedStream> {
let mut recv = sync::lock(ctx, &self.recv).await?.into_async();
recv.recv(ctx).await
}
#[allow(dead_code)]
pub(crate) async fn open(&self, ctx: &ctx::Ctx) -> ctx::OrCanceled<Stream> {
loop {
if let Ok(stream) = self.reserve(ctx).await?.open(ctx).await? {
return Ok(stream);
}
}
}
async fn push(&self, ctx: &ctx::Ctx) -> ctx::OrCanceled<Reservation> {
loop {
let (send, recv) = oneshot::channel();
self.send.send(ctx, ReservedStream(send)).await?;
if let Ok(reservation) = recv.recv_or_disconnected(ctx).await? {
return Ok(reservation);
}
}
}
}
#[derive(Debug)]
pub(super) struct ReadReusableStream {
pub(crate) cache: Option<Frame>,
pub(crate) recv: channel::UnboundedReceiver<Frame>,
pub(crate) close_received: bool,
}
#[derive(Debug)]
pub(super) struct WriteReusableStream {
pub(crate) cfg: Arc<Config>,
pub(crate) stream_id: StreamId,
pub(crate) stream_kind: StreamKind,
pub(crate) buffer: bytes::Buffer,
pub(super) write_send: channel::Sender<WriteCommand>,
pub(super) flush: Arc<sync::Notify>,
}
impl ReadReusableStream {
pub(super) fn new(recv: channel::UnboundedReceiver<Frame>) -> Self {
Self {
cache: None,
recv,
close_received: false,
}
}
pub(super) async fn recv_open(&mut self, ctx: &ctx::Ctx) -> ctx::OrCanceled<()> {
self.cache.take();
self.close_received = false;
while self.recv.recv(ctx).await?.header.frame_kind() != FrameKind::OPEN {}
Ok(())
}
}
impl WriteReusableStream {
pub(super) fn new(
cfg: Arc<Config>,
stream_id: StreamId,
stream_kind: StreamKind,
write_send: channel::Sender<WriteCommand>,
flush: Arc<sync::Notify>,
) -> Self {
Self {
stream_id,
stream_kind,
buffer: bytes::Buffer::new(cfg.write_frame_size as usize),
write_send,
flush,
cfg,
}
}
pub(super) async fn send_data(&mut self, ctx: &ctx::Ctx) -> Result<(), RunError> {
if self.buffer.len() == 0 {
return Ok(());
}
let slot = self
.write_send
.reserve_or_disconnected(ctx)
.await?
.map_err(|_| RunError::Closed)?;
let header = Header::new(FrameKind::DATA, self.stream_kind, self.stream_id);
let frame = Frame {
header,
data: Some(std::mem::replace(
&mut self.buffer,
bytes::Buffer::new(self.cfg.write_frame_size as usize),
)),
_permit: None,
};
slot.send(WriteCommand::Frame(frame));
Ok(())
}
pub(super) async fn send_close(&mut self, ctx: &ctx::Ctx) -> Result<(), RunError> {
self.send_data(ctx).await?;
let header = Header::new(FrameKind::CLOSE, self.stream_kind, self.stream_id);
let frame = Frame {
header,
data: None,
_permit: None,
};
self.write_send
.send(ctx, WriteCommand::Frame(frame))
.await?;
self.flush.notify_one();
Ok(())
}
pub(super) async fn send_open(&mut self, ctx: &ctx::Ctx) -> Result<(), RunError> {
let header = Header::new(FrameKind::OPEN, self.stream_kind, self.stream_id);
let frame = Frame {
header,
data: None,
_permit: None,
};
self.write_send
.send(ctx, WriteCommand::Frame(frame))
.await?;
self.flush.notify_one();
Ok(())
}
}
pub(super) struct ReusableStream {
pub(super) read: ReadReusableStream,
pub(super) write: WriteReusableStream,
pub(super) stream_queue: Arc<StreamQueue>,
}
impl ReusableStream {
pub(super) async fn run(self, ctx: &ctx::Ctx) -> Result<(), RunError> {
scope::run!(ctx, |ctx, s| async {
let (_, mut read_receiver) = sync::ExclusiveLock::new(self.read);
let (_, mut write_receiver) = sync::ExclusiveLock::new(self.write);
loop {
let recv_open_task = s.spawn(async {
let mut read = read_receiver.wait(ctx).await?;
read.recv_open(ctx).await?;
Ok(read)
});
let mut write = write_receiver.wait(ctx).await?;
write.send_close(ctx).await?;
let _open_permit = self.stream_queue.limiter.acquire(ctx, 1).await?;
let (read, reservation) = match write.stream_kind {
StreamKind::ACCEPT => {
let read = recv_open_task.join(ctx).await?;
let reservation = self.stream_queue.push(ctx).await?;
write.send_open(ctx).await?;
(read, reservation)
}
StreamKind::CONNECT => {
let reservation = self.stream_queue.push(ctx).await?;
write.send_open(ctx).await?;
let read = recv_open_task.join(ctx).await?;
(read, reservation)
}
_ => unreachable!("bad StreamKind"),
};
let (read_lock, new_read_receiver) = sync::ExclusiveLock::new(read);
read_receiver = new_read_receiver;
let (write_lock, new_write_receiver) = sync::ExclusiveLock::new(write);
write_receiver = new_write_receiver;
let _ = reservation.send(Stream {
read: ReadStream(read_lock),
write: WriteStream(write_lock),
});
}
})
.await
}
}