use std::io::{Read, Write};
use flate2::read::{GzDecoder, ZlibDecoder};
use flate2::write::{GzEncoder, ZlibEncoder};
use flate2::Compression;
use serde::{Serialize, Deserialize};
use base64::{Engine as _, engine::general_purpose};
use thiserror::Error;
use crate::{QueueMessage, QueueError, QueueResult};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum CompressionAlgorithm {
None,
#[default]
Gzip,
Zlib,
}
impl CompressionAlgorithm {
pub fn parse(s: &str) -> Self {
match s.to_lowercase().as_str() {
"gzip" => CompressionAlgorithm::Gzip,
"zlib" => CompressionAlgorithm::Zlib,
_ => CompressionAlgorithm::None,
}
}
pub fn as_str(&self) -> &'static str {
match self {
CompressionAlgorithm::None => "none",
CompressionAlgorithm::Gzip => "gzip",
CompressionAlgorithm::Zlib => "zlib",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompressionConfig {
pub algorithm: CompressionAlgorithm,
pub level: u32,
pub min_size: usize,
pub force_compression: bool,
pub max_attempts: u32,
}
impl Default for CompressionConfig {
fn default() -> Self {
Self {
algorithm: CompressionAlgorithm::Gzip,
level: 6, min_size: 1024, force_compression: false,
max_attempts: 3,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct CompressionStats {
pub total_messages: u64,
pub compressed_messages: u64,
pub decompressed_messages: u64,
pub original_bytes: u64,
pub compressed_bytes: u64,
pub compression_ratio: f64,
pub compression_time_ms: u64,
pub decompression_time_ms: u64,
}
impl CompressionStats {
pub fn update_ratio(&mut self) {
if self.original_bytes > 0 {
self.compression_ratio = self.compressed_bytes as f64 / self.original_bytes as f64;
}
}
pub fn record_compression(&mut self, original_size: usize, compressed_size: usize, time_ms: u64) {
self.total_messages += 1;
self.compressed_messages += 1;
self.original_bytes += original_size as u64;
self.compressed_bytes += compressed_size as u64;
self.compression_time_ms += time_ms;
self.update_ratio();
}
pub fn record_decompression(&mut self, time_ms: u64) {
self.total_messages += 1;
self.decompressed_messages += 1;
self.decompression_time_ms += time_ms;
}
pub fn avg_compression_time_ms(&self) -> f64 {
if self.compressed_messages > 0 {
self.compression_time_ms as f64 / self.compressed_messages as f64
} else {
0.0
}
}
pub fn avg_decompression_time_ms(&self) -> f64 {
if self.decompressed_messages > 0 {
self.decompression_time_ms as f64 / self.decompressed_messages as f64
} else {
0.0
}
}
}
#[derive(Debug, Error)]
pub enum CompressionError {
#[error("Unsupported compression algorithm: {0}")]
UnsupportedAlgorithm(String),
#[error("Compression failed: {0}")]
CompressionFailed(String),
#[error("Decompression failed: {0}")]
DecompressionFailed(String),
#[error("Invalid compression data: {0}")]
InvalidData(String),
#[error("Message too large for compression: {size} bytes")]
MessageTooLarge { size: usize },
#[error("Compression ratio not beneficial: {ratio:.2}")]
NotBeneficial { ratio: f64 },
}
impl From<CompressionError> for crate::QueueError {
fn from(error: CompressionError) -> Self {
crate::QueueError::Other(error.to_string())
}
}
pub struct MessageCompressor {
config: CompressionConfig,
stats: std::sync::Arc<tokio::sync::RwLock<CompressionStats>>,
}
impl Default for MessageCompressor {
fn default() -> Self {
Self::new(CompressionConfig::default())
}
}
impl MessageCompressor {
pub fn new(config: CompressionConfig) -> Self {
Self {
config,
stats: std::sync::Arc::new(tokio::sync::RwLock::new(CompressionStats::default())),
}
}
pub async fn compress_message(&self, message: &mut QueueMessage) -> QueueResult<()> {
let start_time = std::time::Instant::now();
if !self.should_compress(&message.payload) {
return Ok(());
}
let original_size = serde_json::to_vec(&message.payload).unwrap().len();
let payload_bytes = match serde_json::to_vec(&message.payload) {
Ok(bytes) => bytes,
Err(e) => return Err(QueueError::Serialization(e.to_string())),
};
let compressed_data = self.compress_data(&payload_bytes).await?;
let compressed_size = compressed_data.len();
if !self.config.force_compression && compressed_size >= original_size {
return Err(QueueError::Serialization(
format!("Compression not beneficial: {} -> {} bytes", original_size, compressed_size)
));
}
message.payload = serde_json::Value::String(
general_purpose::STANDARD.encode(&compressed_data)
);
message.compressed = true;
message.original_size = Some(original_size);
message.attributes.insert(
"compression_algorithm".to_string(),
self.config.algorithm.as_str().to_string()
);
message.attributes.insert(
"compression_compressed_size".to_string(),
compressed_size.to_string()
);
let time_ms = start_time.elapsed().as_millis() as u64;
{
let mut stats = self.stats.write().await;
stats.record_compression(original_size, compressed_size, time_ms);
}
Ok(())
}
pub async fn decompress_message(&self, message: &mut QueueMessage) -> QueueResult<()> {
let start_time = std::time::Instant::now();
let compression_info = self.get_compression_info(message)?;
if compression_info.is_none() {
return Ok(()); }
let (algorithm, original_size) = compression_info.unwrap();
let compressed_data = match &message.payload {
serde_json::Value::String(encoded_data) => {
match general_purpose::STANDARD.decode(encoded_data) {
Ok(data) => data,
Err(e) => return Err(QueueError::Deserialization(
format!("Failed to decode compressed data: {}", e)
)),
}
}
_ => {
return Err(QueueError::Deserialization(
"Compressed payload must be a string".to_string()
));
}
};
let decompressed_data = self.decompress_data(&compressed_data, algorithm).await?;
if original_size > 0 && decompressed_data.len() != original_size {
return Err(QueueError::Deserialization(
format!("Decompressed size mismatch: expected {}, got {}",
original_size, decompressed_data.len())
));
}
message.payload = match serde_json::from_slice(&decompressed_data) {
Ok(payload) => payload,
Err(e) => return Err(QueueError::Deserialization(e.to_string())),
};
message.compressed = false;
message.original_size = None;
message.attributes.remove("compression_algorithm");
message.attributes.remove("compression_compressed_size");
let time_ms = start_time.elapsed().as_millis() as u64;
{
let mut stats = self.stats.write().await;
stats.record_decompression(time_ms);
}
Ok(())
}
async fn compress_data(&self, data: &[u8]) -> Result<Vec<u8>, CompressionError> {
let compression = Compression::new(self.config.level);
match self.config.algorithm {
CompressionAlgorithm::Gzip => {
let mut encoder = GzEncoder::new(Vec::new(), compression);
encoder.write_all(data)
.map_err(|e| CompressionError::CompressionFailed(e.to_string()))?;
encoder.finish()
.map_err(|e| CompressionError::CompressionFailed(e.to_string()))
}
CompressionAlgorithm::Zlib => {
let mut encoder = ZlibEncoder::new(Vec::new(), compression);
encoder.write_all(data)
.map_err(|e| CompressionError::CompressionFailed(e.to_string()))?;
encoder.finish()
.map_err(|e| CompressionError::CompressionFailed(e.to_string()))
}
CompressionAlgorithm::None => {
Ok(data.to_vec())
}
}
}
async fn decompress_data(&self, data: &[u8], algorithm: CompressionAlgorithm) -> Result<Vec<u8>, CompressionError> {
match algorithm {
CompressionAlgorithm::Gzip => {
let mut decoder = GzDecoder::new(data);
let mut decompressed = Vec::new();
decoder.read_to_end(&mut decompressed)
.map_err(|e| CompressionError::DecompressionFailed(e.to_string()))?;
Ok(decompressed)
}
CompressionAlgorithm::Zlib => {
let mut decoder = ZlibDecoder::new(data);
let mut decompressed = Vec::new();
decoder.read_to_end(&mut decompressed)
.map_err(|e| CompressionError::DecompressionFailed(e.to_string()))?;
Ok(decompressed)
}
CompressionAlgorithm::None => {
Ok(data.to_vec())
}
}
}
fn should_compress(&self, payload: &serde_json::Value) -> bool {
if self.config.algorithm == CompressionAlgorithm::None {
return false;
}
if payload.is_string() && payload.as_str().unwrap_or("").starts_with("H4sI") { return false;
}
let payload_size = match serde_json::to_vec(payload) {
Ok(vec) => vec.len(),
Err(_) => return false,
};
payload_size >= self.config.min_size || self.config.force_compression
}
fn get_compression_info(&self, message: &QueueMessage) -> Result<Option<(CompressionAlgorithm, usize)>, QueueError> {
if !message.compressed {
return Ok(None);
}
let algorithm_str = match message.attributes.get("compression_algorithm") {
Some(s) => s,
_ => return Ok(None),
};
let algorithm = CompressionAlgorithm::parse(algorithm_str);
if algorithm == CompressionAlgorithm::None {
return Ok(None);
}
let original_size = message.original_size.unwrap_or(0);
Ok(Some((algorithm, original_size)))
}
pub async fn get_stats(&self) -> CompressionStats {
self.stats.read().await.clone()
}
pub async fn reset_stats(&self) {
let mut stats = self.stats.write().await;
*stats = CompressionStats::default();
}
pub fn update_config(&mut self, config: CompressionConfig) {
self.config = config;
}
pub fn get_config(&self) -> &CompressionConfig {
&self.config
}
}
pub struct CompressedMessageBuilder {
inner: crate::QueueMessage,
compressor: std::sync::Arc<MessageCompressor>,
}
impl CompressedMessageBuilder {
pub fn new(compressor: std::sync::Arc<MessageCompressor>) -> Self {
Self {
inner: crate::QueueMessage {
id: String::new(),
payload: serde_json::Value::Null,
priority: crate::QueuePriority::Normal,
receive_count: 0,
max_receive_count: 3,
enqueued_at: chrono::Utc::now(),
created_at: chrono::Utc::now(),
visible_at: chrono::Utc::now(),
expires_at: None,
visibility_timeout: 30,
status: crate::MessageStatus::Pending,
delay_seconds: None,
attributes: std::collections::HashMap::new(),
headers: std::collections::HashMap::new(),
message_group_id: None,
message_deduplication_id: None,
routing_key: None,
compressed: false,
original_size: None,
},
compressor,
}
}
pub fn payload(mut self, payload: serde_json::Value) -> Self {
self.inner.payload = payload;
self
}
pub fn priority(mut self, priority: crate::QueuePriority) -> Self {
self.inner.priority = priority;
self
}
pub fn id(mut self, id: String) -> Self {
self.inner.id = id;
self
}
pub fn metadata(mut self, key: String, value: serde_json::Value) -> Self {
let value_str = match value {
serde_json::Value::String(s) => s,
other => other.to_string(),
};
self.inner.attributes.insert(key, value_str);
self
}
pub async fn build(mut self) -> QueueResult<QueueMessage> {
self.compressor.compress_message(&mut self.inner).await?;
Ok(self.inner)
}
}
pub mod utils {
use super::*;
pub async fn estimate_compression_ratio(data: &[u8], algorithm: CompressionAlgorithm) -> Result<f64, CompressionError> {
if data.is_empty() {
return Ok(1.0);
}
let compressor = MessageCompressor::new(CompressionConfig {
algorithm,
level: 6,
min_size: 0,
force_compression: true,
max_attempts: 1,
});
let compressed = compressor.compress_data(data).await?;
Ok(compressed.len() as f64 / data.len() as f64)
}
pub async fn test_compression(payload: &serde_json::Value) -> Vec<(CompressionAlgorithm, f64, usize)> {
let mut results = Vec::new();
let payload_bytes = serde_json::to_vec(payload).unwrap_or_default();
for algorithm in [CompressionAlgorithm::Gzip, CompressionAlgorithm::Zlib] {
if let Ok(ratio) = estimate_compression_ratio(&payload_bytes, algorithm).await {
let compressed_size = (payload_bytes.len() as f64 * ratio) as usize;
results.push((algorithm, ratio, compressed_size));
}
}
results.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
results
}
pub fn get_compression_recommendations(size: usize) -> CompressionConfig {
let algorithm = if size < 1024 {
CompressionAlgorithm::None } else if size < 10240 {
CompressionAlgorithm::Gzip } else {
CompressionAlgorithm::Zlib };
let level = match size {
0..=1024 => 0, 1025..=10240 => 6, 10241..=102400 => 8, _ => 9, };
CompressionConfig {
algorithm,
level,
min_size: 1024,
force_compression: false,
max_attempts: 3,
}
}
}