use std::io::Cursor;
use std::sync::{Arc, Mutex};
use bytes::Bytes;
use flate2::Decompress;
use std::cell::RefCell;
use crate::error::{BlobError, ErrorKind, Result, new_blob_error, new_error};
use super::blob_wire::{BlobData, MAX_BLOB_MESSAGE_SIZE, WireBlob};
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()
&& 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 {
Err(new_blob_error(BlobError::MessageTooBig { size }))
} else {
buf.extend_from_slice(bytes);
Ok(())
}
}
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
&& 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
&& 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(())
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::cast_possible_truncation)]
mod tests {
use super::{
BlobData, MAX_BLOB_MESSAGE_SIZE, WireBlob, decompress_blob, decompress_parsed_blob_into,
decompress_wire_blob_into,
};
use crate::error::{BlobError, ErrorKind};
use bytes::Bytes;
#[derive(Clone, Copy, Debug)]
enum Kind {
Raw,
Zlib,
Zstd,
}
fn zlib_compress(data: &[u8]) -> Vec<u8> {
use flate2::Compression;
use flate2::write::ZlibEncoder;
use std::io::Write as _;
let mut e = ZlibEncoder::new(Vec::new(), Compression::fast());
e.write_all(data).unwrap();
e.finish().unwrap()
}
fn zstd_compress(data: &[u8]) -> Vec<u8> {
zstd::stream::encode_all(std::io::Cursor::new(data), 1).unwrap()
}
fn make_blob(kind: Kind, payload: &[u8]) -> WireBlob {
let raw_size = Some(i32::try_from(payload.len()).unwrap());
match kind {
Kind::Raw => WireBlob {
data: Some(BlobData::Raw(Bytes::copy_from_slice(payload))),
raw_size,
},
Kind::Zlib => WireBlob {
data: Some(BlobData::Zlib(Bytes::from(zlib_compress(payload)))),
raw_size,
},
Kind::Zstd => WireBlob {
data: Some(BlobData::Zstd(Bytes::from(zstd_compress(payload)))),
raw_size,
},
}
}
fn assert_too_big(result: crate::error::Result<()>, expected_size: u64) {
match result.map_err(crate::error::Error::into_kind) {
Err(ErrorKind::Blob(BlobError::MessageTooBig { size })) => {
assert_eq!(
size, expected_size,
"surfaced size must be the payload size"
);
}
other => panic!("expected MessageTooBig, got {other:?}"),
}
}
#[test]
fn decompress_helpers_agree_at_message_size_boundary() {
let cap = MAX_BLOB_MESSAGE_SIZE as usize;
for kind in [Kind::Raw, Kind::Zlib, Kind::Zstd] {
let at_cap = vec![0u8; cap];
let blob = make_blob(kind, &at_cap);
let mut into = Vec::new();
decompress_wire_blob_into(&blob, &mut into)
.unwrap_or_else(|e| panic!("{kind:?} wire_blob_into at cap: {e:?}"));
assert_eq!(into.len(), cap, "{kind:?} wire_blob_into length at cap");
let mut parsed = Vec::new();
decompress_parsed_blob_into(&blob, &mut parsed)
.unwrap_or_else(|e| panic!("{kind:?} parsed_blob_into at cap: {e:?}"));
assert_eq!(parsed.len(), cap, "{kind:?} parsed_blob_into length at cap");
let bytes = decompress_blob(&blob, None)
.unwrap_or_else(|e| panic!("{kind:?} decompress_blob at cap: {e:?}"));
assert_eq!(bytes.len(), cap, "{kind:?} decompress_blob length at cap");
let over_len = cap + 1;
let over = vec![0u8; over_len];
let blob = make_blob(kind, &over);
let mut into = Vec::new();
assert_too_big(decompress_wire_blob_into(&blob, &mut into), over_len as u64);
let mut parsed = Vec::new();
assert_too_big(
decompress_parsed_blob_into(&blob, &mut parsed),
over_len as u64,
);
assert_too_big(decompress_blob(&blob, None).map(|_| ()), over_len as u64);
}
}
}