1use std::collections::HashMap;
18use std::sync::Arc;
19use tokio::sync::RwLock;
20
21#[derive(Debug, Clone)]
23pub struct CompressionConfig {
24 pub server_max_window_bits: u8,
26 pub client_max_window_bits: u8,
28 pub server_no_context_takeover: bool,
30 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 pub fn new() -> Self {
48 Self::default()
49 }
50
51 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 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 pub fn with_server_no_context_takeover(mut self) -> Self {
65 self.server_no_context_takeover = true;
66 self
67 }
68
69 pub fn with_client_no_context_takeover(mut self) -> Self {
71 self.client_no_context_takeover = true;
72 self
73 }
74
75 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 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#[derive(Debug, Clone, Default)]
109pub struct ClientExtensions {
110 pub raw: String,
112 pub params: HashMap<String, Option<String>>,
114}
115
116impl ClientExtensions {
117 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 pub fn requests_deflate(&self) -> bool {
141 self.raw.contains("permessage-deflate")
142 }
143
144 pub fn get_param(&self, key: &str) -> Option<&str> {
146 self.params.get(key).and_then(|v| v.as_deref())
147 }
148
149 pub fn has_param(&self, key: &str) -> bool {
151 self.params.contains_key(key)
152 }
153}
154
155#[derive(Debug, Clone, PartialEq, Eq)]
157pub enum NegotiationResult {
158 Accepted(String),
160 NotRequested,
162 Rejected(String),
164}
165
166#[derive(Debug, Clone)]
168pub struct CompressionNegotiator {
169 config: CompressionConfig,
170}
171
172impl CompressionNegotiator {
173 pub fn new(config: CompressionConfig) -> Self {
175 Self { config }
176 }
177
178 pub fn config(&self) -> &CompressionConfig {
180 &self.config
181 }
182
183 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 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 let mut response_parts = vec!["permessage-deflate".to_string()];
208
209 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 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 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#[derive(Debug, Clone, Default)]
254pub struct CompressionStats {
255 pub total_uncompressed: u64,
257 pub total_compressed: u64,
259 pub messages_compressed: u64,
261 pub messages_decompressed: u64,
263}
264
265impl CompressionStats {
266 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 pub fn bytes_saved(&self) -> i64 {
276 self.total_uncompressed as i64 - self.total_compressed as i64
277 }
278
279 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 pub fn reset(&mut self) {
290 *self = Self::default();
291 }
292}
293
294fn 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); 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
322fn 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#[derive(Debug)]
346pub struct MessageCompressor {
347 config: CompressionConfig,
348 stats: Arc<RwLock<CompressionStats>>,
349}
350
351impl MessageCompressor {
352 pub fn new(config: CompressionConfig) -> Self {
354 Self {
355 config,
356 stats: Arc::new(RwLock::new(CompressionStats::default())),
357 }
358 }
359
360 pub fn config(&self) -> &CompressionConfig {
362 &self.config
363 }
364
365 pub async fn compress(&self, data: &[u8]) -> Vec<u8> {
367 let uncompressed_size = data.len() as u64;
368
369 let compressed = if data.len() < 32 {
371 data.to_vec()
372 } else {
373 let rle = rle_compress(data);
374 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 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 pub async fn stats(&self) -> CompressionStats {
407 self.stats.read().await.clone()
408 }
409
410 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 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 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 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']; 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 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]; let result = comp.compress(&data).await;
741 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 let data: Vec<u8> = (0..32u8).collect();
755 let result = comp.compress(&data).await;
756 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 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 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}