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
311        && compressed[compressed.len() - 4..] == [0x00, 0x00, 0xFF, 0xFF]
312    {
313        compressed.truncate(compressed.len() - 4);
314    }
315
316    // 若压缩后末尾为 0x00,补一个 0x00 防止与 frame 边界混淆
317    if compressed.last() == Some(&0x00) {
318        compressed.push(0x00);
319    }
320
321    compressed
322}
323
324/// 解压 RFC 7692 permessage-deflate 数据。
325///
326/// 自动补回 4 字节 `00 00 FF FF` 同步标记后调用 zlib 解压。
327/// 同时处理压缩端为防 frame 边界冲突而追加的尾部 0x00(合法的 deflate 流不会以此结尾)。
328fn deflate_decompress(data: &[u8]) -> Result<Vec<u8>, String> {
329    use flate2::read::ZlibDecoder;
330    use std::io::Read;
331
332    if data.is_empty() {
333        return Err("empty compressed data".to_string());
334    }
335
336    // 补回 RFC 7692 移除的 4 字节同步标记
337    let mut full = data.to_vec();
338    full.extend_from_slice(&[0x00, 0x00, 0xFF, 0xFF]);
339
340    let mut decoder = ZlibDecoder::new(&full[..]);
341    let mut out = Vec::new();
342    decoder
343        .read_to_end(&mut out)
344        .map_err(|e| format!("zlib decompress failed: {}", e))?;
345    Ok(out)
346}
347
348/// 消息压缩器,跟踪压缩统计
349#[derive(Debug)]
350pub struct MessageCompressor {
351    config: CompressionConfig,
352    stats: Arc<RwLock<CompressionStats>>,
353}
354
355impl MessageCompressor {
356    /// 创建压缩器
357    pub fn new(config: CompressionConfig) -> Self {
358        Self {
359            config,
360            stats: Arc::new(RwLock::new(CompressionStats::default())),
361        }
362    }
363
364    /// 获取配置
365    pub fn config(&self) -> &CompressionConfig {
366        &self.config
367    }
368
369    /// 压缩消息。对于小消息(<32 字节)不压缩直接返回原文。
370    ///
371    /// 格式:1 字节标记 + 数据
372    /// - 标记 0x00:原始数据(未压缩,后跟原文)
373    /// - 标记 0x01:DEFLATE 压缩数据(后跟 RFC 7692 zlib 流)
374    pub async fn compress(&self, data: &[u8]) -> Vec<u8> {
375        let uncompressed_size = data.len() as u64;
376
377        // 小消息不压缩(DEFLATE 对小数据通常反向膨胀)
378        let compressed = if data.len() < 32 {
379            let mut out = Vec::with_capacity(data.len() + 1);
380            out.push(0x00); // 原始数据标记
381            out.extend_from_slice(data);
382            out
383        } else {
384            let deflated = deflate_compress(data);
385            // 仅当压缩后更小才使用压缩结果
386            if deflated.len() + 1 < data.len() + 1 {
387                let mut out = Vec::with_capacity(deflated.len() + 1);
388                out.push(0x01); // 压缩数据标记
389                out.extend_from_slice(&deflated);
390                out
391            } else {
392                let mut out = Vec::with_capacity(data.len() + 1);
393                out.push(0x00); // 原始数据标记
394                out.extend_from_slice(data);
395                out
396            }
397        };
398
399        let compressed_size = compressed.len() as u64;
400        let mut stats = self.stats.write().await;
401        stats.total_uncompressed += uncompressed_size;
402        stats.total_compressed += compressed_size;
403        stats.messages_compressed += 1;
404
405        compressed
406    }
407
408    /// 解压消息
409    ///
410    /// 根据首字节标记决定解压路径:
411    /// - 0x00:原始数据,跳过标记后直接返回
412    /// - 0x01:DEFLATE 压缩数据,调用 zlib 解压
413    pub async fn decompress(&self, data: &[u8]) -> Result<Vec<u8>, String> {
414        if data.is_empty() {
415            return Err("empty compressed data".to_string());
416        }
417
418        let result = match data[0] {
419            0x00 => data[1..].to_vec(),
420            0x01 => deflate_decompress(&data[1..])?,
421            other => return Err(format!("unknown compression marker: 0x{:02X}", other)),
422        };
423
424        let mut stats = self.stats.write().await;
425        stats.messages_decompressed += 1;
426
427        Ok(result)
428    }
429
430    /// 获取当前压缩统计快照
431    pub async fn stats(&self) -> CompressionStats {
432        self.stats.read().await.clone()
433    }
434
435    /// 重置统计
436    pub async fn reset_stats(&self) {
437        let mut stats = self.stats.write().await;
438        stats.reset();
439    }
440}
441
442#[cfg(test)]
443mod tests {
444    use super::*;
445
446    #[test]
447    fn test_compression_config_default() {
448        let cfg = CompressionConfig::default();
449        assert_eq!(cfg.server_max_window_bits, 15);
450        assert_eq!(cfg.client_max_window_bits, 15);
451        assert!(!cfg.server_no_context_takeover);
452        assert!(!cfg.client_no_context_takeover);
453    }
454
455    #[test]
456    fn test_compression_config_builder() {
457        let cfg = CompressionConfig::new()
458            .with_server_window_bits(10)
459            .with_client_window_bits(12)
460            .with_server_no_context_takeover()
461            .with_client_no_context_takeover();
462        assert_eq!(cfg.server_max_window_bits, 10);
463        assert_eq!(cfg.client_max_window_bits, 12);
464        assert!(cfg.server_no_context_takeover);
465        assert!(cfg.client_no_context_takeover);
466    }
467
468    #[test]
469    fn test_compression_config_clamp_window_bits() {
470        let cfg = CompressionConfig::new()
471            .with_server_window_bits(2)
472            .with_client_window_bits(20);
473        assert_eq!(cfg.server_max_window_bits, 4);
474        assert_eq!(cfg.client_max_window_bits, 15);
475    }
476
477    #[test]
478    fn test_compression_config_validate_ok() {
479        let cfg = CompressionConfig::new();
480        assert!(cfg.validate().is_ok());
481        assert_eq!(cfg.server_max_window_bits, 15, "默认 server_max_window_bits 应为 15");
482        assert_eq!(cfg.client_max_window_bits, 15, "默认 client_max_window_bits 应为 15");
483        assert!(!cfg.server_no_context_takeover, "默认 server_no_context_takeover 应为 false");
484        assert!(!cfg.client_no_context_takeover, "默认 client_no_context_takeover 应为 false");
485    }
486
487    #[test]
488    fn test_compression_config_validate_invalid_server_bits() {
489        let cfg = CompressionConfig {
490            server_max_window_bits: 3,
491            client_max_window_bits: 10,
492            server_no_context_takeover: false,
493            client_no_context_takeover: false,
494        };
495        assert!(cfg.validate().is_err());
496    }
497
498    #[test]
499    fn test_compression_config_validate_invalid_client_bits() {
500        let cfg = CompressionConfig {
501            server_max_window_bits: 10,
502            client_max_window_bits: 16,
503            server_no_context_takeover: false,
504            client_no_context_takeover: false,
505        };
506        assert!(cfg.validate().is_err());
507    }
508
509    #[test]
510    fn test_to_extension_params() {
511        let cfg = CompressionConfig::new()
512            .with_server_window_bits(10)
513            .with_server_no_context_takeover();
514        let params = cfg.to_extension_params();
515        assert!(params.starts_with("permessage-deflate"));
516        assert!(params.contains("server_max_window_bits=10"));
517        assert!(params.contains("server_no_context_takeover"));
518    }
519
520    #[test]
521    fn test_client_extensions_parse_empty() {
522        let ext = ClientExtensions::parse("");
523        assert!(!ext.requests_deflate());
524        assert!(ext.params.is_empty());
525    }
526
527    #[test]
528    fn test_client_extensions_parse_with_values() {
529        let ext = ClientExtensions::parse(
530            "permessage-deflate; server_max_window_bits=10; client_max_window_bits",
531        );
532        assert!(ext.requests_deflate());
533        assert_eq!(ext.get_param("server_max_window_bits"), Some("10"));
534        assert!(ext.has_param("client_max_window_bits"));
535        assert_eq!(ext.get_param("client_max_window_bits"), None);
536    }
537
538    #[test]
539    fn test_client_extensions_parse_quoted_values() {
540        let ext = ClientExtensions::parse("permessage-deflate; param=\"value\"");
541        assert_eq!(ext.get_param("param"), Some("value"));
542    }
543
544    #[test]
545    fn test_negotiator_not_requested_when_empty() {
546        let neg = CompressionNegotiator::new(CompressionConfig::default());
547        let result = neg.negotiate("");
548        assert_eq!(result, NegotiationResult::NotRequested);
549    }
550
551    #[test]
552    fn test_negotiator_not_requested_when_no_deflate() {
553        let neg = CompressionNegotiator::new(CompressionConfig::default());
554        let result = neg.negotiate("other-extension");
555        assert_eq!(result, NegotiationResult::NotRequested);
556    }
557
558    #[test]
559    fn test_negotiator_accepted_basic() {
560        let neg = CompressionNegotiator::new(CompressionConfig::default());
561        let result = neg.negotiate("permessage-deflate");
562        match result {
563            NegotiationResult::Accepted(resp) => {
564                assert!(resp.contains("permessage-deflate"));
565                assert!(resp.contains("server_max_window_bits=15"));
566            }
567            _ => panic!("expected Accepted, got {:?}", result),
568        }
569    }
570
571    #[test]
572    fn test_negotiator_accepted_with_client_window_bits() {
573        let neg = CompressionNegotiator::new(CompressionConfig::default());
574        let result = neg
575            .negotiate("permessage-deflate; server_max_window_bits=10; client_max_window_bits=12");
576        match result {
577            NegotiationResult::Accepted(resp) => {
578                assert!(resp.contains("server_max_window_bits=10"));
579                assert!(resp.contains("client_max_window_bits=12"));
580            }
581            _ => panic!("expected Accepted, got {:?}", result),
582        }
583    }
584
585    #[test]
586    fn test_negotiator_takes_min_window_bits() {
587        let neg = CompressionNegotiator::new(CompressionConfig::new().with_server_window_bits(8));
588        // 客户端请求 server_max_window_bits=10,服务端配置为 8,应取 min=8
589        let result = neg.negotiate("permessage-deflate; server_max_window_bits=10");
590        match result {
591            NegotiationResult::Accepted(resp) => {
592                assert!(resp.contains("server_max_window_bits=8"));
593            }
594            _ => panic!("expected Accepted, got {:?}", result),
595        }
596    }
597
598    #[test]
599    fn test_negotiator_rejected_invalid_window_bits() {
600        let neg = CompressionNegotiator::new(CompressionConfig::default());
601        let result = neg.negotiate("permessage-deflate; server_max_window_bits=99");
602        assert!(matches!(result, NegotiationResult::Rejected(_)));
603    }
604
605    #[test]
606    fn test_negotiator_no_context_takeover_propagated() {
607        let neg =
608            CompressionNegotiator::new(CompressionConfig::new().with_server_no_context_takeover());
609        let result = neg.negotiate("permessage-deflate");
610        match result {
611            NegotiationResult::Accepted(resp) => {
612                assert!(resp.contains("server_no_context_takeover"));
613            }
614            _ => panic!("expected Accepted, got {:?}", result),
615        }
616    }
617
618    #[test]
619    fn test_negotiator_no_context_takeover_from_client() {
620        let neg = CompressionNegotiator::new(CompressionConfig::default());
621        let result = neg.negotiate(
622            "permessage-deflate; server_no_context_takeover; client_no_context_takeover",
623        );
624        match result {
625            NegotiationResult::Accepted(resp) => {
626                assert!(resp.contains("server_no_context_takeover"));
627                assert!(resp.contains("client_no_context_takeover"));
628            }
629            _ => panic!("expected Accepted, got {:?}", result),
630        }
631    }
632
633    #[test]
634    fn test_negotiator_client_window_bits_without_value() {
635        let neg = CompressionNegotiator::new(CompressionConfig::default());
636        let result = neg.negotiate("permessage-deflate; client_max_window_bits");
637        match result {
638            NegotiationResult::Accepted(resp) => {
639                assert!(resp.contains("client_max_window_bits=15"));
640            }
641            _ => panic!("expected Accepted, got {:?}", result),
642        }
643    }
644
645    #[test]
646    fn test_compression_stats_default() {
647        let stats = CompressionStats::default();
648        assert_eq!(stats.total_uncompressed, 0);
649        assert_eq!(stats.total_compressed, 0);
650        assert_eq!(stats.messages_compressed, 0);
651        assert_eq!(stats.messages_decompressed, 0);
652        assert_eq!(stats.ratio(), 1.0);
653        assert_eq!(stats.bytes_saved(), 0);
654        assert_eq!(stats.saved_percent(), 0.0);
655    }
656
657    #[test]
658    fn test_compression_stats_ratio() {
659        let stats = CompressionStats {
660            total_uncompressed: 1000,
661            total_compressed: 400,
662            messages_compressed: 5,
663            messages_decompressed: 0,
664        };
665        assert!((stats.ratio() - 0.4).abs() < 1e-9);
666        assert_eq!(stats.bytes_saved(), 600);
667        assert!((stats.saved_percent() - 60.0).abs() < 1e-9);
668    }
669
670    #[test]
671    fn test_compression_stats_zero_uncompressed() {
672        let stats = CompressionStats {
673            total_uncompressed: 0,
674            total_compressed: 100,
675            messages_compressed: 1,
676            messages_decompressed: 0,
677        };
678        assert_eq!(stats.ratio(), 1.0);
679        assert_eq!(stats.saved_percent(), 0.0);
680    }
681
682    #[test]
683    fn test_compression_stats_reset() {
684        let mut stats = CompressionStats {
685            total_uncompressed: 1000,
686            total_compressed: 400,
687            messages_compressed: 5,
688            messages_decompressed: 3,
689        };
690        stats.reset();
691        assert_eq!(stats.total_uncompressed, 0);
692        assert_eq!(stats.messages_compressed, 0);
693    }
694
695    #[test]
696    fn test_deflate_compress_nonempty() {
697        // 空输入也会产生 zlib header + 同步标记(移除 4 字节后仍非空)
698        let compressed = deflate_compress(b"");
699        assert!(!compressed.is_empty());
700    }
701
702    #[test]
703    fn test_deflate_compress_repeated_bytes() {
704        let data = vec![b'a'; 1000];
705        let compressed = deflate_compress(&data);
706        // DEFLATE 对高度重复数据压缩率极高(应远小于 1000)
707        assert!(compressed.len() < data.len());
708        assert!(compressed.len() < 100);
709    }
710
711    #[test]
712    fn test_deflate_decompress_roundtrip() {
713        let original = b"Hello, permessage-deflate! This is a test message.".repeat(20);
714        let compressed = deflate_compress(&original);
715        let decompressed = deflate_decompress(&compressed).unwrap();
716        assert_eq!(decompressed, original);
717    }
718
719    #[test]
720    fn test_deflate_decompress_empty() {
721        let result = deflate_decompress(b"");
722        assert!(result.is_err());
723    }
724
725    #[test]
726    fn test_deflate_decompress_invalid_data() {
727        // 不是合法 zlib 流
728        let result = deflate_decompress(b"\x00\x01\x02\x03");
729        assert!(result.is_err());
730    }
731
732    #[test]
733    fn test_deflate_roundtrip_random() {
734        let original: Vec<u8> = (0..200u8).collect();
735        let compressed = deflate_compress(&original);
736        let decompressed = deflate_decompress(&compressed).unwrap();
737        assert_eq!(decompressed, original);
738    }
739
740    #[tokio::test]
741    async fn test_compressor_compress_small_message() {
742        let comp = MessageCompressor::new(CompressionConfig::default());
743        // 小消息不压缩,标记为 0x00 + 原文
744        let result = comp.compress(b"hi").await;
745        assert_eq!(result, vec![0x00, b'h', b'i']);
746
747        let stats = comp.stats().await;
748        assert_eq!(stats.messages_compressed, 1);
749        assert_eq!(stats.total_uncompressed, 2);
750        assert_eq!(stats.total_compressed, 3); // 1 字节标记 + 2 字节原文
751    }
752
753    #[tokio::test]
754    async fn test_compressor_compress_large_repeated() {
755        let comp = MessageCompressor::new(CompressionConfig::default());
756        let data = vec![b'a'; 1000]; // 高度重复的数据
757        let result = comp.compress(&data).await;
758        // 应该被压缩(DEFLATE 对重复数据非常高效)
759        assert!(result.len() < data.len());
760        assert_eq!(result[0], 0x01); // 首字节为压缩标记
761
762        let stats = comp.stats().await;
763        assert_eq!(stats.total_uncompressed, 1000);
764        assert!(stats.total_compressed < 1000);
765        assert!(stats.saved_percent() > 50.0);
766    }
767
768    #[tokio::test]
769    async fn test_compressor_compress_incompressible() {
770        let comp = MessageCompressor::new(CompressionConfig::default());
771        // 32 字节顺序递增数据,DEFLATE 可能无法压缩
772        let data: Vec<u8> = (0..32u8).collect();
773        let result = comp.compress(&data).await;
774        // 首字节应为 0x00(原始)或 0x01(压缩);总长度不应远超原文+1
775        assert!(result[0] == 0x00 || result[0] == 0x01);
776        assert!(result.len() <= data.len() + 50); // 允许小幅膨胀但仍合理
777
778        let stats = comp.stats().await;
779        assert_eq!(stats.total_uncompressed, 32);
780    }
781
782    #[tokio::test]
783    async fn test_compressor_decompress_compressed() {
784        let comp = MessageCompressor::new(CompressionConfig::default());
785        let original = vec![b'x'; 100];
786        let compressed = comp.compress(&original).await;
787        let decompressed = comp.decompress(&compressed).await.unwrap();
788        assert_eq!(decompressed, original);
789    }
790
791    #[tokio::test]
792    async fn test_compressor_decompress_uncompressed() {
793        let comp = MessageCompressor::new(CompressionConfig::default());
794        // 小消息不压缩,标记 0x00 + 原文
795        let data = vec![0x00, b'h', b'e', b'l', b'l', b'o'];
796        let decompressed = comp.decompress(&data).await.unwrap();
797        assert_eq!(decompressed, b"hello");
798    }
799
800    #[tokio::test]
801    async fn test_compressor_decompress_invalid_marker() {
802        let comp = MessageCompressor::new(CompressionConfig::default());
803        let data = vec![0x99, 0x01, 0x02]; // 未知标记
804        let result = comp.decompress(&data).await;
805        assert!(result.is_err());
806    }
807
808    #[tokio::test]
809    async fn test_compressor_stats_accumulate() {
810        let comp = MessageCompressor::new(CompressionConfig::default());
811        let data = vec![b'a'; 100];
812        comp.compress(&data).await;
813        comp.compress(&data).await;
814        comp.compress(&data).await;
815
816        let stats = comp.stats().await;
817        assert_eq!(stats.messages_compressed, 3);
818        assert_eq!(stats.total_uncompressed, 300);
819        assert!(stats.total_compressed < 300);
820    }
821
822    #[tokio::test]
823    async fn test_compressor_reset_stats() {
824        let comp = MessageCompressor::new(CompressionConfig::default());
825        let data = vec![b'a'; 100];
826        comp.compress(&data).await;
827        assert!(comp.stats().await.messages_compressed > 0);
828
829        comp.reset_stats().await;
830        let stats = comp.stats().await;
831        assert_eq!(stats.messages_compressed, 0);
832        assert_eq!(stats.total_uncompressed, 0);
833    }
834
835    #[tokio::test]
836    async fn test_compressor_decompress_count() {
837        let comp = MessageCompressor::new(CompressionConfig::default());
838        // 压缩格式要求首字节为标记(0x00 原始 / 0x01 DEFLATE),不能直接传原始字符串
839        let compressed1 = comp.compress(b"hi").await;
840        let compressed2 = comp.compress(b"world").await;
841        comp.decompress(&compressed1).await.unwrap();
842        comp.decompress(&compressed2).await.unwrap();
843
844        let stats = comp.stats().await;
845        assert_eq!(stats.messages_decompressed, 2);
846    }
847
848    #[tokio::test]
849    async fn test_compressor_decompress_invalid_returns_error() {
850        let comp = MessageCompressor::new(CompressionConfig::default());
851        // 0xFF 标记字节后跟单个数据字节但缺少计数字节 —— 截断的压缩数据
852        let result = comp.decompress(&[0xFF, b'a']).await;
853        assert!(result.is_err());
854    }
855
856    #[tokio::test]
857    async fn test_compressor_roundtrip_preserves_data() {
858        let comp = MessageCompressor::new(CompressionConfig::default());
859        let original = vec![b'z'; 500];
860        let compressed = comp.compress(&original).await;
861        let decompressed = comp.decompress(&compressed).await.unwrap();
862        assert_eq!(decompressed, original);
863
864        let stats = comp.stats().await;
865        assert_eq!(stats.messages_compressed, 1);
866        assert_eq!(stats.messages_decompressed, 1);
867    }
868}