use crate::errors::EnkryptitError;
use crate::types::CHUNK_SIZE;
use crate::types::CompressionType::{self, Auto, Lz4, NoComp, Xz, Zstd};
use lz4_flex::block::{compress_into as lz4compress, decompress_into as lz4decompress};
use std::io::Read;
use xz2::read::{XzDecoder, XzEncoder};
use zstd::bulk::compress_to_buffer as zstdcompress;
use zstd::bulk::decompress_to_buffer as zstddecompress;
pub trait EnkryptitCompress {
fn compress(
&self,
output: &mut Vec<u8>,
compression: CompressionType,
) -> Result<(), EnkryptitError>;
}
pub trait EnkryptitDecompress {
fn decompress(
&self,
output: &mut Vec<u8>,
compression: CompressionType,
) -> Result<(), EnkryptitError>;
}
impl EnkryptitCompress for Vec<u8> {
fn compress(
&self,
output: &mut Vec<u8>,
compression: CompressionType,
) -> Result<(), EnkryptitError> {
self.as_slice().compress(output, compression)
}
}
impl EnkryptitCompress for [u8] {
fn compress(
&self,
output: &mut Vec<u8>,
compression: CompressionType,
) -> Result<(), EnkryptitError> {
match compression {
Auto => unreachable!(
"`Auto` should never be reached here, and always infered before reaching this function. There is an error in the code. If you are reading this as an user, please open an Issue."
),
Zstd => compress_with_zstd(self, output),
Lz4 => compress_with_lz4(self, output),
Xz => compress_with_xz(self, output),
NoComp => {
output.clear();
output.extend_from_slice(self);
Ok(())
}
}
}
}
impl EnkryptitDecompress for Vec<u8> {
fn decompress(
&self,
output: &mut Vec<u8>,
compression: CompressionType,
) -> Result<(), EnkryptitError> {
self.as_slice().decompress(output, compression)
}
}
impl EnkryptitDecompress for [u8] {
fn decompress(
&self,
output: &mut Vec<u8>,
compression: CompressionType,
) -> Result<(), EnkryptitError> {
match compression {
Auto => unreachable!(
"`Auto` should never be reached here, and always infered before reaching this function. There is an error in the code. If you are reading this as an user, please open an Issue."
),
Zstd => decompress_with_zstd(self, output),
Lz4 => decompress_with_lz4(self, output),
Xz => decompress_with_xz(self, output),
NoComp => {
output.clear();
output.extend_from_slice(self);
Ok(())
}
}
}
}
fn compress_with_lz4(input: &[u8], output: &mut Vec<u8>) -> Result<(), EnkryptitError> {
let max_compressed_size = lz4_flex::block::get_maximum_output_size(input.len());
output.clear();
output.reserve(max_compressed_size);
output.resize(max_compressed_size, 0);
let actual_size = lz4compress(input, output)?;
output.truncate(actual_size);
Ok(())
}
fn decompress_with_lz4(input: &[u8], output: &mut Vec<u8>) -> Result<(), EnkryptitError> {
output.clear();
let mut current_size = std::cmp::max(input.len() * 4, CHUNK_SIZE / 16);
for _ in 0..5 {
output.resize(current_size, 0);
match lz4decompress(input, output) {
Ok(actual_size) => {
output.truncate(actual_size);
return Ok(());
}
Err(_) => {
current_size *= 2;
if current_size > CHUNK_SIZE * 4 {
output.resize(current_size, 0);
let actual_size = lz4decompress(input, output)?;
output.truncate(actual_size);
return Ok(());
}
}
}
}
output.resize(CHUNK_SIZE * 8, 0);
let actual_size = lz4decompress(input, output)?;
output.truncate(actual_size);
Ok(())
}
fn compress_with_zstd(input: &[u8], output: &mut Vec<u8>) -> Result<(), EnkryptitError> {
output.clear();
let mut current_size = std::cmp::max(input.len() + 20, CHUNK_SIZE / 16);
for _ in 0..5 {
output.resize(current_size, 0);
match zstdcompress(input, output, 6) {
Ok(size) => {
output.truncate(size);
return Ok(());
}
Err(_) => {
current_size *= 2;
if current_size > CHUNK_SIZE * 4 {
output.resize(current_size, 0);
let size = zstdcompress(input, output, 6)?;
output.truncate(size);
return Ok(());
}
}
}
}
output.resize(CHUNK_SIZE * 8, 0);
let size = zstdcompress(input, output, 6)?;
output.truncate(size);
Ok(())
}
fn decompress_with_zstd(input: &[u8], output: &mut Vec<u8>) -> Result<(), EnkryptitError> {
output.clear();
let mut current_size = std::cmp::max(input.len() * 4, CHUNK_SIZE / 16);
for _ in 0..5 {
output.resize(current_size, 0);
match zstddecompress(input, output) {
Ok(size) => {
output.truncate(size);
return Ok(());
}
Err(_) => {
current_size *= 2;
if current_size > CHUNK_SIZE * 4 {
output.resize(current_size, 0);
let size = zstddecompress(input, output)?;
output.truncate(size);
return Ok(());
}
}
}
}
output.resize(CHUNK_SIZE * 8, 0);
let size = zstddecompress(input, output)?;
output.truncate(size);
Ok(())
}
fn compress_with_xz(input: &[u8], output: &mut Vec<u8>) -> Result<(), EnkryptitError> {
output.clear();
let mut encoder = XzEncoder::new(input, 6);
encoder.read_to_end(output)?;
Ok(())
}
fn decompress_with_xz(input: &[u8], output: &mut Vec<u8>) -> Result<(), EnkryptitError> {
output.clear();
let mut decoder = XzDecoder::new(input);
decoder.read_to_end(output)?;
Ok(())
}