use futures_core::ready;
use futures_util::io::AsyncBufRead;
use std::future::Future;
use std::io;
use std::mem;
use std::pin::Pin;
use std::task::{Context, Poll};
pub use crate::{Captures, Needle};
pub trait AsyncUntilNeedleRead: AsyncBufRead {
fn read_until_needle<'a, N>(&'a mut self, needle: N) -> ReadUntilNeedle<'a, Self, N>
where
Self: Unpin + Sized,
N: Needle + 'a;
fn split_read_until_needle<'a, N>(
&'a mut self,
needle: N,
before: &'a mut Vec<u8>,
matched: &'a mut Vec<u8>,
) -> impl Future<Output = io::Result<usize>> + 'a
where
Self: Unpin + Sized,
N: Needle + 'a,
{
async move {
let captures = self.read_until_needle(needle).await?;
let total_bytes_read = captures.total_bytes_read();
let (b, m) = captures.split();
before.extend_from_slice(&b);
matched.extend_from_slice(&m);
Ok(total_bytes_read)
}
}
}
impl<R> AsyncUntilNeedleRead for R
where
R: AsyncBufRead + Unpin,
{
fn read_until_needle<'a, N>(&'a mut self, needle: N) -> ReadUntilNeedle<'a, Self, N>
where
Self: Unpin + Sized,
N: Needle + 'a,
{
ReadUntilNeedle {
reader: self,
needle,
buf: Vec::new(),
}
}
}
pub struct ReadUntilNeedle<'a, R, N>
where
R: Unpin + ?Sized,
{
reader: &'a mut R,
needle: N,
buf: Vec<u8>,
}
impl<R: ?Sized + Unpin, N> Unpin for ReadUntilNeedle<'_, R, N> {}
impl<'a, R, N> Future for ReadUntilNeedle<'a, R, N>
where
R: AsyncBufRead + Unpin + ?Sized,
N: Needle,
{
type Output = io::Result<Captures>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let ReadUntilNeedle {
reader,
needle,
buf,
} = &mut *self;
let reader = Pin::new(reader);
read_until_needle_internal(reader, cx, needle, buf)
}
}
fn read_until_needle_internal<R, N>(
mut reader: Pin<&mut R>,
cx: &mut Context<'_>,
needle: &N,
buf: &mut Vec<u8>,
) -> Poll<io::Result<Captures>>
where
R: AsyncBufRead + ?Sized,
N: Needle,
{
loop {
let available = ready!(reader.as_mut().poll_fill_buf(cx))?;
let len = available.len();
if len == 0 {
return Poll::Ready(Ok(Captures::new(mem::take(buf), buf.len())));
}
buf.extend_from_slice(available);
if let Some(range) = needle.findin(buf) {
let used = len.saturating_sub(buf.len().saturating_sub(range.end));
reader.as_mut().consume(used);
buf.truncate(range.end);
return Poll::Ready(Ok(Captures::new(mem::take(buf), range.start)));
} else {
reader.as_mut().consume(len);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::{
stream::{iter, TryStreamExt as _},
AsyncReadExt as _,
};
#[tokio::test]
async fn test_async_read() {
let mut stream = iter(vec![
Ok(b"hello".to_vec()),
Ok(b" wo".to_vec()),
Ok(b"rld!".to_vec()),
])
.into_async_read();
let mut buf = Vec::new();
stream.read_to_end(&mut buf).await.unwrap();
assert_eq!(buf, b"hello world!");
}
#[tokio::test]
async fn test_split_read_until_needle() {
let mut stream = iter(vec![
Ok(b"hello".to_vec()),
Ok(b" wo".to_vec()),
Ok(b"rld!!".to_vec()),
])
.into_async_read();
let mut before = Vec::new();
let mut matched = Vec::new();
let mut buf = Vec::new();
assert_eq!(
stream
.split_read_until_needle(b"world", &mut before, &mut matched)
.await
.unwrap(),
11
);
assert_eq!(before, b"hello ");
assert_eq!(matched, b"world");
assert_eq!(stream.read_to_end(&mut buf).await.unwrap(), 2);
assert_eq!(buf, b"!!");
}
}