Skip to main content

photon_protocol/compressor/
zstd.rs

1use std::fmt;
2
3#[cfg(not(target_arch = "wasm32"))]
4mod native {
5    use std::sync::Mutex;
6
7    use bytes::BytesMut;
8    use zstd::bulk::{Compressor as Encoder, Decompressor as Decoder};
9
10    use crate::ports::compress::{CompressionError, Compressor};
11
12    pub struct ZstdCompressor {
13        pub level: i32,
14        compressor: Mutex<Encoder<'static>>,
15        decompressor: Mutex<Decoder<'static>>,
16    }
17
18    impl ZstdCompressor {
19        pub const NAME: &str = "zstd";
20
21        pub fn new(level: i32) -> Self {
22            Self {
23                level,
24                compressor: Mutex::new(
25                    Encoder::new(level).expect("failed to create zstd compressor"),
26                ),
27                decompressor: Mutex::new(
28                    Decoder::new().expect("failed to create zstd decompressor"),
29                ),
30            }
31        }
32    }
33
34    impl Compressor for ZstdCompressor {
35        fn compress(&self, input: &[u8], output: &mut BytesMut) -> Result<(), CompressionError> {
36            let compressed = self
37                .compressor
38                .lock()
39                .unwrap()
40                .compress(input)
41                .map_err(|e| CompressionError::Unknown(e.into()))?;
42
43            output.extend_from_slice(&compressed);
44            Ok(())
45        }
46
47        fn decompress(&self, input: &[u8], output: &mut BytesMut) -> Result<(), CompressionError> {
48            let capacity = zstd::zstd_safe::get_frame_content_size(input)
49                .ok()
50                .flatten()
51                .map(|s| s as usize)
52                .unwrap_or(input.len() * 10);
53
54            let decompressed = self
55                .decompressor
56                .lock()
57                .unwrap()
58                .decompress(input, capacity)
59                .map_err(|_| CompressionError::CorruptPayload {
60                    compressor_name: self.name().to_owned(),
61                })?;
62
63            output.extend_from_slice(&decompressed);
64            Ok(())
65        }
66
67        fn name(&self) -> &'static str {
68            Self::NAME
69        }
70    }
71}
72
73#[cfg(target_arch = "wasm32")]
74mod wasm {
75    use std::io::Read;
76
77    use bytes::BytesMut;
78    use ruzstd::decoding::StreamingDecoder;
79    use ruzstd::encoding::{CompressionLevel, compress_to_vec};
80
81    use crate::ports::compress::{CompressionError, Compressor};
82
83    pub struct ZstdCompressor {
84        pub level: i32,
85    }
86
87    impl ZstdCompressor {
88        pub const NAME: &str = "zstd";
89
90        pub fn new(level: i32) -> Self {
91            Self { level }
92        }
93    }
94
95    impl Compressor for ZstdCompressor {
96        fn compress(&self, input: &[u8], output: &mut BytesMut) -> Result<(), CompressionError> {
97            let level = match self.level {
98                0 => CompressionLevel::Uncompressed,
99                1..=3 => CompressionLevel::Fastest,
100                4..=6 => CompressionLevel::Default,
101                7..=9 => CompressionLevel::Better,
102                _ => CompressionLevel::Best,
103            };
104
105            output.extend_from_slice(&compress_to_vec(input, level));
106            Ok(())
107        }
108
109        fn decompress(&self, input: &[u8], output: &mut BytesMut) -> Result<(), CompressionError> {
110            let mut decoder = StreamingDecoder::new(input)
111                .map_err(|e| CompressionError::Unknown(anyhow::anyhow!("zstd frame init: {e}")))?;
112
113            let mut buf = Vec::new();
114            decoder
115                .read_to_end(&mut buf)
116                .map_err(|e| CompressionError::Unknown(anyhow::anyhow!("zstd decompress: {e}")))?;
117
118            output.extend_from_slice(&buf);
119            Ok(())
120        }
121
122        fn name(&self) -> &'static str {
123            Self::NAME
124        }
125    }
126}
127
128#[cfg(not(target_arch = "wasm32"))]
129pub use native::ZstdCompressor;
130#[cfg(target_arch = "wasm32")]
131pub use wasm::ZstdCompressor;
132
133impl Default for ZstdCompressor {
134    fn default() -> Self {
135        Self::new(3)
136    }
137}
138
139impl Clone for ZstdCompressor {
140    fn clone(&self) -> Self {
141        Self::new(self.level)
142    }
143}
144
145impl fmt::Debug for ZstdCompressor {
146    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
147        f.debug_struct("ZstdCompressor")
148            .field("level", &self.level)
149            .finish()
150    }
151}