#![no_std]
extern crate alloc;
#[cfg(any(feature = "std", test))]
extern crate std;
use alloc::boxed::Box;
use alloc::collections::VecDeque;
use alloc::vec::Vec;
use core::future::Future;
use core::pin::Pin;
use core::ptr::copy_nonoverlapping;
mod tree;
use tree::*;
const WORD_SIZE: usize = 8;
const SYMBOL_COUNT: usize = 1 << WORD_SIZE;
static PRESHIFTED7: [u8; 2] = [0b0000_0000, 0b1000_0000];
static PRESHIFTED6: [u8; 2] = [0b0000_0000, 0b0100_0000];
static PRESHIFTED5: [u8; 2] = [0b0000_0000, 0b0010_0000];
static PRESHIFTED4: [u8; 2] = [0b0000_0000, 0b0001_0000];
static PRESHIFTED3: [u8; 2] = [0b0000_0000, 0b0000_1000];
static PRESHIFTED2: [u8; 2] = [0b0000_0000, 0b0000_0100];
static PRESHIFTED1: [u8; 2] = [0b0000_0000, 0b0000_0010];
static PRESHIFTED0: [u8; 2] = [0b0000_0000, 0b0000_0001];
pub struct Encoder {
page_size: usize,
page_threshold: usize,
page_threshold_limit: usize,
page_count: usize,
state: EncodeState,
word_batch: Vec<u8>,
weights: [u32; SYMBOL_COUNT],
code_table: Vec<CodeEntry>,
visit_deque: VecDeque<(Node, Vec<usize>)>,
done: bool,
#[cfg(feature = "ratio")]
bytes_in: usize,
#[cfg(feature = "ratio")]
bytes_out: usize,
}
#[derive(Debug, Copy, Clone)]
enum EncodeState {
Init,
Table,
Data,
Error,
}
impl Encoder {
pub fn new(page_size: usize, page_threshold: usize) -> Encoder {
Encoder {
page_size,
page_threshold: 1, page_threshold_limit: page_threshold,
page_count: 0,
state: EncodeState::Init,
word_batch: Vec::with_capacity(page_size),
weights: [1; SYMBOL_COUNT], code_table: Vec::with_capacity(SYMBOL_COUNT),
visit_deque: VecDeque::with_capacity(SYMBOL_COUNT * 2 - 1),
done: false,
#[cfg(feature = "ratio")]
bytes_in: 0,
#[cfg(feature = "ratio")]
bytes_out: 0,
}
}
pub fn reset(&mut self) {
self.page_count = 0;
self.page_threshold = 1;
self.state = EncodeState::Init;
self.word_batch.clear();
self.weights.fill(1);
self.code_table.clear();
self.done = false;
#[cfg(feature = "ratio")]
{
self.bytes_in = 0;
self.bytes_out = 0;
}
}
#[cfg(feature = "ratio")]
pub fn ratio(&self) -> f32 {
self.bytes_in as f32 / self.bytes_out as f32
}
#[inline(always)]
pub fn batch_fits(&self, bytes: usize) -> bool {
self.word_batch.len() + bytes < self.page_size
}
#[inline(always)]
pub unsafe fn batch_sink(&mut self, bytes: &[u8]) {
#[cfg(feature = "ratio")]
{
self.bytes_in += bytes.len();
}
let weights_ptr = self.weights.as_mut_ptr();
let word_batch_ptr = self.word_batch.as_mut_ptr();
for byte in bytes {
*weights_ptr.add(*byte as usize) += 1;
}
let idx = self.word_batch.len();
let dst = word_batch_ptr.add(idx);
let src = bytes.as_ptr();
copy_nonoverlapping(src, dst, bytes.len());
self.word_batch.set_len(idx + bytes.len());
}
#[inline(always)]
pub async fn sink<E>(&mut self, byte: u8, output: &mut impl PageWriter<E>) -> Result<bool, E> {
self.crank(output, Some(byte)).await
}
pub async fn flush<E>(&mut self, output: &mut impl PageWriter<E>) -> Result<bool, E> {
self.crank(output, None).await
}
#[inline(always)]
async fn crank<E>(
&mut self,
output: &mut impl PageWriter<E>,
byte: Option<u8>,
) -> Result<bool, E> {
let finish = if let Some(byte) = byte {
#[cfg(feature = "ratio")]
{
self.bytes_in += 1;
}
unsafe {
let weights_ptr = self.weights.as_mut_ptr();
let word_batch_ptr = self.word_batch.as_mut_ptr();
*weights_ptr.add(byte as usize) += 1;
let idx = self.word_batch.len();
*word_batch_ptr.add(idx) = byte;
self.word_batch.set_len(idx + 1);
}
if matches!(self.state, EncodeState::Data) & (self.word_batch.len() < self.page_size) {
return Ok(self.done);
}
false
} else {
true
};
loop {
match self.state {
EncodeState::Error => {
unreachable!()
}
EncodeState::Init => {
if self.word_batch.len() == self.page_size || finish {
self.state = EncodeState::Table;
output.reset();
} else {
return Ok(self.done);
}
}
EncodeState::Table => {
for &weight in &self.weights {
output.write_u32le(weight as u32);
}
#[cfg(feature = "ratio")]
{
self.bytes_out += self.page_size;
}
match output.flush().await {
Ok(done) => {
self.done |= done;
}
Err(e) => {
self.state = EncodeState::Error;
return Err(e);
}
}
let root = build_tree(&self.weights);
self.code_table.clear();
self.code_table.resize(SYMBOL_COUNT, Default::default());
build_code_table(root, &mut self.code_table, &mut self.visit_deque);
self.weights.fill(1);
self.state = EncodeState::Data;
}
EncodeState::Data => {
if (self.word_batch.len() < self.page_size) & !finish {
return Ok(self.done);
}
let mut drain = None;
for (idx, byte) in self.word_batch.iter().enumerate() {
let code = unsafe { self.code_table.get_unchecked(*byte as usize) };
if !output.write_code(code) {
drain = Some(idx);
break;
}
}
if let Some(drain) = drain {
self.word_batch.drain(..drain);
#[cfg(feature = "ratio")]
{
self.bytes_out += self.page_size;
}
match output.flush().await {
Ok(done) => {
self.done |= done;
}
Err(e) => {
self.state = EncodeState::Error;
return Err(e);
}
}
self.page_count += 1;
if self.page_count > self.page_threshold {
self.page_count = 0;
self.page_threshold =
self.page_threshold_limit.min(self.page_threshold * 2);
self.state = EncodeState::Table;
} else {
self.state = EncodeState::Data;
}
} else {
self.word_batch.clear();
if finish {
#[cfg(feature = "ratio")]
{
self.bytes_out += self.page_size;
}
return output.flush().await.and_then(|done| {
self.done |= done;
Ok(self.done)
});
}
}
}
}
}
}
}
#[allow(async_fn_in_trait)]
pub trait PageWriter<E> {
async fn flush(&mut self) -> Result<bool, E>;
fn position(&self) -> usize;
fn reset(&mut self);
fn write_header(&mut self, header: u32);
fn write_u32le(&mut self, value: u32);
fn write_code(&mut self, code: &CodeEntry) -> bool;
}
pub struct BufferedPageWriter<E> {
bits_written: usize,
page_size: usize,
bits: Vec<usize>,
bytes: Vec<u8>,
flush_page: WritePageFutureFn<E>,
done: bool,
}
pub type WritePageFutureFn<E> =
Box<dyn for<'a> Fn(&'a [u8]) -> Pin<Box<dyn Future<Output = Result<bool, E>> + 'a>>>;
impl<E> BufferedPageWriter<E> {
pub fn new(page_size: usize, flush: WritePageFutureFn<E>) -> BufferedPageWriter<E> {
let mut buf: Vec<u8> = Vec::with_capacity(page_size);
unsafe { buf.set_len(page_size) };
BufferedPageWriter {
page_size: 8 * page_size,
bits: Vec::with_capacity(8 * 2),
bytes: buf,
bits_written: 0,
flush_page: flush,
done: false,
}
}
}
impl<E> PageWriter<E> for BufferedPageWriter<E> {
#[inline(always)]
fn position(&self) -> usize {
self.bits_written
}
fn reset(&mut self) {
self.bits_written = 32;
}
#[inline(always)]
fn write_header(&mut self, header: u32) {
let bits_written_header = header.to_le_bytes();
unsafe {
copy_nonoverlapping(
bits_written_header.as_ptr(),
self.bytes.as_mut_ptr(),
bits_written_header.len(),
);
}
}
#[inline(always)]
fn write_u32le(&mut self, value: u32) {
debug_assert!(self.bits.len() == 0);
let bytes = value.to_le_bytes();
let offset = self.bits_written / 8;
unsafe {
copy_nonoverlapping(
bytes.as_ptr(),
self.bytes.as_mut_ptr().add(offset),
bytes.len(),
);
}
self.bits_written += bytes.len() * 8;
}
#[inline(always)]
fn write_code(&mut self, code: &CodeEntry) -> bool {
let position = self.position();
let pending = self.bits.len();
debug_assert!(pending < 8); if position + pending + code.bits.len() > self.page_size {
if pending > 0 {
let padding = 8 - pending;
self.bits.extend((0..padding).map(|_| 0));
for byte_bits in self.bits.chunks_exact(8) {
let byte_bits: &[usize; 8] = unsafe { byte_bits.try_into().unwrap_unchecked() };
let byte = unsafe {
PRESHIFTED7.get_unchecked(byte_bits[0])
| PRESHIFTED6.get_unchecked(byte_bits[1])
| PRESHIFTED5.get_unchecked(byte_bits[2])
| PRESHIFTED4.get_unchecked(byte_bits[3])
| PRESHIFTED3.get_unchecked(byte_bits[4])
| PRESHIFTED2.get_unchecked(byte_bits[5])
| PRESHIFTED1.get_unchecked(byte_bits[6])
| PRESHIFTED0.get_unchecked(byte_bits[7])
};
let offset = self.bits_written / 8;
#[cfg(test)]
{
self.bytes[offset] = byte;
}
#[cfg(not(test))]
{
unsafe {
*self.bytes.get_unchecked_mut(offset) = byte;
}
}
self.bits_written += 8;
}
self.bits_written -= padding;
self.bits.clear();
}
return false;
}
self.bits.extend(code.bits.iter());
let mut drained = 0;
for byte_bits in self.bits.chunks_exact(8) {
drained += 8;
let byte_bits: &[usize; 8] = unsafe { byte_bits.try_into().unwrap_unchecked() };
let byte = unsafe {
PRESHIFTED7.get_unchecked(byte_bits[0])
| PRESHIFTED6.get_unchecked(byte_bits[1])
| PRESHIFTED5.get_unchecked(byte_bits[2])
| PRESHIFTED4.get_unchecked(byte_bits[3])
| PRESHIFTED3.get_unchecked(byte_bits[4])
| PRESHIFTED2.get_unchecked(byte_bits[5])
| PRESHIFTED1.get_unchecked(byte_bits[6])
| PRESHIFTED0.get_unchecked(byte_bits[7])
};
let offset = self.bits_written / 8;
#[cfg(test)]
{
self.bytes[offset] = byte;
}
#[cfg(not(test))]
{
unsafe {
*self.bytes.get_unchecked_mut(offset) = byte;
}
}
self.bits_written += 8;
}
self.bits.drain(..drained);
true
}
async fn flush(&mut self) -> Result<bool, E> {
debug_assert!(self.bits.len() < 8); if !self.bits.is_empty() {
let padding = 8 - self.bits.len();
self.bits.extend((0..padding).map(|_| 0));
for byte_bits in self.bits.chunks_exact(8) {
let byte_bits: &[usize; 8] = unsafe { byte_bits.try_into().unwrap_unchecked() };
let byte = unsafe {
PRESHIFTED7.get_unchecked(byte_bits[0])
| PRESHIFTED6.get_unchecked(byte_bits[1])
| PRESHIFTED5.get_unchecked(byte_bits[2])
| PRESHIFTED4.get_unchecked(byte_bits[3])
| PRESHIFTED3.get_unchecked(byte_bits[4])
| PRESHIFTED2.get_unchecked(byte_bits[5])
| PRESHIFTED1.get_unchecked(byte_bits[6])
| PRESHIFTED0.get_unchecked(byte_bits[7])
};
let offset = self.bits_written / 8;
#[cfg(test)]
{
self.bytes[offset] = byte;
}
#[cfg(not(test))]
{
unsafe {
*self.bytes.get_unchecked_mut(offset) = byte;
}
}
self.bits_written += 8;
}
self.bits_written -= padding;
self.bits.clear();
}
self.write_header(self.bits_written as u32);
self.done |= !(*self.flush_page)(&self.bytes).await?;
self.reset();
Ok(self.done)
}
}
pub struct Decoder {
page_size: usize,
page_threshold: usize,
page_threshold_limit: usize,
page_count: usize,
state: DecodeState,
decoder_trie: Option<CodeLookupTrie>,
decoded_bytes: Vec<u8>,
emitted_idx: usize,
}
#[derive(Debug, Copy, Clone)]
enum DecodeState {
Table,
Data,
Done,
Error,
}
impl Decoder {
pub fn new(page_size: usize, page_threshold: usize) -> Decoder {
Decoder {
page_size,
page_threshold: 1,
page_threshold_limit: page_threshold,
page_count: 0,
state: DecodeState::Table,
decoder_trie: None,
decoded_bytes: Vec::with_capacity(page_size),
emitted_idx: 0,
}
}
pub fn reset(&mut self) {
self.page_count = 0;
self.page_threshold = 1;
self.state = DecodeState::Table;
self.decoder_trie = None;
self.decoded_bytes.clear();
self.emitted_idx = 0;
}
pub async fn drain<E>(
&mut self,
input: &mut impl PageReader<E>,
) -> Result<Option<u8>, DecompressionError<E>> {
if self.emitted_idx < self.decoded_bytes.len() {
let byte = if cfg!(test) {
self.decoded_bytes[self.emitted_idx]
} else {
unsafe { *self.decoded_bytes.get_unchecked(self.emitted_idx) }
};
self.emitted_idx += 1;
return Ok(Some(byte));
}
loop {
match self.state {
DecodeState::Done => {
return Ok(None);
}
DecodeState::Error => {
return Err(DecompressionError::Bad);
}
DecodeState::Table => {
let page = input.read_page().await?;
if page[..4] == [0xFF; 4] {
self.state = DecodeState::Done;
return Ok(None);
}
let mut weights = [0u32; SYMBOL_COUNT];
debug_assert!(page.len() >= SYMBOL_COUNT * 4 + 4); unsafe {
let weights_sz = SYMBOL_COUNT * 4; let page_weights_ptr = page.as_ptr().add(4); let weights_ptr = weights.as_mut_ptr() as *mut u8;
core::ptr::copy_nonoverlapping(page_weights_ptr, weights_ptr, weights_sz);
}
let root = build_tree(&weights);
self.decoder_trie = Some(CodeLookupTrie::new(root));
self.state = DecodeState::Data;
}
DecodeState::Data => {
self.emitted_idx = 0;
self.decoded_bytes.clear();
let page = input.read_page().await?;
if page[..4] == [0xFF; 4] {
self.state = DecodeState::Done;
return Ok(None);
}
let symbol_lookup = self.decoder_trie.as_mut().unwrap();
let bits_written = u32::from_le_bytes(page[..4].try_into().unwrap());
if !(32..=self.page_size * 8).contains(&(bits_written as usize)) {
self.state = DecodeState::Error;
return Err(DecompressionError::Bad);
}
let bytes_written = ((bits_written + 7) / 8) as usize;
let page_bytes = &page[4..bytes_written];
if !page_bytes.is_empty() {
let mut bits_read = 32;
let full_bytes = page_bytes.len() - 1;
for &byte in &page_bytes[..full_bytes] {
for i in (0..8).rev() {
let bit = (byte >> i) & 1;
if let Some(symbol) = symbol_lookup.next(bit) {
self.decoded_bytes.push(symbol);
}
}
}
bits_read += full_bytes as u32 * 8;
if let Some(&last_byte) = page_bytes.last() {
let remaining_bits = (bits_written - bits_read) as usize;
for i in (0..8).rev().take(remaining_bits) {
let bit = (last_byte >> i) & 1;
if let Some(symbol) = symbol_lookup.next(bit) {
self.decoded_bytes.push(symbol);
}
}
}
}
self.page_count += 1;
if self.page_count > self.page_threshold {
self.page_count = 0;
self.page_threshold =
self.page_threshold_limit.min(self.page_threshold * 2);
self.state = DecodeState::Table;
} else {
self.state = DecodeState::Data;
}
if self.emitted_idx < self.decoded_bytes.len() {
let byte = if cfg!(test) {
self.decoded_bytes[self.emitted_idx]
} else {
unsafe { *self.decoded_bytes.get_unchecked(self.emitted_idx) }
};
self.emitted_idx += 1;
return Ok(Some(byte));
}
}
}
}
}
}
#[allow(async_fn_in_trait)]
pub trait PageReader<E> {
async fn read_page(&mut self) -> Result<&[u8], E>;
fn reset(&mut self);
}
pub struct BufferedPageReader<E> {
bytes: Vec<u8>,
read_page: ReadPageFutureFn<E>,
done: bool,
}
pub type ReadPageFutureFn<E> =
Box<dyn for<'a> Fn(&'a mut [u8]) -> Pin<Box<dyn Future<Output = Result<bool, E>> + 'a>>>;
impl<E> BufferedPageReader<E> {
pub fn new(page_size: usize, read_page: ReadPageFutureFn<E>) -> BufferedPageReader<E> {
let mut bytes = Vec::with_capacity(page_size);
unsafe { bytes.set_len(page_size) };
BufferedPageReader {
bytes,
read_page,
done: false,
}
}
}
impl<E> PageReader<E> for BufferedPageReader<E> {
async fn read_page(&mut self) -> Result<&[u8], E> {
if self.done {
self.bytes.fill(0xFF);
return Ok(&self.bytes);
}
self.done |= !(*self.read_page)(&mut self.bytes).await?;
Ok(&self.bytes)
}
fn reset(&mut self) {
self.done = false;
}
}
#[derive(Debug, Clone, Copy)]
pub enum DecompressionError<E> {
Bad,
Load(E),
}
impl<E> From<E> for DecompressionError<E> {
fn from(err: E) -> Self {
DecompressionError::Load(err)
}
}
#[cfg(test)]
mod tests {
use super::*;
use core::cell::RefCell;
use std::prelude::v1::*;
use std::rc::Rc;
use std::vec;
use std::vec::Vec;
#[test]
fn test_std_vec() {
let mut vec = Vec::new();
vec.push(1);
vec.push(2);
vec.push(3);
assert_eq!(vec, vec![1, 2, 3]);
}
#[test]
fn test_flush_fn() {
let mut buf = [1, 3, 5, 7];
let flush: WritePageFutureFn<()> = Box::new(|page: &[u8]| {
Box::pin(async move {
std::dbg!("flush", page.len());
Ok(true)
})
});
smol::block_on(async {
(*flush)(&mut buf).await.unwrap();
assert_eq!(buf, [1, 3, 5, 7]);
});
}
#[test]
fn test_page_writer_advance() {
let flush: WritePageFutureFn<()> = Box::new(|page| {
Box::pin(async move {
std::dbg!("flush", page.len());
Ok(true)
})
});
let mut wtr = BufferedPageWriter::new(2048, flush);
smol::block_on(async {
wtr.flush().await.unwrap();
});
}
#[test]
fn test_compress_simple() {
let flush: WritePageFutureFn<()> = Box::new(|page| {
Box::pin(async move {
std::dbg!("flush", page.len());
Ok(true)
})
});
let mut wtr = BufferedPageWriter::new(2048, flush);
let mut encoder = Encoder::new(2048, 4);
smol::block_on(async {
for value in 0..2048 {
encoder.sink(value as u8, &mut wtr).await.unwrap();
}
});
}
#[test]
fn test_compress_multi_page() {
let flush: WritePageFutureFn<()> = Box::new(|_page| Box::pin(async move { Ok(true) }));
let mut wtr = BufferedPageWriter::new(2048, flush);
let mut encoder = Encoder::new(2048, 4);
smol::block_on(async {
for value in 0..2048 * 3 {
encoder.sink(value as u8, &mut wtr).await.unwrap();
}
encoder.flush(&mut wtr).await.unwrap();
});
#[cfg(feature = "ratio")]
{
std::dbg!(
encoder.bytes_in,
encoder.bytes_out,
encoder.bytes_in as f32 / encoder.bytes_out as f32
);
}
}
#[test]
fn test_roundtrip() {
let buf: Vec<u8> = Vec::new();
let buf = Rc::new(RefCell::new(buf));
let wtr_buf = buf.clone();
let rdr_buf = buf.clone();
let flush_page: WritePageFutureFn<()> = Box::new(move |page| {
let buf = wtr_buf.clone();
Box::pin(async move {
let mut buf = buf.borrow_mut();
buf.extend_from_slice(page);
Ok(true)
})
});
const PAGE_SIZE: usize = 2048;
const PAGE_THRESHOLD: usize = 4;
let mut wtr = BufferedPageWriter::new(PAGE_SIZE, flush_page);
let mut encoder = Encoder::new(PAGE_SIZE, PAGE_THRESHOLD);
let read_page: ReadPageFutureFn<()> = Box::new(move |page| {
let buf = rdr_buf.clone();
Box::pin(async move {
let mut buf = buf.borrow_mut();
assert!(buf.len() % PAGE_SIZE == 0);
if buf.is_empty() {
page.fill(0xFF);
Ok(false)
} else {
let drained = buf.drain(..PAGE_SIZE);
page[..drained.len()]
.iter_mut()
.zip(drained)
.for_each(|(p, b)| *p = b);
Ok(true)
}
})
});
let mut rdr = BufferedPageReader::new(PAGE_SIZE, read_page);
let mut decoder = Decoder::new(PAGE_SIZE, PAGE_THRESHOLD);
let bad_rand = (0..100)
.map(|i| vec![0; i])
.collect::<Vec<_>>()
.into_iter()
.map(|v| {
let ptr = v.as_ptr();
(ptr, v.len())
})
.fold(0, |acc, (ptr, len)| acc + ptr as usize * len * 31)
% 9999991
+ 5123457;
std::dbg!(bad_rand);
let test_cases: Vec<Vec<u8>> = vec![
vec![],
(0..10).collect::<Vec<_>>(),
(0..2048).map(|i| i as u8).collect::<Vec<_>>(),
(0..2048 * 3).map(|i| i as u8).collect::<Vec<_>>(),
(0..2048 * 4).map(|i| i as u8).collect::<Vec<_>>(),
(0..1024 * 1024).map(|i| i as u8).collect::<Vec<_>>(),
(0..bad_rand).map(|i| (31 * i) as u8).collect::<Vec<_>>(),
(0..bad_rand)
.map(|i| ((31 * i) % 16) as u8)
.collect::<Vec<_>>(),
];
#[cfg(feature = "ratio")]
let mut compression_ratios = Vec::new();
for (test_case, test_data) in test_cases.into_iter().enumerate() {
std::dbg!(test_case);
buf.borrow_mut().clear();
encoder.reset();
wtr.reset();
decoder.reset();
rdr.reset();
smol::block_on(async {
for value in &test_data {
encoder.sink(*value, &mut wtr).await.unwrap();
}
encoder.flush(&mut wtr).await.unwrap();
});
std::dbg!(buf.borrow().len());
let num_pages = (buf.borrow().len() + PAGE_SIZE - 1) / PAGE_SIZE;
for page in 0..num_pages {
let header_offset = page * PAGE_SIZE;
let header = u32::from_le_bytes(
buf.borrow()[header_offset..header_offset + 4]
.try_into()
.unwrap(),
);
std::dbg!(header);
}
#[cfg(feature = "ratio")]
{
compression_ratios.push((
encoder.bytes_in as f32 / encoder.bytes_out as f32,
humanize_bytes::humanize_bytes_binary!(encoder.bytes_in),
humanize_bytes::humanize_bytes_binary!(encoder.bytes_out),
));
}
smol::block_on(async {
let mut idx = 0;
while let Some(byte) = decoder.drain(&mut rdr).await.unwrap() {
assert_eq!(
byte, test_data[idx],
"test case {} byte {} mismatch",
test_case, idx
);
idx += 1;
}
assert_eq!(idx, test_data.len());
});
}
#[cfg(feature = "ratio")]
{
std::dbg!(compression_ratios);
}
}
}