rustfs_rio/
hardlimit_reader.rs1use 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 let before = buf.filled().len();
44
45 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 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}