use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone)]
pub struct CompressionConfig {
pub server_max_window_bits: u8,
pub client_max_window_bits: u8,
pub server_no_context_takeover: bool,
pub client_no_context_takeover: bool,
}
impl Default for CompressionConfig {
fn default() -> Self {
Self {
server_max_window_bits: 15,
client_max_window_bits: 15,
server_no_context_takeover: false,
client_no_context_takeover: false,
}
}
}
impl CompressionConfig {
pub fn new() -> Self {
Self::default()
}
pub fn with_server_window_bits(mut self, bits: u8) -> Self {
self.server_max_window_bits = bits.clamp(4, 15);
self
}
pub fn with_client_window_bits(mut self, bits: u8) -> Self {
self.client_max_window_bits = bits.clamp(4, 15);
self
}
pub fn with_server_no_context_takeover(mut self) -> Self {
self.server_no_context_takeover = true;
self
}
pub fn with_client_no_context_takeover(mut self) -> Self {
self.client_no_context_takeover = true;
self
}
pub fn validate(&self) -> Result<(), String> {
if !(4..=15).contains(&self.server_max_window_bits) {
return Err("server_max_window_bits must be in [4, 15]".to_string());
}
if !(4..=15).contains(&self.client_max_window_bits) {
return Err("client_max_window_bits must be in [4, 15]".to_string());
}
Ok(())
}
pub fn to_extension_params(&self) -> String {
let mut parts = vec!["permessage-deflate".to_string()];
parts.push(format!(
"server_max_window_bits={}",
self.server_max_window_bits
));
if self.server_no_context_takeover {
parts.push("server_no_context_takeover".to_string());
}
if self.client_no_context_takeover {
parts.push("client_no_context_takeover".to_string());
}
parts.push(format!(
"client_max_window_bits={}",
self.client_max_window_bits
));
parts.join("; ")
}
}
#[derive(Debug, Clone, Default)]
pub struct ClientExtensions {
pub raw: String,
pub params: HashMap<String, Option<String>>,
}
impl ClientExtensions {
pub fn parse(header_value: &str) -> Self {
let mut params = HashMap::new();
for part in header_value.split(';') {
let part = part.trim();
if part.is_empty() {
continue;
}
if let Some((key, value)) = part.split_once('=') {
let key = key.trim().to_string();
let value = value.trim().trim_matches('"').to_string();
params.insert(key, Some(value));
} else {
params.insert(part.to_string(), None);
}
}
Self {
raw: header_value.to_string(),
params,
}
}
pub fn requests_deflate(&self) -> bool {
self.raw.contains("permessage-deflate")
}
pub fn get_param(&self, key: &str) -> Option<&str> {
self.params.get(key).and_then(|v| v.as_deref())
}
pub fn has_param(&self, key: &str) -> bool {
self.params.contains_key(key)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum NegotiationResult {
Accepted(String),
NotRequested,
Rejected(String),
}
#[derive(Debug, Clone)]
pub struct CompressionNegotiator {
config: CompressionConfig,
}
impl CompressionNegotiator {
pub fn new(config: CompressionConfig) -> Self {
Self { config }
}
pub fn config(&self) -> &CompressionConfig {
&self.config
}
pub fn negotiate(&self, client_header: &str) -> NegotiationResult {
if client_header.is_empty() {
return NegotiationResult::NotRequested;
}
let client_ext = ClientExtensions::parse(client_header);
if !client_ext.requests_deflate() {
return NegotiationResult::NotRequested;
}
if let Some(bits_str) = client_ext.get_param("server_max_window_bits") {
if let Ok(bits) = bits_str.parse::<u8>() {
if !(4..=15).contains(&bits) {
return NegotiationResult::Rejected(format!(
"invalid server_max_window_bits: {}",
bits
));
}
}
}
let mut response_parts = vec!["permessage-deflate".to_string()];
let final_server_bits =
if let Some(bits_str) = client_ext.get_param("server_max_window_bits") {
if let Ok(client_bits) = bits_str.parse::<u8>() {
client_bits.min(self.config.server_max_window_bits)
} else {
self.config.server_max_window_bits
}
} else {
self.config.server_max_window_bits
};
response_parts.push(format!("server_max_window_bits={}", final_server_bits));
if self.config.server_no_context_takeover
|| client_ext.has_param("server_no_context_takeover")
{
response_parts.push("server_no_context_takeover".to_string());
}
if let Some(bits_str) = client_ext.get_param("client_max_window_bits") {
if let Ok(client_bits) = bits_str.parse::<u8>() {
let final_client_bits = client_bits.min(self.config.client_max_window_bits);
response_parts.push(format!("client_max_window_bits={}", final_client_bits));
}
} else if client_ext.has_param("client_max_window_bits") {
response_parts.push(format!(
"client_max_window_bits={}",
self.config.client_max_window_bits
));
}
if self.config.client_no_context_takeover
|| client_ext.has_param("client_no_context_takeover")
{
response_parts.push("client_no_context_takeover".to_string());
}
NegotiationResult::Accepted(response_parts.join("; "))
}
}
#[derive(Debug, Clone, Default)]
pub struct CompressionStats {
pub total_uncompressed: u64,
pub total_compressed: u64,
pub messages_compressed: u64,
pub messages_decompressed: u64,
}
impl CompressionStats {
pub fn ratio(&self) -> f64 {
if self.total_uncompressed == 0 {
return 1.0;
}
self.total_compressed as f64 / self.total_uncompressed as f64
}
pub fn bytes_saved(&self) -> i64 {
self.total_uncompressed as i64 - self.total_compressed as i64
}
pub fn saved_percent(&self) -> f64 {
if self.total_uncompressed == 0 {
return 0.0;
}
let saved = self.total_uncompressed - self.total_compressed;
(saved as f64 / self.total_uncompressed as f64) * 100.0
}
pub fn reset(&mut self) {
*self = Self::default();
}
}
fn deflate_compress(data: &[u8]) -> Vec<u8> {
use flate2::write::ZlibEncoder;
use std::io::Write;
let mut encoder = ZlibEncoder::new(Vec::new(), flate2::Compression::default());
let _ = encoder.write_all(data);
let mut compressed = encoder
.finish()
.expect("zlib compression of in-memory buffer cannot fail");
if compressed.len() >= 4
&& compressed[compressed.len() - 4..] == [0x00, 0x00, 0xFF, 0xFF]
{
compressed.truncate(compressed.len() - 4);
}
if compressed.last() == Some(&0x00) {
compressed.push(0x00);
}
compressed
}
fn deflate_decompress(data: &[u8]) -> Result<Vec<u8>, String> {
use flate2::read::ZlibDecoder;
use std::io::Read;
if data.is_empty() {
return Err("empty compressed data".to_string());
}
let mut full = data.to_vec();
full.extend_from_slice(&[0x00, 0x00, 0xFF, 0xFF]);
let mut decoder = ZlibDecoder::new(&full[..]);
let mut out = Vec::new();
decoder
.read_to_end(&mut out)
.map_err(|e| format!("zlib decompress failed: {}", e))?;
Ok(out)
}
#[derive(Debug)]
pub struct MessageCompressor {
config: CompressionConfig,
stats: Arc<RwLock<CompressionStats>>,
}
impl MessageCompressor {
pub fn new(config: CompressionConfig) -> Self {
Self {
config,
stats: Arc::new(RwLock::new(CompressionStats::default())),
}
}
pub fn config(&self) -> &CompressionConfig {
&self.config
}
pub async fn compress(&self, data: &[u8]) -> Vec<u8> {
let uncompressed_size = data.len() as u64;
let compressed = if data.len() < 32 {
let mut out = Vec::with_capacity(data.len() + 1);
out.push(0x00); out.extend_from_slice(data);
out
} else {
let deflated = deflate_compress(data);
if deflated.len() + 1 < data.len() + 1 {
let mut out = Vec::with_capacity(deflated.len() + 1);
out.push(0x01); out.extend_from_slice(&deflated);
out
} else {
let mut out = Vec::with_capacity(data.len() + 1);
out.push(0x00); out.extend_from_slice(data);
out
}
};
let compressed_size = compressed.len() as u64;
let mut stats = self.stats.write().await;
stats.total_uncompressed += uncompressed_size;
stats.total_compressed += compressed_size;
stats.messages_compressed += 1;
compressed
}
pub async fn decompress(&self, data: &[u8]) -> Result<Vec<u8>, String> {
if data.is_empty() {
return Err("empty compressed data".to_string());
}
let result = match data[0] {
0x00 => data[1..].to_vec(),
0x01 => deflate_decompress(&data[1..])?,
other => return Err(format!("unknown compression marker: 0x{:02X}", other)),
};
let mut stats = self.stats.write().await;
stats.messages_decompressed += 1;
Ok(result)
}
pub async fn stats(&self) -> CompressionStats {
self.stats.read().await.clone()
}
pub async fn reset_stats(&self) {
let mut stats = self.stats.write().await;
stats.reset();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_compression_config_default() {
let cfg = CompressionConfig::default();
assert_eq!(cfg.server_max_window_bits, 15);
assert_eq!(cfg.client_max_window_bits, 15);
assert!(!cfg.server_no_context_takeover);
assert!(!cfg.client_no_context_takeover);
}
#[test]
fn test_compression_config_builder() {
let cfg = CompressionConfig::new()
.with_server_window_bits(10)
.with_client_window_bits(12)
.with_server_no_context_takeover()
.with_client_no_context_takeover();
assert_eq!(cfg.server_max_window_bits, 10);
assert_eq!(cfg.client_max_window_bits, 12);
assert!(cfg.server_no_context_takeover);
assert!(cfg.client_no_context_takeover);
}
#[test]
fn test_compression_config_clamp_window_bits() {
let cfg = CompressionConfig::new()
.with_server_window_bits(2)
.with_client_window_bits(20);
assert_eq!(cfg.server_max_window_bits, 4);
assert_eq!(cfg.client_max_window_bits, 15);
}
#[test]
fn test_compression_config_validate_ok() {
let cfg = CompressionConfig::new();
assert!(cfg.validate().is_ok());
assert_eq!(cfg.server_max_window_bits, 15, "默认 server_max_window_bits 应为 15");
assert_eq!(cfg.client_max_window_bits, 15, "默认 client_max_window_bits 应为 15");
assert!(!cfg.server_no_context_takeover, "默认 server_no_context_takeover 应为 false");
assert!(!cfg.client_no_context_takeover, "默认 client_no_context_takeover 应为 false");
}
#[test]
fn test_compression_config_validate_invalid_server_bits() {
let cfg = CompressionConfig {
server_max_window_bits: 3,
client_max_window_bits: 10,
server_no_context_takeover: false,
client_no_context_takeover: false,
};
assert!(cfg.validate().is_err());
}
#[test]
fn test_compression_config_validate_invalid_client_bits() {
let cfg = CompressionConfig {
server_max_window_bits: 10,
client_max_window_bits: 16,
server_no_context_takeover: false,
client_no_context_takeover: false,
};
assert!(cfg.validate().is_err());
}
#[test]
fn test_to_extension_params() {
let cfg = CompressionConfig::new()
.with_server_window_bits(10)
.with_server_no_context_takeover();
let params = cfg.to_extension_params();
assert!(params.starts_with("permessage-deflate"));
assert!(params.contains("server_max_window_bits=10"));
assert!(params.contains("server_no_context_takeover"));
}
#[test]
fn test_client_extensions_parse_empty() {
let ext = ClientExtensions::parse("");
assert!(!ext.requests_deflate());
assert!(ext.params.is_empty());
}
#[test]
fn test_client_extensions_parse_with_values() {
let ext = ClientExtensions::parse(
"permessage-deflate; server_max_window_bits=10; client_max_window_bits",
);
assert!(ext.requests_deflate());
assert_eq!(ext.get_param("server_max_window_bits"), Some("10"));
assert!(ext.has_param("client_max_window_bits"));
assert_eq!(ext.get_param("client_max_window_bits"), None);
}
#[test]
fn test_client_extensions_parse_quoted_values() {
let ext = ClientExtensions::parse("permessage-deflate; param=\"value\"");
assert_eq!(ext.get_param("param"), Some("value"));
}
#[test]
fn test_negotiator_not_requested_when_empty() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg.negotiate("");
assert_eq!(result, NegotiationResult::NotRequested);
}
#[test]
fn test_negotiator_not_requested_when_no_deflate() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg.negotiate("other-extension");
assert_eq!(result, NegotiationResult::NotRequested);
}
#[test]
fn test_negotiator_accepted_basic() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg.negotiate("permessage-deflate");
match result {
NegotiationResult::Accepted(resp) => {
assert!(resp.contains("permessage-deflate"));
assert!(resp.contains("server_max_window_bits=15"));
}
_ => panic!("expected Accepted, got {:?}", result),
}
}
#[test]
fn test_negotiator_accepted_with_client_window_bits() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg
.negotiate("permessage-deflate; server_max_window_bits=10; client_max_window_bits=12");
match result {
NegotiationResult::Accepted(resp) => {
assert!(resp.contains("server_max_window_bits=10"));
assert!(resp.contains("client_max_window_bits=12"));
}
_ => panic!("expected Accepted, got {:?}", result),
}
}
#[test]
fn test_negotiator_takes_min_window_bits() {
let neg = CompressionNegotiator::new(CompressionConfig::new().with_server_window_bits(8));
let result = neg.negotiate("permessage-deflate; server_max_window_bits=10");
match result {
NegotiationResult::Accepted(resp) => {
assert!(resp.contains("server_max_window_bits=8"));
}
_ => panic!("expected Accepted, got {:?}", result),
}
}
#[test]
fn test_negotiator_rejected_invalid_window_bits() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg.negotiate("permessage-deflate; server_max_window_bits=99");
assert!(matches!(result, NegotiationResult::Rejected(_)));
}
#[test]
fn test_negotiator_no_context_takeover_propagated() {
let neg =
CompressionNegotiator::new(CompressionConfig::new().with_server_no_context_takeover());
let result = neg.negotiate("permessage-deflate");
match result {
NegotiationResult::Accepted(resp) => {
assert!(resp.contains("server_no_context_takeover"));
}
_ => panic!("expected Accepted, got {:?}", result),
}
}
#[test]
fn test_negotiator_no_context_takeover_from_client() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg.negotiate(
"permessage-deflate; server_no_context_takeover; client_no_context_takeover",
);
match result {
NegotiationResult::Accepted(resp) => {
assert!(resp.contains("server_no_context_takeover"));
assert!(resp.contains("client_no_context_takeover"));
}
_ => panic!("expected Accepted, got {:?}", result),
}
}
#[test]
fn test_negotiator_client_window_bits_without_value() {
let neg = CompressionNegotiator::new(CompressionConfig::default());
let result = neg.negotiate("permessage-deflate; client_max_window_bits");
match result {
NegotiationResult::Accepted(resp) => {
assert!(resp.contains("client_max_window_bits=15"));
}
_ => panic!("expected Accepted, got {:?}", result),
}
}
#[test]
fn test_compression_stats_default() {
let stats = CompressionStats::default();
assert_eq!(stats.total_uncompressed, 0);
assert_eq!(stats.total_compressed, 0);
assert_eq!(stats.messages_compressed, 0);
assert_eq!(stats.messages_decompressed, 0);
assert_eq!(stats.ratio(), 1.0);
assert_eq!(stats.bytes_saved(), 0);
assert_eq!(stats.saved_percent(), 0.0);
}
#[test]
fn test_compression_stats_ratio() {
let stats = CompressionStats {
total_uncompressed: 1000,
total_compressed: 400,
messages_compressed: 5,
messages_decompressed: 0,
};
assert!((stats.ratio() - 0.4).abs() < 1e-9);
assert_eq!(stats.bytes_saved(), 600);
assert!((stats.saved_percent() - 60.0).abs() < 1e-9);
}
#[test]
fn test_compression_stats_zero_uncompressed() {
let stats = CompressionStats {
total_uncompressed: 0,
total_compressed: 100,
messages_compressed: 1,
messages_decompressed: 0,
};
assert_eq!(stats.ratio(), 1.0);
assert_eq!(stats.saved_percent(), 0.0);
}
#[test]
fn test_compression_stats_reset() {
let mut stats = CompressionStats {
total_uncompressed: 1000,
total_compressed: 400,
messages_compressed: 5,
messages_decompressed: 3,
};
stats.reset();
assert_eq!(stats.total_uncompressed, 0);
assert_eq!(stats.messages_compressed, 0);
}
#[test]
fn test_deflate_compress_nonempty() {
let compressed = deflate_compress(b"");
assert!(!compressed.is_empty());
}
#[test]
fn test_deflate_compress_repeated_bytes() {
let data = vec![b'a'; 1000];
let compressed = deflate_compress(&data);
assert!(compressed.len() < data.len());
assert!(compressed.len() < 100);
}
#[test]
fn test_deflate_decompress_roundtrip() {
let original = b"Hello, permessage-deflate! This is a test message.".repeat(20);
let compressed = deflate_compress(&original);
let decompressed = deflate_decompress(&compressed).unwrap();
assert_eq!(decompressed, original);
}
#[test]
fn test_deflate_decompress_empty() {
let result = deflate_decompress(b"");
assert!(result.is_err());
}
#[test]
fn test_deflate_decompress_invalid_data() {
let result = deflate_decompress(b"\x00\x01\x02\x03");
assert!(result.is_err());
}
#[test]
fn test_deflate_roundtrip_random() {
let original: Vec<u8> = (0..200u8).collect();
let compressed = deflate_compress(&original);
let decompressed = deflate_decompress(&compressed).unwrap();
assert_eq!(decompressed, original);
}
#[tokio::test]
async fn test_compressor_compress_small_message() {
let comp = MessageCompressor::new(CompressionConfig::default());
let result = comp.compress(b"hi").await;
assert_eq!(result, vec![0x00, b'h', b'i']);
let stats = comp.stats().await;
assert_eq!(stats.messages_compressed, 1);
assert_eq!(stats.total_uncompressed, 2);
assert_eq!(stats.total_compressed, 3); }
#[tokio::test]
async fn test_compressor_compress_large_repeated() {
let comp = MessageCompressor::new(CompressionConfig::default());
let data = vec![b'a'; 1000]; let result = comp.compress(&data).await;
assert!(result.len() < data.len());
assert_eq!(result[0], 0x01);
let stats = comp.stats().await;
assert_eq!(stats.total_uncompressed, 1000);
assert!(stats.total_compressed < 1000);
assert!(stats.saved_percent() > 50.0);
}
#[tokio::test]
async fn test_compressor_compress_incompressible() {
let comp = MessageCompressor::new(CompressionConfig::default());
let data: Vec<u8> = (0..32u8).collect();
let result = comp.compress(&data).await;
assert!(result[0] == 0x00 || result[0] == 0x01);
assert!(result.len() <= data.len() + 50);
let stats = comp.stats().await;
assert_eq!(stats.total_uncompressed, 32);
}
#[tokio::test]
async fn test_compressor_decompress_compressed() {
let comp = MessageCompressor::new(CompressionConfig::default());
let original = vec![b'x'; 100];
let compressed = comp.compress(&original).await;
let decompressed = comp.decompress(&compressed).await.unwrap();
assert_eq!(decompressed, original);
}
#[tokio::test]
async fn test_compressor_decompress_uncompressed() {
let comp = MessageCompressor::new(CompressionConfig::default());
let data = vec![0x00, b'h', b'e', b'l', b'l', b'o'];
let decompressed = comp.decompress(&data).await.unwrap();
assert_eq!(decompressed, b"hello");
}
#[tokio::test]
async fn test_compressor_decompress_invalid_marker() {
let comp = MessageCompressor::new(CompressionConfig::default());
let data = vec![0x99, 0x01, 0x02]; let result = comp.decompress(&data).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_compressor_stats_accumulate() {
let comp = MessageCompressor::new(CompressionConfig::default());
let data = vec![b'a'; 100];
comp.compress(&data).await;
comp.compress(&data).await;
comp.compress(&data).await;
let stats = comp.stats().await;
assert_eq!(stats.messages_compressed, 3);
assert_eq!(stats.total_uncompressed, 300);
assert!(stats.total_compressed < 300);
}
#[tokio::test]
async fn test_compressor_reset_stats() {
let comp = MessageCompressor::new(CompressionConfig::default());
let data = vec![b'a'; 100];
comp.compress(&data).await;
assert!(comp.stats().await.messages_compressed > 0);
comp.reset_stats().await;
let stats = comp.stats().await;
assert_eq!(stats.messages_compressed, 0);
assert_eq!(stats.total_uncompressed, 0);
}
#[tokio::test]
async fn test_compressor_decompress_count() {
let comp = MessageCompressor::new(CompressionConfig::default());
let compressed1 = comp.compress(b"hi").await;
let compressed2 = comp.compress(b"world").await;
comp.decompress(&compressed1).await.unwrap();
comp.decompress(&compressed2).await.unwrap();
let stats = comp.stats().await;
assert_eq!(stats.messages_decompressed, 2);
}
#[tokio::test]
async fn test_compressor_decompress_invalid_returns_error() {
let comp = MessageCompressor::new(CompressionConfig::default());
let result = comp.decompress(&[0xFF, b'a']).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_compressor_roundtrip_preserves_data() {
let comp = MessageCompressor::new(CompressionConfig::default());
let original = vec![b'z'; 500];
let compressed = comp.compress(&original).await;
let decompressed = comp.decompress(&compressed).await.unwrap();
assert_eq!(decompressed, original);
let stats = comp.stats().await;
assert_eq!(stats.messages_compressed, 1);
assert_eq!(stats.messages_decompressed, 1);
}
}