use super::atmost::AtMost;
use super::{Encode, EncodingStrategy};
use crate::{Incompressible, Small, Sorted};
#[cfg(test)]
use expect_test::expect;
impl Encode for u8 {
type Context = <AtMost<255> as Encode>::Context;
#[inline]
fn encode<E: super::EntropyCoder>(&self, writer: &mut E, ctx: &mut Self::Context) {
AtMost::<255>::new(*self as usize).encode(writer, ctx)
}
#[inline]
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<Self, std::io::Error> {
Ok(usize::from(AtMost::<255>::decode(reader, ctx)?) as u8)
}
}
impl Encode for i8 {
type Context = <u8 as Encode>::Context;
#[inline]
fn encode<E: super::EntropyCoder>(&self, writer: &mut E, ctx: &mut Self::Context) {
(*self as u8).encode(writer, ctx)
}
#[inline]
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<Self, std::io::Error> {
<u8 as Encode>::decode(reader, ctx).map(|v| v as i8)
}
}
#[derive(Default, Clone)]
pub struct SmallContext {
nonzero: <AtMost<7> as Encode>::Context,
b1: <AtMost<1> as Encode>::Context,
b2: <AtMost<3> as Encode>::Context,
b3: <AtMost<7> as Encode>::Context,
b4: <AtMost<15> as Encode>::Context,
b5: <AtMost<31> as Encode>::Context,
need_seven_bits: <bool as Encode>::Context,
b6: <AtMost<63> as Encode>::Context,
b7: <AtMost<127> as Encode>::Context,
}
impl EncodingStrategy<u8> for Small {
type Context = SmallContext;
fn encode<E: super::EntropyCoder>(value: &u8, writer: &mut E, ctx: &mut Self::Context) {
let bucket = |code: usize| AtMost::<7>::new(code);
let rest = |first: u8| (*value - first) as usize;
match *value {
0 => bucket(0).encode(writer, &mut ctx.nonzero),
1 => bucket(1).encode(writer, &mut ctx.nonzero),
2..4 => {
bucket(2).encode(writer, &mut ctx.nonzero);
AtMost::<1>::new(rest(2)).encode(writer, &mut ctx.b1)
}
4..8 => {
bucket(3).encode(writer, &mut ctx.nonzero);
AtMost::<3>::new(rest(4)).encode(writer, &mut ctx.b2)
}
8..16 => {
bucket(4).encode(writer, &mut ctx.nonzero);
AtMost::<7>::new(rest(8)).encode(writer, &mut ctx.b3)
}
16..32 => {
bucket(5).encode(writer, &mut ctx.nonzero);
AtMost::<15>::new(rest(16)).encode(writer, &mut ctx.b4)
}
32..64 => {
bucket(6).encode(writer, &mut ctx.nonzero);
AtMost::<31>::new(rest(32)).encode(writer, &mut ctx.b5)
}
64..128 => {
bucket(7).encode(writer, &mut ctx.nonzero);
false.encode(writer, &mut ctx.need_seven_bits);
AtMost::<63>::new(rest(64)).encode(writer, &mut ctx.b6)
}
128..=255 => {
bucket(7).encode(writer, &mut ctx.nonzero);
true.encode(writer, &mut ctx.need_seven_bits);
AtMost::<127>::new(rest(128)).encode(writer, &mut ctx.b7)
}
}
}
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<u8, std::io::Error> {
fn rest<const MAX: usize, D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut <AtMost<MAX> as Encode>::Context,
) -> Result<u8, std::io::Error> {
Ok(usize::from(AtMost::<MAX>::decode(reader, ctx)?) as u8)
}
match usize::from(AtMost::<7>::decode(reader, &mut ctx.nonzero)?) {
0 => Ok(0),
1 => Ok(1),
2 => Ok(rest::<1, D>(reader, &mut ctx.b1)? + 2),
3 => Ok(rest::<3, D>(reader, &mut ctx.b2)? + 4),
4 => Ok(rest::<7, D>(reader, &mut ctx.b3)? + 8),
5 => Ok(rest::<15, D>(reader, &mut ctx.b4)? + 16),
6 => Ok(rest::<31, D>(reader, &mut ctx.b5)? + 32),
7 => {
if <bool as Encode>::decode(reader, &mut ctx.need_seven_bits)? {
Ok(rest::<127, D>(reader, &mut ctx.b7)? + 128)
} else {
Ok(rest::<63, D>(reader, &mut ctx.b6)? + 64)
}
}
_ => unreachable!(),
}
}
}
impl EncodingStrategy<i8> for Small {
type Context = SmallContext;
fn encode<E: super::EntropyCoder>(value: &i8, writer: &mut E, ctx: &mut Self::Context) {
let v = *value as u8;
let zigzag = (v << 1) ^ (0u8.wrapping_sub(v >> 7));
Small::encode(&zigzag, writer, ctx)
}
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<i8, std::io::Error> {
let z = <Small as EncodingStrategy<u8>>::decode(reader, ctx)?;
Ok(((z >> 1) as i8) ^ (-((z & 1) as i8)))
}
}
impl EncodingStrategy<u8> for Incompressible {
type Context = ();
fn encode<E: super::EntropyCoder>(value: &u8, writer: &mut E, _ctx: &mut Self::Context) {
writer.encode_incompressible_bytes(&[*value])
}
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
_ctx: &mut Self::Context,
) -> Result<u8, std::io::Error> {
let mut byte = [0u8];
reader.decode_incompressible_bytes(&mut byte)?;
Ok(byte[0])
}
}
#[derive(Default, Clone)]
pub struct SortedU8Context {
previous: Option<u8>,
delta: <Small as EncodingStrategy<i8>>::Context,
}
impl EncodingStrategy<u8> for Sorted {
type Context = SortedU8Context;
fn encode<E: super::EntropyCoder>(value: &u8, writer: &mut E, ctx: &mut Self::Context) {
if let Some(previous) = ctx.previous.take() {
Small::encode(
&(value.wrapping_sub(previous) as i8),
writer,
&mut ctx.delta,
);
} else {
writer.encode_incompressible_bytes(&[*value]);
}
ctx.previous = Some(*value);
}
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<u8, std::io::Error> {
let out = if let Some(previous) = ctx.previous.take() {
let delta: i8 = Small::decode(reader, &mut ctx.delta)?;
previous.wrapping_add(delta as u8)
} else {
let mut byte = [0u8];
reader.decode_incompressible_bytes(&mut byte)?;
byte[0]
};
ctx.previous = Some(out);
Ok(out)
}
}
impl EncodingStrategy<i8> for Sorted {
type Context = SortedU8Context;
fn encode<E: super::EntropyCoder>(value: &i8, writer: &mut E, ctx: &mut Self::Context) {
Sorted::encode(&(*value as u8), writer, ctx)
}
fn decode<D: super::EntropyDecoder>(
reader: &mut D,
ctx: &mut Self::Context,
) -> Result<i8, std::io::Error> {
<Sorted as EncodingStrategy<u8>>::decode(reader, ctx).map(|v| v as i8)
}
}
#[test]
fn size() {
use super::{assert_bits_all, estimated_bits};
expect!["8"].assert_eq(&estimated_bits!(u8::MAX));
expect!["8"].assert_eq(&estimated_bits!(0_u8));
assert_bits_all!(3_u8..255, expect!["8"]);
expect!["31"].assert_eq(&estimated_bits!(*b"hello"));
expect!["68"].assert_eq(&estimated_bits!(*b"hello world"));
expect!["129"].assert_eq(&estimated_bits!(*b"hello world, hello world"));
expect!["111"].assert_eq(&estimated_bits!(*b"hello hello, hello hello"));
expect!["195"].assert_eq(&estimated_bits!(
*b"hello hello, hello hello, hello hello, hello hello"
));
expect!["37"].assert_eq(&estimated_bits!(*b"hhhhhhhhhhhhhhhhhhhhhhhh"));
expect!["44"].assert_eq(&estimated_bits!(
*b"hhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhhh"
));
expect!["8"].assert_eq(&estimated_bits!(*b"\0"));
expect!["8"].assert_eq(&estimated_bits!(*b"\x01"));
expect!["13"].assert_eq(&estimated_bits!(*b"\x01\x01"));
expect!["19"].assert_eq(&estimated_bits!(*b"\x01\x01\x01\x01"));
expect!["21"].assert_eq(&estimated_bits!(*b"\x01\x01\x01\x01\x01"));
expect!["22"].assert_eq(&estimated_bits!(*b"\x01\x01\x01\x01\x01\x01"));
expect!["25"].assert_eq(&estimated_bits!(*b"\x01\x02\x03\x04"));
expect!["30"].assert_eq(&estimated_bits!(*b"\x01\x02\x03\x04\x05"));
expect!["35"].assert_eq(&estimated_bits!(*b"\x01\x02\x03\x04\x05\x06"));
expect!["40"].assert_eq(&estimated_bits!(*b"\x01\x02\x03\x04\x05\x06\x07"));
expect!["47"].assert_eq(&estimated_bits!(*b"\x01\x02\x03\x04\x05\x06\x07\x08"));
expect!["8"].assert_eq(&estimated_bits!(i8::MAX));
expect!["8"].assert_eq(&estimated_bits!(0_i8));
}
#[test]
fn small() {
use super::Small;
use crate::Encoded;
fn size_of(vals: impl IntoIterator<Item = u8>) -> String {
let mut sizes = vals.into_iter().map(|v| {
println!("Checking {v}");
let bits = super::encoded_bits!(Encoded::<u8, Small>::new(v));
assert_eq!(
Encoded::<u8, Small>::new(v).millibits(),
super::Millibits::bits(bits.parse().unwrap()),
"millibits estimate disagrees for {v}"
);
(v, bits)
});
let (_, bits) = sizes.next().expect("size_of needs at least one value");
for (v, other) in sizes {
assert_eq!(other, bits, "encoded size differs for {v}");
}
bits
}
expect!["3"].assert_eq(&size_of(0..2));
expect!["4"].assert_eq(&size_of(2..4));
expect!["5"].assert_eq(&size_of(4..8));
expect!["6"].assert_eq(&size_of(8..16));
expect!["7"].assert_eq(&size_of(16..32));
expect!["8"].assert_eq(&size_of(32..64));
expect!["10"].assert_eq(&size_of(64..128));
expect!["11"].assert_eq(&size_of(128..255));
assert_eq!(
Encoded::<u8, Small>::new(255u8).millibits(),
super::Millibits::bits(11)
);
}
#[test]
fn small_i8() {
use super::Small;
use crate::Encoded;
for v in i8::MIN..=i8::MAX {
let enc = super::encode(&Encoded::<i8, Small>::new(v));
let dec = super::decode::<Encoded<i8, Small>>(&enc).unwrap().value();
assert_eq!(v, dec, "round-trip failed for {v}");
}
fn size_of(vals: impl IntoIterator<Item = i8>) -> String {
let mut sizes = vals.into_iter().map(|v| {
println!("Checking {v}");
(v, super::estimated_bits!(Encoded::<i8, Small>::new(v)))
});
let (_, bits) = sizes.next().expect("size_of needs at least one value");
for (v, other) in sizes {
assert_eq!(other, bits, "encoded size differs for {v}");
}
bits
}
expect!["3"].assert_eq(&size_of([0]));
expect!["3"].assert_eq(&size_of([-1]));
expect!["4"].assert_eq(&size_of([1]));
expect!["4"].assert_eq(&size_of([-2]));
expect!["5"].assert_eq(&size_of([2i8, 3, -3, -4]));
expect!["6"].assert_eq(&size_of([4i8, 7, -5, -8]));
expect!["7"].assert_eq(&size_of([8i8, 15, -9, -16]));
expect!["8"].assert_eq(&size_of([16i8, 31, -17, -32]));
expect!["10"].assert_eq(&size_of([32i8, 63, -33, -64]));
expect!["11"].assert_eq(&size_of([64i8, 127, -65]));
assert_eq!(
crate::Encoded::<i8, Small>::new(-128).millibits(),
super::Millibits::bits(11)
);
}
#[test]
fn sorted_u8_roundtrip() {
use crate::Encoded;
for prev in 0u8..=255 {
for cur in 0u8..=255 {
let data = [
Encoded::<u8, Sorted>::new(prev),
Encoded::<u8, Sorted>::new(cur),
];
let enc = super::encode(&data);
let dec: [Encoded<u8, Sorted>; 2] = super::decode(&enc).unwrap();
assert_eq!(
[dec[0].value(), dec[1].value()],
[prev, cur],
"round-trip failed for [{prev}, {cur}]"
);
}
}
for v in 0u8..=255 {
let enc = super::encode_with(Sorted, &v);
let dec: u8 = super::decode_with(Sorted, &enc).unwrap();
assert_eq!(dec, v);
}
for v in i8::MIN..=i8::MAX {
let enc = super::encode_with(Sorted, &v);
let dec: i8 = super::decode_with(Sorted, &enc).unwrap();
assert_eq!(dec, v);
}
}
#[test]
fn sorted_u8_ascii() {
use super::estimated_bits;
use crate::Encoded;
expect!["28"].assert_eq(&estimated_bits!([
Encoded::<u8, Sorted>::new(b'h'),
Encoded::<u8, Sorted>::new(b'e'),
Encoded::<u8, Sorted>::new(b'l'),
Encoded::<u8, Sorted>::new(b'l'),
Encoded::<u8, Sorted>::new(b'o'),
]));
}