use std::pin::Pin;
use std::task::{Context, Poll, ready};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
use weida_runtime::Exec;
const CHUNK_DATA: u8 = 0x00;
const CHUNK_FIN: u8 = 0x01;
const CHUNK_RESET: u8 = 0x02;
const HEADER_MAX: usize = 9;
const HEADER_DATA: usize = 5;
const DRAIN_BUF: usize = 8 * 1024;
#[derive(Debug)]
pub(crate) struct PeerReset(pub(crate) u64);
impl std::fmt::Display for PeerReset {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "stream reset by the peer with code {}", self.0)
}
}
impl std::error::Error for PeerReset {}
pub(crate) enum Marker {
Fin,
Reset(u64),
}
impl Marker {
fn encode(&self) -> ([u8; HEADER_MAX], usize) {
let mut bytes = [0u8; HEADER_MAX];
match self {
Marker::Fin => {
bytes[0] = CHUNK_FIN;
(bytes, 1)
}
Marker::Reset(code) => {
bytes[0] = CHUNK_RESET;
bytes[1..].copy_from_slice(&code.to_le_bytes());
(bytes, HEADER_MAX)
}
}
}
}
fn closed() -> std::io::Error {
std::io::Error::new(std::io::ErrorKind::BrokenPipe, "stream already closed")
}
pub(crate) struct ChunkWriter<W> {
io: Option<W>,
exec: Exec,
header: [u8; HEADER_MAX],
header_len: usize,
header_written: usize,
body_left: usize,
finished: bool,
}
impl<W: AsyncWrite + Unpin + Send + 'static> ChunkWriter<W> {
pub(crate) fn new(io: W, exec: Exec) -> ChunkWriter<W> {
ChunkWriter {
io: Some(io),
exec,
header: [0; HEADER_MAX],
header_len: 0,
header_written: 0,
body_left: 0,
finished: false,
}
}
pub(crate) fn end(mut self, marker: Marker) {
let Some(mut io) = self.io.take() else {
return;
};
if self.finished || self.body_left > 0 || self.header_written != self.header_len {
return;
}
let (bytes, len) = marker.encode();
self.exec.spawn(async move {
if let Err(e) = io.write_all(&bytes[..len]).await {
tracing::debug!(error = %e, "local stream end marker not written");
return;
}
let _ = io.flush().await;
});
}
fn poll_header(&mut self, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
let Some(io) = self.io.as_mut() else {
return Poll::Ready(Err(closed()));
};
while self.header_written < self.header_len {
let n = ready!(
Pin::new(&mut *io)
.poll_write(cx, &self.header[self.header_written..self.header_len])
)?;
if n == 0 {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::WriteZero,
"the local stream accepted no bytes",
)));
}
self.header_written += n;
}
Poll::Ready(Ok(()))
}
fn poll_flush_inner(&mut self, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.io.as_mut() {
Some(io) => Pin::new(io).poll_flush(cx),
None => Poll::Ready(Ok(())),
}
}
}
impl<W: AsyncWrite + Unpin + Send + 'static> AsyncWrite for ChunkWriter<W> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
let this = self.get_mut();
if this.io.is_none() || this.finished {
return Poll::Ready(Err(closed()));
}
if buf.is_empty() {
return Poll::Ready(Ok(0));
}
if this.body_left == 0 && this.header_written == this.header_len {
let len = buf.len().min(u32::MAX as usize);
this.header[0] = CHUNK_DATA;
this.header[1..HEADER_DATA].copy_from_slice(&(len as u32).to_le_bytes());
this.header_len = HEADER_DATA;
this.header_written = 0;
this.body_left = len;
}
ready!(this.poll_header(cx))?;
let want = buf.len().min(this.body_left);
let io = this.io.as_mut().expect("checked above");
let n = ready!(Pin::new(io).poll_write(cx, &buf[..want]))?;
this.body_left -= n;
Poll::Ready(Ok(n))
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
self.get_mut().poll_flush_inner(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
if this.io.is_none() {
return Poll::Ready(Ok(()));
}
if this.finished {
return this.poll_flush_inner(cx);
}
if this.body_left > 0 {
return Poll::Ready(Err(std::io::Error::other(
"shutdown in the middle of a write",
)));
}
if this.header_written == this.header_len {
let (bytes, len) = Marker::Fin.encode();
this.header = bytes;
this.header_len = len;
this.header_written = 0;
}
ready!(this.poll_header(cx))?;
this.finished = true;
this.poll_flush_inner(cx)
}
}
enum ReadState {
Header {
buf: [u8; HEADER_MAX],
filled: usize,
},
Body { left: usize },
Ended,
}
enum Chunk {
Data(usize),
Fin,
Reset(u64),
}
pub(crate) struct ChunkReader<R> {
io: Option<R>,
exec: Exec,
state: ReadState,
}
impl<R: AsyncRead + Unpin + Send + 'static> ChunkReader<R> {
pub(crate) fn new(io: R, exec: Exec) -> ChunkReader<R> {
ChunkReader {
io: Some(io),
exec,
state: ReadState::Header {
buf: [0; HEADER_MAX],
filled: 0,
},
}
}
pub(crate) fn drain(mut self) {
if matches!(self.state, ReadState::Ended) || self.io.is_none() {
return;
}
let exec = self.exec.clone();
exec.spawn(async move {
let mut scratch = vec![0u8; DRAIN_BUF];
loop {
match self.read(&mut scratch).await {
Ok(0) | Err(_) => return,
Ok(_) => {}
}
}
});
}
fn header_len(kind: u8) -> Option<usize> {
match kind {
CHUNK_DATA => Some(HEADER_DATA),
CHUNK_FIN => Some(1),
CHUNK_RESET => Some(HEADER_MAX),
_ => None,
}
}
fn poll_read_inner(
&mut self,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let Some(io) = self.io.as_mut() else {
return Poll::Ready(Ok(()));
};
loop {
match &mut self.state {
ReadState::Ended => return Poll::Ready(Ok(())),
ReadState::Header {
buf: header,
filled,
} => {
let need = if *filled == 0 {
1
} else {
match ChunkReader::<R>::header_len(header[0]) {
Some(need) => need,
None => {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("unknown local chunk kind {:#04x}", header[0]),
)));
}
}
};
if *filled < need {
let mut slice = ReadBuf::new(&mut header[*filled..need]);
ready!(Pin::new(&mut *io).poll_read(cx, &mut slice))?;
let n = slice.filled().len();
if n == 0 {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"the local stream closed before the end of the payload",
)));
}
*filled += n;
continue;
}
let chunk = match header[0] {
CHUNK_DATA => {
let len =
u32::from_le_bytes([header[1], header[2], header[3], header[4]]);
Chunk::Data(len as usize)
}
CHUNK_FIN => Chunk::Fin,
_ => {
let mut code = [0u8; 8];
code.copy_from_slice(&header[1..HEADER_MAX]);
Chunk::Reset(u64::from_le_bytes(code))
}
};
match chunk {
Chunk::Data(0) => {
self.state = ReadState::Header {
buf: [0; HEADER_MAX],
filled: 0,
};
}
Chunk::Data(len) => self.state = ReadState::Body { left: len },
Chunk::Fin => {
self.state = ReadState::Ended;
return Poll::Ready(Ok(()));
}
Chunk::Reset(code) => {
self.state = ReadState::Ended;
return Poll::Ready(Err(std::io::Error::other(PeerReset(code))));
}
}
}
ReadState::Body { left } => {
let want = buf.remaining().min(*left);
if want == 0 {
return Poll::Ready(Ok(()));
}
buf.initialize_unfilled_to(want);
let mut slice = buf.take(want);
ready!(Pin::new(&mut *io).poll_read(cx, &mut slice))?;
let n = slice.filled().len();
if n == 0 {
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"the local stream closed in the middle of a chunk",
)));
}
buf.advance(n);
*left -= n;
if *left == 0 {
self.state = ReadState::Header {
buf: [0; HEADER_MAX],
filled: 0,
};
}
return Poll::Ready(Ok(()));
}
}
}
}
}
impl<R: AsyncRead + Unpin + Send + 'static> AsyncRead for ChunkReader<R> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
let polled = this.poll_read_inner(cx, buf);
if let Poll::Ready(Err(_)) = &polled {
this.state = ReadState::Ended;
}
polled
}
}
pub(crate) fn read_error(error: std::io::Error) -> weida_core::Error {
match error
.get_ref()
.and_then(|inner| inner.downcast_ref::<PeerReset>())
{
Some(reset) => weida_protocol::codes::stop_reason(reset.0).into(),
None => weida_core::Error::Transport(format!("local stream read failed: {error}")),
}
}