Skip to main content

rustfs_rio/
compress_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};
17use crate::{HashReaderMut, Reader};
18use pin_project_lite::pin_project;
19use rustfs_utils::compress::{CompressionAlgorithm, compress_block, decompress_block};
20use rustfs_utils::{put_uvarint, uvarint};
21use std::cmp::min;
22use std::io::{self};
23use std::pin::Pin;
24use std::task::{Context, Poll};
25use tokio::io::{AsyncRead, ReadBuf};
26// use tracing::error;
27
28const COMPRESS_TYPE_COMPRESSED: u8 = 0x00;
29const COMPRESS_TYPE_UNCOMPRESSED: u8 = 0x01;
30const COMPRESS_TYPE_END: u8 = 0xFF;
31
32const DEFAULT_BLOCK_SIZE: usize = 1 << 20; // 1MB
33const HEADER_LEN: usize = 8;
34
35pin_project! {
36    #[derive(Debug)]
37    /// A reader wrapper that compresses data on the fly using DEFLATE algorithm.
38    pub struct CompressReader<R> {
39        #[pin]
40        pub inner: R,
41        buffer: Vec<u8>,
42        pos: usize,
43        done: bool,
44        block_size: usize,
45        compression_algorithm: CompressionAlgorithm,
46        index: Index,
47        written: usize,
48        uncomp_written: usize,
49        temp_buffer: Vec<u8>,
50        temp_pos: usize,
51    }
52}
53
54impl<R> CompressReader<R>
55where
56    R: Reader,
57{
58    pub fn new(inner: R, compression_algorithm: CompressionAlgorithm) -> Self {
59        Self {
60            inner,
61            buffer: Vec::new(),
62            pos: 0,
63            done: false,
64            compression_algorithm,
65            block_size: DEFAULT_BLOCK_SIZE,
66            index: Index::new(),
67            written: 0,
68            uncomp_written: 0,
69            temp_buffer: Vec::with_capacity(DEFAULT_BLOCK_SIZE), // Pre-allocate capacity
70            temp_pos: 0,
71        }
72    }
73
74    /// Optional: allow users to customize block_size
75    pub fn with_block_size(inner: R, block_size: usize, compression_algorithm: CompressionAlgorithm) -> Self {
76        Self {
77            inner,
78            buffer: Vec::new(),
79            pos: 0,
80            done: false,
81            compression_algorithm,
82            block_size,
83            index: Index::new(),
84            written: 0,
85            uncomp_written: 0,
86            temp_buffer: Vec::with_capacity(block_size),
87            temp_pos: 0,
88        }
89    }
90}
91
92impl<R> TryGetIndex for CompressReader<R>
93where
94    R: Reader,
95{
96    fn try_get_index(&self) -> Option<&Index> {
97        Some(&self.index)
98    }
99}
100
101impl<R> AsyncRead for CompressReader<R>
102where
103    R: AsyncRead + Unpin + Send + Sync,
104{
105    fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
106        let mut this = self.project();
107        // Copy from buffer first if available
108        if *this.pos < this.buffer.len() {
109            let to_copy = min(buf.remaining(), this.buffer.len() - *this.pos);
110            buf.put_slice(&this.buffer[*this.pos..*this.pos + to_copy]);
111            *this.pos += to_copy;
112            if *this.pos == this.buffer.len() {
113                this.buffer.clear();
114                *this.pos = 0;
115            }
116            return Poll::Ready(Ok(()));
117        }
118        if *this.done {
119            return Poll::Ready(Ok(()));
120        }
121        // Fill temporary buffer
122        while this.temp_buffer.len() < *this.block_size {
123            let remaining = *this.block_size - this.temp_buffer.len();
124            let mut temp = vec![0u8; remaining];
125            let mut temp_buf = ReadBuf::new(&mut temp);
126            match this.inner.as_mut().poll_read(cx, &mut temp_buf) {
127                Poll::Pending => {
128                    if this.temp_buffer.is_empty() {
129                        return Poll::Pending;
130                    }
131                    break;
132                }
133                Poll::Ready(Ok(())) => {
134                    let n = temp_buf.filled().len();
135                    if n == 0 {
136                        if this.temp_buffer.is_empty() {
137                            return Poll::Ready(Ok(()));
138                        }
139                        break;
140                    }
141                    this.temp_buffer.extend_from_slice(&temp[..n]);
142                }
143                Poll::Ready(Err(e)) => {
144                    // error!("CompressReader poll_read: read inner error: {e}");
145                    return Poll::Ready(Err(e));
146                }
147            }
148        }
149        // Process accumulated data
150        if !this.temp_buffer.is_empty() {
151            let uncompressed_data = &this.temp_buffer;
152            let out = build_compressed_block(uncompressed_data, *this.compression_algorithm);
153            *this.written += out.len();
154            *this.uncomp_written += uncompressed_data.len();
155            if let Err(e) = this.index.add(*this.written as i64, *this.uncomp_written as i64) {
156                // error!("CompressReader index add error: {e}");
157                return Poll::Ready(Err(e));
158            }
159            *this.buffer = out;
160            *this.pos = 0;
161            this.temp_buffer.truncate(0); // More efficient way to clear
162            let to_copy = min(buf.remaining(), this.buffer.len());
163            buf.put_slice(&this.buffer[..to_copy]);
164            *this.pos += to_copy;
165            if *this.pos == this.buffer.len() {
166                this.buffer.clear();
167                *this.pos = 0;
168            }
169            Poll::Ready(Ok(()))
170        } else {
171            Poll::Pending
172        }
173    }
174}
175
176impl<R> EtagResolvable for CompressReader<R>
177where
178    R: EtagResolvable,
179{
180    fn try_resolve_etag(&mut self) -> Option<String> {
181        self.inner.try_resolve_etag()
182    }
183}
184
185impl<R> HashReaderDetector for CompressReader<R>
186where
187    R: HashReaderDetector,
188{
189    fn is_hash_reader(&self) -> bool {
190        self.inner.is_hash_reader()
191    }
192
193    fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> {
194        self.inner.as_hash_reader_mut()
195    }
196}
197
198pin_project! {
199    /// A reader wrapper that decompresses data on the fly using DEFLATE algorithm.
200    /// Header format:
201    /// - First byte: compression type (00 = compressed, 01 = uncompressed, FF = end)
202    /// - Bytes 1-3: length of compressed data (little-endian)
203    /// - Bytes 4-7: CRC32 checksum of uncompressed data (little-endian)
204    #[derive(Debug)]
205    pub struct DecompressReader<R> {
206        #[pin]
207        pub inner: R,
208        buffer: Vec<u8>,
209        buffer_pos: usize,
210        finished: bool,
211        // Fields for saving header read progress across polls
212        header_buf: [u8; 8],
213        header_read: usize,
214        header_done: bool,
215        // Fields for saving compressed block read progress across polls
216        compressed_buf: Option<Vec<u8>>,
217        compressed_read: usize,
218        compressed_len: usize,
219        compression_algorithm: CompressionAlgorithm,
220    }
221}
222
223impl<R> DecompressReader<R>
224where
225    R: AsyncRead + Unpin + Send + Sync,
226{
227    pub fn new(inner: R, compression_algorithm: CompressionAlgorithm) -> Self {
228        Self {
229            inner,
230            buffer: Vec::new(),
231            buffer_pos: 0,
232            finished: false,
233            header_buf: [0u8; 8],
234            header_read: 0,
235            header_done: false,
236            compressed_buf: None,
237            compressed_read: 0,
238            compressed_len: 0,
239            compression_algorithm,
240        }
241    }
242}
243
244impl<R> AsyncRead for DecompressReader<R>
245where
246    R: AsyncRead + Unpin + Send + Sync,
247{
248    fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
249        let mut this = self.project();
250        // Copy from buffer first if available
251        if *this.buffer_pos < this.buffer.len() {
252            let to_copy = min(buf.remaining(), this.buffer.len() - *this.buffer_pos);
253            buf.put_slice(&this.buffer[*this.buffer_pos..*this.buffer_pos + to_copy]);
254            *this.buffer_pos += to_copy;
255            if *this.buffer_pos == this.buffer.len() {
256                this.buffer.clear();
257                *this.buffer_pos = 0;
258            }
259            return Poll::Ready(Ok(()));
260        }
261        if *this.finished {
262            return Poll::Ready(Ok(()));
263        }
264        // Read header
265        while !*this.header_done && *this.header_read < HEADER_LEN {
266            let mut temp = [0u8; HEADER_LEN];
267            let mut temp_buf = ReadBuf::new(&mut temp[0..HEADER_LEN - *this.header_read]);
268            match this.inner.as_mut().poll_read(cx, &mut temp_buf) {
269                Poll::Pending => return Poll::Pending,
270                Poll::Ready(Ok(())) => {
271                    let n = temp_buf.filled().len();
272                    if n == 0 {
273                        break;
274                    }
275                    this.header_buf[*this.header_read..*this.header_read + n].copy_from_slice(&temp_buf.filled()[..n]);
276                    *this.header_read += n;
277                }
278                Poll::Ready(Err(e)) => {
279                    // error!("DecompressReader poll_read: read header error: {e}");
280                    return Poll::Ready(Err(e));
281                }
282            }
283            if *this.header_read < HEADER_LEN {
284                return Poll::Pending;
285            }
286        }
287        if !*this.header_done && *this.header_read == 0 {
288            return Poll::Ready(Ok(()));
289        }
290        let typ = this.header_buf[0];
291        let len = (this.header_buf[1] as usize) | ((this.header_buf[2] as usize) << 8) | ((this.header_buf[3] as usize) << 16);
292        let crc = (this.header_buf[4] as u32)
293            | ((this.header_buf[5] as u32) << 8)
294            | ((this.header_buf[6] as u32) << 16)
295            | ((this.header_buf[7] as u32) << 24);
296        *this.header_read = 0;
297        *this.header_done = true;
298        if this.compressed_buf.is_none() {
299            *this.compressed_len = len;
300            *this.compressed_buf = Some(vec![0u8; *this.compressed_len]);
301            *this.compressed_read = 0;
302        }
303        let compressed_buf = this.compressed_buf.as_mut().unwrap();
304        while *this.compressed_read < *this.compressed_len {
305            let mut temp_buf = ReadBuf::new(&mut compressed_buf[*this.compressed_read..]);
306            match this.inner.as_mut().poll_read(cx, &mut temp_buf) {
307                Poll::Pending => return Poll::Pending,
308                Poll::Ready(Ok(())) => {
309                    let n = temp_buf.filled().len();
310                    if n == 0 {
311                        break;
312                    }
313                    *this.compressed_read += n;
314                }
315                Poll::Ready(Err(e)) => {
316                    // error!("DecompressReader poll_read: read compressed block error: {e}");
317                    this.compressed_buf.take();
318                    *this.compressed_read = 0;
319                    *this.compressed_len = 0;
320                    return Poll::Ready(Err(e));
321                }
322            }
323        }
324        let (uncompress_len, uvarint) = uvarint(&compressed_buf[0..16]);
325        let compressed_data = &compressed_buf[uvarint as usize..];
326        let decompressed = if typ == COMPRESS_TYPE_COMPRESSED {
327            match decompress_block(compressed_data, *this.compression_algorithm) {
328                Ok(out) => out,
329                Err(e) => {
330                    // error!("DecompressReader decompress_block error: {e}");
331                    this.compressed_buf.take();
332                    *this.compressed_read = 0;
333                    *this.compressed_len = 0;
334                    return Poll::Ready(Err(e));
335                }
336            }
337        } else if typ == COMPRESS_TYPE_UNCOMPRESSED {
338            compressed_data.to_vec()
339        } else if typ == COMPRESS_TYPE_END {
340            this.compressed_buf.take();
341            *this.compressed_read = 0;
342            *this.compressed_len = 0;
343            *this.finished = true;
344            return Poll::Ready(Ok(()));
345        } else {
346            // error!("DecompressReader unknown compression type: {typ}");
347            this.compressed_buf.take();
348            *this.compressed_read = 0;
349            *this.compressed_len = 0;
350            return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, "Unknown compression type")));
351        };
352        if decompressed.len() != uncompress_len as usize {
353            // error!("DecompressReader decompressed length mismatch: {} != {}", decompressed.len(), uncompress_len);
354            this.compressed_buf.take();
355            *this.compressed_read = 0;
356            *this.compressed_len = 0;
357            return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, "Decompressed length mismatch")));
358        }
359        let actual_crc = crc32fast::hash(&decompressed);
360        if actual_crc != crc {
361            // error!("DecompressReader CRC32 mismatch: actual {actual_crc} != expected {crc}");
362            this.compressed_buf.take();
363            *this.compressed_read = 0;
364            *this.compressed_len = 0;
365            return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, "CRC32 mismatch")));
366        }
367        *this.buffer = decompressed;
368        *this.buffer_pos = 0;
369        this.compressed_buf.take();
370        *this.compressed_read = 0;
371        *this.compressed_len = 0;
372        *this.header_done = false;
373        let to_copy = min(buf.remaining(), this.buffer.len());
374        buf.put_slice(&this.buffer[..to_copy]);
375        *this.buffer_pos += to_copy;
376        if *this.buffer_pos == this.buffer.len() {
377            this.buffer.clear();
378            *this.buffer_pos = 0;
379        }
380        Poll::Ready(Ok(()))
381    }
382}
383
384impl<R> EtagResolvable for DecompressReader<R>
385where
386    R: EtagResolvable,
387{
388    fn try_resolve_etag(&mut self) -> Option<String> {
389        self.inner.try_resolve_etag()
390    }
391}
392
393impl<R> HashReaderDetector for DecompressReader<R>
394where
395    R: HashReaderDetector,
396{
397    fn is_hash_reader(&self) -> bool {
398        self.inner.is_hash_reader()
399    }
400    fn as_hash_reader_mut(&mut self) -> Option<&mut dyn HashReaderMut> {
401        self.inner.as_hash_reader_mut()
402    }
403}
404
405/// Build compressed block with header + uvarint + compressed data
406fn build_compressed_block(uncompressed_data: &[u8], compression_algorithm: CompressionAlgorithm) -> Vec<u8> {
407    let crc = crc32fast::hash(uncompressed_data);
408    let compressed_data = compress_block(uncompressed_data, compression_algorithm);
409    let uncompressed_len = uncompressed_data.len();
410    let mut uncompressed_len_buf = [0u8; 10];
411    let int_len = put_uvarint(&mut uncompressed_len_buf[..], uncompressed_len as u64);
412    let len = compressed_data.len() + int_len;
413    let mut header = [0u8; HEADER_LEN];
414    header[0] = COMPRESS_TYPE_COMPRESSED;
415    header[1] = (len & 0xFF) as u8;
416    header[2] = ((len >> 8) & 0xFF) as u8;
417    header[3] = ((len >> 16) & 0xFF) as u8;
418    header[4] = (crc & 0xFF) as u8;
419    header[5] = ((crc >> 8) & 0xFF) as u8;
420    header[6] = ((crc >> 16) & 0xFF) as u8;
421    header[7] = ((crc >> 24) & 0xFF) as u8;
422    let mut out = Vec::with_capacity(len + HEADER_LEN);
423    out.extend_from_slice(&header);
424    out.extend_from_slice(&uncompressed_len_buf[..int_len]);
425    out.extend_from_slice(&compressed_data);
426    out
427}
428
429#[cfg(test)]
430mod tests {
431    use crate::WarpReader;
432
433    use super::*;
434    use std::io::Cursor;
435    use tokio::io::{AsyncReadExt, BufReader};
436
437    #[tokio::test]
438    async fn test_compress_reader_basic() {
439        let data = b"hello world, hello world, hello world!";
440        let reader = Cursor::new(&data[..]);
441        let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip);
442
443        let mut compressed = Vec::new();
444        compress_reader.read_to_end(&mut compressed).await.unwrap();
445
446        // DecompressReader解包
447        let mut decompress_reader = DecompressReader::new(Cursor::new(compressed.clone()), CompressionAlgorithm::Gzip);
448        let mut decompressed = Vec::new();
449        decompress_reader.read_to_end(&mut decompressed).await.unwrap();
450
451        assert_eq!(&decompressed, data);
452    }
453
454    #[tokio::test]
455    async fn test_compress_reader_basic_deflate() {
456        let data = b"hello world, hello world, hello world!";
457        let reader = BufReader::new(&data[..]);
458        let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Deflate);
459
460        let mut compressed = Vec::new();
461        compress_reader.read_to_end(&mut compressed).await.unwrap();
462
463        // DecompressReader解包
464        let mut decompress_reader = DecompressReader::new(Cursor::new(compressed.clone()), CompressionAlgorithm::Deflate);
465        let mut decompressed = Vec::new();
466        decompress_reader.read_to_end(&mut decompressed).await.unwrap();
467
468        assert_eq!(&decompressed, data);
469    }
470
471    #[tokio::test]
472    async fn test_compress_reader_empty() {
473        let data = b"";
474        let reader = BufReader::new(&data[..]);
475        let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip);
476
477        let mut compressed = Vec::new();
478        compress_reader.read_to_end(&mut compressed).await.unwrap();
479
480        let mut decompress_reader = DecompressReader::new(Cursor::new(compressed.clone()), CompressionAlgorithm::Gzip);
481        let mut decompressed = Vec::new();
482        decompress_reader.read_to_end(&mut decompressed).await.unwrap();
483
484        assert_eq!(&decompressed, data);
485    }
486
487    #[tokio::test]
488    async fn test_compress_reader_large() {
489        use rand::Rng;
490        // Generate 1MB of random bytes
491        let mut data = vec![0u8; 1024 * 1024 * 32];
492        rand::rng().fill(&mut data[..]);
493        let reader = Cursor::new(data.clone());
494        let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::Gzip);
495
496        let mut compressed = Vec::new();
497        compress_reader.read_to_end(&mut compressed).await.unwrap();
498
499        let mut decompress_reader = DecompressReader::new(Cursor::new(compressed.clone()), CompressionAlgorithm::Gzip);
500        let mut decompressed = Vec::new();
501        decompress_reader.read_to_end(&mut decompressed).await.unwrap();
502
503        assert_eq!(&decompressed, &data);
504    }
505
506    #[tokio::test]
507    async fn test_compress_reader_large_deflate() {
508        use rand::Rng;
509        // Generate 1MB of random bytes
510        let mut data = vec![0u8; 1024 * 1024 * 3 + 512];
511        rand::rng().fill(&mut data[..]);
512        let reader = Cursor::new(data.clone());
513        let mut compress_reader = CompressReader::new(WarpReader::new(reader), CompressionAlgorithm::default());
514
515        let mut compressed = Vec::new();
516        compress_reader.read_to_end(&mut compressed).await.unwrap();
517
518        let mut decompress_reader = DecompressReader::new(Cursor::new(compressed.clone()), CompressionAlgorithm::default());
519        let mut decompressed = Vec::new();
520        decompress_reader.read_to_end(&mut decompressed).await.unwrap();
521
522        assert_eq!(&decompressed, &data);
523    }
524}