Skip to main content

sz_orm_websocket/
compression.rs

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