Skip to main content

khive_runtime/
bounded_read.rs

1//! Size-capped reads of streams whose length the caller does not control.
2//!
3//! Each reader takes at most `max + 1` bytes from the stream: one byte past the
4//! cap is enough to tell a stream of exactly `max` bytes from a longer one.
5//! `None` means more than `max` bytes were available. How that is reported is
6//! left to the caller, so each call site keeps its own error.
7
8use std::io::{self, Read};
9
10use tokio::io::{AsyncRead, AsyncReadExt};
11
12/// The most bytes a read with cap `max` takes. Saturates so `u64::MAX` is a valid cap.
13fn read_limit(max: u64) -> u64 {
14    max.saturating_add(1)
15}
16
17fn within_bound(bytes: Vec<u8>, max: u64) -> Option<Vec<u8>> {
18    if bytes.len() as u64 > max {
19        return None;
20    }
21    Some(bytes)
22}
23
24/// Read `reader` to its end. `Some` holds every byte when the stream has at most
25/// `max` bytes, `None` means it had more, and a read error is returned unchanged.
26pub fn read_to_end_bounded(reader: impl Read, max: u64) -> io::Result<Option<Vec<u8>>> {
27    let mut bytes = Vec::new();
28    let mut limited = reader.take(read_limit(max));
29    limited.read_to_end(&mut bytes)?;
30    Ok(within_bound(bytes, max))
31}
32
33/// Async counterpart of [`read_to_end_bounded`] over a `tokio` reader.
34pub async fn read_to_end_bounded_async(
35    reader: impl AsyncRead + Unpin,
36    max: u64,
37) -> io::Result<Option<Vec<u8>>> {
38    let mut bytes = Vec::new();
39    let mut limited = reader.take(read_limit(max));
40    limited.read_to_end(&mut bytes).await?;
41    Ok(within_bound(bytes, max))
42}
43
44#[cfg(test)]
45mod tests {
46    use super::*;
47    use std::pin::Pin;
48    use std::task::{Context, Poll};
49    use tokio::io::ReadBuf;
50
51    /// Yields `prefix`, then fails every later read.
52    struct Flaky {
53        prefix: &'static [u8],
54    }
55
56    impl Read for Flaky {
57        fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
58            if self.prefix.is_empty() {
59                return Err(io::Error::other("flaky reader failed"));
60            }
61            let len = self.prefix.len().min(buf.len());
62            buf[..len].copy_from_slice(&self.prefix[..len]);
63            self.prefix = &self.prefix[len..];
64            Ok(len)
65        }
66    }
67
68    impl AsyncRead for Flaky {
69        fn poll_read(
70            self: Pin<&mut Self>,
71            _cx: &mut Context<'_>,
72            buf: &mut ReadBuf<'_>,
73        ) -> Poll<io::Result<()>> {
74            let this = self.get_mut();
75            if this.prefix.is_empty() {
76                return Poll::Ready(Err(io::Error::other("flaky reader failed")));
77            }
78            let len = this.prefix.len().min(buf.remaining());
79            buf.put_slice(&this.prefix[..len]);
80            this.prefix = &this.prefix[len..];
81            Poll::Ready(Ok(()))
82        }
83    }
84
85    #[test]
86    fn blocking_exact_max_bytes_are_returned_whole() {
87        let data = [7_u8; 16];
88        let bytes = read_to_end_bounded(&data[..], 16).unwrap();
89        assert_eq!(bytes, Some(data.to_vec()));
90    }
91
92    #[test]
93    fn blocking_one_byte_past_max_is_refused() {
94        let data = [7_u8; 17];
95        let bytes = read_to_end_bounded(&data[..], 16).unwrap();
96        assert_eq!(bytes, None);
97    }
98
99    #[test]
100    fn blocking_empty_reader_yields_an_empty_vector() {
101        let bytes = read_to_end_bounded(io::empty(), 16).unwrap();
102        assert_eq!(bytes, Some(Vec::new()));
103    }
104
105    #[test]
106    fn blocking_maximum_cap_does_not_overflow() {
107        let data = [7_u8; 4];
108        let bytes = read_to_end_bounded(&data[..], u64::MAX).unwrap();
109        assert_eq!(bytes, Some(data.to_vec()));
110    }
111
112    #[test]
113    fn blocking_read_error_is_returned() {
114        let reader = Flaky { prefix: b"abc" };
115        let error = read_to_end_bounded(reader, 16).unwrap_err();
116        assert_eq!(error.kind(), io::ErrorKind::Other);
117        assert_eq!(error.to_string(), "flaky reader failed");
118    }
119
120    #[test]
121    fn blocking_stops_reading_one_byte_past_max() {
122        let data = [7_u8; 100];
123        let mut reader = &data[..];
124        let bytes = read_to_end_bounded(&mut reader, 10).unwrap();
125        assert_eq!(bytes, None);
126        assert_eq!(reader.len(), 89);
127    }
128
129    #[tokio::test]
130    async fn async_exact_max_bytes_are_returned_whole() {
131        let data = [7_u8; 16];
132        let bytes = read_to_end_bounded_async(&data[..], 16).await.unwrap();
133        assert_eq!(bytes, Some(data.to_vec()));
134    }
135
136    #[tokio::test]
137    async fn async_one_byte_past_max_is_refused() {
138        let data = [7_u8; 17];
139        let bytes = read_to_end_bounded_async(&data[..], 16).await.unwrap();
140        assert_eq!(bytes, None);
141    }
142
143    #[tokio::test]
144    async fn async_empty_reader_yields_an_empty_vector() {
145        let reader = tokio::io::empty();
146        let bytes = read_to_end_bounded_async(reader, 16).await.unwrap();
147        assert_eq!(bytes, Some(Vec::new()));
148    }
149
150    #[tokio::test]
151    async fn async_maximum_cap_does_not_overflow() {
152        let data = [7_u8; 4];
153        let max = u64::MAX;
154        let bytes = read_to_end_bounded_async(&data[..], max).await.unwrap();
155        assert_eq!(bytes, Some(data.to_vec()));
156    }
157
158    #[tokio::test]
159    async fn async_read_error_is_returned() {
160        let flaky = Flaky { prefix: b"abc" };
161        let error = read_to_end_bounded_async(flaky, 16).await.unwrap_err();
162        assert_eq!(error.kind(), io::ErrorKind::Other);
163        assert_eq!(error.to_string(), "flaky reader failed");
164    }
165
166    #[tokio::test]
167    async fn async_stops_reading_one_byte_past_max() {
168        let data = [7_u8; 100];
169        let mut rest = &data[..];
170        let bytes = read_to_end_bounded_async(&mut rest, 10).await.unwrap();
171        assert_eq!(bytes, None);
172        assert_eq!(rest.len(), 89);
173    }
174}