use tokio::io::{AsyncRead, AsyncSeek, SeekFrom, AsyncSeekExt, AsyncReadExt, AsyncWriteExt};
use std::pin::Pin;
use core::task::Poll;
use tokio::io::ReadBuf;
use std::error::Error;
use std::io::Read;
use std::io::Write;
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub enum ObserverDecision {
Continue,
Abort,
}
pub trait StreamObserver {
fn begin(&mut self) {
}
fn before_read(&mut self) -> ObserverDecision {
return ObserverDecision::Continue;
}
fn before_write(&mut self, _: &[u8]) -> ObserverDecision {
return ObserverDecision::Continue;
}
fn after_write(&mut self, _:&[u8]) -> ObserverDecision {
return ObserverDecision::Continue;
}
fn end(&mut self, _:usize, _:Option<Box<&dyn Error>>) {
}
}
struct DumbObserver;
impl StreamObserver for DumbObserver {
}
pub trait EhRead:Read {
fn try_read_exact(&mut self, buffer: &mut [u8]) -> Result<usize, Box<dyn Error>> {
let wanted = buffer.len();
let mut copied:usize = 0;
loop {
let rr = self.read(&mut buffer[copied..])?;
if rr == 0 {
return Ok(copied);
}
copied = copied + rr;
if copied >= wanted {
return Ok(copied);
}
}
}
fn skip(&mut self, bytes: usize) -> Result<usize, Box<dyn Error>> {
if bytes == 0 {
return Ok(0);
}
let mut buffer = [0u8; 4096];
let mut remaining = bytes;
while remaining > 0 {
let rr = self.try_read_exact(&mut buffer[..remaining])?;
if rr == 0 {
return Err(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "Insufficient bytes to skip").into());
}
remaining -= rr;
}
Ok(bytes)
}
fn stream_to<W>(&mut self, w:&mut W, buffer_size:Option<usize>, observer:Option<Box<dyn StreamObserver>>) -> Result<usize, Box<dyn Error>>
where W:Write + Sized
{
let mut buffer_size = buffer_size.unwrap_or(4096);
if buffer_size == 0 {
buffer_size = 4096;
}
let mut buffer = vec![0u8; buffer_size];
return self.stream_to_with_buffer(w, &mut buffer, observer);
}
fn stream_to_with_buffer<W>(&mut self, w:&mut W, buffer:&mut[u8], observer:Option<Box<dyn StreamObserver>>) -> Result<usize, Box<dyn Error>>
where W:Write+Sized {
let mut observer = observer;
let default_ob: Box<dyn StreamObserver> = Box::new(DumbObserver);
let mut obs = observer.take().unwrap_or(default_ob);
let mut copied:usize = 0;
loop {
let decision = obs.before_read();
if decision == ObserverDecision::Abort {
break;
}
let rr = self.read(buffer);
if rr.is_err() {
let err = rr.err().unwrap();
obs.end(copied, Some(Box::new(&err)));
return Err(err.into());
}
let rr = rr.unwrap();
if rr == 0 {
break;
}
let decision = obs.before_write(&buffer[0..rr]);
if decision == ObserverDecision::Abort {
break;
}
let wr = w.write_all(&buffer[0..rr]);
if wr.is_err() {
let err = wr.err().unwrap();
obs.end(copied, Some(Box::new(&err)));
return Err(err.into());
}
let decision = obs.after_write(&buffer[0..rr]);
if decision == ObserverDecision::Abort {
break;
}
copied += rr;
}
return Ok(copied);
}
}
impl <T> EhRead for T where T:Read{}
pub struct UndoReader<T>
where T:AsyncRead + Unpin
{
src: T,
read_count: usize,
limit: usize,
buffer: Vec<Vec<u8>>
}
impl<T> UndoReader<T>
where T:AsyncRead + Unpin
{
pub fn destruct(self) -> (Vec<u8>, T) {
let count = self.count_unread();
let mut resultv = vec![0u8; count];
self.copy_into(&mut resultv);
return (resultv, self.src)
}
fn copy_into(&self, buf:&mut [u8]) -> usize{
let mut copied = 0;
for i in 0.. self.buffer.len() {
let v = &self.buffer[ self.buffer.len() - i - 1];
for i in 0..v.len() {
buf[copied + i] = v[i];
}
copied += v.len();
}
return copied;
}
pub fn limit(&self)->usize {
self.limit
}
pub fn count_unread(&self) -> usize {
let mut result:usize = 0;
for v in &self.buffer {
result += v.len();
}
return result;
}
pub fn new(src:T, limit:Option<usize>) -> UndoReader<T> {
UndoReader {
src,
limit: match limit {
None => std::usize::MAX,
Some(actual) => actual
},
read_count: 0,
buffer: Vec::new()
}
}
pub fn unread(&mut self, data:&[u8]) -> &mut Self {
if data.len() > 0 {
let mut new = vec![0u8;data.len()];
for (index, payload) in data.iter().enumerate() {
new[index] = *payload;
}
self.buffer.push(new);
}
return self;
}
}
impl<T> AsyncRead for UndoReader<T>
where T:AsyncRead + Unpin
{
fn poll_read(mut self: Pin<&mut Self>, ctx: &mut std::task::Context<'_>,
data: &mut ReadBuf<'_>) -> Poll<Result<(), std::io::Error>> {
loop {
let next = self.buffer.pop();
match next {
Some(bufdata) => {
if bufdata.len() == 0 {
continue;
}
let available = bufdata.len();
let remaining = data.remaining();
if available <= remaining {
data.put_slice(&bufdata);
} else {
data.put_slice(&bufdata[0..remaining]);
let left_over = &bufdata[remaining..];
let mut new_vec = vec![0u8;left_over.len()];
for (index, payload) in left_over.iter().enumerate() {
new_vec[index] = *payload;
}
self.buffer.push(new_vec);
}
return Poll::Ready(Ok(()));
},
None => {
break;
}
}
}
if self.read_count >= self.limit {
return Poll::Ready(Ok(()));
}
let ms = &mut *self;
let p = Pin::new(&mut ms.src);
let before_filled = data.filled().len();
let result = p.poll_read(ctx, data);
let after_filled = data.filled().len();
let this_read = after_filled - before_filled;
self.read_count += this_read;
let overread = self.read_count > self.limit;
if overread {
let overread_count = self.read_count - self.limit;
data.set_filled(after_filled - overread_count);
self.read_count = self.limit;
}
return result;
}
}
pub struct LimitSeekerReader<T>
where T:AsyncRead + AsyncSeek + Unpin
{
src: T,
read_count: usize,
limit: usize,
}
impl<T> LimitSeekerReader<T>
where T:AsyncRead + AsyncSeek + Unpin
{
pub fn destruct(self) -> (usize, T) {
(self.read_count, self.src)
}
pub fn new(src:T, limit:Option<usize>) -> LimitSeekerReader<T> {
LimitSeekerReader {
src,
limit: {
match limit {
None => std::usize::MAX,
Some(actual_limit) => actual_limit
}
},
read_count: 0
}
}
}
impl<T> AsyncRead for LimitSeekerReader<T>
where T:AsyncRead + AsyncSeek + Unpin
{
fn poll_read(mut self: Pin<&mut Self>, ctx: &mut std::task::Context<'_>,
data: &mut ReadBuf<'_>) -> Poll<Result<(), std::io::Error>> {
if self.read_count >= self.limit {
return Poll::Ready(Ok(()));
}
let ms = &mut *self;
let p = Pin::new(&mut ms.src);
let before_filled = data.filled().len();
let result = p.poll_read(ctx, data);
let after_filled = data.filled().len();
let this_read = after_filled - before_filled;
self.read_count += this_read;
let overread = self.read_count > self.limit;
if overread {
let overread_count = self.read_count - self.limit;
data.set_filled(after_filled - overread_count);
self.read_count = self.limit;
}
return result;
}
}
impl<T> AsyncSeek for LimitSeekerReader<T>
where T:AsyncRead + AsyncSeek + Unpin
{
fn start_seek(mut self: Pin<&mut Self>, from: SeekFrom) -> Result<(), std::io::Error> {
let ms = &mut *self;
let p = Pin::new(&mut ms.src);
return p.start_seek(from);
}
fn poll_complete(mut self: Pin<&mut Self>, ctx: &mut std::task::Context<'_>) -> Poll<Result<u64, std::io::Error>> {
let ms = &mut *self;
let p = Pin::new(&mut ms.src);
return p.poll_complete(ctx);
}
}
pub struct LimitReader<T>
where T:AsyncRead + Unpin
{
src: T,
read_count: usize,
limit: usize,
}
impl<T> LimitReader<T>
where T:AsyncRead + Unpin
{
pub fn new(src:T, limit:Option<usize>) -> LimitReader<T> {
LimitReader {
src,
limit: {
match limit {
None => std::usize::MAX,
Some(actual_limit) => actual_limit
}
},
read_count: 0
}
}
pub fn destruct(self) -> (usize, T) {
(self.read_count, self.src)
}
}
impl<T> AsyncRead for LimitReader<T>
where T:AsyncRead + Unpin
{
fn poll_read(mut self: Pin<&mut Self>, ctx: &mut std::task::Context<'_>,
data: &mut ReadBuf<'_>) -> Poll<Result<(), std::io::Error>> {
if self.read_count >= self.limit {
return Poll::Ready(Ok(()));
}
let ms = &mut *self;
let p = Pin::new(&mut ms.src);
let before_filled = data.filled().len();
let result = p.poll_read(ctx, data);
let after_filled = data.filled().len();
let this_read = after_filled - before_filled;
self.read_count += this_read;
let overread = self.read_count > self.limit;
if overread {
let overread_count = self.read_count - self.limit;
data.set_filled(after_filled - overread_count);
self.read_count = self.limit;
}
return result;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_undo() {
let file = tokio::fs::File::open("test.data").await.unwrap();
let mut undor = UndoReader::new(file, Some(10));
let mut buf = [0u8; 1024];
let undo = "XXX".as_bytes();
let undo1 = "YYY".as_bytes();
undor.unread(undo);
undor.unread(undo1);
let rr = undor.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "YYY".as_bytes()); let rr = undor.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "XXX".as_bytes()); let rr = undor.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "123456789\n".as_bytes());
}
#[tokio::test]
async fn test_limit() {
let file = tokio::fs::File::open("test.data").await.unwrap();
let mut limitr = LimitReader::new(file, Some(10));
let mut buf = [0u8; 1024];
let rr = limitr.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "123456789\n".as_bytes());
let rr = limitr.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "".as_bytes());
}
#[tokio::test]
async fn test_seek() {
let file = tokio::fs::File::open("test.data").await.unwrap();
let mut limitr = LimitSeekerReader::new(file, Some(10));
limitr.seek(SeekFrom::Current(13)).await.unwrap();
let mut buf = [0u8; 1024];
let rr = limitr.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "456789\n123".as_bytes());
let rr = limitr.read(&mut buf).await.unwrap();
assert_eq!(&buf[0..rr], "".as_bytes());
}
}