rustfs_rio/
limit_reader.rs1use 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
52impl<R> LimitReader<R>
54where
55 R: AsyncRead + Unpin + Send + Sync,
56{
57 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 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 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}