photon_protocol/compressor/
zstd.rs1use 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}