use std::io;
use std::pin::Pin;
#[cfg(feature = "tls")]
use std::sync::Arc;
use std::task::{Context, Poll};
use futures_core::Stream;
use minarrow::structs::shared_buffer::SharedBuffer;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
#[cfg(feature = "tls")]
use tokio::net::TcpStream;
use tokio::net::ToSocketAddrs;
use crate::enums::BufferChunkSize;
use crate::models::streams::stream_arena::StreamArena;
use crate::models::transports::tcp::TcpTransport;
pub enum TcpReadHalf {
Plain(OwnedReadHalf),
#[cfg(feature = "tls")]
Tls(Box<dyn AsyncRead + Send + Unpin + 'static>),
}
impl AsyncRead for TcpReadHalf {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
match self.get_mut() {
TcpReadHalf::Plain(s) => Pin::new(s).poll_read(cx, buf),
#[cfg(feature = "tls")]
TcpReadHalf::Tls(s) => Pin::new(&mut **s).poll_read(cx, buf),
}
}
}
pub enum TcpWriteHalf {
Plain(OwnedWriteHalf),
#[cfg(feature = "tls")]
Tls(Box<dyn AsyncWrite + Send + Sync + Unpin + 'static>),
}
impl AsyncWrite for TcpWriteHalf {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match self.get_mut() {
TcpWriteHalf::Plain(s) => Pin::new(s).poll_write(cx, buf),
#[cfg(feature = "tls")]
TcpWriteHalf::Tls(s) => Pin::new(&mut **s).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
TcpWriteHalf::Plain(s) => Pin::new(s).poll_flush(cx),
#[cfg(feature = "tls")]
TcpWriteHalf::Tls(s) => Pin::new(&mut **s).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
TcpWriteHalf::Plain(s) => Pin::new(s).poll_shutdown(cx),
#[cfg(feature = "tls")]
TcpWriteHalf::Tls(s) => Pin::new(&mut **s).poll_shutdown(cx),
}
}
}
pub struct TcpByteStream {
reader: TcpReadHalf,
eof: bool,
chunk_size: usize,
arena: StreamArena,
}
impl TcpByteStream {
pub async fn connect(addr: impl ToSocketAddrs) -> io::Result<Self> {
let (read_half, _write_half) = TcpTransport::connect(addr).await?;
Ok(Self::from_read_half(read_half, BufferChunkSize::Http))
}
pub fn from_read_half(read_half: OwnedReadHalf, size: BufferChunkSize) -> Self {
Self {
reader: TcpReadHalf::Plain(read_half),
eof: false,
chunk_size: size.chunk_size(),
arena: StreamArena::new(),
}
}
#[cfg(feature = "tls")]
pub async fn connect_tls(
addr: impl ToSocketAddrs,
server_name: rustls_pki_types::ServerName<'static>,
config: Arc<tokio_rustls::rustls::ClientConfig>,
) -> io::Result<Self> {
let tcp = TcpStream::connect(addr).await?;
let connector = tokio_rustls::TlsConnector::from(config);
let tls = connector.connect(server_name, tcp).await?;
let (read_half, _write_half) = tokio::io::split(tls);
Ok(Self::from_tls_read_half(read_half, BufferChunkSize::Http))
}
#[cfg(feature = "tls")]
pub fn from_tls_read_half<R>(read_half: R, size: BufferChunkSize) -> Self
where
R: AsyncRead + Send + Unpin + 'static,
{
Self {
reader: TcpReadHalf::Tls(Box::new(read_half)),
eof: false,
chunk_size: size.chunk_size(),
arena: StreamArena::new(),
}
}
}
impl Stream for TcpByteStream {
type Item = Result<SharedBuffer, io::Error>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let me = self.get_mut();
if me.eof {
return Poll::Ready(None);
}
if me.arena.remaining() < me.chunk_size {
me.arena.recycle_or_reset();
}
let chunk_start = me.arena.write_pos();
let n = {
let spare = me.arena.spare_uninit();
let read_len = spare.len().min(me.chunk_size);
let mut read_buf = ReadBuf::uninit(&mut spare[..read_len]);
match Pin::new(&mut me.reader).poll_read(cx, &mut read_buf) {
Poll::Ready(Ok(())) => read_buf.filled().len(),
Poll::Ready(Err(e)) => {
me.eof = true;
return Poll::Ready(Some(Err(e)));
}
Poll::Pending => return Poll::Pending,
}
};
if n == 0 {
me.eof = true;
return Poll::Ready(None);
}
unsafe { me.arena.advance(n) };
let shared = me.arena.window(chunk_start, n);
me.arena.align();
Poll::Ready(Some(Ok(shared)))
}
}
impl AsyncRead for TcpByteStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let me = self.get_mut();
Pin::new(&mut me.reader).poll_read(cx, buf)
}
}