Skip to main content

sz_orm_websocket/
compression.rs

1//! # 消息压缩(RFC 7692 permessage-deflate)
2//!
3//! 实现 RFC 7692 `permessage-deflate` 扩展的协商与压缩/解压缩流程。
4//! 基于 `flate2` 的 zlib DEFLATE 算法,与浏览器/服务端互通:
5//!
6//! - 协商阶段:解析客户端 `Sec-WebSocket-Extensions` 头并匹配服务端支持的参数
7//! - 压缩阶段:使用 zlib DEFLATE 压缩消息,按 RFC 7692 移除尾部 4 字节同步标记
8//!
9//! ## 主要类型
10//!
11//! - [`CompressionConfig`] — 压缩配置
12//! - [`CompressionNegotiator`] — 扩展协商器
13//! - [`CompressionStats`] — 压缩统计
14//! - [`MessageCompressor`] — 消息压缩器
15
16use std::collections::HashMap;
17use std::sync::Arc;
18use tokio::sync::RwLock;
19
20/// 压缩配置,对应 permessage-deflate 扩展参数
21#[derive(Debug, Clone)]
22pub struct CompressionConfig {
23    /// 服务端窗口大小指数(4-15,对应 2^server_no_context_takeover_bits)
24    pub server_max_window_bits: u8,
25    /// 客户端窗口大小指数(4-15)
26    pub client_max_window_bits: u8,
27    /// 服务端不保留上下文(每条消息独立压缩)
28    pub server_no_context_takeover: bool,
29    /// 客户端不保留上下文
30    pub client_no_context_takeover: bool,
31}
32
33impl Default for CompressionConfig {
34    fn default() -> Self {
35        Self {
36            server_max_window_bits: 15,
37            client_max_window_bits: 15,
38            server_no_context_takeover: false,
39            client_no_context_takeover: false,
40        }
41    }
42}
43
44impl CompressionConfig {
45    /// 创建默认配置
46    pub fn new() -> Self {
47        Self::default()
48    }
49
50    /// 设置服务端窗口大小
51    pub fn with_server_window_bits(mut self, bits: u8) -> Self {
52        self.server_max_window_bits = bits.clamp(4, 15);
53        self
54    }
55
56    /// 设置客户端窗口大小
57    pub fn with_client_window_bits(mut self, bits: u8) -> Self {
58        self.client_max_window_bits = bits.clamp(4, 15);
59        self
60    }
61
62    /// 启用服务端无上下文接管
63    pub fn with_server_no_context_takeover(mut self) -> Self {
64        self.server_no_context_takeover = true;
65        self
66    }
67
68    /// 启用客户端无上下文接管
69    pub fn with_client_no_context_takeover(mut self) -> Self {
70        self.client_no_context_takeover = true;
71        self
72    }
73
74    /// 校验配置合法性
75    pub fn validate(&self) -> Result<(), String> {
76        if !(4..=15).contains(&self.server_max_window_bits) {
77            return Err("server_max_window_bits must be in [4, 15]".to_string());
78        }
79        if !(4..=15).contains(&self.client_max_window_bits) {
80            return Err("client_max_window_bits must be in [4, 15]".to_string());
81        }
82        Ok(())
83    }
84
85    /// 将配置渲染为 permessage-deflate 扩展参数字符串
86    pub fn to_extension_params(&self) -> String {
87        let mut parts = vec!["permessage-deflate".to_string()];
88        parts.push(format!(
89            "server_max_window_bits={}",
90            self.server_max_window_bits
91        ));
92        if self.server_no_context_takeover {
93            parts.push("server_no_context_takeover".to_string());
94        }
95        if self.client_no_context_takeover {
96            parts.push("client_no_context_takeover".to_string());
97        }
98        parts.push(format!(
99            "client_max_window_bits={}",
100            self.client_max_window_bits
101        ));
102        parts.join("; ")
103    }
104}
105
106/// 客户端提供的扩展参数
107#[derive(Debug, Clone, Default)]
108pub struct ClientExtensions {
109    /// 原始扩展头值
110    pub raw: String,
111    /// 解析出的参数键值对
112    pub params: HashMap<String, Option<String>>,
113}
114
115impl ClientExtensions {
116    /// 从 `Sec-WebSocket-Extensions` 头值解析
117    pub fn parse(header_value: &str) -> Self {
118        let mut params = HashMap::new();
119        for part in header_value.split(';') {
120            let part = part.trim();
121            if part.is_empty() {
122                continue;
123            }
124            if let Some((key, value)) = part.split_once('=') {
125                let key = key.trim().to_string();
126                let value = value.trim().trim_matches('"').to_string();
127                params.insert(key, Some(value));
128            } else {
129                params.insert(part.to_string(), None);
130            }
131        }
132        Self {
133            raw: header_value.to_string(),
134            params,
135        }
136    }
137
138    /// 是否请求了 permessage-deflate
139    pub fn requests_deflate(&self) -> bool {
140        self.raw.contains("permessage-deflate")
141    }
142
143    /// 获取参数值
144    pub fn get_param(&self, key: &str) -> Option<&str> {
145        self.params.get(key).and_then(|v| v.as_deref())
146    }
147
148    /// 是否包含某参数(无论有无值)
149    pub fn has_param(&self, key: &str) -> bool {
150        self.params.contains_key(key)
151    }
152}
153
154/// 扩展协商结果
155#[derive(Debug, Clone, PartialEq, Eq)]
156pub enum NegotiationResult {
157    /// 协商成功,返回服务端选定的扩展头值
158    Accepted(String),
159    /// 客户端未请求压缩,服务端不启用
160    NotRequested,
161    /// 客户端请求的参数不合法,拒绝压缩
162    Rejected(String),
163}
164
165/// 扩展协商器
166#[derive(Debug, Clone)]
167pub struct CompressionNegotiator {
168    config: CompressionConfig,
169}
170
171impl CompressionNegotiator {
172    /// 创建协商器
173    pub fn new(config: CompressionConfig) -> Self {
174        Self { config }
175    }
176
177    /// 获取配置
178    pub fn config(&self) -> &CompressionConfig {
179        &self.config
180    }
181
182    /// 根据客户端请求协商压缩参数
183    pub fn negotiate(&self, client_header: &str) -> NegotiationResult {
184        if client_header.is_empty() {
185            return NegotiationResult::NotRequested;
186        }
187
188        let client_ext = ClientExtensions::parse(client_header);
189        if !client_ext.requests_deflate() {
190            return NegotiationResult::NotRequested;
191        }
192
193        // 校验客户端请求的窗口大小是否在合法范围
194        if let Some(bits_str) = client_ext.get_param("server_max_window_bits") {
195            if let Ok(bits) = bits_str.parse::<u8>() {
196                if !(4..=15).contains(&bits) {
197                    return NegotiationResult::Rejected(format!(
198                        "invalid server_max_window_bits: {}",
199                        bits
200                    ));
201                }
202            }
203        }
204
205        // 服务端选定最终参数(取客户端与服务端的较小值)
206        let mut response_parts = vec!["permessage-deflate".to_string()];
207
208        // server_max_window_bits:若客户端指定了,取 min(客户端, 服务端)
209        let final_server_bits =
210            if let Some(bits_str) = client_ext.get_param("server_max_window_bits") {
211                if let Ok(client_bits) = bits_str.parse::<u8>() {
212                    client_bits.min(self.config.server_max_window_bits)
213                } else {
214                    self.config.server_max_window_bits
215                }
216            } else {
217                self.config.server_max_window_bits
218            };
219        response_parts.push(format!("server_max_window_bits={}", final_server_bits));
220
221        if self.config.server_no_context_takeover
222            || client_ext.has_param("server_no_context_takeover")
223        {
224            response_parts.push("server_no_context_takeover".to_string());
225        }
226
227        // client_max_window_bits:仅当客户端提供了值时才回应
228        if let Some(bits_str) = client_ext.get_param("client_max_window_bits") {
229            if let Ok(client_bits) = bits_str.parse::<u8>() {
230                let final_client_bits = client_bits.min(self.config.client_max_window_bits);
231                response_parts.push(format!("client_max_window_bits={}", final_client_bits));
232            }
233        } else if client_ext.has_param("client_max_window_bits") {
234            // 客户端声明支持但未指定值
235            response_parts.push(format!(
236                "client_max_window_bits={}",
237                self.config.client_max_window_bits
238            ));
239        }
240
241        if self.config.client_no_context_takeover
242            || client_ext.has_param("client_no_context_takeover")
243        {
244            response_parts.push("client_no_context_takeover".to_string());
245        }
246
247        NegotiationResult::Accepted(response_parts.join("; "))
248    }
249}
250
251/// 压缩统计
252#[derive(Debug, Clone, Default)]
253pub struct CompressionStats {
254    /// 压缩前的总字节数
255    pub total_uncompressed: u64,
256    /// 压缩后的总字节数
257    pub total_compressed: u64,
258    /// 已压缩的消息数
259    pub messages_compressed: u64,
260    /// 已解压的消息数
261    pub messages_decompressed: u64,
262}
263
264impl CompressionStats {
265    /// 平均压缩比(compressed / uncompressed,越小越好)
266    pub fn ratio(&self) -> f64 {
267        if self.total_uncompressed == 0 {
268            return 1.0;
269        }
270        self.total_compressed as f64 / self.total_uncompressed as f64
271    }
272
273    /// 总节省字节数
274    pub fn bytes_saved(&self) -> i64 {
275        self.total_uncompressed as i64 - self.total_compressed as i64
276    }
277
278    /// 节省百分比(0.0-100.0)
279    pub fn saved_percent(&self) -> f64 {
280        if self.total_uncompressed == 0 {
281            return 0.0;
282        }
283        let saved = self.total_uncompressed - self.total_compressed;
284        (saved as f64 / self.total_uncompressed as f64) * 100.0
285    }
286
287    /// 重置统计
288    pub fn reset(&mut self) {
289        *self = Self::default();
290    }
291}
292
293/// 使用 zlib DEFLATE 压缩字节序列,并按 RFC 7692 处理尾部同步标记。
294///
295/// RFC 7692 规定:1. 使用 zlib 压缩;2. 移除尾部的 4 字节 `00 00 FF FF` 同步标记;
296/// 3. 若结果以 0x00 结尾且位于 WebSocket frame 的最后,需补一个 0x00 防止 frame 边界冲突。
297///
298/// 返回值始终非空(空输入也会产生至少 1 字节的 zlib header 信息)。
299fn deflate_compress(data: &[u8]) -> Vec<u8> {
300    use flate2::write::ZlibEncoder;
301    use std::io::Write;
302    let mut encoder = ZlibEncoder::new(Vec::new(), flate2::Compression::default());
303    // 写入数据;空数据也能压缩(zlib header + empty block + trailer)
304    let _ = encoder.write_all(data);
305    let mut compressed = encoder
306        .finish()
307        .expect("zlib compression of in-memory buffer cannot fail");
308
309    // RFC 7692: 移除尾部 4 字节 00 00 FF FF 同步标记
310    if compressed.len() >= 4 && compressed[compressed.len() - 4..] == [0x00, 0x00, 0xFF, 0xFF] {
311        compressed.truncate(compressed.len() - 4);
312    }
313
314    // 若压缩后末尾为 0x00,补一个 0x00 防止与 frame 边界混淆
315    if compressed.last() == Some(&0x00) {
316        compressed.push(0x00);
317    }
318
319    compressed
320}
321
322/// 解压 RFC 7692 permessage-deflate 数据。
323///
324/// 自动补回 4 字节 `00 00 FF FF` 同步标记后调用 zlib 解压。
325/// 同时处理压缩端为防 frame 边界冲突而追加的尾部 0x00(合法的 deflate 流不会以此结尾)。
326fn deflate_decompress(data: &[u8]) -> Result<Vec<u8>, String> {
327    use flate2::read::ZlibDecoder;
328    use std::io::Read;
329
330    if data.is_empty() {
331        return Err("empty compressed data".to_string());
332    }
333
334    // 补回 RFC 7692 移除的 4 字节同步标记
335    let mut full = data.to_vec();
336    full.extend_from_slice(&[0x00, 0x00, 0xFF, 0xFF]);
337
338    let mut decoder = ZlibDecoder::new(&full[..]);
339    let mut out = Vec::new();
340    decoder
341        .read_to_end(&mut out)
342        .map_err(|e| format!("zlib decompress failed: {}", e))?;
343    Ok(out)
344}
345
346/// 消息压缩器,跟踪压缩统计
347#[derive(Debug)]
348pub struct MessageCompressor {
349    config: CompressionConfig,
350    stats: Arc<RwLock<CompressionStats>>,
351}
352
353impl MessageCompressor {
354    /// 创建压缩器
355    pub fn new(config: CompressionConfig) -> Self {
356        Self {
357            config,
358            stats: Arc::new(RwLock::new(CompressionStats::default())),
359        }
360    }
361
362    /// 获取配置
363    pub fn config(&self) -> &CompressionConfig {
364        &self.config
365    }
366
367    /// 压缩消息。对于小消息(<32 字节)不压缩直接返回原文。
368    ///
369    /// 格式:1 字节标记 + 数据
370    /// - 标记 0x00:原始数据(未压缩,后跟原文)
371    /// - 标记 0x01:DEFLATE 压缩数据(后跟 RFC 7692 zlib 流)
372    pub async fn compress(&self, data: &[u8]) -> Vec<u8> {
373        let uncompressed_size = data.len() as u64;
374
375        // 小消息不压缩(DEFLATE 对小数据通常反向膨胀)
376        let compressed = if data.len() < 32 {
377            let mut out = Vec::with_capacity(data.len() + 1);
378            out.push(0x00); // 原始数据标记
379            out.extend_from_slice(data);
380            out
381        } else {
382            let deflated = deflate_compress(data);
383            // 仅当压缩后更小才使用压缩结果
384            if deflated.len() + 1 < data.len() + 1 {
385                let mut out = Vec::with_capacity(deflated.len() + 1);
386                out.push(0x01); // 压缩数据标记
387                out.extend_from_slice(&deflated);
388                out
389            } else {
390                let mut out = Vec::with_capacity(data.len() + 1);
391                out.push(0x00); // 原始数据标记
392                out.extend_from_slice(data);
393                out
394            }
395        };
396
397        let compressed_size = compressed.len() as u64;
398        let mut stats = self.stats.write().await;
399        stats.total_uncompressed += uncompressed_size;
400        stats.total_compressed += compressed_size;
401        stats.messages_compressed += 1;
402
403        compressed
404    }
405
406    /// 解压消息
407    ///
408    /// 根据首字节标记决定解压路径:
409    /// - 0x00:原始数据,跳过标记后直接返回
410    /// - 0x01:DEFLATE 压缩数据,调用 zlib 解压
411    pub async fn decompress(&self, data: &[u8]) -> Result<Vec<u8>, String> {
412        if data.is_empty() {
413            return Err("empty compressed data".to_string());
414        }
415
416        let result = match data[0] {
417            0x00 => data[1..].to_vec(),
418            0x01 => deflate_decompress(&data[1..])?,
419            other => return Err(format!("unknown compression marker: 0x{:02X}", other)),
420        };
421
422        let mut stats = self.stats.write().await;
423        stats.messages_decompressed += 1;
424
425        Ok(result)
426    }
427
428    /// 获取当前压缩统计快照
429    pub async fn stats(&self) -> CompressionStats {
430        self.stats.read().await.clone()
431    }
432
433    /// 重置统计
434    pub async fn reset_stats(&self) {
435        let mut stats = self.stats.write().await;
436        stats.reset();
437    }
438}
439
440#[cfg(test)]
441mod tests {
442    use super::*;
443
444    #[test]
445    fn test_compression_config_default() {
446        let cfg = CompressionConfig::default();
447        assert_eq!(cfg.server_max_window_bits, 15);
448        assert_eq!(cfg.client_max_window_bits, 15);
449        assert!(!cfg.server_no_context_takeover);
450        assert!(!cfg.client_no_context_takeover);
451    }
452
453    #[test]
454    fn test_compression_config_builder() {
455        let cfg = CompressionConfig::new()
456            .with_server_window_bits(10)
457            .with_client_window_bits(12)
458            .with_server_no_context_takeover()
459            .with_client_no_context_takeover();
460        assert_eq!(cfg.server_max_window_bits, 10);
461        assert_eq!(cfg.client_max_window_bits, 12);
462        assert!(cfg.server_no_context_takeover);
463        assert!(cfg.client_no_context_takeover);
464    }
465
466    #[test]
467    fn test_compression_config_clamp_window_bits() {
468        let cfg = CompressionConfig::new()
469            .with_server_window_bits(2)
470            .with_client_window_bits(20);
471        assert_eq!(cfg.server_max_window_bits, 4);
472        assert_eq!(cfg.client_max_window_bits, 15);
473    }
474
475    #[test]
476    fn test_compression_config_validate_ok() {
477        let cfg = CompressionConfig::new();
478        assert!(cfg.validate().is_ok());
479        assert_eq!(
480            cfg.server_max_window_bits, 15,
481            "默认 server_max_window_bits 应为 15"
482        );
483        assert_eq!(
484            cfg.client_max_window_bits, 15,
485            "默认 client_max_window_bits 应为 15"
486        );
487        assert!(
488            !cfg.server_no_context_takeover,
489            "默认 server_no_context_takeover 应为 false"
490        );
491        assert!(
492            !cfg.client_no_context_takeover,
493            "默认 client_no_context_takeover 应为 false"
494        );
495    }
496
497    #[test]
498    fn test_compression_config_validate_invalid_server_bits() {
499        let cfg = CompressionConfig {
500            server_max_window_bits: 3,
501            client_max_window_bits: 10,
502            server_no_context_takeover: false,
503            client_no_context_takeover: false,
504        };
505        assert!(cfg.validate().is_err());
506    }
507
508    #[test]
509    fn test_compression_config_validate_invalid_client_bits() {
510        let cfg = CompressionConfig {
511            server_max_window_bits: 10,
512            client_max_window_bits: 16,
513            server_no_context_takeover: false,
514            client_no_context_takeover: false,
515        };
516        assert!(cfg.validate().is_err());
517    }
518
519    #[test]
520    fn test_to_extension_params() {
521        let cfg = CompressionConfig::new()
522            .with_server_window_bits(10)
523            .with_server_no_context_takeover();
524        let params = cfg.to_extension_params();
525        assert!(params.starts_with("permessage-deflate"));
526        assert!(params.contains("server_max_window_bits=10"));
527        assert!(params.contains("server_no_context_takeover"));
528    }
529
530    #[test]
531    fn test_client_extensions_parse_empty() {
532        let ext = ClientExtensions::parse("");
533        assert!(!ext.requests_deflate());
534        assert!(ext.params.is_empty());
535    }
536
537    #[test]
538    fn test_client_extensions_parse_with_values() {
539        let ext = ClientExtensions::parse(
540            "permessage-deflate; server_max_window_bits=10; client_max_window_bits",
541        );
542        assert!(ext.requests_deflate());
543        assert_eq!(ext.get_param("server_max_window_bits"), Some("10"));
544        assert!(ext.has_param("client_max_window_bits"));
545        assert_eq!(ext.get_param("client_max_window_bits"), None);
546    }
547
548    #[test]
549    fn test_client_extensions_parse_quoted_values() {
550        let ext = ClientExtensions::parse("permessage-deflate; param=\"value\"");
551        assert_eq!(ext.get_param("param"), Some("value"));
552    }
553
554    #[test]
555    fn test_negotiator_not_requested_when_empty() {
556        let neg = CompressionNegotiator::new(CompressionConfig::default());
557        let result = neg.negotiate("");
558        assert_eq!(result, NegotiationResult::NotRequested);
559    }
560
561    #[test]
562    fn test_negotiator_not_requested_when_no_deflate() {
563        let neg = CompressionNegotiator::new(CompressionConfig::default());
564        let result = neg.negotiate("other-extension");
565        assert_eq!(result, NegotiationResult::NotRequested);
566    }
567
568    #[test]
569    fn test_negotiator_accepted_basic() {
570        let neg = CompressionNegotiator::new(CompressionConfig::default());
571        let result = neg.negotiate("permessage-deflate");
572        match result {
573            NegotiationResult::Accepted(resp) => {
574                assert!(resp.contains("permessage-deflate"));
575                assert!(resp.contains("server_max_window_bits=15"));
576            }
577            _ => panic!("expected Accepted, got {:?}", result),
578        }
579    }
580
581    #[test]
582    fn test_negotiator_accepted_with_client_window_bits() {
583        let neg = CompressionNegotiator::new(CompressionConfig::default());
584        let result = neg
585            .negotiate("permessage-deflate; server_max_window_bits=10; client_max_window_bits=12");
586        match result {
587            NegotiationResult::Accepted(resp) => {
588                assert!(resp.contains("server_max_window_bits=10"));
589                assert!(resp.contains("client_max_window_bits=12"));
590            }
591            _ => panic!("expected Accepted, got {:?}", result),
592        }
593    }
594
595    #[test]
596    fn test_negotiator_takes_min_window_bits() {
597        let neg = CompressionNegotiator::new(CompressionConfig::new().with_server_window_bits(8));
598        // 客户端请求 server_max_window_bits=10,服务端配置为 8,应取 min=8
599        let result = neg.negotiate("permessage-deflate; server_max_window_bits=10");
600        match result {
601            NegotiationResult::Accepted(resp) => {
602                assert!(resp.contains("server_max_window_bits=8"));
603            }
604            _ => panic!("expected Accepted, got {:?}", result),
605        }
606    }
607
608    #[test]
609    fn test_negotiator_rejected_invalid_window_bits() {
610        let neg = CompressionNegotiator::new(CompressionConfig::default());
611        let result = neg.negotiate("permessage-deflate; server_max_window_bits=99");
612        assert!(matches!(result, NegotiationResult::Rejected(_)));
613    }
614
615    #[test]
616    fn test_negotiator_no_context_takeover_propagated() {
617        let neg =
618            CompressionNegotiator::new(CompressionConfig::new().with_server_no_context_takeover());
619        let result = neg.negotiate("permessage-deflate");
620        match result {
621            NegotiationResult::Accepted(resp) => {
622                assert!(resp.contains("server_no_context_takeover"));
623            }
624            _ => panic!("expected Accepted, got {:?}", result),
625        }
626    }
627
628    #[test]
629    fn test_negotiator_no_context_takeover_from_client() {
630        let neg = CompressionNegotiator::new(CompressionConfig::default());
631        let result = neg.negotiate(
632            "permessage-deflate; server_no_context_takeover; client_no_context_takeover",
633        );
634        match result {
635            NegotiationResult::Accepted(resp) => {
636                assert!(resp.contains("server_no_context_takeover"));
637                assert!(resp.contains("client_no_context_takeover"));
638            }
639            _ => panic!("expected Accepted, got {:?}", result),
640        }
641    }
642
643    #[test]
644    fn test_negotiator_client_window_bits_without_value() {
645        let neg = CompressionNegotiator::new(CompressionConfig::default());
646        let result = neg.negotiate("permessage-deflate; client_max_window_bits");
647        match result {
648            NegotiationResult::Accepted(resp) => {
649                assert!(resp.contains("client_max_window_bits=15"));
650            }
651            _ => panic!("expected Accepted, got {:?}", result),
652        }
653    }
654
655    #[test]
656    fn test_compression_stats_default() {
657        let stats = CompressionStats::default();
658        assert_eq!(stats.total_uncompressed, 0);
659        assert_eq!(stats.total_compressed, 0);
660        assert_eq!(stats.messages_compressed, 0);
661        assert_eq!(stats.messages_decompressed, 0);
662        assert_eq!(stats.ratio(), 1.0);
663        assert_eq!(stats.bytes_saved(), 0);
664        assert_eq!(stats.saved_percent(), 0.0);
665    }
666
667    #[test]
668    fn test_compression_stats_ratio() {
669        let stats = CompressionStats {
670            total_uncompressed: 1000,
671            total_compressed: 400,
672            messages_compressed: 5,
673            messages_decompressed: 0,
674        };
675        assert!((stats.ratio() - 0.4).abs() < 1e-9);
676        assert_eq!(stats.bytes_saved(), 600);
677        assert!((stats.saved_percent() - 60.0).abs() < 1e-9);
678    }
679
680    #[test]
681    fn test_compression_stats_zero_uncompressed() {
682        let stats = CompressionStats {
683            total_uncompressed: 0,
684            total_compressed: 100,
685            messages_compressed: 1,
686            messages_decompressed: 0,
687        };
688        assert_eq!(stats.ratio(), 1.0);
689        assert_eq!(stats.saved_percent(), 0.0);
690    }
691
692    #[test]
693    fn test_compression_stats_reset() {
694        let mut stats = CompressionStats {
695            total_uncompressed: 1000,
696            total_compressed: 400,
697            messages_compressed: 5,
698            messages_decompressed: 3,
699        };
700        stats.reset();
701        assert_eq!(stats.total_uncompressed, 0);
702        assert_eq!(stats.messages_compressed, 0);
703    }
704
705    #[test]
706    fn test_deflate_compress_nonempty() {
707        // 空输入也会产生 zlib header + 同步标记(移除 4 字节后仍非空)
708        let compressed = deflate_compress(b"");
709        assert!(!compressed.is_empty());
710    }
711
712    #[test]
713    fn test_deflate_compress_repeated_bytes() {
714        let data = vec![b'a'; 1000];
715        let compressed = deflate_compress(&data);
716        // DEFLATE 对高度重复数据压缩率极高(应远小于 1000)
717        assert!(compressed.len() < data.len());
718        assert!(compressed.len() < 100);
719    }
720
721    #[test]
722    fn test_deflate_decompress_roundtrip() {
723        let original = b"Hello, permessage-deflate! This is a test message.".repeat(20);
724        let compressed = deflate_compress(&original);
725        let decompressed = deflate_decompress(&compressed).unwrap();
726        assert_eq!(decompressed, original);
727    }
728
729    #[test]
730    fn test_deflate_decompress_empty() {
731        let result = deflate_decompress(b"");
732        assert!(result.is_err());
733    }
734
735    #[test]
736    fn test_deflate_decompress_invalid_data() {
737        // 不是合法 zlib 流
738        let result = deflate_decompress(b"\x00\x01\x02\x03");
739        assert!(result.is_err());
740    }
741
742    #[test]
743    fn test_deflate_roundtrip_random() {
744        let original: Vec<u8> = (0..200u8).collect();
745        let compressed = deflate_compress(&original);
746        let decompressed = deflate_decompress(&compressed).unwrap();
747        assert_eq!(decompressed, original);
748    }
749
750    #[tokio::test]
751    async fn test_compressor_compress_small_message() {
752        let comp = MessageCompressor::new(CompressionConfig::default());
753        // 小消息不压缩,标记为 0x00 + 原文
754        let result = comp.compress(b"hi").await;
755        assert_eq!(result, vec![0x00, b'h', b'i']);
756
757        let stats = comp.stats().await;
758        assert_eq!(stats.messages_compressed, 1);
759        assert_eq!(stats.total_uncompressed, 2);
760        assert_eq!(stats.total_compressed, 3); // 1 字节标记 + 2 字节原文
761    }
762
763    #[tokio::test]
764    async fn test_compressor_compress_large_repeated() {
765        let comp = MessageCompressor::new(CompressionConfig::default());
766        let data = vec![b'a'; 1000]; // 高度重复的数据
767        let result = comp.compress(&data).await;
768        // 应该被压缩(DEFLATE 对重复数据非常高效)
769        assert!(result.len() < data.len());
770        assert_eq!(result[0], 0x01); // 首字节为压缩标记
771
772        let stats = comp.stats().await;
773        assert_eq!(stats.total_uncompressed, 1000);
774        assert!(stats.total_compressed < 1000);
775        assert!(stats.saved_percent() > 50.0);
776    }
777
778    #[tokio::test]
779    async fn test_compressor_compress_incompressible() {
780        let comp = MessageCompressor::new(CompressionConfig::default());
781        // 32 字节顺序递增数据,DEFLATE 可能无法压缩
782        let data: Vec<u8> = (0..32u8).collect();
783        let result = comp.compress(&data).await;
784        // 首字节应为 0x00(原始)或 0x01(压缩);总长度不应远超原文+1
785        assert!(result[0] == 0x00 || result[0] == 0x01);
786        assert!(result.len() <= data.len() + 50); // 允许小幅膨胀但仍合理
787
788        let stats = comp.stats().await;
789        assert_eq!(stats.total_uncompressed, 32);
790    }
791
792    #[tokio::test]
793    async fn test_compressor_decompress_compressed() {
794        let comp = MessageCompressor::new(CompressionConfig::default());
795        let original = vec![b'x'; 100];
796        let compressed = comp.compress(&original).await;
797        let decompressed = comp.decompress(&compressed).await.unwrap();
798        assert_eq!(decompressed, original);
799    }
800
801    #[tokio::test]
802    async fn test_compressor_decompress_uncompressed() {
803        let comp = MessageCompressor::new(CompressionConfig::default());
804        // 小消息不压缩,标记 0x00 + 原文
805        let data = vec![0x00, b'h', b'e', b'l', b'l', b'o'];
806        let decompressed = comp.decompress(&data).await.unwrap();
807        assert_eq!(decompressed, b"hello");
808    }
809
810    #[tokio::test]
811    async fn test_compressor_decompress_invalid_marker() {
812        let comp = MessageCompressor::new(CompressionConfig::default());
813        let data = vec![0x99, 0x01, 0x02]; // 未知标记
814        let result = comp.decompress(&data).await;
815        assert!(result.is_err());
816    }
817
818    #[tokio::test]
819    async fn test_compressor_stats_accumulate() {
820        let comp = MessageCompressor::new(CompressionConfig::default());
821        let data = vec![b'a'; 100];
822        comp.compress(&data).await;
823        comp.compress(&data).await;
824        comp.compress(&data).await;
825
826        let stats = comp.stats().await;
827        assert_eq!(stats.messages_compressed, 3);
828        assert_eq!(stats.total_uncompressed, 300);
829        assert!(stats.total_compressed < 300);
830    }
831
832    #[tokio::test]
833    async fn test_compressor_reset_stats() {
834        let comp = MessageCompressor::new(CompressionConfig::default());
835        let data = vec![b'a'; 100];
836        comp.compress(&data).await;
837        assert!(comp.stats().await.messages_compressed > 0);
838
839        comp.reset_stats().await;
840        let stats = comp.stats().await;
841        assert_eq!(stats.messages_compressed, 0);
842        assert_eq!(stats.total_uncompressed, 0);
843    }
844
845    #[tokio::test]
846    async fn test_compressor_decompress_count() {
847        let comp = MessageCompressor::new(CompressionConfig::default());
848        // 压缩格式要求首字节为标记(0x00 原始 / 0x01 DEFLATE),不能直接传原始字符串
849        let compressed1 = comp.compress(b"hi").await;
850        let compressed2 = comp.compress(b"world").await;
851        comp.decompress(&compressed1).await.unwrap();
852        comp.decompress(&compressed2).await.unwrap();
853
854        let stats = comp.stats().await;
855        assert_eq!(stats.messages_decompressed, 2);
856    }
857
858    #[tokio::test]
859    async fn test_compressor_decompress_invalid_returns_error() {
860        let comp = MessageCompressor::new(CompressionConfig::default());
861        // 0xFF 标记字节后跟单个数据字节但缺少计数字节 —— 截断的压缩数据
862        let result = comp.decompress(&[0xFF, b'a']).await;
863        assert!(result.is_err());
864    }
865
866    #[tokio::test]
867    async fn test_compressor_roundtrip_preserves_data() {
868        let comp = MessageCompressor::new(CompressionConfig::default());
869        let original = vec![b'z'; 500];
870        let compressed = comp.compress(&original).await;
871        let decompressed = comp.decompress(&compressed).await.unwrap();
872        assert_eq!(decompressed, original);
873
874        let stats = comp.stats().await;
875        assert_eq!(stats.messages_compressed, 1);
876        assert_eq!(stats.messages_decompressed, 1);
877    }
878}