use crate::{
collections::Vector,
misc::{Lease, LeaseMut},
stream::StreamReader,
};
use core::{fmt::Debug, hint::cold_path, mem::MaybeUninit};
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum BufStreamReaderError {
AbruptDisconnect,
CapacityOverflow,
ForbiddenClear,
}
pub struct BufStreamReader {
antecedent_end_idx: usize,
buffer: Vector<u8>,
capacity_ub: usize,
current_end_idx: usize,
forbid_clear: bool,
}
impl BufStreamReader {
#[inline]
pub const fn new() -> Self {
Self {
antecedent_end_idx: 0,
buffer: Vector::new(),
capacity_ub: 1024 * 1024 * 32,
current_end_idx: 0,
forbid_clear: false,
}
}
#[inline]
pub fn antecedent(&self) -> &[u8] {
let range = 0..self.antecedent_end_idx;
unsafe { self.buffer.get(range).unwrap_unchecked() }
}
#[inline]
pub const fn antecedent_end_idx(&self) -> usize {
self.antecedent_end_idx
}
#[inline]
pub const fn capacity_ub(&self) -> usize {
self.capacity_ub
}
#[inline]
pub fn clear_if_exhausted(&mut self) {
if self.current_end_idx == self.buffer.len() {
self.clear();
}
}
#[inline]
pub fn current(&self) -> &[u8] {
let range = self.antecedent_end_idx..self.current_end_idx;
unsafe { self.buffer.get(range).unwrap_unchecked() }
}
#[inline]
pub fn current_mut(&mut self) -> &mut [u8] {
let range = self.antecedent_end_idx..self.current_end_idx;
unsafe { self.buffer.get_mut(range).unwrap_unchecked() }
}
#[inline]
pub const fn current_end_idx(&self) -> usize {
self.current_end_idx
}
#[inline]
pub fn filled(&self) -> &[u8] {
&self.buffer
}
#[inline]
pub fn following(&self) -> &[u8] {
unsafe { self.buffer.get(self.current_end_idx..).unwrap_unchecked() }
}
#[inline]
pub const fn forbid_clear(&self) -> bool {
self.forbid_clear
}
#[inline]
pub const fn forbid_clear_mut(&mut self) -> &mut bool {
&mut self.forbid_clear
}
#[inline]
pub async fn read_header<SR, const LEN: usize>(
&mut self,
stream_reader: &mut SR,
) -> crate::Result<Option<[u8; LEN]>>
where
SR: StreamReader,
{
self.manage_capacity(LEN)?;
let Self { antecedent_end_idx, buffer, current_end_idx, .. } = self;
let read_fut = async move {
let local_current_end_idx = *current_end_idx;
loop {
let (init, uninit) = buffer.split_at_spare_mut();
let following = unsafe { init.get(local_current_end_idx..).unwrap_unchecked() };
if let Some(slice) = following.get(..LEN) {
let rslt = slice.try_into().unwrap_or([0; LEN]);
Self::remove_current(antecedent_end_idx, current_end_idx, LEN);
return Ok(Some(rslt));
}
let Some(len) = stream_reader.read(uninit.into()).await? else {
cold_path();
return Ok(None);
};
let new_len = init.len().wrapping_add(len.get());
unsafe {
buffer.set_len(new_len);
}
}
};
read_fut.await
}
#[inline]
pub async fn read_payload<SR>(
&mut self,
payload_len: usize,
stream_reader: &mut SR,
) -> crate::Result<()>
where
SR: StreamReader,
{
self.manage_capacity(payload_len)?;
let current_end_idx = self.current_end_idx;
loop {
let (init, uninit) = self.split_at_spare_mut();
let following_len = init.len().wrapping_sub(current_end_idx);
if following_len >= payload_len {
self.current_end_idx = current_end_idx.wrapping_add(payload_len);
return Ok(());
}
let Some(len) = stream_reader.read(uninit.into()).await? else {
cold_path();
return Err(BufStreamReaderError::AbruptDisconnect.into());
};
let new_len = init.len().wrapping_add(len.get());
unsafe {
self.buffer.set_len(new_len);
}
}
}
#[inline]
pub fn set_capacity_ub(&mut self, value: usize) {
self.capacity_ub = self.capacity_ub.max(value);
}
#[inline]
pub fn split_at_spare_mut(&mut self) -> (&mut [u8], &mut [MaybeUninit<u8>]) {
self.buffer.split_at_spare_mut()
}
#[cfg(any(feature = "tls", feature = "postgres"))]
#[inline]
pub(crate) fn buffer_mut(&mut self) -> &mut Vector<u8> {
&mut self.buffer
}
#[inline]
pub(crate) fn clear(&mut self) {
let Self { antecedent_end_idx, buffer, capacity_ub: _, current_end_idx, forbid_clear } = self;
if *forbid_clear {
return;
}
*antecedent_end_idx = 0;
buffer.clear();
*current_end_idx = 0;
*forbid_clear = false;
}
#[cfg(feature = "web-socket")]
pub(crate) async fn read_arbitrary<SR>(
&mut self,
reserve_len: usize,
stream_reader: &mut SR,
) -> crate::Result<Option<core::num::NonZeroUsize>>
where
SR: StreamReader,
{
self.manage_capacity(reserve_len)?;
let (init, uninit) = self.buffer.split_at_spare_mut();
let Some(len) = stream_reader.read(uninit.into()).await? else {
cold_path();
return Ok(None);
};
let new_len = init.len().wrapping_add(len.get());
unsafe {
self.buffer.set_len(new_len);
}
self.current_end_idx = new_len;
Ok(Some(len))
}
#[cfg(feature = "web-socket")]
pub(crate) fn set_indices(&mut self, antecedent_end_idx: usize, current_end_idx: usize) {
self.current_end_idx = current_end_idx.min(self.buffer.len());
self.antecedent_end_idx = antecedent_end_idx.min(self.current_end_idx);
}
#[cfg(any(feature = "postgres", feature = "web-socket"))]
pub(crate) fn suffix_pusher(&mut self) -> crate::collections::SuffixGuardVectorMut<'_, u8> {
crate::collections::SuffixGuardVectorMut::from(&mut self.buffer)
}
#[inline]
fn manage_capacity(&mut self, additional: usize) -> crate::Result<()> {
let buffer_len = self.buffer.len();
let capacity_ub = self.capacity_ub;
let current_end_idx = self.current_end_idx;
let following_len = buffer_len.wrapping_sub(current_end_idx);
if additional > capacity_ub {
cold_path();
return Err(BufStreamReaderError::CapacityOverflow.into());
}
if following_len == 0 && !self.forbid_clear {
self.clear();
self.buffer.reserve(additional)?;
return Ok(());
}
let required_capacity = current_end_idx.wrapping_add(additional);
if self.buffer.capacity() >= required_capacity {
return Ok(());
}
if required_capacity <= capacity_ub {
self.buffer.reserve(required_capacity.wrapping_sub(buffer_len))?;
return Ok(());
}
cold_path();
if self.forbid_clear {
return Err(BufStreamReaderError::ForbiddenClear.into());
}
self.antecedent_end_idx = 0;
self.current_end_idx = 0;
self.buffer.copy_within(current_end_idx.., 0);
self.buffer.truncate(following_len);
self.buffer.reserve(additional.wrapping_sub(following_len))?;
Ok(())
}
#[inline]
fn remove_current(antecedent_end_idx: &mut usize, current_end_idx: &mut usize, offset: usize) {
let idx = current_end_idx.wrapping_add(offset);
*antecedent_end_idx = idx;
*current_end_idx = idx;
}
}
impl Lease<BufStreamReader> for BufStreamReader {
#[inline]
fn lease(&self) -> &BufStreamReader {
self
}
}
impl LeaseMut<BufStreamReader> for BufStreamReader {
#[inline]
fn lease_mut(&mut self) -> &mut BufStreamReader {
self
}
}
impl Debug for BufStreamReader {
#[inline]
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("NetReadBuffer").finish()
}
}
impl Default for BufStreamReader {
#[inline]
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use crate::stream::{BufStreamReader, BytesStream, StreamWriter};
#[wtx::test]
async fn read_header_and_payload() {
let mut stream = BytesStream::default();
stream.write_all(&[0, 2, 1, 2]).await.unwrap();
let mut nrb = BufStreamReader::default();
let header = nrb.read_header::<_, 2>(&mut stream).await.unwrap().unwrap();
let len = u16::from_be_bytes(header);
nrb.read_payload(len.into(), &mut stream).await.unwrap();
assert_eq!(nrb.current(), &[1, 2][..]);
}
#[wtx::test]
async fn zero_payload() {
let mut stream = BytesStream::default();
stream.write_all(&[0, 0]).await.unwrap();
let mut nrb = BufStreamReader::default();
let header = nrb.read_header::<_, 2>(&mut stream).await.unwrap().unwrap();
let len = u16::from_be_bytes(header);
nrb.read_payload(len.into(), &mut stream).await.unwrap();
assert!(nrb.current().is_empty());
}
}