use crate::format::{FormatError, FormatResult};
pub(crate) const NBIT_ATOMIC: u32 = 1;
const NBIT_ARRAY: u32 = 2;
const NBIT_COMPOUND: u32 = 3;
const NBIT_NOOPTYPE: u32 = 4;
pub(crate) const NBIT_ORDER_LE: u32 = 0;
pub(crate) const NBIT_ORDER_BE: u32 = 1;
#[derive(Clone, Copy)]
struct NbitAtomic {
size: u32,
order: u32,
precision: u32,
offset: u32,
}
struct BitWriter<'a> {
buf: &'a mut [u8],
j: usize,
acc: u64,
nacc: u32,
overrun: bool,
}
impl<'a> BitWriter<'a> {
fn new(buf: &'a mut [u8]) -> Self {
Self {
buf,
j: 0,
acc: 0,
nacc: 0,
overrun: false,
}
}
#[inline]
fn store(&mut self, b: u8) {
match self.buf.get_mut(self.j) {
Some(slot) => *slot = b,
None => self.overrun = true,
}
self.j += 1;
}
#[inline]
fn put(&mut self, v: u64, n: u32) {
if n > 32 {
self.put_half(v >> 32, n - 32);
self.put_half(v, 32);
} else {
self.put_half(v, n);
}
}
#[inline]
fn put_half(&mut self, v: u64, n: u32) {
self.acc = (self.acc << n) | (v & mask_u64(n as usize));
self.nacc += n;
while self.nacc >= 8 {
self.nacc -= 8;
self.store((self.acc >> self.nacc) as u8);
}
}
fn finish(mut self) -> FormatResult<usize> {
if self.nacc > 0 {
let b = (self.acc << (8 - self.nacc)) as u8;
match self.buf.get_mut(self.j) {
Some(slot) => *slot = b,
None => self.overrun = true,
}
}
if self.overrun {
return Err(FormatError::InvalidData(
"packed stream longer than the buffer it was sized for".into(),
));
}
Ok(self.j)
}
}
struct BitReader<'a> {
buf: &'a [u8],
j: usize,
acc: u64,
nacc: u32,
short: &'static str,
}
impl<'a> BitReader<'a> {
fn new(buf: &'a [u8], short: &'static str) -> Self {
Self {
buf,
j: 0,
acc: 0,
nacc: 0,
short,
}
}
#[inline]
fn get(&mut self, n: u32) -> FormatResult<u64> {
if n > 32 {
let hi = self.get_half(n - 32)?;
let lo = self.get_half(32)?;
Ok((hi << 32) | lo)
} else {
self.get_half(n)
}
}
#[inline]
fn get_half(&mut self, n: u32) -> FormatResult<u64> {
while self.nacc < n {
let Some(&b) = self.buf.get(self.j) else {
return Err(FormatError::InvalidData(self.short.into()));
};
self.acc = (self.acc << 8) | u64::from(b);
self.j += 1;
self.nacc += 8;
}
self.nacc -= n;
Ok((self.acc >> self.nacc) & mask_u64(n as usize))
}
}
const NBIT_SHORT: &str = "nbit: buffer too short";
fn nbit_bytes(p: &NbitAtomic) -> impl Iterator<Item = (usize, u32, u32)> {
let len = p.size * 8;
let top = p.precision + p.offset;
let (begin, end, step): (i64, i64, i64) = if p.order == NBIT_ORDER_LE {
let begin = if top.is_multiple_of(8) {
top / 8 - 1
} else {
top / 8
};
(i64::from(begin), i64::from(p.offset / 8), -1)
} else {
let end = if p.offset.is_multiple_of(8) {
(len - p.offset) / 8 - 1
} else {
(len - p.offset) / 8
};
(i64::from((len - top) / 8), i64::from(end), 1)
};
let p = *p;
std::iter::successors(Some(begin), move |&k| (k != end).then(|| k + step)).map(move |k| {
let (bits, shift) = if begin == end {
(p.precision, p.offset % 8)
} else if k == begin {
(8 - (len - top) % 8, 0)
} else if k == end {
let bits = 8 - p.offset % 8;
(bits, 8 - bits)
} else {
(8, 0)
};
(k as usize, bits, shift)
})
}
#[inline]
fn load_uint<const N: usize>(bytes: [u8; N], le: bool) -> u64 {
let mut padded = [0u8; 8];
if le {
padded[..N].copy_from_slice(&bytes);
u64::from_le_bytes(padded)
} else {
padded[8 - N..].copy_from_slice(&bytes);
u64::from_be_bytes(padded)
}
}
#[inline]
fn store_uint<const N: usize>(v: u64, le: bool) -> [u8; N] {
let mut out = [0u8; N];
if le {
out.copy_from_slice(&v.to_le_bytes()[..N]);
} else {
out.copy_from_slice(&v.to_be_bytes()[8 - N..]);
}
out
}
macro_rules! by_width {
($size:expr, $f:ident($($arg:expr),* $(,)?)) => {
match $size {
1 => $f::<1>($($arg),*),
2 => $f::<2>($($arg),*),
4 => $f::<4>($($arg),*),
8 => $f::<8>($($arg),*),
n => unreachable!("element width {n} is not 1, 2, 4 or 8"),
}
};
}
#[inline]
fn nbit_field<const N: usize>(elem: &[u8], le: bool, offset: u32) -> u64 {
load_uint::<N>(elem.try_into().expect("elem is N bytes"), le) >> offset
}
#[inline]
fn nbit_place<const N: usize>(elem: &mut [u8], le: bool, offset: u32, field: u64) {
elem.copy_from_slice(&store_uint::<N>(field << offset, le));
}
fn nbit_compress_atomics<const N: usize>(data: &[u8], w: &mut BitWriter, p: &NbitAtomic) {
let le = p.order == NBIT_ORDER_LE;
let (elems, _) = data.as_chunks::<N>();
for &e in elems {
w.put(load_uint(e, le) >> p.offset, p.precision);
}
}
fn nbit_decompress_atomics<const N: usize>(
out: &mut [u8],
r: &mut BitReader,
p: &NbitAtomic,
) -> FormatResult<()> {
let le = p.order == NBIT_ORDER_LE;
let (elems, _) = out.as_chunks_mut::<N>();
for e in elems {
*e = store_uint(r.get(p.precision)? << p.offset, le);
}
Ok(())
}
struct Parms<'a> {
list: &'a [u32],
pos: usize,
}
impl<'a> Parms<'a> {
fn at(list: &'a [u32], pos: usize) -> Self {
Self { list, pos }
}
fn next(&mut self) -> FormatResult<u32> {
let v = *self
.list
.get(self.pos)
.ok_or_else(|| FormatError::InvalidData("nbit: parameter list truncated".into()))?;
self.pos += 1;
Ok(v)
}
fn position(&self) -> usize {
self.pos
}
fn seek(&mut self, pos: usize) {
self.pos = pos;
}
}
fn element(data: &[u8], offset: usize, size: u32) -> FormatResult<&[u8]> {
offset
.checked_add(size as usize)
.and_then(|end| data.get(offset..end))
.ok_or_else(|| FormatError::InvalidData("nbit: element extends past buffer".into()))
}
fn element_mut(data: &mut [u8], offset: usize, size: u32) -> FormatResult<&mut [u8]> {
offset
.checked_add(size as usize)
.and_then(|end| data.get_mut(offset..end))
.ok_or_else(|| FormatError::InvalidData("nbit: element extends past buffer".into()))
}
fn repeat_count(total_size: u32, base_size: u32) -> FormatResult<usize> {
if base_size == 0 {
return Err(FormatError::InvalidData(
"nbit: zero-sized array base type".into(),
));
}
Ok((total_size / base_size) as usize)
}
fn nbit_decompress_one_nooptype(
data: &mut [u8],
data_offset: usize,
r: &mut BitReader,
size: u32,
) -> FormatResult<()> {
for b in element_mut(data, data_offset, size)? {
*b = r.get(8)? as u8;
}
Ok(())
}
fn nbit_compress_one_nooptype(
data: &[u8],
data_offset: usize,
w: &mut BitWriter,
size: u32,
) -> FormatResult<()> {
for &b in element(data, data_offset, size)? {
w.put(u64::from(b), 8);
}
Ok(())
}
fn nbit_decompress_one_atomic(
data: &mut [u8],
data_offset: usize,
r: &mut BitReader,
p: &NbitAtomic,
) -> FormatResult<()> {
let elem = element_mut(data, data_offset, p.size)?;
if matches!(p.size, 1 | 2 | 4 | 8) {
let field = r.get(p.precision)?;
let le = p.order == NBIT_ORDER_LE;
by_width!(p.size, nbit_place(elem, le, p.offset, field));
return Ok(());
}
for (k, bits, shift) in nbit_bytes(p) {
elem[k] = (r.get(bits)? << shift) as u8;
}
Ok(())
}
fn nbit_compress_one_atomic(
data: &[u8],
data_offset: usize,
w: &mut BitWriter,
p: &NbitAtomic,
) -> FormatResult<()> {
let elem = element(data, data_offset, p.size)?;
if matches!(p.size, 1 | 2 | 4 | 8) {
let le = p.order == NBIT_ORDER_LE;
w.put(
by_width!(p.size, nbit_field(elem, le, p.offset)),
p.precision,
);
return Ok(());
}
for (k, bits, shift) in nbit_bytes(p) {
w.put(u64::from(elem[k] >> shift), bits);
}
Ok(())
}
fn read_atomic(parms: &mut Parms) -> FormatResult<NbitAtomic> {
let p = NbitAtomic {
size: parms.next()?,
order: parms.next()?,
precision: parms.next()?,
offset: parms.next()?,
};
let bits = p.size.checked_mul(8);
let span = p.precision.checked_add(p.offset);
match (bits, span) {
(Some(bits), Some(span))
if p.size > 0 && p.precision > 0 && p.precision <= bits && span <= bits => {}
_ => {
return Err(FormatError::InvalidData(format!(
"nbit: invalid atomic datatype (size={}, precision={}, offset={})",
p.size, p.precision, p.offset
)));
}
}
Ok(p)
}
fn nbit_decompress_one_array(
data: &mut [u8],
data_offset: usize,
r: &mut BitReader,
parms: &mut Parms,
) -> FormatResult<()> {
let total_size = parms.next()?;
let base_class = parms.next()?;
match base_class {
NBIT_ATOMIC => {
let p = read_atomic(parms)?;
let n = repeat_count(total_size, p.size)?;
for i in 0..n {
nbit_decompress_one_atomic(data, data_offset + i * p.size as usize, r, &p)?;
}
}
NBIT_ARRAY => {
let begin = parms.position();
let base_size = parms.next()?;
let n = repeat_count(total_size, base_size)?;
for i in 0..n {
parms.seek(begin);
nbit_decompress_one_array(data, data_offset + i * base_size as usize, r, parms)?;
}
}
NBIT_COMPOUND => {
let begin = parms.position();
let base_size = parms.next()?;
let n = repeat_count(total_size, base_size)?;
for i in 0..n {
parms.seek(begin);
nbit_decompress_one_compound(data, data_offset + i * base_size as usize, r, parms)?;
}
}
NBIT_NOOPTYPE => {
parms.next()?; nbit_decompress_one_nooptype(data, data_offset, r, total_size)?;
}
_ => {
return Err(FormatError::InvalidData(format!(
"nbit: bad base class {}",
base_class
)))
}
}
Ok(())
}
fn nbit_decompress_one_compound(
data: &mut [u8],
data_offset: usize,
r: &mut BitReader,
parms: &mut Parms,
) -> FormatResult<()> {
parms.next()?; let nmembers = parms.next()?;
for _ in 0..nmembers {
let member_offset = parms.next()? as usize;
let member_class = parms.next()?;
match member_class {
NBIT_ATOMIC => {
let p = read_atomic(parms)?;
nbit_decompress_one_atomic(data, data_offset + member_offset, r, &p)?;
}
NBIT_ARRAY => {
nbit_decompress_one_array(data, data_offset + member_offset, r, parms)?;
}
NBIT_COMPOUND => {
nbit_decompress_one_compound(data, data_offset + member_offset, r, parms)?;
}
NBIT_NOOPTYPE => {
let size = parms.next()?;
nbit_decompress_one_nooptype(data, data_offset + member_offset, r, size)?;
}
_ => {
return Err(FormatError::InvalidData(format!(
"nbit: bad member class {}",
member_class
)))
}
}
}
Ok(())
}
fn nbit_compress_one_array(
data: &[u8],
data_offset: usize,
w: &mut BitWriter,
parms: &mut Parms,
) -> FormatResult<()> {
let total_size = parms.next()?;
let base_class = parms.next()?;
match base_class {
NBIT_ATOMIC => {
let p = read_atomic(parms)?;
let n = repeat_count(total_size, p.size)?;
for i in 0..n {
nbit_compress_one_atomic(data, data_offset + i * p.size as usize, w, &p)?;
}
}
NBIT_ARRAY => {
let begin = parms.position();
let base_size = parms.next()?;
let n = repeat_count(total_size, base_size)?;
for i in 0..n {
parms.seek(begin);
nbit_compress_one_array(data, data_offset + i * base_size as usize, w, parms)?;
}
}
NBIT_COMPOUND => {
let begin = parms.position();
let base_size = parms.next()?;
let n = repeat_count(total_size, base_size)?;
for i in 0..n {
parms.seek(begin);
nbit_compress_one_compound(data, data_offset + i * base_size as usize, w, parms)?;
}
}
NBIT_NOOPTYPE => {
parms.next()?;
nbit_compress_one_nooptype(data, data_offset, w, total_size)?;
}
_ => {
return Err(FormatError::InvalidData(format!(
"nbit: bad base class {}",
base_class
)))
}
}
Ok(())
}
fn nbit_compress_one_compound(
data: &[u8],
data_offset: usize,
w: &mut BitWriter,
parms: &mut Parms,
) -> FormatResult<()> {
parms.next()?;
let nmembers = parms.next()?;
for _ in 0..nmembers {
let member_offset = parms.next()? as usize;
let member_class = parms.next()?;
match member_class {
NBIT_ATOMIC => {
let p = read_atomic(parms)?;
nbit_compress_one_atomic(data, data_offset + member_offset, w, &p)?;
}
NBIT_ARRAY => {
nbit_compress_one_array(data, data_offset + member_offset, w, parms)?;
}
NBIT_COMPOUND => {
nbit_compress_one_compound(data, data_offset + member_offset, w, parms)?;
}
NBIT_NOOPTYPE => {
let size = parms.next()?;
nbit_compress_one_nooptype(data, data_offset + member_offset, w, size)?;
}
_ => {
return Err(FormatError::InvalidData(format!(
"nbit: bad member class {}",
member_class
)))
}
}
}
Ok(())
}
const NBIT_HEADER_NPARMS: usize = 5;
pub fn apply_nbit(data: &[u8], cd_values: &[u32], compress: bool) -> FormatResult<Vec<u8>> {
if cd_values.len() < NBIT_HEADER_NPARMS {
return Err(FormatError::InvalidData("nbit: cd_values too short".into()));
}
if cd_values[0] as usize != cd_values.len() {
return Err(FormatError::InvalidData(format!(
"nbit: cd_values[0] names {} parameters but {} are stored",
cd_values[0],
cd_values.len()
)));
}
if cd_values[1] != 0 {
return Ok(data.to_vec());
}
let d_nelmts = cd_values[2] as usize;
let dtype_size = cd_values[4] as usize;
if dtype_size == 0 {
return Err(FormatError::InvalidData("nbit: zero datatype size".into()));
}
let unpacked_size = d_nelmts.checked_mul(dtype_size).ok_or_else(|| {
FormatError::InvalidData("nbit: (de)compression buffer size overflow".into())
})?;
if compress {
if data.len() != unpacked_size {
return Err(FormatError::InvalidData(format!(
"nbit: input size {} != expected {}",
data.len(),
unpacked_size
)));
}
let mut buffer = vec![0u8; unpacked_size + 1];
let mut w = BitWriter::new(&mut buffer);
match cd_values[3] {
NBIT_ATOMIC => {
let p = read_atomic(&mut Parms::at(cd_values, 4))?;
if matches!(p.size, 1 | 2 | 4 | 8) {
by_width!(p.size, nbit_compress_atomics(data, &mut w, &p));
} else {
for i in 0..d_nelmts {
nbit_compress_one_atomic(data, i * p.size as usize, &mut w, &p)?;
}
}
}
NBIT_ARRAY => {
for i in 0..d_nelmts {
let mut parms = Parms::at(cd_values, 4);
nbit_compress_one_array(data, i * dtype_size, &mut w, &mut parms)?;
}
}
NBIT_COMPOUND => {
for i in 0..d_nelmts {
let mut parms = Parms::at(cd_values, 4);
nbit_compress_one_compound(data, i * dtype_size, &mut w, &mut parms)?;
}
}
other => {
return Err(FormatError::InvalidData(format!(
"nbit: unsupported top class {}",
other
)))
}
}
let j = w.finish()?;
buffer.truncate(j + 1);
Ok(buffer)
} else {
let mut out = Vec::new();
out.try_reserve_exact(unpacked_size).map_err(|_| {
FormatError::InvalidData(format!(
"nbit: cannot allocate {unpacked_size} bytes for decompression"
))
})?;
out.resize(unpacked_size, 0);
let mut r = BitReader::new(data, NBIT_SHORT);
match cd_values[3] {
NBIT_ATOMIC => {
let p = read_atomic(&mut Parms::at(cd_values, 4))?;
if matches!(p.size, 1 | 2 | 4 | 8) {
by_width!(p.size, nbit_decompress_atomics(&mut out, &mut r, &p))?;
} else {
for i in 0..d_nelmts {
nbit_decompress_one_atomic(&mut out, i * p.size as usize, &mut r, &p)?;
}
}
}
NBIT_ARRAY => {
for i in 0..d_nelmts {
let mut parms = Parms::at(cd_values, 4);
nbit_decompress_one_array(&mut out, i * dtype_size, &mut r, &mut parms)?;
}
}
NBIT_COMPOUND => {
for i in 0..d_nelmts {
let mut parms = Parms::at(cd_values, 4);
nbit_decompress_one_compound(&mut out, i * dtype_size, &mut r, &mut parms)?;
}
}
other => {
return Err(FormatError::InvalidData(format!(
"nbit: unsupported top class {}",
other
)))
}
}
Ok(out)
}
}
const SO_PARM_SCALETYPE: usize = 0;
const SO_PARM_SCALEFACTOR: usize = 1;
const SO_PARM_NELMTS: usize = 2;
const SO_PARM_CLASS: usize = 3;
const SO_PARM_SIZE: usize = 4;
const SO_PARM_SIGN: usize = 5;
const SO_PARM_ORDER: usize = 6;
const SO_PARM_FILAVAIL: usize = 7;
const SO_PARM_FILVAL: usize = 8;
pub(crate) const SO_CLS_INTEGER: u32 = 0;
pub(crate) const SO_CLS_FLOAT: u32 = 1;
pub(crate) const SO_ORDER_LE: u32 = 0;
const SO_FILL_DEFINED: u32 = 1;
pub(crate) const SO_FLOAT_DSCALE: u32 = 0;
pub(crate) const SO_INT: u32 = 2;
pub(crate) const SO_SGN_NONE: u32 = 0;
pub(crate) const SO_SGN_2: u32 = 1;
pub(crate) const SO_ORDER_BE: u32 = 1;
pub(crate) const SO_TOTAL_NPARMS: usize = 20;
const SO_BUF_OFFSET: usize = 21;
const SO_SHORT: &str = "scaleoffset: buffer too short";
fn so_log2(num: u64) -> u32 {
let mut v = 0u32;
let mut lower_bound: u64 = 1;
let mut val = num;
while {
val >>= 1;
val != 0
} {
v += 1;
lower_bound <<= 1;
}
if num == lower_bound {
v
} else {
v + 1
}
}
#[derive(Clone, Copy)]
struct SoParams {
scale_factor: i32,
d_nelmts: usize,
dtype_class: u32,
size: usize,
dtype_sign: u32,
order: u32,
fill_defined: bool,
filval: u64,
}
impl SoParams {
fn parse(cd_values: &[u32]) -> FormatResult<Self> {
if cd_values.len() < 8 {
return Err(FormatError::InvalidData(
"scaleoffset: cd_values too short".into(),
));
}
let scale_type = cd_values[SO_PARM_SCALETYPE];
let dtype_class = cd_values[SO_PARM_CLASS];
let size = cd_values[SO_PARM_SIZE] as usize;
let fill_defined = cd_values[SO_PARM_FILAVAIL] == SO_FILL_DEFINED;
if !matches!(size, 1 | 2 | 4 | 8) {
return Err(FormatError::InvalidData(format!(
"scaleoffset: unsupported datatype size {}",
size
)));
}
if dtype_class == SO_CLS_FLOAT && scale_type != SO_FLOAT_DSCALE {
return Err(FormatError::UnsupportedFeature(
"scaleoffset E-scaling method is not supported".into(),
));
}
let filval: u64 = if fill_defined {
let mut v: u64 = 0;
let n_cd = size.div_ceil(4);
if cd_values.len() < SO_PARM_FILVAL + n_cd {
return Err(FormatError::InvalidData(
"scaleoffset: cd_values missing fill value".into(),
));
}
for (w, cd) in cd_values[SO_PARM_FILVAL..SO_PARM_FILVAL + n_cd]
.iter()
.enumerate()
{
v |= (*cd as u64) << (w * 32);
}
v & mask_u64(size * 8)
} else {
0
};
let mut scale_factor = cd_values[SO_PARM_SCALEFACTOR] as i32;
if dtype_class == SO_CLS_INTEGER && scale_factor < 0 {
scale_factor = 0;
}
Ok(Self {
scale_factor,
d_nelmts: cd_values[SO_PARM_NELMTS] as usize,
dtype_class,
size,
dtype_sign: cd_values[SO_PARM_SIGN],
order: cd_values[SO_PARM_ORDER],
fill_defined,
filval,
})
}
fn is_noop(&self) -> bool {
self.dtype_class == SO_CLS_INTEGER && self.scale_factor as usize == self.size * 8
}
fn dtype_len(&self) -> u32 {
(self.size * 8) as u32
}
fn width_mask(&self) -> u64 {
mask_u64(self.size * 8)
}
}
fn mask_u64(n: usize) -> u64 {
if n >= 64 {
u64::MAX
} else {
!(u64::MAX << n)
}
}
fn so_pack<const N: usize>(buf: &[u8], le: bool, minbits: u32, w: &mut BitWriter) {
let (elems, _) = buf.as_chunks::<N>();
for &elem in elems {
w.put(load_uint(elem, le), minbits);
}
}
fn so_unpack<const N: usize>(
out: &mut [u8],
le: bool,
minbits: u32,
r: &mut BitReader,
) -> FormatResult<()> {
let (elems, _) = out.as_chunks_mut::<N>();
for elem in elems {
*elem = store_uint(r.get(minbits)?, le);
}
Ok(())
}
pub fn reverse_scaleoffset(data: &[u8], cd_values: &[u32]) -> FormatResult<Vec<u8>> {
let p = SoParams::parse(cd_values)?;
let (d_nelmts, size, order) = (p.d_nelmts, p.size, p.order);
let size_out = d_nelmts * size;
if p.is_noop() {
if data.len() < size_out {
return Err(FormatError::InvalidData(SO_SHORT.into()));
}
return Ok(data[..size_out].to_vec());
}
if data.len() < SO_BUF_OFFSET {
return Err(FormatError::InvalidData(
"scaleoffset: buffer too short for header".into(),
));
}
let mut minbits: u32 = 0;
for (i, &b) in data[..4].iter().enumerate() {
minbits |= (b as u32) << (i * 8);
}
if minbits as usize > size * 8 {
return Err(FormatError::InvalidData(
"scaleoffset: minbits exceeds datatype size".into(),
));
}
let minval_size = std::cmp::min(8usize, data[4] as usize);
let mut minval: u64 = 0;
for i in 0..minval_size {
minval |= (data[5 + i] as u64) << (i * 8);
}
if minbits as usize == size * 8 {
if data.len() < SO_BUF_OFFSET + size_out {
return Err(FormatError::InvalidData(SO_SHORT.into()));
}
return Ok(data[SO_BUF_OFFSET..SO_BUF_OFFSET + size_out].to_vec());
}
let mut out = vec![0u8; size_out];
if minbits != 0 {
if data.len() < SO_BUF_OFFSET {
return Err(FormatError::InvalidData(SO_SHORT.into()));
}
let mut r = BitReader::new(&data[SO_BUF_OFFSET..], SO_SHORT);
let le = order == SO_ORDER_LE;
by_width!(size, so_unpack(&mut out, le, minbits, &mut r))?;
}
postdecompress(&mut out, &p, minbits, minval);
Ok(out)
}
pub fn forward_scaleoffset(data: &[u8], cd_values: &[u32]) -> FormatResult<Vec<u8>> {
let p = SoParams::parse(cd_values)?;
let nbytes = p.d_nelmts * p.size;
if data.len() != nbytes {
return Err(FormatError::InvalidData(format!(
"scaleoffset: chunk is {} bytes, but the filter parameters describe {} elements of \
{} bytes",
data.len(),
p.d_nelmts,
p.size
)));
}
if p.is_noop() {
return Ok(data.to_vec());
}
if p.dtype_class == SO_CLS_INTEGER && p.scale_factor as usize > p.size * 8 {
return Err(FormatError::InvalidData(
"scaleoffset: minimum number of bits exceeds the datatype".into(),
));
}
let mut buf = data.to_vec();
let (minbits, minval) = if p.dtype_class == SO_CLS_INTEGER {
by_width!(p.size, precompress_int(&mut buf, &p))
} else {
precompress_float(&mut buf, &p)?
};
debug_assert!(minbits <= p.dtype_len());
let size_out = SO_BUF_OFFSET + nbytes * minbits as usize / (p.size * 8) + 1;
let mut out = vec![0u8; size_out];
out[..4].copy_from_slice(&minbits.to_le_bytes());
out[4] = 8;
out[5..13].copy_from_slice(&minval.to_le_bytes());
if minbits as usize == p.size * 8 {
out.truncate(SO_BUF_OFFSET + nbytes);
out[SO_BUF_OFFSET..].copy_from_slice(&buf);
return Ok(out);
}
if minbits != 0 {
let mut w = BitWriter::new(&mut out[SO_BUF_OFFSET..]);
let le = p.order == SO_ORDER_LE;
by_width!(p.size, so_pack(&buf, le, minbits, &mut w));
w.finish()?;
}
Ok(out)
}
fn precompress_int<const N: usize>(buf: &mut [u8], p: &SoParams) -> (u32, u64) {
let signed = p.dtype_sign == SO_SGN_2;
let le = p.order == SO_ORDER_LE;
let width_mask = p.width_mask();
let key = |raw: u64| -> i128 {
if signed {
i128::from(sign_extend::<N>(raw))
} else {
i128::from(raw)
}
};
let (elems, _) = buf.as_chunks_mut::<N>();
let mut minbits = p.scale_factor as u32;
let mut min: i128 = 0;
let mut max: i128 = 0;
if p.fill_defined {
let first = elems.iter().position(|&e| load_uint(e, le) != p.filval);
if let Some(f) = first {
min = key(load_uint(elems[f], le));
max = min;
for &e in &elems[f..] {
let raw = load_uint(e, le);
if raw == p.filval {
continue;
}
let v = key(raw);
max = max.max(v);
min = min.min(v);
}
}
if minbits == 0 {
let span_minus_1 = (max - min) as u64;
if span_minus_1 > width_mask - 2 {
return (p.dtype_len(), 0);
}
minbits = so_log2(span_minus_1 + 2);
}
if minbits != p.dtype_len() {
let sentinel = mask_u64(minbits as usize);
for e in elems.iter_mut() {
let raw = load_uint(*e, le);
let v = if raw == p.filval {
sentinel
} else {
(key(raw) - min) as u64 & width_mask
};
*e = store_uint(v, le);
}
}
} else {
if let Some(&e0) = elems.first() {
min = key(load_uint(e0, le));
max = min;
}
for &e in elems.iter() {
let v = key(load_uint(e, le));
max = max.max(v);
min = min.min(v);
}
if minbits == 0 {
let span_minus_1 = (max - min) as u64;
if span_minus_1 > width_mask - 2 {
return (p.dtype_len(), 0);
}
minbits = so_log2(span_minus_1 + 1);
}
if minbits != p.dtype_len() {
for e in elems.iter_mut() {
let v = (key(load_uint(*e, le)) - min) as u64 & width_mask;
*e = store_uint(v, le);
}
}
}
(
minbits,
min as i64 as u64 & if signed { u64::MAX } else { width_mask },
)
}
trait SoFloat:
Copy
+ PartialOrd
+ std::ops::Mul<Output = Self>
+ std::ops::Sub<Output = Self>
+ std::ops::Div<Output = Self>
+ std::ops::Add<Output = Self>
{
const ZERO: Self;
fn from_stored(v: u64) -> Self;
fn to_stored(self) -> u64;
fn widen(self) -> f64;
fn narrow(v: f64) -> Self;
fn from_int(v: i64) -> Self;
fn pow(base: f64, exp: f64) -> Self;
fn abs(self) -> Self;
fn round(self) -> Self;
fn lround(self) -> i64;
}
impl SoFloat for f32 {
const ZERO: Self = 0.0;
fn from_stored(v: u64) -> Self {
f32::from_bits(v as u32)
}
fn to_stored(self) -> u64 {
self.to_bits() as u64
}
fn widen(self) -> f64 {
self as f64
}
fn narrow(v: f64) -> Self {
v as f32
}
fn from_int(v: i64) -> Self {
v as f32
}
fn pow(base: f64, exp: f64) -> Self {
(base as f32).powf(exp as f32)
}
fn abs(self) -> Self {
f32::abs(self)
}
fn round(self) -> Self {
f32::round(self)
}
fn lround(self) -> i64 {
f32::round(self) as i64
}
}
impl SoFloat for f64 {
const ZERO: Self = 0.0;
fn from_stored(v: u64) -> Self {
f64::from_bits(v)
}
fn to_stored(self) -> u64 {
self.to_bits()
}
fn widen(self) -> f64 {
self
}
fn narrow(v: f64) -> Self {
v
}
fn from_int(v: i64) -> Self {
v as f64
}
fn pow(base: f64, exp: f64) -> Self {
base.powf(exp)
}
fn abs(self) -> Self {
f64::abs(self)
}
fn round(self) -> Self {
f64::round(self)
}
fn lround(self) -> i64 {
f64::round(self) as i64
}
}
fn precompress_float(buf: &mut [u8], p: &SoParams) -> FormatResult<(u32, u64)> {
match p.size {
4 => Ok(precompress_float_typed::<f32, 4>(buf, p)),
8 => Ok(precompress_float_typed::<f64, 8>(buf, p)),
n => Err(FormatError::InvalidData(format!(
"scaleoffset: no floating-point type of {n} bytes"
))),
}
}
fn precompress_float_typed<T: SoFloat, const N: usize>(buf: &mut [u8], p: &SoParams) -> (u32, u64) {
let d_val = p.scale_factor as f64;
let pow10 = T::pow(10.0, d_val);
let filval = T::from_stored(p.filval);
let le = p.order == SO_ORDER_LE;
let get = |e: [u8; N]| T::from_stored(load_uint(e, le));
let (elems, _) = buf.as_chunks_mut::<N>();
let scan_epsilon = 10f64.powf(-d_val);
let is_fill_scan = |v: T| (v - filval).widen().abs() < scan_epsilon;
let modify_epsilon = T::pow(10.0, -d_val);
let is_fill_modify = |v: T| (v - filval).abs() < modify_epsilon;
let mut min = T::ZERO;
let mut max = T::ZERO;
if p.fill_defined {
if let Some(f) = elems.iter().position(|&e| !is_fill_scan(get(e))) {
min = get(elems[f]);
max = min;
for &e in &elems[f..] {
let v = get(e);
if is_fill_scan(v) {
continue;
}
if v > max {
max = v;
}
if v < min {
min = v;
}
}
}
} else if let Some(&e0) = elems.first() {
min = get(e0);
max = min;
for &e in elems.iter() {
let v = get(e);
if v > max {
max = v;
}
if v < min {
min = v;
}
}
}
let dtype_len = p.dtype_len();
let scaled = max * pow10 - min * pow10;
if scaled.round() > T::pow(2.0, (dtype_len - 1) as f64) {
return (dtype_len, 0);
}
let span = scaled.lround() as u64 + 1;
let minbits = if p.fill_defined {
so_log2(span + 1)
} else {
so_log2(span)
};
if minbits != dtype_len {
let sentinel = mask_u64(minbits as usize);
for e in elems.iter_mut() {
let v = get(*e);
let stored = if p.fill_defined && is_fill_modify(v) {
sentinel
} else {
(v * pow10 - min * pow10).lround() as u64 & p.width_mask()
};
*e = store_uint(stored, le);
}
}
(minbits, min.to_stored())
}
#[inline]
fn sign_extend<const N: usize>(v: u64) -> i64 {
let shift = 64 - 8 * N as u32;
((v << shift) as i64) >> shift
}
fn postdecompress(out: &mut [u8], p: &SoParams, minbits: u32, minval: u64) {
let sentinel = mask_u64(minbits as usize);
if p.dtype_class == SO_CLS_INTEGER {
by_width!(p.size, postdecompress_int(out, p, sentinel, minval));
} else {
match p.size {
4 => postdecompress_float::<f32, 4>(out, p, sentinel, minval),
8 => postdecompress_float::<f64, 8>(out, p, sentinel, minval),
_ => {}
}
}
}
fn postdecompress_int<const N: usize>(out: &mut [u8], p: &SoParams, sentinel: u64, minval: u64) {
let le = p.order == SO_ORDER_LE;
let width_mask = p.width_mask();
let (elems, _) = out.as_chunks_mut::<N>();
for e in elems {
let v = load_uint(*e, le);
let result = if p.fill_defined && v == sentinel {
p.filval
} else {
v.wrapping_add(minval) & width_mask
};
*e = store_uint(result, le);
}
}
fn postdecompress_float<T: SoFloat, const N: usize>(
out: &mut [u8],
p: &SoParams,
sentinel: u64,
minval: u64,
) {
let le = p.order == SO_ORDER_LE;
let divisor = T::narrow(10f64.powf(p.scale_factor as f64));
let min = T::from_stored(minval);
let filval = T::from_stored(p.filval);
let (elems, _) = out.as_chunks_mut::<N>();
for e in elems {
let raw = load_uint(*e, le);
let val = if p.fill_defined && raw == sentinel {
filval
} else {
T::from_int(sign_extend::<N>(raw)) / divisor + min
};
*e = store_uint(val.to_stored(), le);
}
}
use crate::format::messages::datatype::{ByteOrder, DatatypeMessage};
fn is_standard_ieee_float(dt: &DatatypeMessage) -> bool {
dt.ieee_format().is_some()
}
pub fn datatype_needs_bit_conversion(dt: &DatatypeMessage) -> bool {
match dt {
DatatypeMessage::FixedPoint {
size,
bit_offset,
bit_precision,
..
} => *bit_offset != 0 || (*bit_precision as u32) < *size * 8,
DatatypeMessage::FloatingPoint { .. } => !is_standard_ieee_float(dt),
_ => false,
}
}
pub fn apply_datatype_conversion(buffer: &mut [u8], dt: &DatatypeMessage) -> FormatResult<()> {
match dt {
DatatypeMessage::FixedPoint {
size,
byte_order,
signed,
bit_offset,
bit_precision,
} => {
let size = *size as usize;
let precision = *bit_precision as usize;
let offset = *bit_offset as usize;
if offset == 0 && precision == size * 8 {
return Ok(());
}
if size == 0 || size > 8 {
return Err(FormatError::InvalidData(format!(
"datatype conversion: unsupported FixedPoint size {size}"
)));
}
if precision == 0 || offset + precision > size * 8 {
return Err(FormatError::InvalidData(format!(
"datatype conversion: invalid bit layout (offset {offset}, \
precision {precision}, size {size})"
)));
}
if !buffer.len().is_multiple_of(size) {
return Err(FormatError::InvalidData(format!(
"datatype conversion: buffer length {} not a multiple of \
element size {size}",
buffer.len()
)));
}
let big_endian = matches!(byte_order, ByteOrder::BigEndian);
let precision_mask: u64 = if precision == 64 {
u64::MAX
} else {
(1u64 << precision) - 1
};
let sign_bit: u64 = 1u64 << (precision - 1);
for elem in buffer.chunks_exact_mut(size) {
let mut raw: u64 = 0;
if big_endian {
for &b in elem.iter() {
raw = (raw << 8) | b as u64;
}
} else {
for (i, &b) in elem.iter().enumerate() {
raw |= (b as u64) << (8 * i);
}
}
let mut value = (raw >> offset) & precision_mask;
if *signed && (value & sign_bit) != 0 {
value |= !precision_mask;
}
if big_endian {
for i in 0..size {
elem[size - 1 - i] = (value >> (8 * i)) as u8;
}
} else {
for (i, b) in elem.iter_mut().enumerate() {
*b = (value >> (8 * i)) as u8;
}
}
}
Ok(())
}
DatatypeMessage::FloatingPoint { .. } => {
if is_standard_ieee_float(dt) {
Ok(())
} else {
Err(FormatError::InvalidData(
"datatype conversion: non-standard floating-point bit \
layout cannot be converted"
.into(),
))
}
}
_ => Ok(()),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn nbit_atomic_cd(d_nelmts: u32, size: u32, precision: u32, offset: u32) -> Vec<u32> {
let need_not_compress = if offset == 0 && precision == size * 8 {
1
} else {
0
};
vec![
8,
need_not_compress,
d_nelmts,
NBIT_ATOMIC,
size,
NBIT_ORDER_LE,
precision,
offset,
]
}
#[test]
fn nbit_roundtrip_u16_precision12() {
let values: Vec<u16> = (0..40u16).map(|i| (i * 71) & 0x0FFF).collect();
let mut raw = Vec::new();
for &v in &values {
raw.extend_from_slice(&v.to_le_bytes());
}
let cd = nbit_atomic_cd(values.len() as u32, 2, 12, 0);
let packed = apply_nbit(&raw, &cd, true).unwrap();
assert!(packed.len() <= raw.len());
let unpacked = apply_nbit(&packed, &cd, false).unwrap();
assert_eq!(unpacked, raw);
}
#[test]
fn nbit_roundtrip_u32_precision20_offset4() {
let values: Vec<u32> = (0..32u32).map(|i| ((i * 9999) & 0xFFFFF) << 4).collect();
let mut raw = Vec::new();
for &v in &values {
raw.extend_from_slice(&v.to_le_bytes());
}
let cd = nbit_atomic_cd(values.len() as u32, 4, 20, 4);
let packed = apply_nbit(&raw, &cd, true).unwrap();
let unpacked = apply_nbit(&packed, &cd, false).unwrap();
assert_eq!(unpacked, raw);
}
#[test]
fn nbit_passthrough_full_precision() {
let raw: Vec<u8> = (0..64).collect();
let cd = nbit_atomic_cd(16, 4, 32, 0); let packed = apply_nbit(&raw, &cd, true).unwrap();
assert_eq!(packed, raw);
let unpacked = apply_nbit(&packed, &cd, false).unwrap();
assert_eq!(unpacked, raw);
}
#[test]
fn nbit_roundtrip_big_endian() {
let values: Vec<u16> = (0..24u16).map(|i| (i * 53) & 0x03FF).collect();
let mut raw = Vec::new();
for &v in &values {
raw.extend_from_slice(&v.to_be_bytes());
}
let mut cd = nbit_atomic_cd(values.len() as u32, 2, 10, 0);
cd[5] = NBIT_ORDER_BE;
let packed = apply_nbit(&raw, &cd, true).unwrap();
let unpacked = apply_nbit(&packed, &cd, false).unwrap();
assert_eq!(unpacked, raw);
}
fn refusal(data: &[u8], cd: &[u32], compress: bool) -> String {
match apply_nbit(data, cd, compress) {
Ok(out) => panic!("accepted {cd:?}: {} bytes out", out.len()),
Err(e) => e.to_string(),
}
}
#[test]
fn a_list_without_the_datatype_size_is_refused() {
let err = refusal(&[0; 8], &[4, 0, 1, NBIT_ATOMIC], false);
assert!(err.contains("cd_values too short"), "{err}");
}
#[test]
fn a_count_that_disagrees_with_the_list_is_refused() {
let mut cd = nbit_atomic_cd(1, 2, 12, 0);
cd[0] = 9;
let err = refusal(&[0; 2], &cd, false);
assert!(err.contains("names 9 parameters but 8"), "{err}");
}
#[test]
fn an_array_whose_base_lies_past_the_list_is_refused() {
let cd = [6, 0, 1, NBIT_ARRAY, 8, NBIT_ARRAY];
for compress in [false, true] {
let err = refusal(&[0; 8], &cd, compress);
assert!(err.contains("parameter list truncated"), "{err}");
}
}
#[test]
fn an_array_over_a_zero_sized_base_is_refused() {
let cd = [8, 0, 1, NBIT_ARRAY, 8, NBIT_COMPOUND, 0, 0];
for compress in [false, true] {
let err = refusal(&[0; 8], &cd, compress);
assert!(err.contains("zero-sized array base type"), "{err}");
}
}
#[test]
fn a_member_past_the_element_is_refused() {
let cd = [12, 0, 1, NBIT_COMPOUND, 4, 1, 8, NBIT_ATOMIC, 4, 0, 32, 0];
for compress in [false, true] {
let err = refusal(&[0; 4], &cd, compress);
assert!(err.contains("element extends past buffer"), "{err}");
}
}
#[test]
fn a_tree_that_packs_more_bits_than_the_element_holds_is_refused() {
let cd = [
18,
0,
1,
NBIT_COMPOUND,
4,
2,
0,
NBIT_ATOMIC,
4,
0,
32,
0,
0,
NBIT_ATOMIC,
4,
0,
32,
0,
];
let err = refusal(&[0xAB; 4], &cd, true);
assert!(err.contains("packed stream longer"), "{err}");
}
#[test]
fn an_element_count_no_machine_can_hold_is_refused() {
let cd = [8, 0, u32::MAX, NBIT_ATOMIC, u32::MAX, NBIT_ORDER_LE, 1, 0];
let err = refusal(&[0; 8], &cd, false);
assert!(err.contains("cannot allocate"), "{err}");
}
fn fixed(size: u32, signed: bool, offset: u16, precision: u16) -> DatatypeMessage {
DatatypeMessage::FixedPoint {
size,
byte_order: ByteOrder::LittleEndian,
signed,
bit_offset: offset,
bit_precision: precision,
}
}
#[test]
fn conversion_noop_for_full_width_types() {
let dt = fixed(4, false, 0, 32);
assert!(!datatype_needs_bit_conversion(&dt));
let mut buf = vec![0x78, 0x56, 0x34, 0x12, 0xFF, 0xFF, 0xFF, 0xFF];
let before = buf.clone();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(buf, before);
}
#[test]
fn conversion_noop_for_non_numeric_types() {
let dt = DatatypeMessage::fixed_string(8);
assert!(!datatype_needs_bit_conversion(&dt));
let mut buf = b"hello!!\0".to_vec();
let before = buf.clone();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(buf, before);
}
#[test]
fn conversion_unsigned_offset_shifts_right() {
let dt = fixed(2, false, 3, 10);
assert!(datatype_needs_bit_conversion(&dt));
let mut buf = (0x1528u16).to_le_bytes().to_vec();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(u16::from_le_bytes([buf[0], buf[1]]), 0x2A5);
}
#[test]
fn conversion_signed_negative_sign_extends() {
let dt = fixed(2, true, 4, 8);
let mut buf = (0x0FD0u16).to_le_bytes().to_vec();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(i16::from_le_bytes([buf[0], buf[1]]), -3);
}
#[test]
fn conversion_signed_positive_stays_positive() {
let dt = fixed(2, true, 4, 8);
let mut buf = (0x0050u16).to_le_bytes().to_vec();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(i16::from_le_bytes([buf[0], buf[1]]), 5);
}
#[test]
fn conversion_reduced_precision_offset_zero() {
let dt = fixed(4, true, 0, 20);
assert!(datatype_needs_bit_conversion(&dt));
let mut buf = (0x000FFFFFu32).to_le_bytes().to_vec();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(i32::from_le_bytes(buf.clone().try_into().unwrap()), -1);
}
#[test]
fn conversion_big_endian_signed() {
let dt = DatatypeMessage::FixedPoint {
size: 2,
byte_order: ByteOrder::BigEndian,
signed: true,
bit_offset: 4,
bit_precision: 8,
};
let mut buf = (0x0FD0u16).to_be_bytes().to_vec();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(i16::from_be_bytes([buf[0], buf[1]]), -3);
}
#[test]
fn conversion_multiple_elements() {
let dt = fixed(4, false, 5, 16);
let vals: [u32; 3] = [0x1234, 0xABCD, 0x0001];
let mut buf = Vec::new();
for v in vals {
buf.extend_from_slice(&(v << 5).to_le_bytes());
}
apply_datatype_conversion(&mut buf, &dt).unwrap();
for (i, v) in vals.iter().enumerate() {
let e = u32::from_le_bytes(buf[i * 4..i * 4 + 4].try_into().unwrap());
assert_eq!(e, *v);
}
}
#[test]
fn conversion_rejects_non_standard_float() {
let dt = DatatypeMessage::FloatingPoint {
size: 4,
byte_order: ByteOrder::LittleEndian,
sign_location: 30,
bit_offset: 1,
bit_precision: 31,
exponent_location: 22,
exponent_size: 8,
mantissa_location: 0,
mantissa_size: 22,
exponent_bias: 127,
};
assert!(datatype_needs_bit_conversion(&dt));
let mut buf = vec![0u8; 4];
assert!(apply_datatype_conversion(&mut buf, &dt).is_err());
}
#[test]
fn conversion_standard_float_is_noop() {
let dt = DatatypeMessage::f64_type();
assert!(!datatype_needs_bit_conversion(&dt));
let mut buf = 12.5f64.to_le_bytes().to_vec();
let before = buf.clone();
apply_datatype_conversion(&mut buf, &dt).unwrap();
assert_eq!(buf, before);
}
#[test]
fn conversion_rejects_bad_buffer_length() {
let dt = fixed(4, false, 3, 16);
let mut buf = vec![0u8; 5]; assert!(apply_datatype_conversion(&mut buf, &dt).is_err());
}
}