use crate::DecodeError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ValidationError {
TruncatedStream,
ExcessiveRenormalization,
ZeroFrequency,
CumulativeOverflow,
InvalidScaleBits,
RangeOverflow,
TrailingData,
}
impl core::fmt::Display for ValidationError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
ValidationError::TruncatedStream => write!(f, "compressed stream is truncated"),
ValidationError::ExcessiveRenormalization => {
write!(f, "excessive renormalization steps – probable corruption")
}
ValidationError::ZeroFrequency => write!(f, "symbol frequency is zero"),
ValidationError::CumulativeOverflow => {
write!(f, "cumulative frequency exceeds allowed range")
}
ValidationError::InvalidScaleBits => {
write!(f, "scale_bits parameter is out of valid range")
}
ValidationError::RangeOverflow => {
write!(f, "start + frequency exceeds allowed range")
}
ValidationError::TrailingData => write!(f, "trailing data after complete decode"),
}
}
}
#[cfg(feature = "std")]
impl std::error::Error for ValidationError {}
const MAX_RENORM_STEPS_BYTE: u32 = 16;
const MAX_RENORM_STEPS_R64: u32 = 8;
const MAX_RENORM_STEPS_WORD: u32 = 8;
#[inline]
pub fn validate_byte_compressed(compressed: &[u8]) -> Result<(), ValidationError> {
if compressed.len() < 4 {
return Err(ValidationError::TruncatedStream);
}
Ok(())
}
#[inline]
pub fn validate_r64_compressed(compressed: &[u8]) -> Result<(), ValidationError> {
if compressed.len() < 8 {
return Err(ValidationError::TruncatedStream);
}
Ok(())
}
#[inline]
pub fn validate_word_compressed(compressed: &[u16]) -> Result<(), ValidationError> {
if compressed.len() < 2 {
return Err(ValidationError::TruncatedStream);
}
Ok(())
}
#[inline]
pub fn validate_byte_scale_bits(scale_bits: u32) -> Result<(), ValidationError> {
if !(1..=16).contains(&scale_bits) {
return Err(ValidationError::InvalidScaleBits);
}
Ok(())
}
#[inline]
pub fn validate_r64_scale_bits(scale_bits: u32) -> Result<(), ValidationError> {
if !(1..=31).contains(&scale_bits) {
return Err(ValidationError::InvalidScaleBits);
}
Ok(())
}
#[inline]
pub fn validate_freq_model(
cum_freqs: &[u32],
freqs: &[u32],
scale_bits: u32,
) -> Result<(), ValidationError> {
let total = 1u64 << scale_bits;
let n = freqs.len().min(cum_freqs.len().saturating_sub(1));
for i in 0..n {
if freqs[i] == 0 {
continue;
}
let start = cum_freqs[i] as u64;
let freq = freqs[i] as u64;
if start + freq > total {
return Err(ValidationError::RangeOverflow);
}
}
let m = cum_freqs.len().min(256);
for i in 1..m {
if cum_freqs[i] < cum_freqs[i - 1] {
return Err(ValidationError::CumulativeOverflow);
}
}
Ok(())
}
pub struct RenormGuard {
remaining: u32,
}
impl RenormGuard {
#[inline]
pub fn new_byte() -> Self {
Self {
remaining: MAX_RENORM_STEPS_BYTE,
}
}
#[inline]
pub fn new_r64() -> Self {
Self {
remaining: MAX_RENORM_STEPS_R64,
}
}
#[inline]
pub fn new_word() -> Self {
Self {
remaining: MAX_RENORM_STEPS_WORD,
}
}
#[inline]
pub fn check(&mut self) -> Result<(), ValidationError> {
if self.remaining == 0 {
return Err(ValidationError::ExcessiveRenormalization);
}
self.remaining -= 1;
Ok(())
}
#[inline]
pub fn reset(&mut self) {
self.remaining = Self::default_remaining_for(&self);
}
}
impl RenormGuard {
fn default_remaining_for(&self) -> u32 {
MAX_RENORM_STEPS_BYTE }
}
#[inline]
pub fn has_dominant_symbol(freqs: &[u32], total: u32) -> bool {
freqs.iter().any(|&f| f as u64 * 2 > total as u64)
}
#[inline]
pub fn is_single_symbol(freqs: &[u32]) -> bool {
freqs.iter().filter(|&&f| f > 0).count() == 1
}
#[inline]
pub fn has_freq_one(freqs: &[u32]) -> bool {
freqs.iter().any(|&f| f == 1)
}
#[inline]
pub fn validation_to_decode_error(e: ValidationError) -> DecodeError {
match e {
ValidationError::TruncatedStream => DecodeError::InputTooShort,
ValidationError::ExcessiveRenormalization => DecodeError::InputTooShort,
ValidationError::ZeroFrequency => DecodeError::InputTooShort,
ValidationError::CumulativeOverflow => DecodeError::InputTooShort,
ValidationError::InvalidScaleBits => DecodeError::InputTooShort,
ValidationError::RangeOverflow => DecodeError::InputTooShort,
ValidationError::TrailingData => DecodeError::InputTooShort,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validate_byte_compressed_short() {
assert_eq!(
validate_byte_compressed(&[]),
Err(ValidationError::TruncatedStream)
);
assert_eq!(
validate_byte_compressed(&[0; 3]),
Err(ValidationError::TruncatedStream)
);
assert!(validate_byte_compressed(&[0; 4]).is_ok());
}
#[test]
fn test_validate_r64_compressed_short() {
assert_eq!(
validate_r64_compressed(&[0; 7]),
Err(ValidationError::TruncatedStream)
);
assert!(validate_r64_compressed(&[0; 8]).is_ok());
}
#[test]
fn test_validate_word_compressed_short() {
assert_eq!(
validate_word_compressed(&[0; 1]),
Err(ValidationError::TruncatedStream)
);
assert!(validate_word_compressed(&[0; 2]).is_ok());
}
#[test]
fn test_validate_byte_scale_bits() {
assert_eq!(
validate_byte_scale_bits(0),
Err(ValidationError::InvalidScaleBits)
);
assert_eq!(
validate_byte_scale_bits(17),
Err(ValidationError::InvalidScaleBits)
);
assert!(validate_byte_scale_bits(14).is_ok());
assert!(validate_byte_scale_bits(1).is_ok());
assert!(validate_byte_scale_bits(16).is_ok());
}
#[test]
fn test_validate_r64_scale_bits() {
assert_eq!(
validate_r64_scale_bits(0),
Err(ValidationError::InvalidScaleBits)
);
assert_eq!(
validate_r64_scale_bits(32),
Err(ValidationError::InvalidScaleBits)
);
assert!(validate_r64_scale_bits(14).is_ok());
assert!(validate_r64_scale_bits(31).is_ok());
}
#[test]
fn test_validate_freq_model_valid() {
let total = 1u32 << 14;
let freqs = [total / 3, total / 3, total - 2 * (total / 3)];
let cum = [0u32, freqs[0], freqs[0] + freqs[1], total];
assert!(validate_freq_model(&cum, &freqs, 14).is_ok());
}
#[test]
fn test_validate_freq_model_range_overflow() {
let total = 1u32 << 14;
let freqs = [total + 1];
let cum = [0u32, 0];
assert_eq!(
validate_freq_model(&cum, &freqs, 14),
Err(ValidationError::RangeOverflow)
);
}
#[test]
fn test_validate_freq_model_non_monotonic() {
let freqs = [100u32, 50];
let cum = [0u32, 100, 50]; assert_eq!(
validate_freq_model(&cum, &freqs, 14),
Err(ValidationError::CumulativeOverflow)
);
}
#[test]
fn test_renorm_guard_byte() {
let mut guard = RenormGuard::new_byte();
for _ in 0..MAX_RENORM_STEPS_BYTE - 1 {
assert!(guard.check().is_ok());
}
assert!(guard.check().is_ok());
assert_eq!(
guard.check(),
Err(ValidationError::ExcessiveRenormalization)
);
}
#[test]
fn test_renorm_guard_r64() {
let mut guard = RenormGuard::new_r64();
for _ in 0..MAX_RENORM_STEPS_R64 {
assert!(guard.check().is_ok());
}
assert_eq!(
guard.check(),
Err(ValidationError::ExcessiveRenormalization)
);
}
#[test]
fn test_has_dominant_symbol() {
let freqs = [1000u32, 1, 1];
assert!(has_dominant_symbol(&freqs, 1002));
let freqs = [500u32, 500];
assert!(!has_dominant_symbol(&freqs, 1000)); let freqs = [501u32, 499];
assert!(has_dominant_symbol(&freqs, 1000));
}
#[test]
fn test_is_single_symbol() {
assert!(is_single_symbol(&[100u32, 0, 0]));
assert!(!is_single_symbol(&[50u32, 50]));
assert!(!is_single_symbol(&[0u32, 0]));
}
#[test]
fn test_has_freq_one() {
assert!(has_freq_one(&[1u32, 100]));
assert!(!has_freq_one(&[2u32, 100]));
}
}