1use std::collections::HashMap;
17use std::sync::Arc;
18use tokio::sync::RwLock;
19
20#[derive(Debug, Clone)]
22pub struct CompressionConfig {
23 pub server_max_window_bits: u8,
25 pub client_max_window_bits: u8,
27 pub server_no_context_takeover: bool,
29 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 pub fn new() -> Self {
47 Self::default()
48 }
49
50 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 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 pub fn with_server_no_context_takeover(mut self) -> Self {
64 self.server_no_context_takeover = true;
65 self
66 }
67
68 pub fn with_client_no_context_takeover(mut self) -> Self {
70 self.client_no_context_takeover = true;
71 self
72 }
73
74 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 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#[derive(Debug, Clone, Default)]
108pub struct ClientExtensions {
109 pub raw: String,
111 pub params: HashMap<String, Option<String>>,
113}
114
115impl ClientExtensions {
116 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 pub fn requests_deflate(&self) -> bool {
140 self.raw.contains("permessage-deflate")
141 }
142
143 pub fn get_param(&self, key: &str) -> Option<&str> {
145 self.params.get(key).and_then(|v| v.as_deref())
146 }
147
148 pub fn has_param(&self, key: &str) -> bool {
150 self.params.contains_key(key)
151 }
152}
153
154#[derive(Debug, Clone, PartialEq, Eq)]
156pub enum NegotiationResult {
157 Accepted(String),
159 NotRequested,
161 Rejected(String),
163}
164
165#[derive(Debug, Clone)]
167pub struct CompressionNegotiator {
168 config: CompressionConfig,
169}
170
171impl CompressionNegotiator {
172 pub fn new(config: CompressionConfig) -> Self {
174 Self { config }
175 }
176
177 pub fn config(&self) -> &CompressionConfig {
179 &self.config
180 }
181
182 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 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 let mut response_parts = vec!["permessage-deflate".to_string()];
207
208 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 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 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#[derive(Debug, Clone, Default)]
253pub struct CompressionStats {
254 pub total_uncompressed: u64,
256 pub total_compressed: u64,
258 pub messages_compressed: u64,
260 pub messages_decompressed: u64,
262}
263
264impl CompressionStats {
265 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 pub fn bytes_saved(&self) -> i64 {
275 self.total_uncompressed as i64 - self.total_compressed as i64
276 }
277
278 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 pub fn reset(&mut self) {
289 *self = Self::default();
290 }
291}
292
293fn 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 let _ = encoder.write_all(data);
305 let mut compressed = encoder
306 .finish()
307 .expect("zlib compression of in-memory buffer cannot fail");
308
309 if compressed.len() >= 4
311 && compressed[compressed.len() - 4..] == [0x00, 0x00, 0xFF, 0xFF]
312 {
313 compressed.truncate(compressed.len() - 4);
314 }
315
316 if compressed.last() == Some(&0x00) {
318 compressed.push(0x00);
319 }
320
321 compressed
322}
323
324fn 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 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#[derive(Debug)]
350pub struct MessageCompressor {
351 config: CompressionConfig,
352 stats: Arc<RwLock<CompressionStats>>,
353}
354
355impl MessageCompressor {
356 pub fn new(config: CompressionConfig) -> Self {
358 Self {
359 config,
360 stats: Arc::new(RwLock::new(CompressionStats::default())),
361 }
362 }
363
364 pub fn config(&self) -> &CompressionConfig {
366 &self.config
367 }
368
369 pub async fn compress(&self, data: &[u8]) -> Vec<u8> {
375 let uncompressed_size = data.len() as u64;
376
377 let compressed = if data.len() < 32 {
379 let mut out = Vec::with_capacity(data.len() + 1);
380 out.push(0x00); out.extend_from_slice(data);
382 out
383 } else {
384 let deflated = deflate_compress(data);
385 if deflated.len() + 1 < data.len() + 1 {
387 let mut out = Vec::with_capacity(deflated.len() + 1);
388 out.push(0x01); out.extend_from_slice(&deflated);
390 out
391 } else {
392 let mut out = Vec::with_capacity(data.len() + 1);
393 out.push(0x00); 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 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 pub async fn stats(&self) -> CompressionStats {
432 self.stats.read().await.clone()
433 }
434
435 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 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 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 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 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 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); }
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]; let result = comp.compress(&data).await;
758 assert!(result.len() < data.len());
760 assert_eq!(result[0], 0x01); 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 let data: Vec<u8> = (0..32u8).collect();
773 let result = comp.compress(&data).await;
774 assert!(result[0] == 0x00 || result[0] == 0x01);
776 assert!(result.len() <= data.len() + 50); 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 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]; 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 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 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}