Skip to main content

rustfs_rio/
hardlimit_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
15use crate::compress_index::{Index, TryGetIndex};
16use crate::{EtagResolvable, HashReaderDetector, HashReaderMut, Reader};
17use pin_project_lite::pin_project;
18use std::io::{Error, Result};
19use std::pin::Pin;
20use std::task::{Context, Poll};
21use tokio::io::{AsyncRead, ReadBuf};
22
23pin_project! {
24    pub struct HardLimitReader {
25        #[pin]
26        pub inner: Box<dyn Reader>,
27        remaining: i64,
28    }
29}
30
31impl HardLimitReader {
32    pub fn new(inner: Box<dyn Reader>, limit: i64) -> Self {
33        HardLimitReader { inner, remaining: limit }
34    }
35}
36
37impl AsyncRead for HardLimitReader {
38    fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<Result<()>> {
39        if self.remaining < 0 {
40            return Poll::Ready(Err(Error::other("input provided more bytes than specified")));
41        }
42        // Save the initial length
43        let before = buf.filled().len();
44
45        // Poll the inner reader
46        let this = self.as_mut().project();
47        let poll = this.inner.poll_read(cx, buf);
48
49        if let Poll::Ready(Ok(())) = &poll {
50            let after = buf.filled().len();
51            let read = (after - before) as i64;
52            self.remaining -= read;
53            if self.remaining < 0 {
54                return Poll::Ready(Err(Error::other("input provided more bytes than specified")));
55            }
56        }
57        poll
58    }
59}
60
61impl EtagResolvable for HardLimitReader {
62    fn try_resolve_etag(&mut self) -> Option<String> {
63        self.inner.try_resolve_etag()
64    }
65}
66
67impl HashReaderDetector for HardLimitReader {
68    fn is_hash_reader(&self) -> bool {
69        self.inner.is_hash_reader()
70    }
71    fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> {
72        self.inner.as_hash_reader_mut()
73    }
74}
75
76impl TryGetIndex for HardLimitReader {
77    fn try_get_index(&self) -> Option<&Index> {
78        self.inner.try_get_index()
79    }
80}
81
82#[cfg(test)]
83mod tests {
84    use std::vec;
85
86    use crate::WarpReader;
87
88    use super::*;
89    use rustfs_utils::read_full;
90    use tokio::io::{AsyncReadExt, BufReader};
91
92    #[tokio::test]
93    async fn test_hardlimit_reader_normal() {
94        let data = b"hello world";
95        let reader = BufReader::new(&data[..]);
96        let reader = Box::new(WarpReader::new(reader));
97        let hardlimit = HardLimitReader::new(reader, 20);
98        let mut r = hardlimit;
99        let mut buf = Vec::new();
100        let n = r.read_to_end(&mut buf).await.unwrap();
101        assert_eq!(n, data.len());
102        assert_eq!(&buf, data);
103    }
104
105    #[tokio::test]
106    async fn test_hardlimit_reader_exact_limit() {
107        let data = b"1234567890";
108        let reader = BufReader::new(&data[..]);
109        let reader = Box::new(WarpReader::new(reader));
110        let hardlimit = HardLimitReader::new(reader, 10);
111        let mut r = hardlimit;
112        let mut buf = Vec::new();
113        let n = r.read_to_end(&mut buf).await.unwrap();
114        assert_eq!(n, 10);
115        assert_eq!(&buf, data);
116    }
117
118    #[tokio::test]
119    async fn test_hardlimit_reader_exceed_limit() {
120        let data = b"abcdef";
121        let reader = BufReader::new(&data[..]);
122        let reader = Box::new(WarpReader::new(reader));
123        let hardlimit = HardLimitReader::new(reader, 3);
124        let mut r = hardlimit;
125        let mut buf = vec![0u8; 10];
126        // 读取超限,应该返回错误
127        let err = match read_full(&mut r, &mut buf).await {
128            Ok(n) => {
129                println!("Read {n} bytes");
130                assert_eq!(n, 3);
131                assert_eq!(&buf[..n], b"abc");
132                None
133            }
134            Err(e) => Some(e),
135        };
136
137        assert!(err.is_some());
138
139        let err = err.unwrap();
140        assert_eq!(err.kind(), std::io::ErrorKind::Other);
141    }
142
143    #[tokio::test]
144    async fn test_hardlimit_reader_empty() {
145        let data = b"";
146        let reader = BufReader::new(&data[..]);
147        let reader = Box::new(WarpReader::new(reader));
148        let hardlimit = HardLimitReader::new(reader, 5);
149        let mut r = hardlimit;
150        let mut buf = Vec::new();
151        let n = r.read_to_end(&mut buf).await.unwrap();
152        assert_eq!(n, 0);
153        assert_eq!(&buf, data);
154    }
155}