use std::{
io,
pin::Pin,
task::{Context, Poll},
};
use bytes::{Buf, Bytes};
use http_body::Body;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use super::{request::RequestBodyWriter, response::H3ResponseBody};
use crate::h3::common::H3Error;
pub struct H3DuplexStream {
writer: RequestBodyWriter,
body: H3ResponseBody,
read_leftover: Bytes,
}
impl H3DuplexStream {
pub fn new(writer: RequestBodyWriter, body: H3ResponseBody) -> Self {
Self {
writer,
body,
read_leftover: Bytes::new(),
}
}
}
impl AsyncRead for H3DuplexStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
loop {
if !this.read_leftover.is_empty() {
let n = this.read_leftover.len().min(buf.remaining());
buf.put_slice(&this.read_leftover[..n]);
this.read_leftover.advance(n);
return Poll::Ready(Ok(()));
}
match Pin::new(&mut this.body).poll_frame(cx) {
Poll::Ready(Some(Ok(frame))) => {
match frame.into_data() {
Ok(data) => {
this.read_leftover = data;
}
Err(_non_data) => return Poll::Ready(Ok(())),
}
}
Poll::Ready(Some(Err(err))) => return Poll::Ready(Err(h3_to_io(err))),
Poll::Ready(None) => return Poll::Ready(Ok(())),
Poll::Pending => return Poll::Pending,
}
}
}
}
impl AsyncWrite for H3DuplexStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.get_mut().writer).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().writer).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().writer).poll_shutdown(cx)
}
}
fn h3_to_io(err: H3Error) -> io::Error {
match err {
H3Error::Reset(code) => {
io::Error::new(
io::ErrorKind::ConnectionReset,
format!("stream reset by peer (code {code:#x})"),
)
}
H3Error::ConnectionClosed => {
io::Error::new(io::ErrorKind::NotConnected, "connection closed")
}
H3Error::H3(err) => io::Error::other(format!("h3 error: {err}")),
}
}