use std::io;
use std::io::prelude::*;
use tokio::io::AsyncRead;
use tokio::io::AsyncReadExt;
const PACKET_BUFFER_SIZE: usize = 4_096;
const PACKET_LARGE_BUFFER_SIZE: usize = 1_048_576;
pub struct PacketReader<R> {
bytes: Vec<u8>,
start: usize,
remaining: usize,
pub r: R,
}
impl<R> PacketReader<R> {
pub fn new(r: R) -> Self {
PacketReader {
bytes: Vec::new(),
start: 0,
remaining: 0,
r,
}
}
}
impl<R: Read> PacketReader<R> {
#[allow(dead_code)]
pub fn next(&mut self) -> io::Result<Option<(u8, Packet<'_>)>> {
self.start = self.bytes.len() - self.remaining;
loop {
if self.remaining != 0 {
let bytes = {
let bytes = &self.bytes[self.start..];
unsafe { ::std::slice::from_raw_parts(bytes.as_ptr(), bytes.len()) }
};
match packet(bytes) {
Ok((rest, p)) => {
self.remaining = rest.len();
return Ok(Some(p));
}
Err(nom::Err::Incomplete(_)) | Err(nom::Err::Error(_)) => {}
Err(nom::Err::Failure(ctx)) => {
let err = Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("{:?}", ctx),
));
self.bytes.truncate(self.remaining);
return err;
}
}
}
self.bytes.drain(0..self.start);
self.start = 0;
let end = self.bytes.len();
self.bytes.resize(std::cmp::max(4096, end * 2), 0);
let read = {
let buf = &mut self.bytes[end..];
self.r.read(buf)?
};
self.bytes.truncate(end + read);
self.remaining = self.bytes.len();
if read == 0 {
if self.bytes.is_empty() {
return Ok(None);
} else {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
format!("{} unhandled bytes", self.bytes.len()),
));
}
}
}
}
}
impl<R: AsyncRead + Unpin> AsyncRead for PacketReader<R> {
fn poll_read(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<io::Result<()>> {
if self.remaining != 0 {
buf.put_slice(&self.bytes[self.start..]);
self.bytes.clear();
self.start = 0;
self.remaining = 0;
std::task::Poll::Ready(Ok(()))
} else {
std::pin::Pin::new(&mut self.r).poll_read(cx, buf)
}
}
}
impl<R: AsyncRead + Unpin> PacketReader<R> {
pub async fn next_async(&mut self) -> io::Result<Option<(u8, Packet<'_>)>> {
self.start = self.bytes.len() - self.remaining;
let mut buffer_size = PACKET_BUFFER_SIZE;
loop {
if self.remaining != 0 {
let bytes = {
let bytes = &self.bytes[self.start..];
unsafe { ::std::slice::from_raw_parts(bytes.as_ptr(), self.remaining) }
};
match packet(bytes) {
Ok((rest, p)) => {
self.remaining = rest.len();
if self.remaining > 0 {
self.bytes = rest.to_vec();
self.start = 0;
}
return Ok(Some(p));
}
Err(nom::Err::Incomplete(_)) | Err(nom::Err::Error(_)) => {}
Err(nom::Err::Failure(ctx)) => {
self.bytes.truncate(self.remaining);
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("{:?}", ctx),
));
}
}
}
self.bytes.drain(0..self.start);
self.start = 0;
let end = self.remaining;
if self.bytes.len() - end < buffer_size {
let new_len = std::cmp::max(buffer_size, end * 2);
self.bytes.resize(new_len, 0);
}
let read = {
let buf = &mut self.bytes[end..];
self.r.read(buf).await?
};
self.remaining = end + read;
buffer_size = PACKET_LARGE_BUFFER_SIZE;
if read == 0 {
self.bytes.truncate(self.remaining);
if self.bytes.is_empty() {
return Ok(None);
} else {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
format!("{} unhandled bytes", self.bytes.len()),
));
}
}
}
}
}
pub fn fullpacket(i: &[u8]) -> nom::IResult<&[u8], (u8, &[u8])> {
let (i, _) = nom::bytes::complete::tag(&[0xff, 0xff, 0xff])(i)?;
let (i, seq) = nom::bytes::complete::take(1u8)(i)?;
let (i, bytes) = nom::bytes::complete::take(U24_MAX)(i)?;
Ok((i, (seq[0], bytes)))
}
pub fn onepacket(i: &[u8]) -> nom::IResult<&[u8], (u8, &[u8])> {
let (i, length) = nom::number::complete::le_u24(i)?;
let (i, seq) = nom::bytes::complete::take(1u8)(i)?;
let (i, bytes) = nom::bytes::complete::take(length)(i)?;
Ok((i, (seq[0], bytes)))
}
#[derive(Clone)]
pub struct Packet<'a>(&'a [u8], Vec<u8>);
impl<'a> Packet<'a> {
fn extend(&mut self, bytes: &'a [u8]) {
if self.0.is_empty() {
if self.1.is_empty() {
self.0 = bytes;
} else {
self.1.extend(bytes);
}
} else {
assert!(self.1.is_empty());
let mut v = self.0.to_vec();
v.extend(bytes);
self.1 = v;
self.0 = &[];
}
}
}
impl<'a> AsRef<[u8]> for Packet<'a> {
fn as_ref(&self) -> &[u8] {
if self.1.is_empty() {
self.0
} else {
&self.1
}
}
}
use crate::U24_MAX;
use std::ops::Deref;
impl<'a> Deref for Packet<'a> {
type Target = [u8];
fn deref(&self) -> &Self::Target {
self.as_ref()
}
}
pub(crate) fn packet(i: &[u8]) -> nom::IResult<&[u8], (u8, Packet<'_>)> {
nom::combinator::map(
nom::sequence::pair(
nom::multi::fold_many0(
fullpacket,
|| (0, None),
|(seq, pkt): (_, Option<Packet<'_>>), (nseq, p)| {
let pkt = if let Some(mut pkt) = pkt {
assert_eq!(nseq, seq + 1);
pkt.extend(p);
Some(pkt)
} else {
Some(Packet(p, Vec::new()))
};
(nseq, pkt)
},
),
onepacket,
),
move |(full, last)| {
let seq = last.0;
let pkt = if let Some(mut pkt) = full.1 {
assert_eq!(last.0, full.0 + 1);
pkt.extend(last.1);
pkt
} else {
Packet(last.1, Vec::new())
};
(seq, pkt)
},
)(i)
}