use core::num::NonZeroUsize;
use std::time::Duration;
use tokio::{
io::{AsyncRead, AsyncReadExt as _},
time::Instant,
};
#[derive(Debug)]
pub struct PeekOutput<D> {
pub data: Option<D>,
pub peek_size: usize,
}
#[inline(always)]
pub fn peek_input_until<R, O, P>(
reader: &mut R,
buffer: &mut [u8],
timeout: Option<Duration>,
predicate: P,
) -> impl Future<Output = PeekOutput<O>>
where
R: AsyncRead + Unpin,
P: Fn(&[u8]) -> Option<O>,
{
peek_input_until_with_offset(reader, buffer, 0, timeout, predicate)
}
#[inline]
pub fn peek_input_until_with_offset<R, O, P>(
reader: &mut R,
buffer: &mut [u8],
offset: usize,
timeout: Option<Duration>,
predicate: P,
) -> impl Future<Output = PeekOutput<O>>
where
R: AsyncRead + Unpin,
P: Fn(&[u8]) -> Option<O>,
{
let default_budget =
NonZeroUsize::new(buffer.len().saturating_div(4).max(1) + 1).unwrap_or(NonZeroUsize::MIN);
peek_input_until_with_options(
reader,
buffer,
offset,
timeout,
Some(default_budget),
predicate,
)
}
pub async fn peek_input_until_with_options<R, O, P>(
reader: &mut R,
buffer: &mut [u8],
offset: usize,
timeout: Option<Duration>,
max_attempts: Option<NonZeroUsize>,
predicate: P,
) -> PeekOutput<O>
where
R: AsyncRead + Unpin,
P: Fn(&[u8]) -> Option<O>,
{
let mut output = PeekOutput {
data: None,
peek_size: offset.min(buffer.len()),
};
if buffer[output.peek_size..].is_empty() {
return output;
}
let peek_deadline = timeout.map(|d| Instant::now() + d);
let attempt_cap = max_attempts.map(NonZeroUsize::get).unwrap_or(usize::MAX);
for _ in 0..attempt_cap {
let read_fut = reader.read(&mut buffer[output.peek_size..]);
let n = match peek_deadline {
Some(deadline) => {
let now = Instant::now();
if now >= deadline {
tracing::debug!("I/O peek: abort: deadline reached");
return output;
}
let remaining = deadline - now;
match tokio::time::timeout(remaining, read_fut).await {
Err(err) => {
tracing::debug!("I/O peek: time-fenced peek read timeout error: {err}");
return output;
}
Ok(Err(err)) => {
tracing::debug!("I/O peek: time-fenced peek read error: {err}");
return output;
}
Ok(Ok(n)) => n,
}
}
None => match read_fut.await {
Err(err) => {
tracing::debug!("I/O peek: peek read error: {err}");
return output;
}
Ok(n) => n,
},
};
if n == 0 {
tracing::trace!("I/O peek: break loop: no new data read...");
return output;
}
output.peek_size = (output.peek_size + n).min(buffer.len());
if let Some(data) = predicate(&buffer[..output.peek_size]) {
output.data = Some(data);
tracing::trace!("I/O peek: data found using predicate: return it...");
return output;
}
}
output
}
#[cfg(test)]
mod tests {
#![expect(
clippy::unreachable,
reason = "test fixture: closure is wired up but never invoked on the tested path"
)]
use super::*;
use std::{
io,
pin::Pin,
task::{Context, Poll},
};
use tokio::io::ReadBuf;
#[tokio::test]
async fn returns_immediately_for_empty_buffer() {
let mut reader = tokio_test::io::Builder::new().build();
let mut buffer = [];
let output =
peek_input_until::<_, (), _>(&mut reader, &mut buffer, None, |_| unreachable!()).await;
assert!(output.data.is_none());
assert_eq!(output.peek_size, 0);
}
#[tokio::test]
async fn returns_data_when_predicate_matches_on_first_read() {
let mut reader = tokio_test::io::Builder::new().read(b"hello").build();
let mut buffer = [0_u8; 8];
let output = peek_input_until(&mut reader, &mut buffer, None, |buf| {
(buf == b"hello").then_some("hello")
})
.await;
assert_eq!(output.data, Some("hello"));
assert_eq!(output.peek_size, 5);
assert_eq!(&buffer[..output.peek_size], b"hello");
}
#[tokio::test]
async fn accumulates_across_multiple_reads_until_predicate_matches() {
let mut reader = tokio_test::io::Builder::new()
.read(b"he")
.read(b"llo")
.build();
let mut buffer = [0_u8; 8];
let output = peek_input_until(&mut reader, &mut buffer, None, |buf| {
(buf == b"hello").then_some(buf.len())
})
.await;
assert_eq!(output.data, Some(5));
assert_eq!(output.peek_size, 5);
assert_eq!(&buffer[..output.peek_size], b"hello");
}
#[tokio::test]
async fn returns_partial_bytes_when_reader_hits_eof_before_match() {
let mut reader = tokio_test::io::Builder::new().read(b"he").build();
let mut buffer = [0_u8; 8];
let output = peek_input_until(&mut reader, &mut buffer, None, |buf| {
(buf == b"hello").then_some(())
})
.await;
assert!(output.data.is_none());
assert_eq!(output.peek_size, 2);
assert_eq!(&buffer[..output.peek_size], b"he");
}
#[tokio::test]
async fn returns_partial_bytes_when_reader_errors_after_progress() {
let mut reader = tokio_test::io::Builder::new()
.read(b"he")
.read_error(io::Error::new(io::ErrorKind::BrokenPipe, "boom"))
.build();
let mut buffer = [0_u8; 8];
let output = peek_input_until(&mut reader, &mut buffer, None, |buf| {
(buf == b"hello").then_some(())
})
.await;
assert!(output.data.is_none());
assert_eq!(output.peek_size, 2);
assert_eq!(&buffer[..output.peek_size], b"he");
}
#[tokio::test]
async fn returns_no_data_when_first_read_errors() {
let mut reader = tokio_test::io::Builder::new()
.read_error(io::Error::new(io::ErrorKind::BrokenPipe, "boom"))
.build();
let mut buffer = [0_u8; 8];
let output = peek_input_until(&mut reader, &mut buffer, None, |buf| {
(!buf.is_empty()).then_some(())
})
.await;
assert!(output.data.is_none());
assert_eq!(output.peek_size, 0);
}
#[tokio::test]
async fn timeout_returns_partial_bytes_already_peeked() {
let mut reader = TwoPhaseReader {
first_chunk: Some(b"he".to_vec()),
sleep: Some(Box::pin(tokio::time::sleep(Duration::from_millis(50)))),
};
let mut buffer = [0_u8; 8];
let output = peek_input_until(
&mut reader,
&mut buffer,
Some(Duration::from_millis(10)),
|buf| (buf == b"hello").then_some(()),
)
.await;
assert!(output.data.is_none());
assert_eq!(output.peek_size, 2);
assert_eq!(&buffer[..output.peek_size], b"he");
}
#[tokio::test]
async fn stops_when_attempt_budget_is_exhausted() {
let mut reader = tokio_test::io::Builder::new().read(b"h").read(b"e").build();
let mut buffer = [0_u8; 8];
let output = peek_input_until(&mut reader, &mut buffer, None, |buf| {
(buf == b"hel").then_some(())
})
.await;
assert!(output.data.is_none());
assert_eq!(output.peek_size, 2);
assert_eq!(&buffer[..output.peek_size], b"he");
}
#[tokio::test]
async fn explicit_max_attempts_overrides_buffer_derived_budget() {
let mut reader = tokio_test::io::Builder::new()
.read(b"h")
.read(b"e")
.read(b"l")
.build();
let mut buffer = [0_u8; 8];
let output = peek_input_until_with_options(
&mut reader,
&mut buffer,
0,
None,
NonZeroUsize::new(5),
|buf| (buf == b"hel").then_some(()),
)
.await;
assert_eq!(output.data, Some(()));
assert_eq!(output.peek_size, 3);
}
#[tokio::test]
async fn explicit_max_attempts_none_means_no_cap() {
let mut reader = tokio_test::io::Builder::new()
.read(b"h")
.read(b"e")
.read(b"l")
.read(b"l")
.read(b"o")
.build();
let mut buffer = [0_u8; 8];
let output =
peek_input_until_with_options(&mut reader, &mut buffer, 0, None, None, |buf| {
(buf == b"hello").then_some(buf.len())
})
.await;
assert_eq!(output.data, Some(5));
assert_eq!(output.peek_size, 5);
}
struct TwoPhaseReader {
first_chunk: Option<Vec<u8>>,
sleep: Option<Pin<Box<tokio::time::Sleep>>>,
}
impl AsyncRead for TwoPhaseReader {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if let Some(chunk) = self.first_chunk.take() {
buf.put_slice(&chunk[..chunk.len().min(buf.remaining())]);
return Poll::Ready(Ok(()));
}
match self.sleep.as_mut() {
Some(sleep) => sleep.as_mut().poll(cx).map(|_| Ok(())),
None => Poll::Ready(Ok(())),
}
}
}
}