Skip to main content

rustfs_rio/
limit_reader.rs

1// Copyright 2024 RustFS Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! LimitReader: a wrapper for AsyncRead that limits the total number of bytes read.
16//!
17//! # Example
18//! ```
19//! use tokio::io::{AsyncReadExt, BufReader};
20//! use rustfs_rio::LimitReader;
21//!
22//! #[tokio::main]
23//! async fn main() {
24//!  let data = b"hello world";
25//!       let reader = BufReader::new(&data[..]);
26//!      let mut limit_reader = LimitReader::new(reader, data.len());
27//!
28//!      let mut buf = Vec::new();
29//!      let n = limit_reader.read_to_end(&mut buf).await.unwrap();
30//!      assert_eq!(n, data.len());
31//!      assert_eq!(&buf, data);
32//! }
33//! ```
34
35use pin_project_lite::pin_project;
36use std::pin::Pin;
37use std::task::{Context, Poll};
38use tokio::io::{AsyncRead, ReadBuf};
39
40use crate::{EtagResolvable, HashReaderDetector, HashReaderMut};
41
42pin_project! {
43    #[derive(Debug)]
44    pub struct LimitReader<R> {
45        #[pin]
46        pub inner: R,
47        limit: usize,
48        read: usize,
49    }
50}
51
52/// A wrapper for AsyncRead that limits the total number of bytes read.
53impl<R> LimitReader<R>
54where
55    R: AsyncRead + Unpin + Send + Sync,
56{
57    /// Create a new LimitReader wrapping `inner`, with a total read limit of `limit` bytes.
58    pub fn new(inner: R, limit: usize) -> Self {
59        Self { inner, limit, read: 0 }
60    }
61}
62
63impl<R> AsyncRead for LimitReader<R>
64where
65    R: AsyncRead + Unpin + Send + Sync,
66{
67    fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<std::io::Result<()>> {
68        let mut this = self.project();
69        let remaining = this.limit.saturating_sub(*this.read);
70        if remaining == 0 {
71            return Poll::Ready(Ok(()));
72        }
73        let orig_remaining = buf.remaining();
74        let allowed = remaining.min(orig_remaining);
75        if allowed == 0 {
76            return Poll::Ready(Ok(()));
77        }
78        if allowed == orig_remaining {
79            let before_size = buf.filled().len();
80            let poll = this.inner.as_mut().poll_read(cx, buf);
81            if let Poll::Ready(Ok(())) = &poll {
82                let n = buf.filled().len() - before_size;
83                *this.read += n;
84            }
85            poll
86        } else {
87            let mut temp = vec![0u8; allowed];
88            let mut temp_buf = ReadBuf::new(&mut temp);
89            let poll = this.inner.as_mut().poll_read(cx, &mut temp_buf);
90            if let Poll::Ready(Ok(())) = &poll {
91                let n = temp_buf.filled().len();
92                buf.put_slice(temp_buf.filled());
93                *this.read += n;
94            }
95            poll
96        }
97    }
98}
99
100impl<R> EtagResolvable for LimitReader<R>
101where
102    R: EtagResolvable,
103{
104    fn try_resolve_etag(&mut self) -> Option<String> {
105        self.inner.try_resolve_etag()
106    }
107}
108
109impl<R> HashReaderDetector for LimitReader<R>
110where
111    R: HashReaderDetector,
112{
113    fn is_hash_reader(&self) -> bool {
114        self.inner.is_hash_reader()
115    }
116    fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> {
117        self.inner.as_hash_reader_mut()
118    }
119}
120
121#[cfg(test)]
122mod tests {
123    use std::io::Cursor;
124
125    use super::*;
126    use tokio::io::{AsyncReadExt, BufReader};
127
128    #[tokio::test]
129    async fn test_limit_reader_exact() {
130        let data = b"hello world";
131        let reader = BufReader::new(&data[..]);
132        let mut limit_reader = LimitReader::new(reader, data.len());
133
134        let mut buf = Vec::new();
135        let n = limit_reader.read_to_end(&mut buf).await.unwrap();
136        assert_eq!(n, data.len());
137        assert_eq!(&buf, data);
138    }
139
140    #[tokio::test]
141    async fn test_limit_reader_less_than_data() {
142        let data = b"hello world";
143        let reader = BufReader::new(&data[..]);
144        let mut limit_reader = LimitReader::new(reader, 5);
145
146        let mut buf = Vec::new();
147        let n = limit_reader.read_to_end(&mut buf).await.unwrap();
148        assert_eq!(n, 5);
149        assert_eq!(&buf, b"hello");
150    }
151
152    #[tokio::test]
153    async fn test_limit_reader_zero() {
154        let data = b"hello world";
155        let reader = BufReader::new(&data[..]);
156        let mut limit_reader = LimitReader::new(reader, 0);
157
158        let mut buf = Vec::new();
159        let n = limit_reader.read_to_end(&mut buf).await.unwrap();
160        assert_eq!(n, 0);
161        assert!(buf.is_empty());
162    }
163
164    #[tokio::test]
165    async fn test_limit_reader_multiple_reads() {
166        let data = b"abcdefghij";
167        let reader = BufReader::new(&data[..]);
168        let mut limit_reader = LimitReader::new(reader, 7);
169
170        let mut buf1 = [0u8; 3];
171        let n1 = limit_reader.read(&mut buf1).await.unwrap();
172        assert_eq!(n1, 3);
173        assert_eq!(&buf1, b"abc");
174
175        let mut buf2 = [0u8; 5];
176        let n2 = limit_reader.read(&mut buf2).await.unwrap();
177        assert_eq!(n2, 4);
178        assert_eq!(&buf2[..n2], b"defg");
179
180        let mut buf3 = [0u8; 2];
181        let n3 = limit_reader.read(&mut buf3).await.unwrap();
182        assert_eq!(n3, 0);
183    }
184
185    #[tokio::test]
186    async fn test_limit_reader_large_file() {
187        use rand::Rng;
188        // Generate a 3MB random byte array for testing
189        let size = 3 * 1024 * 1024;
190        let mut data = vec![0u8; size];
191        rand::rng().fill(&mut data[..]);
192        let reader = Cursor::new(data.clone());
193        let mut limit_reader = LimitReader::new(reader, size);
194
195        // Read data into buffer
196        let mut buf = Vec::new();
197        let n = limit_reader.read_to_end(&mut buf).await.unwrap();
198        assert_eq!(n, size);
199        assert_eq!(buf.len(), size);
200        assert_eq!(&buf, &data);
201    }
202}