use super::super::bit_context::BitContext;
use super::super::model::{SymbolCoder, SymbolDecoder, SymbolRange};
use super::super::{EntropyCoder, EntropyDecoder};
use super::{AtMost, AtMostContext};
#[inline(always)]
pub(crate) fn encode_symbol_or_bitwise<C: SymbolCoder, const MAX: usize>(
coder: &mut C,
ctx: &mut AtMostContext<MAX>,
value: AtMost<MAX>,
) {
let value = usize::from(value);
if let Some(walk) = Walk::production::<MAX>(false) {
encode_atmost_walk(walk, coder, ctx, value);
}
}
#[inline(always)]
pub(crate) fn decode_symbol_or_bitwise<D: SymbolDecoder, const MAX: usize>(
reader: &mut D,
ctx: &mut AtMostContext<MAX>,
) -> AtMost<MAX> {
let value = match Walk::production::<MAX>(D::SPECULATES) {
Some(walk) => decode_atmost_walk(walk, reader, ctx),
None => 0,
};
debug_assert!(value <= MAX);
AtMost(value)
}
#[inline]
pub(super) const fn half(i: usize) -> usize {
let half = i / 2;
if half > 1 {
1 << half.ilog(2)
} else {
half
}
}
const fn tree_depth(len: usize) -> u32 {
if len <= 1 {
0
} else {
let vc = half(len);
let lo = vc.ilog2();
let hi = tree_depth(len - vc);
1 + if lo > hi { lo } else { hi }
}
}
const SPECULATE_MIN_MAX: usize = 3;
const SPECULATE_MAX_DEPTH: u32 = 6;
const SPECULATE_HUGE_MIN_MAX: usize = 700;
const fn speculation_pays(max: usize) -> bool {
max >= SPECULATE_MIN_MAX
&& (tree_depth(max + 1) <= SPECULATE_MAX_DEPTH || max >= SPECULATE_HUGE_MIN_MAX)
}
#[doc(hidden)] #[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum Walk {
Complete,
CompleteSpeculating,
Uneven,
UnevenSpeculating,
CompleteBitwise,
UnevenBitwise,
}
impl Walk {
#[inline(always)]
#[doc(hidden)]
pub const fn production<const MAX: usize>(speculate: bool) -> Option<Walk> {
if MAX == 0 {
None
} else if MAX >= SymbolRange::M as usize {
Some(if (MAX + 1).is_power_of_two() {
Walk::CompleteBitwise
} else {
Walk::UnevenBitwise
})
} else if (MAX + 1).is_power_of_two() {
Some(Walk::CompleteSpeculating)
} else if speculate && speculation_pays(MAX) {
Some(Walk::UnevenSpeculating)
} else {
Some(Walk::Uneven)
}
}
#[doc(hidden)]
pub const fn applies_to<const MAX: usize>(self) -> bool {
match self {
Walk::Complete | Walk::CompleteSpeculating => {
(MAX + 1).is_power_of_two() && MAX < SymbolRange::M as usize
}
Walk::Uneven | Walk::UnevenSpeculating => MAX < SymbolRange::M as usize,
Walk::CompleteBitwise => (MAX + 1).is_power_of_two(),
Walk::UnevenBitwise => true,
}
}
#[doc(hidden)]
pub const fn encode_with(self) -> Walk {
match self {
Walk::CompleteSpeculating => Walk::Complete,
Walk::UnevenSpeculating => Walk::Uneven,
other => other,
}
}
}
#[doc(hidden)]
pub const WALKS: [Walk; 6] = [
Walk::Complete,
Walk::CompleteSpeculating,
Walk::Uneven,
Walk::UnevenSpeculating,
Walk::CompleteBitwise,
Walk::UnevenBitwise,
];
#[inline(always)]
pub(crate) fn encode_atmost_walk<C: SymbolCoder, const MAX: usize>(
walk: Walk,
coder: &mut C,
ctx: &mut AtMostContext<MAX>,
value: usize,
) {
match walk.encode_with() {
Walk::Complete => coder.encode_symbol(complete::for_value(&mut ctx.bits, value)),
Walk::Uneven => coder.encode_symbol(uneven::for_value(&mut ctx.bits, value)),
Walk::CompleteBitwise => complete::encode_bitwise(coder, &mut ctx.bits, value),
Walk::UnevenBitwise => uneven::encode_bitwise(coder, &mut ctx.bits, value),
Walk::CompleteSpeculating | Walk::UnevenSpeculating => {
unreachable!("Walk::encode_with never returns a speculating variant")
}
}
}
#[inline(always)]
pub(crate) fn decode_atmost_walk<D: SymbolDecoder, const MAX: usize>(
walk: Walk,
reader: &mut D,
ctx: &mut AtMostContext<MAX>,
) -> usize {
let contexts = &mut ctx.bits;
match walk {
Walk::Complete => reader.decode_symbol_step(|slot| complete::from_slot(contexts, slot)),
Walk::CompleteSpeculating => {
reader.decode_symbol_step(|slot| complete::from_slot_speculating(contexts, slot))
}
Walk::Uneven => reader.decode_symbol_step(|slot| uneven::from_slot(contexts, slot)),
Walk::UnevenSpeculating => {
reader.decode_symbol_step(|slot| uneven::from_slot_speculating(contexts, slot))
}
Walk::CompleteBitwise => complete::decode_bitwise(reader, contexts),
Walk::UnevenBitwise => uneven::decode_bitwise(reader, contexts),
}
}
pub(crate) fn encode_atmost_batch<C: SymbolCoder, const MAX: usize, const WHICH_WALK: usize>(
mut coder: C,
values: &[AtMost<MAX>],
) -> C {
let walk = const { WALKS[WHICH_WALK] };
let mut ctx = AtMostContext::<MAX>::default();
for &value in values {
encode_atmost_walk(walk, &mut coder, &mut ctx, usize::from(value));
}
coder
}
pub(crate) fn decode_atmost_batch<D: SymbolDecoder, const MAX: usize, const WHICH_WALK: usize>(
mut reader: D,
n: usize,
) -> Vec<AtMost<MAX>> {
let walk = const { WALKS[WHICH_WALK] };
let mut ctx = AtMostContext::<MAX>::default();
(0..n)
.map(|_| AtMost::new(decode_atmost_walk(walk, &mut reader, &mut ctx)))
.collect()
}
#[inline]
pub(crate) fn encode_bitwise<E: EntropyCoder, const MAX: usize>(
writer: &mut E,
contexts: &mut [BitContext; MAX],
value: usize,
) {
if (MAX + 1).is_power_of_two() {
complete::encode_bitwise(writer, contexts, value)
} else {
uneven::encode_bitwise(writer, contexts, value)
}
}
#[inline]
pub(crate) fn decode_bitwise<D: EntropyDecoder, const MAX: usize>(
reader: &mut D,
contexts: &mut [BitContext; MAX],
) -> usize {
if (MAX + 1).is_power_of_two() {
complete::decode_bitwise(reader, contexts)
} else {
uneven::decode_bitwise(reader, contexts)
}
}
mod complete {
use super::*;
#[inline]
pub(super) fn for_value<const MAX: usize>(
contexts: &mut [BitContext; MAX],
value: usize,
) -> SymbolRange {
let mut range = SymbolRange::full();
let n_bits = (MAX + 1).ilog2();
debug_assert_eq!(1 << n_bits, MAX + 1);
debug_assert!(value <= MAX);
let mut node = 0usize;
for i in (0..n_bits).rev() {
let cur = contexts[node].model();
let reserve = 1u32 << i;
let split = range.split_reserving(cur.prob, reserve, reserve);
let bit = (value >> i) & 1 == 1;
range = if bit {
range.upper(split)
} else {
range.lower(split)
};
contexts[node] = cur.next[bit as usize];
node = (node << 1) + 1 + bit as usize;
}
range
}
#[inline]
pub(super) fn from_slot<const MAX: usize>(
contexts: &mut [BitContext; MAX],
slot: u32,
) -> (SymbolRange, usize) {
let mut range = SymbolRange::full();
let n_bits = (MAX + 1).ilog2();
debug_assert_eq!(1 << n_bits, MAX + 1);
debug_assert!(slot < SymbolRange::M);
let mut node = 0usize;
for i in (0..n_bits).rev() {
let cur = contexts[node].model();
let reserve = 1u32 << i;
let split = range.split_reserving(cur.prob, reserve, reserve);
let lower = range.lower(split);
let bit = !lower.contains(slot);
contexts[node] = cur.next[bit as usize];
range = if bit { range.upper(split) } else { lower };
node = (node << 1) + 1 + bit as usize;
}
(range, node - MAX)
}
#[inline]
pub(super) fn from_slot_speculating<const MAX: usize>(
contexts: &mut [BitContext; MAX],
slot: u32,
) -> (SymbolRange, usize) {
let mut range = SymbolRange::full();
if MAX == 0 {
return (range, 0);
}
let n_bits = (MAX + 1).ilog2();
debug_assert_eq!(1 << n_bits, MAX + 1);
debug_assert!(slot < SymbolRange::M);
let mut node = 0usize;
let mut cur = contexts[0].model();
for i in (0..n_bits).rev() {
let reserve = 1u32 << i;
let split = range.split_reserving(cur.prob, reserve, reserve);
let lower = range.lower(split);
let bit = !lower.contains(slot);
let adapted = cur.next[bit as usize];
if i > 0 {
let left = contexts[2 * node + 1].model();
let right = contexts[2 * node + 2].model();
cur = if bit { right } else { left };
}
contexts[node] = adapted;
range = if bit { range.upper(split) } else { lower };
node = (node << 1) + 1 + bit as usize;
}
(range, node - MAX)
}
#[inline]
pub(super) fn encode_bitwise<E: EntropyCoder, const MAX: usize>(
writer: &mut E,
contexts: &mut [BitContext; MAX],
value: usize,
) {
debug_assert!(value <= MAX);
let n_bits = (MAX + 1).ilog2();
let mut node = 0usize;
for i in (0..n_bits).rev() {
let bit = (value >> i) & 1 == 1;
writer.encode_bit(&mut contexts[node], bit);
node = (node << 1) + 1 + bit as usize;
}
}
#[inline]
pub(super) fn decode_bitwise<D: EntropyDecoder, const MAX: usize>(
reader: &mut D,
contexts: &mut [BitContext; MAX],
) -> usize {
let n_bits = (MAX + 1).ilog2();
let mut node = 0usize;
for _ in 0..n_bits {
let bit = reader.decode_bit(&mut contexts[node]);
node = (node << 1) + 1 + bit as usize;
}
node - MAX
}
}
mod uneven {
use super::*;
#[inline]
pub(super) fn for_value<const MAX: usize>(
contexts: &mut [BitContext; MAX],
value: usize,
) -> SymbolRange {
debug_assert!(MAX < SymbolRange::M as usize);
debug_assert!(value <= MAX);
let mut range = SymbolRange::full();
let mut accumulated_value = 0;
let mut possible_values_left = MAX + 1;
for _ in 0..const { tree_depth(MAX + 1) } {
if possible_values_left <= 1 {
break;
}
let value_considered = half(possible_values_left);
let split = accumulated_value + value_considered;
let cur = contexts[split - 1].model();
let slot_split = range.split_reserving(
cur.prob,
value_considered as u32,
(possible_values_left - value_considered) as u32,
);
let bit = value >= split;
contexts[split - 1] = cur.next[bit as usize];
if bit {
range = range.upper(slot_split);
accumulated_value = split;
possible_values_left -= value_considered;
} else {
range = range.lower(slot_split);
possible_values_left = value_considered;
}
}
range
}
#[inline]
pub(super) fn from_slot<const MAX: usize>(
contexts: &mut [BitContext; MAX],
slot: u32,
) -> (SymbolRange, usize) {
debug_assert!(MAX < SymbolRange::M as usize);
debug_assert!(slot < SymbolRange::M);
let mut range = SymbolRange::full();
let mut accumulated_value = 0;
let mut possible_values_left = MAX + 1;
for _ in 0..const { tree_depth(MAX + 1) } {
if possible_values_left <= 1 {
break;
}
let value_considered = half(possible_values_left);
let split = accumulated_value + value_considered;
let cur = contexts[split - 1].model();
let slot_split = range.split_reserving(
cur.prob,
value_considered as u32,
(possible_values_left - value_considered) as u32,
);
let lower = range.lower(slot_split);
let bit = !lower.contains(slot);
contexts[split - 1] = cur.next[bit as usize];
if bit {
range = range.upper(slot_split);
accumulated_value = split;
possible_values_left -= value_considered;
} else {
range = lower;
possible_values_left = value_considered;
}
}
(range, accumulated_value)
}
#[inline]
pub(super) fn from_slot_speculating<const MAX: usize>(
contexts: &mut [BitContext; MAX],
slot: u32,
) -> (SymbolRange, usize) {
let mut range = SymbolRange::full();
if MAX == 0 {
return (range, 0);
}
let mut accumulated_value = 0;
let mut possible_values_left = MAX + 1;
let mut value_considered = half(MAX + 1);
let mut split = value_considered;
let mut cur = contexts[split - 1].model();
for _ in 0..const { tree_depth(MAX + 1) } {
let lo_len = value_considered;
let hi_len = possible_values_left - value_considered;
let slot_split = range.split_reserving(cur.prob, lo_len as u32, hi_len as u32);
let lo_vc = half(lo_len);
let hi_vc = half(hi_len);
let lo_split = accumulated_value + lo_vc;
let hi_split = split + hi_vc;
let lo_cur = contexts[if lo_len > 1 { lo_split - 1 } else { 0 }].model();
let hi_cur = contexts[if hi_len > 1 { hi_split - 1 } else { 0 }].model();
let lower = range.lower(slot_split);
let bit = !lower.contains(slot);
contexts[split - 1] = cur.next[bit as usize];
range = if bit { range.upper(slot_split) } else { lower };
if bit {
accumulated_value = split;
possible_values_left = hi_len;
if hi_len <= 1 {
break;
}
value_considered = hi_vc;
split = hi_split;
cur = hi_cur;
} else {
possible_values_left = lo_len;
if lo_len <= 1 {
break;
}
value_considered = lo_vc;
split = lo_split;
cur = lo_cur;
}
}
(range, accumulated_value)
}
#[inline]
pub(super) fn encode_bitwise<E: EntropyCoder, const MAX: usize>(
writer: &mut E,
contexts: &mut [BitContext; MAX],
value: usize,
) {
debug_assert!(value <= MAX);
let mut accumulated_value = 0;
let mut possible_values_left = MAX + 1;
while possible_values_left > 1 {
let value_considered = half(possible_values_left);
let split = accumulated_value + value_considered;
let bit = value >= split;
writer.encode_bit(&mut contexts[split - 1], bit);
if bit {
accumulated_value = split;
possible_values_left -= value_considered;
} else {
possible_values_left = value_considered;
}
}
}
#[inline]
pub(super) fn decode_bitwise<D: EntropyDecoder, const MAX: usize>(
reader: &mut D,
contexts: &mut [BitContext; MAX],
) -> usize {
let mut accumulated_value = 0;
let mut possible_values_left = MAX + 1;
while possible_values_left > 1 {
let value_considered = half(possible_values_left);
let split = accumulated_value + value_considered;
let bit = reader.decode_bit(&mut contexts[split - 1]);
if bit {
accumulated_value = split;
possible_values_left -= value_considered;
} else {
possible_values_left = value_considered;
}
}
accumulated_value
}
}
#[cfg(test)]
mod tests {
use super::*;
fn reference_for_value<const MAX: usize>(
contexts: &mut [BitContext; MAX],
index_of: impl Fn(usize, usize) -> usize + Copy,
start: usize,
len: usize,
range: SymbolRange,
value: usize,
) -> SymbolRange {
if len <= 1 {
return range;
}
let vc = half(len);
let node = index_of(start, len);
let split =
range.split_reserving(contexts[node].probability(), vc as u32, (len - vc) as u32);
let bit = value >= start + vc;
contexts[node] = contexts[node].adapt(bit);
if bit {
reference_for_value(
contexts,
index_of,
start + vc,
len - vc,
range.upper(split),
value,
)
} else {
reference_for_value(contexts, index_of, start, vc, range.lower(split), value)
}
}
fn heap_index<const MAX: usize>(start: usize, len: usize) -> usize {
(MAX + 1) / len - 1 + start / len
}
fn split_index(start: usize, len: usize) -> usize {
start + half(len) - 1
}
fn check_complete_determinism<const MAX: usize>(contexts: [BitContext; MAX]) {
let mut total = 0u32;
for value in 0..=MAX {
let mut ref_ctx = contexts;
let range = reference_for_value(
&mut ref_ctx,
heap_index::<MAX>,
0,
MAX + 1,
SymbolRange::full(),
value,
);
let mut enc_ctx = contexts;
assert_eq!(
complete::for_value(&mut enc_ctx, value),
range,
"interval must match the bit-at-a-time reference for value {value}"
);
assert_eq!(enc_ctx, ref_ctx, "contexts must adapt like the reference");
assert!(range.width() >= 1, "leaf lost its slot for value {value}");
assert_eq!(
range.start(),
total,
"intervals must tile [0, M) in value order"
);
total += range.width();
for slot in [range.start(), range.start() + range.width() - 1] {
let mut dec_ctx = contexts;
let (dec_range, decoded) = complete::from_slot(&mut dec_ctx, slot);
assert_eq!(dec_range, range);
assert_eq!(decoded, value);
assert_eq!(dec_ctx, enc_ctx, "contexts must adapt identically");
let mut spec_ctx = contexts;
let (spec_range, spec_decoded) =
complete::from_slot_speculating(&mut spec_ctx, slot);
assert_eq!(spec_range, range);
assert_eq!(spec_decoded, value);
assert_eq!(spec_ctx, enc_ctx, "speculating walk must adapt identically");
}
}
assert_eq!(total, SymbolRange::M, "intervals must cover all of M");
}
fn check_uneven_determinism<const MAX: usize>(contexts: [BitContext; MAX]) {
let mut total = 0u32;
for value in 0..=MAX {
let mut ref_ctx = contexts;
let range = reference_for_value(
&mut ref_ctx,
split_index,
0,
MAX + 1,
SymbolRange::full(),
value,
);
let mut enc_ctx = contexts;
assert_eq!(
uneven::for_value(&mut enc_ctx, value),
range,
"interval must match the bit-at-a-time reference for value {value}"
);
assert_eq!(enc_ctx, ref_ctx, "contexts must adapt like the reference");
assert!(range.width() >= 1, "leaf lost its slot for value {value}");
assert_eq!(
range.start(),
total,
"intervals must tile [0, M) in value order"
);
total += range.width();
for slot in [range.start(), range.start() + range.width() - 1] {
let mut dec_ctx = contexts;
let (dec_range, decoded) = uneven::from_slot(&mut dec_ctx, slot);
assert_eq!(dec_range, range);
assert_eq!(decoded, value);
assert_eq!(dec_ctx, enc_ctx, "contexts must adapt identically");
let mut spec_ctx = contexts;
let (spec_range, spec_decoded) = uneven::from_slot_speculating(&mut spec_ctx, slot);
assert_eq!(spec_range, range);
assert_eq!(spec_decoded, value);
assert_eq!(spec_ctx, enc_ctx, "speculating walk must adapt identically");
}
}
assert_eq!(total, SymbolRange::M, "intervals must cover all of M");
}
fn check_bitwise_matches_reference<const MAX: usize>(contexts: [BitContext; MAX]) {
let index_of = |start: usize, len: usize| {
if (MAX + 1).is_power_of_two() {
heap_index::<MAX>(start, len)
} else {
split_index(start, len)
}
};
for value in 0..=MAX {
let mut ref_ctx = contexts;
reference_for_value(
&mut ref_ctx,
index_of,
0,
MAX + 1,
SymbolRange::full(),
value,
);
let mut enc_ctx = contexts;
let mut coder = crate::v2::Range::default();
encode_bitwise(&mut coder, &mut enc_ctx, value);
assert_eq!(
enc_ctx, ref_ctx,
"bitwise encode must adapt like the reference for value {value}"
);
let bytes = coder.into_vec();
let mut decoder = crate::v2::arith::Decoder::new(&bytes);
let mut dec_ctx = contexts;
let decoded = decode_bitwise(&mut decoder, &mut dec_ctx);
assert_eq!(decoded, value);
assert_eq!(
dec_ctx, ref_ctx,
"bitwise decode must adapt like the reference for value {value}"
);
}
}
fn extreme(bit: bool) -> BitContext {
let mut ctx = BitContext::default();
for _ in 0..2000 {
ctx = ctx.adapt(bit);
}
ctx
}
#[test]
fn complete_deterministic_and_lossless() {
check_complete_determinism::<0>([]);
check_complete_determinism::<1>([BitContext::default(); 1]);
check_complete_determinism::<15>([BitContext::default(); 15]);
check_complete_determinism::<255>([BitContext::default(); 255]);
check_complete_determinism::<255>([extreme(true); 255]);
check_complete_determinism::<255>([extreme(false); 255]);
check_complete_determinism::<127>([extreme(true); 127]);
let mut alternating = [extreme(false); 255];
for (i, ctx) in alternating.iter_mut().enumerate() {
if i % 2 == 0 {
*ctx = extreme(true);
}
}
check_complete_determinism::<255>(alternating);
for _ in 0..50 {
let mut random = [BitContext::default(); 63];
for ctx in random.iter_mut() {
*ctx = rand::random();
}
check_complete_determinism::<63>(random);
}
}
#[test]
fn uneven_deterministic_and_lossless() {
check_uneven_determinism::<0>([]);
check_uneven_determinism::<1>([BitContext::default(); 1]);
check_uneven_determinism::<2>([BitContext::default(); 2]);
check_uneven_determinism::<4>([BitContext::default(); 4]);
check_uneven_determinism::<5>([BitContext::default(); 5]);
check_uneven_determinism::<6>([BitContext::default(); 6]);
check_uneven_determinism::<9>([BitContext::default(); 9]);
check_uneven_determinism::<254>([BitContext::default(); 254]);
check_uneven_determinism::<255>([BitContext::default(); 255]);
check_uneven_determinism::<256>([BitContext::default(); 256]);
check_uneven_determinism::<256>([extreme(true); 256]);
check_uneven_determinism::<256>([extreme(false); 256]);
check_uneven_determinism::<99>([extreme(true); 99]);
let mut alternating = [extreme(false); 256];
for (i, ctx) in alternating.iter_mut().enumerate() {
if i % 2 == 0 {
*ctx = extreme(true);
}
}
check_uneven_determinism::<256>(alternating);
for _ in 0..50 {
let mut random = [BitContext::default(); 76];
for ctx in random.iter_mut() {
*ctx = rand::random();
}
check_uneven_determinism::<76>(random);
}
}
#[test]
fn bitwise_matches_reference() {
check_bitwise_matches_reference::<0>([]);
check_bitwise_matches_reference::<1>([BitContext::default(); 1]);
check_bitwise_matches_reference::<5>([BitContext::default(); 5]);
check_bitwise_matches_reference::<255>([BitContext::default(); 255]);
check_bitwise_matches_reference::<256>([BitContext::default(); 256]);
check_bitwise_matches_reference::<255>([extreme(true); 255]);
check_bitwise_matches_reference::<99>([extreme(false); 99]);
for _ in 0..20 {
let mut random = [BitContext::default(); 76];
for ctx in random.iter_mut() {
*ctx = rand::random();
}
check_bitwise_matches_reference::<76>(random);
}
}
#[test]
fn complete_and_uneven_agree_on_balanced_trees() {
fn check<const MAX: usize>() {
for value in 0..=MAX {
let mut heap_ctx = [BitContext::default(); MAX];
let mut split_ctx = [BitContext::default(); MAX];
assert_eq!(
complete::for_value(&mut heap_ctx, value),
uneven::for_value(&mut split_ctx, value),
"implementations disagree for value {value} of 0..={MAX}"
);
}
}
check::<1>();
check::<7>();
check::<255>();
}
#[test]
fn matches_per_bit_probabilities_when_unclamped() {
let mut ctx = [BitContext::default(); 255];
let range = complete::for_value(&mut ctx, 0x5a);
assert_eq!(range.width(), SymbolRange::M >> 8);
}
fn check_walk_round_trips<const MAX: usize>() {
for &walk in WALKS.iter() {
if !walk.applies_to::<MAX>() {
continue;
}
for value in 0..=MAX {
let mut enc_ctx = AtMostContext::<MAX>::default();
let mut coder = crate::v2::Range::default();
encode_atmost_walk(walk, &mut coder, &mut enc_ctx, value);
let bytes = coder.into_vec();
let mut decoder = crate::v2::arith::Decoder::new(&bytes);
let mut dec_ctx = AtMostContext::<MAX>::default();
let decoded = decode_atmost_walk(walk, &mut decoder, &mut dec_ctx);
assert_eq!(
decoded, value,
"{walk:?} failed to round-trip {value} of 0..={MAX}"
);
}
}
}
#[test]
fn every_applicable_walk_round_trips() {
check_walk_round_trips::<0>();
check_walk_round_trips::<1>();
check_walk_round_trips::<2>();
check_walk_round_trips::<3>();
check_walk_round_trips::<7>();
check_walk_round_trips::<8>();
check_walk_round_trips::<9>();
check_walk_round_trips::<255>();
check_walk_round_trips::<256>();
}
fn check_batch_round_trips<const MAX: usize, const WHICH_WALK: usize>() {
if !WALKS[WHICH_WALK].applies_to::<MAX>() {
return;
}
let values: Vec<AtMost<MAX>> = (0..=MAX).chain((0..=MAX).rev()).map(AtMost::new).collect();
let ans_bytes = crate::v2::Ans::encode_atmost_batch::<MAX, WHICH_WALK>(&values);
let ans_decoded =
crate::v2::Ans::decode_atmost_batch::<MAX, WHICH_WALK>(&ans_bytes, values.len());
assert_eq!(
ans_decoded, values,
"Ans batch round-trip failed for MAX={MAX}, {:?}",
WALKS[WHICH_WALK]
);
let range_bytes = crate::v2::Range::encode_atmost_batch::<MAX, WHICH_WALK>(&values);
let range_decoded =
crate::v2::Range::decode_atmost_batch::<MAX, WHICH_WALK>(&range_bytes, values.len());
assert_eq!(
range_decoded, values,
"Range batch round-trip failed for MAX={MAX}, {:?}",
WALKS[WHICH_WALK]
);
}
#[test]
fn every_applicable_walk_batch_round_trips() {
macro_rules! check_all_walks {
($max:expr) => {
check_batch_round_trips::<$max, 0>();
check_batch_round_trips::<$max, 1>();
check_batch_round_trips::<$max, 2>();
check_batch_round_trips::<$max, 3>();
check_batch_round_trips::<$max, 4>();
check_batch_round_trips::<$max, 5>();
};
}
check_all_walks!(0);
check_all_walks!(1);
check_all_walks!(2);
check_all_walks!(3);
check_all_walks!(7);
check_all_walks!(8);
check_all_walks!(9);
check_all_walks!(255);
check_all_walks!(256);
}
#[test]
fn speculation_window_boundaries() {
assert_eq!(Walk::production::<2>(true), Some(Walk::Uneven));
assert_eq!(Walk::production::<4>(true), Some(Walk::UnevenSpeculating));
assert_eq!(Walk::production::<33>(true), Some(Walk::UnevenSpeculating));
assert_eq!(Walk::production::<34>(true), Some(Walk::Uneven));
assert_eq!(Walk::production::<512>(true), Some(Walk::Uneven));
assert_eq!(Walk::production::<700>(true), Some(Walk::UnevenSpeculating));
assert_eq!(Walk::production::<33>(false), Some(Walk::Uneven));
assert_eq!(Walk::production::<700>(false), Some(Walk::Uneven));
}
#[test]
fn production_matches_walks_array() {
fn check<const MAX: usize>() {
for speculate in [false, true] {
if let Some(walk) = Walk::production::<MAX>(speculate) {
assert!(
WALKS.contains(&walk),
"{walk:?} missing from WALKS for MAX={MAX}"
);
assert!(
walk.applies_to::<MAX>(),
"{walk:?} does not apply_to MAX={MAX}"
);
}
}
}
check::<0>();
check::<1>();
check::<2>();
check::<3>();
check::<7>();
check::<8>();
check::<9>();
check::<255>();
check::<256>();
}
}