use std::io::Cursor;
use std::sync::{Arc, Mutex};
use bytes::Bytes;
use flate2::Decompress;
use std::cell::RefCell;
use crate::error::{new_blob_error, new_error, BlobError, ErrorKind, Result};
use super::blob_wire::{BlobData, WireBlob, MAX_BLOB_MESSAGE_SIZE};
thread_local! {
static ZLIB_DECOMPRESS: RefCell<Decompress> = RefCell::new(Decompress::new(true));
}
const MAX_RETAINED_CAPACITY: usize = 4 * 1024 * 1024;
const MAX_POOL_SIZE: usize = 64;
pub(crate) struct DecompressPool {
buffers: Mutex<Vec<Vec<u8>>>,
}
impl DecompressPool {
pub fn new() -> Arc<Self> {
Arc::new(Self {
buffers: Mutex::new(Vec::new()),
})
}
pub fn get(&self) -> Vec<u8> {
self.buffers
.lock()
.ok()
.and_then(|mut v| v.pop())
.unwrap_or_default()
}
fn put(&self, mut buf: Vec<u8>) {
if buf.capacity() > MAX_RETAINED_CAPACITY {
return;
}
buf.clear();
if let Ok(mut v) = self.buffers.lock() {
if v.len() < MAX_POOL_SIZE {
v.push(buf);
}
}
}
}
struct PooledBuffer {
vec: Vec<u8>,
pool: Arc<DecompressPool>,
}
impl AsRef<[u8]> for PooledBuffer {
fn as_ref(&self) -> &[u8] {
&self.vec
}
}
impl Drop for PooledBuffer {
fn drop(&mut self) {
let v = std::mem::take(&mut self.vec);
self.pool.put(v);
}
}
pub(super) fn pool_get(pool: Option<&Arc<DecompressPool>>, capacity: usize) -> Vec<u8> {
match pool {
Some(p) => {
let mut buf = p.get();
buf.reserve(capacity.saturating_sub(buf.capacity()));
buf
}
None => Vec::with_capacity(capacity),
}
}
pub(crate) fn pool_get_pub(pool: &Arc<DecompressPool>, capacity: usize) -> Vec<u8> {
let mut buf = pool.get();
buf.reserve(capacity.saturating_sub(buf.capacity()));
buf
}
pub(crate) fn pool_wrap(decoded: Vec<u8>, pool: Option<&Arc<DecompressPool>>) -> Bytes {
match pool {
Some(p) => Bytes::from_owner(PooledBuffer {
vec: decoded,
pool: Arc::clone(p),
}),
None => Bytes::from(decoded),
}
}
#[allow(clippy::cast_possible_truncation)] fn zlib_decompress_into(compressed: &[u8], buf: &mut Vec<u8>) -> Result<()> {
ZLIB_DECOMPRESS.with_borrow_mut(|decompress| {
decompress.reset(true);
let mut input = compressed;
loop {
if buf.len() == buf.capacity() {
buf.reserve(input.len().max(4096));
}
let before_in = decompress.total_in();
let status = decompress
.decompress_vec(input, buf, flate2::FlushDecompress::None)
.map_err(|e| {
new_error(ErrorKind::Io(std::io::Error::other(format!(
"zlib decompress error: {e}"
))))
})?;
let consumed = (decompress.total_in() - before_in) as usize;
input = &input[consumed..];
if matches!(status, flate2::Status::StreamEnd) {
break;
}
if buf.len() == buf.capacity() {
buf.reserve(buf.len().max(4096));
}
}
let size = buf.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
Ok(())
})
}
pub(crate) fn decompress_blob_data_into(blob_bytes: &[u8], buf: &mut Vec<u8>) -> Result<()> {
let blob = WireBlob::parse_slice(blob_bytes)?;
decompress_parsed_blob_into(&blob, buf)
}
#[allow(clippy::cast_sign_loss)]
pub(super) fn decompress_parsed_blob_into(blob: &WireBlob, buf: &mut Vec<u8>) -> Result<()> {
buf.clear();
match &blob.data {
Some(BlobData::Raw(bytes)) => {
let size = bytes.len() as u64;
if size < MAX_BLOB_MESSAGE_SIZE {
buf.extend_from_slice(bytes);
Ok(())
} else {
Err(new_blob_error(BlobError::MessageTooBig { size }))
}
}
Some(BlobData::Zlib(bytes)) => {
let cap = blob.estimated_capacity();
if cap > 0 {
buf.reserve(cap.saturating_sub(buf.capacity()));
}
zlib_decompress_into(bytes, buf)
}
Some(BlobData::Zstd(bytes)) => {
let cap = blob.estimated_capacity();
if cap > 0 {
buf.reserve(cap.saturating_sub(buf.capacity()));
}
zstd::stream::copy_decode(Cursor::new(&**bytes), &mut *buf)?;
let size = buf.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
Ok(())
}
None => Err(new_blob_error(BlobError::Empty)),
}
}
#[hotpath::measure]
pub(crate) fn decompress_blob_raw(raw_blob: &[u8], buf: &mut Vec<u8>) -> Result<()> {
use super::wire::Cursor;
buf.clear();
let mut cursor = Cursor::new(raw_blob);
let mut raw_size: Option<i32> = None;
let mut found = false;
while let Some((field, wire_type)) = cursor.read_tag()? {
match field {
1 => {
let slice = cursor.read_len_delimited()?;
let size = slice.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
buf.extend_from_slice(slice);
return Ok(());
}
2 => {
#[allow(clippy::cast_possible_truncation)]
{ raw_size = Some(cursor.read_varint()? as i32); }
}
3 => {
let slice = cursor.read_len_delimited()?;
if let Some(rs) = raw_size {
if rs > 0 {
#[allow(clippy::cast_sign_loss)]
buf.reserve((rs as usize).saturating_sub(buf.capacity()));
}
}
zlib_decompress_into(slice, buf)?;
found = true;
}
7 => {
let slice = cursor.read_len_delimited()?;
if let Some(rs) = raw_size {
if rs > 0 {
#[allow(clippy::cast_sign_loss)]
buf.reserve((rs as usize).saturating_sub(buf.capacity()));
}
}
zstd::stream::copy_decode(std::io::Cursor::new(slice), &mut *buf)?;
let size = buf.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
found = true;
}
_ => cursor.skip_field(wire_type)?,
}
}
if found { Ok(()) } else { Err(new_blob_error(BlobError::Empty)) }
}
pub(crate) fn decompress_blob(
blob: &WireBlob,
pool: Option<&Arc<DecompressPool>>,
) -> Result<Bytes> {
match &blob.data {
Some(BlobData::Raw(bytes)) => {
let size = bytes.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
Err(new_blob_error(BlobError::MessageTooBig { size }))
} else {
Ok(bytes.clone())
}
}
Some(BlobData::Zlib(bytes)) => {
let est = blob.estimated_capacity();
let capacity = if est > 0 { est } else { bytes.len() * 4 };
let mut decoded_bytes = pool_get(pool, capacity);
zlib_decompress_into(bytes, &mut decoded_bytes)?;
Ok(pool_wrap(decoded_bytes, pool))
}
Some(BlobData::Zstd(bytes)) => {
let est = blob.estimated_capacity();
let capacity = if est > 0 { est } else { bytes.len() * 4 };
let mut decoded_bytes = pool_get(pool, capacity);
zstd::stream::copy_decode(Cursor::new(&**bytes), &mut decoded_bytes)?;
let size = decoded_bytes.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
Ok(pool_wrap(decoded_bytes, pool))
}
None => Err(new_blob_error(BlobError::Empty)),
}
}
pub(crate) fn decompress_wire_blob_into(blob: &WireBlob, buf: &mut Vec<u8>) -> Result<()> {
buf.clear();
match &blob.data {
Some(BlobData::Raw(bytes)) => {
let size = bytes.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
buf.extend_from_slice(bytes);
}
Some(BlobData::Zlib(bytes)) => {
let est = blob.estimated_capacity();
if est > 0 { buf.reserve(est); }
zlib_decompress_into(bytes, buf)?;
}
Some(BlobData::Zstd(bytes)) => {
let est = blob.estimated_capacity();
if est > 0 { buf.reserve(est); }
zstd::stream::copy_decode(Cursor::new(&**bytes), &mut *buf)?;
let size = buf.len() as u64;
if size > MAX_BLOB_MESSAGE_SIZE {
return Err(new_blob_error(BlobError::MessageTooBig { size }));
}
}
None => return Err(new_blob_error(BlobError::Empty)),
}
Ok(())
}