Skip to main content

rust_hdf5/format/
szip.rs

1// Copyright 2024 Mathis Rosenhauer, Moritz Hanke, Joerg Behrens, Luis Kornblueh
2// All rights reserved.
3//
4// Redistribution and use in source and binary forms, with or without
5// modification, are permitted provided that the following conditions
6// are met:
7//
8// 1. Redistributions of source code must retain the above copyright
9//    notice, this list of conditions and the following disclaimer.
10// 2. Redistributions in binary form must reproduce the above
11//    copyright notice, this list of conditions and the following
12//    disclaimer in the documentation and/or other materials provided
13//    with the distribution.
14//
15// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
16// "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
17// LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS
18// FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE
19// COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT,
20// INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES
21// (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
22// SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION)
23// HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
24// STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
25// ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED
26// OF THE POSSIBILITY OF SUCH DAMAGE.
27//
28// Pure Rust port of libaec (Adaptive Entropy Coding) for SZIP/HDF5.
29// Based on the CCSDS recommended standard 121.0-B-3.
30
31// ---------------------------------------------------------------------------
32// SZIP option masks (HDF5 filter interface, see H5Zpublic.h / szlib.h)
33// ---------------------------------------------------------------------------
34#[allow(dead_code)]
35const SZ_ALLOW_K13_OPTION_MASK: u32 = 1;
36#[allow(dead_code)]
37const SZ_CHIP_OPTION_MASK: u32 = 2;
38#[allow(dead_code)]
39const SZ_EC_OPTION_MASK: u32 = 4;
40#[allow(dead_code)]
41const SZ_LSB_OPTION_MASK: u32 = 8;
42const SZ_MSB_OPTION_MASK: u32 = 16;
43const SZ_NN_OPTION_MASK: u32 = 32;
44// RAW: libaec's szlib wrapper accepts but does not act on this mask for the
45// interleave decision (see sz_compat.c). Kept for documentation only.
46#[allow(dead_code)]
47const SZ_RAW_OPTION_MASK: u32 = 128;
48
49// ---------------------------------------------------------------------------
50// AEC flags
51// ---------------------------------------------------------------------------
52const AEC_DATA_SIGNED: u32 = 1;
53#[allow(dead_code)]
54const AEC_DATA_3BYTE: u32 = 2;
55const AEC_DATA_MSB: u32 = 4;
56const AEC_DATA_PREPROCESS: u32 = 8;
57const AEC_RESTRICTED: u32 = 16;
58const AEC_NOT_ENFORCE: u32 = 64;
59
60// Marker for Remainder Of Segment in zero block encoding
61const ROS_ENC: i32 = -1;
62const ROS_DEC: u32 = 5;
63
64const SE_TABLE_SIZE: usize = 90;
65
66// ---------------------------------------------------------------------------
67// Convert SZIP option mask to AEC flags
68// ---------------------------------------------------------------------------
69fn convert_options(sz_opts: u32) -> u32 {
70    let mut flags: u32 = 0;
71    if sz_opts & SZ_MSB_OPTION_MASK != 0 {
72        flags |= AEC_DATA_MSB;
73    }
74    if sz_opts & SZ_NN_OPTION_MASK != 0 {
75        flags |= AEC_DATA_PREPROCESS;
76    }
77    flags
78}
79
80fn bits_to_bytes(bits: u32) -> u32 {
81    if bits > 16 {
82        4
83    } else if bits > 8 {
84        2
85    } else {
86        1
87    }
88}
89
90// ---------------------------------------------------------------------------
91// SZIP interleaving (for 32/64-bit samples)
92// ---------------------------------------------------------------------------
93fn interleave_buffer(src: &[u8], wordsize: usize) -> Vec<u8> {
94    let n = src.len();
95    let count = n / wordsize;
96    let mut dest = vec![0u8; n];
97    for i in 0..count {
98        for j in 0..wordsize {
99            dest[j * count + i] = src[i * wordsize + j];
100        }
101    }
102    dest
103}
104
105fn deinterleave_buffer(src: &[u8], wordsize: usize) -> Vec<u8> {
106    let n = src.len();
107    let count = n / wordsize;
108    let mut dest = vec![0u8; n];
109    for i in 0..count {
110        for j in 0..wordsize {
111            dest[i * wordsize + j] = src[j * count + i];
112        }
113    }
114    dest
115}
116
117// ---------------------------------------------------------------------------
118// Scanline padding helpers
119// ---------------------------------------------------------------------------
120fn add_padding(
121    src: &[u8],
122    line_size: usize,
123    padding_size: usize,
124    pixel_size: usize,
125    pp: bool,
126) -> Vec<u8> {
127    let padded_line = line_size + padding_size;
128    let num_lines = src.len().div_ceil(line_size);
129    let mut dest = vec![0u8; num_lines * padded_line];
130    let mut si = 0;
131    let mut di = 0;
132    while si < src.len() {
133        let ls = std::cmp::min(src.len() - si, line_size);
134        dest[di..di + ls].copy_from_slice(&src[si..si + ls]);
135        di += ls;
136        si += ls;
137        let pad_pixels = padded_line - ls;
138        let pixel: &[u8] = if pp && si >= pixel_size {
139            &src[si - pixel_size..si]
140        } else {
141            &[0u8; 4][..pixel_size]
142        };
143        for k in (0..pad_pixels).step_by(pixel_size) {
144            let end = std::cmp::min(k + pixel_size, pad_pixels);
145            dest[di + k..di + end].copy_from_slice(&pixel[..end - k]);
146        }
147        di += pad_pixels;
148    }
149    dest.truncate(di);
150    dest
151}
152
153fn remove_padding(buf: &mut Vec<u8>, line_size: usize, padding_size: usize) {
154    let padded = line_size + padding_size;
155    if padded == 0 || padding_size == 0 {
156        return;
157    }
158    let mut dst = line_size;
159    let mut src_off = padded;
160    while src_off < buf.len() {
161        let copy_len = std::cmp::min(line_size, buf.len() - src_off);
162        buf.copy_within(src_off..src_off + copy_len, dst);
163        dst += copy_len;
164        src_off += padded;
165    }
166    buf.truncate(dst);
167}
168
169// ---------------------------------------------------------------------------
170// Determine id_len from bits_per_sample and flags
171// ---------------------------------------------------------------------------
172fn compute_id_len(bits_per_sample: u32, flags: u32) -> Result<u32, String> {
173    if bits_per_sample > 16 {
174        Ok(5)
175    } else if bits_per_sample > 8 {
176        Ok(4)
177    } else if flags & AEC_RESTRICTED != 0 {
178        if bits_per_sample <= 2 {
179            Ok(1)
180        } else if bits_per_sample <= 4 {
181            Ok(2)
182        } else {
183            Err("restricted mode only supports <= 4 bits".into())
184        }
185    } else {
186        Ok(3)
187    }
188}
189
190// ===========================================================================
191//  ENCODER
192// ===========================================================================
193
194struct BitWriter {
195    buf: Vec<u8>,
196    bits: i32, // free bits in current byte (1..8)
197}
198
199impl BitWriter {
200    fn new() -> Self {
201        Self {
202            buf: vec![0u8],
203            bits: 8,
204        }
205    }
206
207    fn emit(&mut self, data: u32, mut nbits: i32) {
208        if nbits == 0 {
209            return;
210        }
211        if nbits <= self.bits {
212            self.bits -= nbits;
213            *self.buf.last_mut().unwrap() |= (data << self.bits) as u8;
214        } else {
215            nbits -= self.bits;
216            *self.buf.last_mut().unwrap() |=
217                ((data as u64 >> nbits) as u8) & ((1u16 << self.bits) - 1) as u8;
218            while nbits > 8 {
219                nbits -= 8;
220                self.buf.push((data >> nbits) as u8);
221            }
222            self.bits = 8 - nbits;
223            self.buf.push((data << self.bits) as u8);
224        }
225    }
226
227    fn emitfs(&mut self, fs: u32) {
228        // fs zero bits followed by one 1 bit
229        let mut remaining = fs as i32;
230        loop {
231            if remaining < self.bits {
232                self.bits -= remaining + 1;
233                *self.buf.last_mut().unwrap() |= 1u8 << self.bits;
234                break;
235            } else {
236                remaining -= self.bits;
237                self.buf.push(0);
238                self.bits = 8;
239            }
240        }
241    }
242
243    fn emit_block_fs(&mut self, block: &[u32], k: u32, ref_skip: usize) {
244        // Emit fundamental sequences for each sample's high bits (sample >> k)
245        for &s in &block[ref_skip..] {
246            self.emitfs(s >> k);
247        }
248    }
249
250    fn emit_block(&mut self, block: &[u32], k: u32, ref_skip: usize) {
251        // Emit k LSBs of each sample
252        if k == 0 {
253            return;
254        }
255        let mask = (1u64 << k) - 1;
256        for &s in &block[ref_skip..] {
257            self.emit((s as u64 & mask) as u32, k as i32);
258        }
259    }
260
261    fn flush_to_byte(&mut self) {
262        // Pad remaining bits with zeros to byte boundary
263        if self.bits < 8 {
264            self.bits = 8;
265            self.buf.push(0);
266        }
267    }
268
269    fn finish(mut self) -> Vec<u8> {
270        // Remove trailing empty byte if fully aligned
271        if self.bits == 8 && self.buf.len() > 1 {
272            // The last byte is fully empty (no bits used)
273            if *self.buf.last().unwrap() == 0 {
274                self.buf.pop();
275            }
276        }
277        self.buf
278    }
279}
280
281fn preprocess_unsigned(raw: &[u32], xmax: u32) -> Vec<u32> {
282    let n = raw.len();
283    let mut d = vec![0u32; n];
284    d[0] = 0; // placeholder for ref sample
285    for i in 0..n - 1 {
286        if raw[i + 1] >= raw[i] {
287            let diff = raw[i + 1] - raw[i];
288            if diff <= raw[i] {
289                d[i + 1] = 2 * diff;
290            } else {
291                d[i + 1] = raw[i + 1];
292            }
293        } else {
294            let diff = raw[i] - raw[i + 1];
295            if diff <= xmax - raw[i] {
296                d[i + 1] = 2 * diff - 1;
297            } else {
298                d[i + 1] = xmax - raw[i + 1];
299            }
300        }
301    }
302    d
303}
304
305fn preprocess_signed(raw: &[u32], bits_per_sample: u32, xmax: u32) -> Vec<u32> {
306    let n = raw.len();
307    let mut d = vec![0u32; n];
308    let m = 1u32 << (bits_per_sample - 1);
309
310    // Sign-extend all samples into i32 values, then compute deltas
311    let mut sx = vec![0i32; n];
312    for i in 0..n {
313        sx[i] = ((raw[i] ^ m).wrapping_sub(m)) as i32;
314    }
315
316    d[0] = 0;
317    for i in 0..n - 1 {
318        let cur = sx[i];
319        let nxt = sx[i + 1];
320        if nxt < cur {
321            let diff = (cur as u32).wrapping_sub(nxt as u32);
322            if diff <= xmax.wrapping_add((cur as u32).wrapping_add(1)) {
323                // half_d style: 2*D - 1
324                d[i + 1] = 2u32.wrapping_mul(diff).wrapping_sub(1);
325            } else {
326                d[i + 1] = xmax.wrapping_sub(nxt as u32);
327            }
328        } else {
329            let diff = (nxt as u32).wrapping_sub(cur as u32);
330            let xmin_val = (!xmax) as i32; // ~xmax gives xmin for signed
331            if diff <= (cur as u32).wrapping_sub(xmin_val as u32) {
332                d[i + 1] = 2u32.wrapping_mul(diff);
333            } else {
334                d[i + 1] = (nxt as u32).wrapping_sub(xmin_val as u32);
335            }
336        }
337    }
338    d
339}
340
341fn assess_splitting(
342    block: &[u32],
343    block_size: usize,
344    has_ref: bool,
345    prev_k: u32,
346    kmax: u32,
347) -> (u32, u32) {
348    let this_bs = if has_ref { block_size - 1 } else { block_size } as u64;
349    let effective = if has_ref { &block[1..] } else { block };
350
351    let mut len_min = u64::MAX;
352    let mut k = prev_k;
353    let mut k_min = k;
354    let mut no_turn = k == 0;
355    let mut dir = true; // true = increasing k
356
357    loop {
358        let fs_len: u64 = effective.iter().map(|&s| (s >> k) as u64).sum();
359        let len = fs_len + this_bs * (k as u64 + 1);
360
361        if len < len_min {
362            if len_min < u64::MAX {
363                no_turn = true;
364            }
365            len_min = len;
366            k_min = k;
367
368            if dir {
369                if fs_len < this_bs || k >= kmax {
370                    if no_turn {
371                        break;
372                    }
373                    if prev_k == 0 {
374                        break;
375                    }
376                    k = prev_k - 1;
377                    dir = false;
378                    no_turn = true;
379                } else {
380                    k += 1;
381                }
382            } else {
383                if fs_len >= this_bs || k == 0 {
384                    break;
385                }
386                k -= 1;
387            }
388        } else {
389            if no_turn {
390                break;
391            }
392            if prev_k == 0 {
393                break;
394            }
395            k = prev_k - 1;
396            dir = false;
397            no_turn = true;
398        }
399    }
400    (k_min, len_min as u32)
401}
402
403fn assess_se(block: &[u32], block_size: usize, uncomp_len: u32) -> u32 {
404    let mut len = 1u64;
405    let mut i = 0;
406    while i < block_size {
407        // `d` can reach ~2*u32::MAX (~2^33) when bits_per_sample is large, so
408        // `d * (d + 1)` (~2^66) overflows u64. SE's triangular code length is
409        // only ever compared against `uncomp_len`; any term that big means SE
410        // already lost, so saturate instead of wrapping (which would panic in
411        // debug and silently corrupt the cost estimate in release).
412        let d = block[i] as u64 + block[i + 1] as u64;
413        let triangular = d.saturating_mul(d + 1) / 2;
414        len = len
415            .saturating_add(triangular)
416            .saturating_add(block[i + 1] as u64)
417            .saturating_add(1);
418        if len > uncomp_len as u64 {
419            return u32::MAX;
420        }
421        i += 2;
422    }
423    len as u32
424}
425
426struct Encoder {
427    bits_per_sample: u32,
428    block_size: u32,
429    rsi: u32,
430    flags: u32,
431    id_len: u32,
432    kmax: u32,
433    xmax: u32,
434    bytes_per_sample: u32,
435}
436
437impl Encoder {
438    fn new(bits_per_sample: u32, block_size: u32, rsi: u32, flags: u32) -> Result<Self, String> {
439        let id_len = compute_id_len(bits_per_sample, flags)?;
440        let kmax = (1u32 << id_len) - 3;
441        let xmax = if flags & AEC_DATA_SIGNED != 0 {
442            ((1u64 << (bits_per_sample - 1)) - 1) as u32
443        } else {
444            ((1u64 << bits_per_sample) - 1) as u32
445        };
446        let bytes_per_sample = bits_to_bytes(bits_per_sample);
447        Ok(Self {
448            bits_per_sample,
449            block_size,
450            rsi,
451            flags,
452            id_len,
453            kmax,
454            xmax,
455            bytes_per_sample,
456        })
457    }
458
459    fn read_samples(&self, data: &[u8]) -> Vec<u32> {
460        let bps = self.bytes_per_sample as usize;
461        let msb = self.flags & AEC_DATA_MSB != 0;
462        let n = data.len() / bps;
463        let mut samples = Vec::with_capacity(n);
464        for i in 0..n {
465            let off = i * bps;
466            let s = match bps {
467                1 => data[off] as u32,
468                2 => {
469                    if msb {
470                        ((data[off] as u32) << 8) | data[off + 1] as u32
471                    } else {
472                        (data[off] as u32) | ((data[off + 1] as u32) << 8)
473                    }
474                }
475                3 => {
476                    if msb {
477                        ((data[off] as u32) << 16)
478                            | ((data[off + 1] as u32) << 8)
479                            | data[off + 2] as u32
480                    } else {
481                        (data[off] as u32)
482                            | ((data[off + 1] as u32) << 8)
483                            | ((data[off + 2] as u32) << 16)
484                    }
485                }
486                4 => {
487                    if msb {
488                        ((data[off] as u32) << 24)
489                            | ((data[off + 1] as u32) << 16)
490                            | ((data[off + 2] as u32) << 8)
491                            | data[off + 3] as u32
492                    } else {
493                        (data[off] as u32)
494                            | ((data[off + 1] as u32) << 8)
495                            | ((data[off + 2] as u32) << 16)
496                            | ((data[off + 3] as u32) << 24)
497                    }
498                }
499                _ => unreachable!(),
500            };
501            // Mask to bits_per_sample
502            let mask = if self.bits_per_sample == 32 {
503                u32::MAX
504            } else {
505                (1u32 << self.bits_per_sample) - 1
506            };
507            samples.push(s & mask);
508        }
509        samples
510    }
511
512    fn encode(&self, input_data: &[u8]) -> Result<Vec<u8>, String> {
513        let rsi_samples = (self.rsi * self.block_size) as usize;
514        let samples = self.read_samples(input_data);
515        let total_samples = samples.len();
516
517        let mut writer = BitWriter::new();
518        let mut offset = 0;
519        let mut prev_k = 0u32;
520
521        while offset < total_samples {
522            // Read one RSI worth of samples (pad with last if short)
523            let avail = std::cmp::min(rsi_samples, total_samples - offset);
524            let mut rsi_buf = Vec::with_capacity(rsi_samples);
525            rsi_buf.extend_from_slice(&samples[offset..offset + avail]);
526            if avail < rsi_samples {
527                let last = *rsi_buf.last().unwrap_or(&0);
528                rsi_buf.resize(rsi_samples, last);
529            }
530
531            // Number of actual blocks to encode
532            let blocks_to_encode = if avail < rsi_samples {
533                let b = avail.div_ceil(self.block_size as usize);
534                if b == 0 {
535                    1
536                } else {
537                    b
538                }
539            } else {
540                self.rsi as usize
541            };
542
543            // Preprocess
544            let pp = if self.flags & AEC_DATA_PREPROCESS != 0 {
545                if self.flags & AEC_DATA_SIGNED != 0 {
546                    preprocess_signed(&rsi_buf, self.bits_per_sample, self.xmax)
547                } else {
548                    preprocess_unsigned(&rsi_buf, self.xmax)
549                }
550            } else {
551                rsi_buf.clone()
552            };
553
554            let ref_sample = rsi_buf[0];
555            let has_preprocess = self.flags & AEC_DATA_PREPROCESS != 0;
556
557            // Encode each block
558            let mut zero_blocks: i32 = 0;
559            let mut zero_ref = false;
560            let mut zero_ref_sample = 0u32;
561
562            let bs = self.block_size as usize;
563
564            for b in 0..blocks_to_encode {
565                let block_start = b * bs;
566                let block = &pp[block_start..block_start + bs];
567                let is_first = b == 0;
568                let has_ref = has_preprocess && is_first;
569
570                let uncomp_len = if has_ref {
571                    (bs as u32 - 1) * self.bits_per_sample
572                } else {
573                    bs as u32 * self.bits_per_sample
574                };
575
576                // Check if block is all zeros
577                let all_zero = block.iter().all(|&x| x == 0);
578
579                if all_zero {
580                    zero_blocks += 1;
581                    if zero_blocks == 1 {
582                        zero_ref = has_ref;
583                        zero_ref_sample = ref_sample;
584                    }
585                    // Check if we need to flush zero blocks:
586                    // at end of RSI, or every 64 blocks
587                    let is_last = b + 1 >= blocks_to_encode;
588                    let at_boundary = (b + 1) % 64 == 0;
589                    if is_last || at_boundary {
590                        if zero_blocks > 4 {
591                            zero_blocks = ROS_ENC;
592                        }
593                        // Encode zero block
594                        writer.emit(0, self.id_len as i32 + 1);
595                        if zero_ref {
596                            writer.emit(zero_ref_sample, self.bits_per_sample as i32);
597                        }
598                        if zero_blocks == ROS_ENC {
599                            writer.emitfs(4);
600                        } else if zero_blocks >= 5 {
601                            writer.emitfs(zero_blocks as u32);
602                        } else {
603                            writer.emitfs((zero_blocks - 1) as u32);
604                        }
605                        zero_blocks = 0;
606                    }
607                    continue;
608                }
609
610                // Non-zero block: first flush any pending zero blocks
611                if zero_blocks > 0 {
612                    writer.emit(0, self.id_len as i32 + 1);
613                    if zero_ref {
614                        writer.emit(zero_ref_sample, self.bits_per_sample as i32);
615                    }
616                    if zero_blocks == ROS_ENC {
617                        writer.emitfs(4);
618                    } else if zero_blocks >= 5 {
619                        writer.emitfs(zero_blocks as u32);
620                    } else {
621                        writer.emitfs((zero_blocks - 1) as u32);
622                    }
623                    zero_blocks = 0;
624                }
625
626                // Assess coding options
627                let (split_len, best_k) = if self.id_len > 1 {
628                    let (k, len) = assess_splitting(block, bs, has_ref, prev_k, self.kmax);
629                    prev_k = k;
630                    (len, k)
631                } else {
632                    (u32::MAX, 0)
633                };
634
635                // SE always operates on the full (even-sized) block
636                let se_len = if bs >= 2 {
637                    assess_se(block, bs, uncomp_len)
638                } else {
639                    u32::MAX
640                };
641
642                if split_len < uncomp_len {
643                    // libaec's m_select_code_option uses a strict `<` here:
644                    // when splitting and SE produce equal-length CDS, SE wins.
645                    if split_len < se_len {
646                        // Splitting (Golomb-Rice)
647                        writer.emit(best_k + 1, self.id_len as i32);
648                        if has_ref {
649                            writer.emit(ref_sample, self.bits_per_sample as i32);
650                        }
651                        writer.emit_block_fs(block, best_k, if has_ref { 1 } else { 0 });
652                        if best_k > 0 {
653                            writer.emit_block(block, best_k, if has_ref { 1 } else { 0 });
654                        }
655                    } else {
656                        // Second extension
657                        encode_se(
658                            &mut writer,
659                            block,
660                            bs,
661                            has_ref,
662                            ref_sample,
663                            self.id_len,
664                            self.bits_per_sample,
665                        );
666                    }
667                } else if uncomp_len <= se_len {
668                    // Uncompressed
669                    writer.emit((1u32 << self.id_len) - 1, self.id_len as i32);
670                    if has_ref {
671                        // For uncompressed, first sample is the raw reference
672                        let mut ublock = block.to_vec();
673                        ublock[0] = ref_sample;
674                        writer.emit_block(&ublock, self.bits_per_sample, 0);
675                    } else {
676                        writer.emit_block(block, self.bits_per_sample, 0);
677                    }
678                } else {
679                    // Second extension
680                    encode_se(
681                        &mut writer,
682                        block,
683                        bs,
684                        has_ref,
685                        ref_sample,
686                        self.id_len,
687                        self.bits_per_sample,
688                    );
689                }
690            }
691
692            offset += avail;
693        }
694
695        // Pad final byte
696        writer.flush_to_byte();
697        Ok(writer.finish())
698    }
699}
700
701fn encode_se(
702    writer: &mut BitWriter,
703    block: &[u32],
704    block_size: usize,
705    has_ref: bool,
706    ref_sample: u32,
707    id_len: u32,
708    bits_per_sample: u32,
709) {
710    // SE uses id_len + 1 bits: 0 followed by 1
711    writer.emit(1, id_len as i32 + 1);
712    if has_ref {
713        writer.emit(ref_sample, bits_per_sample as i32);
714    }
715    // Always encode the full block (block[0] is 0 for ref blocks from preprocessing)
716    let mut i = 0;
717    while i < block_size {
718        let a = block[i] as u64;
719        let b = block[i + 1] as u64;
720        let d = a + b;
721        let fs = d * (d + 1) / 2 + b;
722        writer.emitfs(fs as u32);
723        i += 2;
724    }
725}
726
727// ===========================================================================
728//  DECODER
729// ===========================================================================
730
731struct BitReader<'a> {
732    data: &'a [u8],
733    pos: usize,
734    acc: u64,
735    bitp: i32,
736}
737
738impl<'a> BitReader<'a> {
739    fn new(data: &'a [u8]) -> Self {
740        Self {
741            data,
742            pos: 0,
743            acc: 0,
744            bitp: 0,
745        }
746    }
747
748    fn fill(&mut self) {
749        while self.bitp <= 56 && self.pos < self.data.len() {
750            self.acc = (self.acc << 8) | self.data[self.pos] as u64;
751            self.pos += 1;
752            self.bitp += 8;
753        }
754    }
755
756    fn get_bits(&mut self, n: i32) -> u32 {
757        while self.bitp < n {
758            if self.pos < self.data.len() {
759                self.acc = (self.acc << 8) | self.data[self.pos] as u64;
760                self.pos += 1;
761                self.bitp += 8;
762            } else {
763                // pad with zeros
764                self.acc <<= 8;
765                self.bitp += 8;
766            }
767        }
768        self.bitp -= n;
769        ((self.acc >> self.bitp) & ((1u64 << n) - 1)) as u32
770    }
771
772    fn get_fs(&mut self) -> u32 {
773        let mut fs = 0u32;
774
775        // Mask accumulator to valid bits
776        if self.bitp > 0 {
777            self.acc &= (1u64 << self.bitp) - 1;
778        } else {
779            self.acc = 0;
780        }
781
782        while self.acc == 0 {
783            fs += self.bitp as u32;
784            self.acc = 0;
785            self.bitp = 0;
786            // read more bytes
787            let to_read = std::cmp::min(7, self.data.len() - self.pos);
788            if to_read == 0 {
789                return fs;
790            }
791            for _ in 0..to_read {
792                self.acc = (self.acc << 8) | self.data[self.pos] as u64;
793                self.pos += 1;
794                self.bitp += 8;
795            }
796        }
797
798        // Find highest set bit
799        let highest = 63 - self.acc.leading_zeros() as i32;
800        fs += (self.bitp - highest - 1) as u32;
801        self.bitp = highest; // consume the 1 bit
802        fs
803    }
804}
805
806fn create_se_table() -> [i32; 2 * (SE_TABLE_SIZE + 1)] {
807    let mut table = [0i32; 2 * (SE_TABLE_SIZE + 1)];
808    let mut k = 0usize;
809    for i in 0..13i32 {
810        let ms = k as i32;
811        for _j in 0..=i {
812            if k <= SE_TABLE_SIZE {
813                table[2 * k] = i;
814                table[2 * k + 1] = ms;
815            }
816            k += 1;
817        }
818    }
819    table
820}
821
822fn postprocess_unsigned(rsi_buf: &[u32], xmax: u32) -> Vec<u32> {
823    let n = rsi_buf.len();
824    if n == 0 {
825        return vec![];
826    }
827    let mut out = vec![0u32; n];
828    out[0] = rsi_buf[0]; // reference sample
829    let med = xmax / 2 + 1;
830
831    let mut data = out[0];
832    for i in 1..n {
833        let d = rsi_buf[i];
834        let half_d = (d >> 1) + (d & 1);
835        let mask = if data >= med { xmax } else { 0 };
836
837        if half_d <= (mask ^ data) {
838            data = data.wrapping_add((d >> 1) ^ (!((d & 1).wrapping_sub(1))));
839        } else {
840            data = mask ^ d;
841        }
842        out[i] = data;
843    }
844    out
845}
846
847fn postprocess_signed(rsi_buf: &[u32], bits_per_sample: u32, xmax: u32) -> Vec<u32> {
848    let n = rsi_buf.len();
849    if n == 0 {
850        return vec![];
851    }
852    let mut out = vec![0u32; n];
853    let m = 1u32 << (bits_per_sample - 1);
854    // Sign-extend the reference sample
855    let ref_val = (rsi_buf[0] ^ m).wrapping_sub(m);
856    out[0] = ref_val;
857
858    let mut data = ref_val;
859    for i in 1..n {
860        let d = rsi_buf[i];
861        let half_d = (d >> 1) + (d & 1);
862
863        if (data as i32) < 0 {
864            if half_d <= xmax.wrapping_add(data).wrapping_add(1) {
865                data = data.wrapping_add((d >> 1) ^ (!((d & 1).wrapping_sub(1))));
866            } else {
867                data = d.wrapping_sub(xmax).wrapping_sub(1);
868            }
869        } else {
870            if half_d <= xmax.wrapping_sub(data) {
871                data = data.wrapping_add((d >> 1) ^ (!((d & 1).wrapping_sub(1))));
872            } else {
873                data = xmax.wrapping_sub(d);
874            }
875        }
876        out[i] = data;
877    }
878    out
879}
880
881struct Decoder {
882    bits_per_sample: u32,
883    block_size: u32,
884    rsi: u32,
885    flags: u32,
886    id_len: u32,
887    xmax: u32,
888    bytes_per_sample: u32,
889}
890
891impl Decoder {
892    fn new(bits_per_sample: u32, block_size: u32, rsi: u32, flags: u32) -> Result<Self, String> {
893        let id_len = compute_id_len(bits_per_sample, flags)?;
894        let xmax = if flags & AEC_DATA_SIGNED != 0 {
895            ((1u64 << (bits_per_sample - 1)) - 1) as u32
896        } else {
897            ((1u64 << bits_per_sample) - 1) as u32
898        };
899        let bytes_per_sample = bits_to_bytes(bits_per_sample);
900        Ok(Self {
901            bits_per_sample,
902            block_size,
903            rsi,
904            flags,
905            id_len,
906            xmax,
907            bytes_per_sample,
908        })
909    }
910
911    fn decode(&self, compressed: &[u8], output_samples: usize) -> Result<Vec<u32>, String> {
912        let mut reader = BitReader::new(compressed);
913        reader.fill();
914
915        let se_table = create_se_table();
916        let rsi_samples = (self.rsi * self.block_size) as usize;
917        let pp = self.flags & AEC_DATA_PREPROCESS != 0;
918
919        let mut all_output: Vec<u32> = Vec::with_capacity(output_samples);
920
921        while all_output.len() < output_samples {
922            // Decode one RSI
923            let mut rsi_buf: Vec<u32> = Vec::with_capacity(rsi_samples);
924            let mut first_block_in_rsi = true;
925
926            while rsi_buf.len() < rsi_samples
927                && all_output.len() + rsi_buf.len() < output_samples + rsi_samples
928            {
929                let has_ref = pp && first_block_in_rsi;
930                let encoded_block_size = if has_ref {
931                    self.block_size - 1
932                } else {
933                    self.block_size
934                } as usize;
935
936                // Read ID
937                let id = reader.get_bits(self.id_len as i32);
938
939                if id == 0 {
940                    // Low entropy
941                    let sub_id = reader.get_bits(1);
942                    if sub_id == 1 {
943                        // Second extension
944                        if has_ref {
945                            rsi_buf.push(reader.get_bits(self.bits_per_sample as i32));
946                        }
947                        // SE decoding: i starts at ref (0 or 1), runs to block_size
948                        // Each iteration reads one FS and produces 1 or 2 samples
949                        let ref_offset = if has_ref { 1usize } else { 0 };
950                        let mut i = ref_offset;
951                        while i < self.block_size as usize {
952                            let m = reader.get_fs();
953                            if m as usize > SE_TABLE_SIZE {
954                                return Err("SE table overflow".into());
955                            }
956                            let d1 = m as i32 - se_table[2 * m as usize + 1];
957
958                            if (i & 1) == 0 {
959                                rsi_buf.push((se_table[2 * m as usize] - d1) as u32);
960                                i += 1;
961                            }
962                            rsi_buf.push(d1 as u32);
963                            i += 1;
964                        }
965                    } else {
966                        // Zero block
967                        if has_ref {
968                            rsi_buf.push(reader.get_bits(self.bits_per_sample as i32));
969                        }
970                        let fs = reader.get_fs();
971                        let mut zero_blocks = fs + 1;
972
973                        if zero_blocks == ROS_DEC {
974                            let b = rsi_buf.len() / self.block_size as usize;
975                            let remaining = self.rsi as usize - b;
976                            let boundary = 64 - (b % 64);
977                            zero_blocks = std::cmp::min(remaining, boundary) as u32;
978                        } else if zero_blocks > ROS_DEC {
979                            zero_blocks -= 1;
980                        }
981
982                        // `fs` (hence `zero_blocks`) is bitstream-derived; a
983                        // corrupt run of zero bits could ask for a huge
984                        // expansion. Clamp to what remains in the RSI so a
985                        // bad stream cannot drive an unbounded allocation.
986                        let zero_samples = (zero_blocks as usize * self.block_size as usize)
987                            .saturating_sub(if has_ref { 1 } else { 0 })
988                            .min(
989                                (self.rsi as usize * self.block_size as usize)
990                                    .saturating_sub(rsi_buf.len()),
991                            );
992                        rsi_buf.extend(std::iter::repeat_n(0, zero_samples));
993                    }
994                } else if id == (1u32 << self.id_len) - 1 {
995                    // Uncompressed
996                    for _ in 0..self.block_size {
997                        rsi_buf.push(reader.get_bits(self.bits_per_sample as i32));
998                    }
999                } else {
1000                    // Split (Golomb-Rice) with k = id - 1
1001                    let k = id - 1;
1002
1003                    if has_ref {
1004                        rsi_buf.push(reader.get_bits(self.bits_per_sample as i32));
1005                    }
1006
1007                    // Read FS parts
1008                    let base = rsi_buf.len();
1009                    for _ in 0..encoded_block_size {
1010                        let fs = reader.get_fs();
1011                        rsi_buf.push(fs << k);
1012                    }
1013
1014                    // Read binary parts and add
1015                    if k > 0 {
1016                        for j in 0..encoded_block_size {
1017                            let bits = reader.get_bits(k as i32);
1018                            rsi_buf[base + j] += bits;
1019                        }
1020                    }
1021                }
1022
1023                first_block_in_rsi = false;
1024
1025                // Check if RSI is complete
1026                if rsi_buf.len() >= rsi_samples {
1027                    break;
1028                }
1029            }
1030
1031            // Postprocess RSI
1032            if pp {
1033                let processed = if self.flags & AEC_DATA_SIGNED != 0 {
1034                    postprocess_signed(&rsi_buf, self.bits_per_sample, self.xmax)
1035                } else {
1036                    postprocess_unsigned(&rsi_buf, self.xmax)
1037                };
1038                all_output.extend_from_slice(&processed);
1039            } else {
1040                all_output.extend_from_slice(&rsi_buf);
1041            }
1042        }
1043
1044        all_output.truncate(output_samples);
1045        Ok(all_output)
1046    }
1047
1048    fn write_samples(&self, samples: &[u32], output_size: usize) -> Vec<u8> {
1049        let bps = self.bytes_per_sample as usize;
1050        let msb = self.flags & AEC_DATA_MSB != 0;
1051        let mut out = Vec::with_capacity(output_size);
1052        for &s in samples {
1053            match bps {
1054                1 => out.push(s as u8),
1055                2 => {
1056                    if msb {
1057                        out.push((s >> 8) as u8);
1058                        out.push(s as u8);
1059                    } else {
1060                        out.push(s as u8);
1061                        out.push((s >> 8) as u8);
1062                    }
1063                }
1064                3 => {
1065                    if msb {
1066                        out.push((s >> 16) as u8);
1067                        out.push((s >> 8) as u8);
1068                        out.push(s as u8);
1069                    } else {
1070                        out.push(s as u8);
1071                        out.push((s >> 8) as u8);
1072                        out.push((s >> 16) as u8);
1073                    }
1074                }
1075                4 => {
1076                    if msb {
1077                        out.push((s >> 24) as u8);
1078                        out.push((s >> 16) as u8);
1079                        out.push((s >> 8) as u8);
1080                        out.push(s as u8);
1081                    } else {
1082                        out.push(s as u8);
1083                        out.push((s >> 8) as u8);
1084                        out.push((s >> 16) as u8);
1085                        out.push((s >> 24) as u8);
1086                    }
1087                }
1088                _ => unreachable!(),
1089            }
1090            if out.len() >= output_size {
1091                break;
1092            }
1093        }
1094        out.truncate(output_size);
1095        out
1096    }
1097}
1098
1099// ===========================================================================
1100//  Public API
1101// ===========================================================================
1102
1103/// Compress data using the SZIP (AEC) algorithm.
1104///
1105/// Parameters match the HDF5 SZIP filter interface:
1106/// - `data`: raw uncompressed bytes
1107/// - `bits_per_pixel`: sample width (1-32, or 64 for double interleaving)
1108/// - `pixels_per_block`: block size (must be even, typically 8/16/32)
1109/// - `pixels_per_scanline`: scanline width in pixels
1110/// - `options_mask`: SZIP option flags
1111pub fn compress(
1112    data: &[u8],
1113    bits_per_pixel: u32,
1114    pixels_per_block: u32,
1115    pixels_per_scanline: u32,
1116    options_mask: u32,
1117) -> Result<Vec<u8>, String> {
1118    if pixels_per_scanline == 0
1119        || pixels_per_block == 0
1120        || pixels_per_block & 1 != 0
1121        || bits_per_pixel == 0
1122        || (bits_per_pixel > 32 && bits_per_pixel != 64)
1123    {
1124        return Err("invalid SZIP parameters".into());
1125    }
1126
1127    let flags = AEC_NOT_ENFORCE | convert_options(options_mask);
1128    let block_size = pixels_per_block;
1129    let rsi = pixels_per_scanline.div_ceil(pixels_per_block);
1130
1131    // libaec's SZ_BufftoBuffCompress treats 32- and 64-bit pixels by
1132    // byte-interleaving them into 8-bit samples. The RAW option mask does
1133    // NOT influence this decision (see sz_compat.c).
1134    let interleave = bits_per_pixel == 32 || bits_per_pixel == 64;
1135
1136    // The true input pixel width in bytes, BEFORE any interleaving. libhdf5's
1137    // H5Zszip.c only ever feeds SZ_BufftoBuffCompress buffers whose length is a
1138    // whole multiple of the on-disk type size, and libaec's interleave_buffer
1139    // computes `count = n / wordsize`, silently discarding the trailing
1140    // `n % wordsize` bytes. Validate against the real pixel width here -- the
1141    // post-interleave sample size is always 8 bits, so the old check against
1142    // `bits_to_bytes(bits_per_sample)` collapsed to `is_multiple_of(1)` and
1143    // never rejected anything on the 32/64-bit path. `bits_to_bytes` caps at 4,
1144    // so it cannot express the 8-byte width of a 64-bit pixel; compute the
1145    // ceiling-to-byte width directly.
1146    let input_pixel_size = bits_per_pixel.div_ceil(8) as usize;
1147    if !data.len().is_multiple_of(input_pixel_size) {
1148        return Err(format!(
1149            "input size {} is not a multiple of pixel size {} (bits_per_pixel={})",
1150            data.len(),
1151            input_pixel_size,
1152            bits_per_pixel
1153        ));
1154    }
1155
1156    let bits_per_sample;
1157    let input_buf: Vec<u8>;
1158
1159    if interleave {
1160        bits_per_sample = 8;
1161        input_buf = interleave_buffer(data, (bits_per_pixel / 8) as usize);
1162    } else {
1163        bits_per_sample = bits_per_pixel;
1164        input_buf = data.to_vec();
1165    }
1166
1167    let pixel_size = bits_to_bytes(bits_per_sample) as usize;
1168
1169    let line_size_bytes = pixels_per_scanline as usize * pixel_size;
1170    let padded_line_pixels = rsi * block_size;
1171    let padding_pixels = padded_line_pixels as usize - pixels_per_scanline as usize;
1172    let padding_size = padding_pixels * pixel_size;
1173
1174    // libaec's add_padding is always applied: besides filling the
1175    // scanline-vs-block gap (padding_size), it also pads a short final
1176    // scanline up to a full padded scanline. Skipping it when
1177    // padding_size == 0 would under-pad an input whose length is not a
1178    // whole multiple of the (padded) scanline size.
1179    let padded_input = add_padding(
1180        &input_buf,
1181        line_size_bytes,
1182        padding_size,
1183        pixel_size,
1184        flags & AEC_DATA_PREPROCESS != 0,
1185    );
1186
1187    let encoder = Encoder::new(bits_per_sample, block_size, rsi, flags)?;
1188    encoder.encode(&padded_input)
1189}
1190
1191/// Decompress SZIP (AEC) compressed data.
1192///
1193/// Parameters match the HDF5 SZIP filter interface:
1194/// - `data`: compressed bytes
1195/// - `output_size`: expected size of decompressed data in bytes
1196/// - `bits_per_pixel`: sample width (1-32, or 64 for double interleaving)
1197/// - `pixels_per_block`: block size (must be even)
1198/// - `pixels_per_scanline`: scanline width in pixels
1199/// - `options_mask`: SZIP option flags
1200pub fn decompress(
1201    data: &[u8],
1202    output_size: usize,
1203    bits_per_pixel: u32,
1204    pixels_per_block: u32,
1205    pixels_per_scanline: u32,
1206    options_mask: u32,
1207) -> Result<Vec<u8>, String> {
1208    if pixels_per_scanline == 0
1209        || pixels_per_block == 0
1210        || pixels_per_block & 1 != 0
1211        || bits_per_pixel == 0
1212        || (bits_per_pixel > 32 && bits_per_pixel != 64)
1213    {
1214        return Err("invalid SZIP parameters".into());
1215    }
1216
1217    let flags = convert_options(options_mask);
1218    let block_size = pixels_per_block;
1219    let rsi = pixels_per_scanline.div_ceil(pixels_per_block);
1220
1221    // Symmetric to `compress`: `output_size` is the requested uncompressed
1222    // length. On the 32/64-bit path it is fed through `deinterleave_buffer`,
1223    // whose `count = n / wordsize` would silently drop the trailing
1224    // `n % wordsize` bytes if `output_size` is not a whole pixel multiple.
1225    // Reject such requests up front instead of returning a short buffer.
1226    let output_pixel_size = bits_per_pixel.div_ceil(8) as usize;
1227    if !output_size.is_multiple_of(output_pixel_size) {
1228        return Err(format!(
1229            "output size {} is not a multiple of pixel size {} (bits_per_pixel={})",
1230            output_size, output_pixel_size, bits_per_pixel
1231        ));
1232    }
1233
1234    // libaec's SZ_BufftoBuffDecompress byte-deinterleaves 32- and 64-bit
1235    // pixels back from 8-bit samples. The RAW option mask does NOT influence
1236    // this decision (see sz_compat.c).
1237    let deinterleave = bits_per_pixel == 32 || bits_per_pixel == 64;
1238    let bits_per_sample = if deinterleave { 8 } else { bits_per_pixel };
1239    let pixel_size = bits_to_bytes(bits_per_sample) as usize;
1240
1241    let pad_scanline = !pixels_per_scanline.is_multiple_of(pixels_per_block);
1242    let _extra_buffer = pad_scanline || deinterleave;
1243
1244    let decode_output_size = if pad_scanline {
1245        let scanlines = (output_size / pixel_size).div_ceil(pixels_per_scanline as usize);
1246        rsi as usize * block_size as usize * pixel_size * scanlines
1247    } else {
1248        output_size
1249    };
1250
1251    let decoder = Decoder::new(bits_per_sample, block_size, rsi, flags)?;
1252    let output_samples = decode_output_size / pixel_size;
1253    let samples = decoder.decode(data, output_samples)?;
1254    let mut raw_bytes = decoder.write_samples(&samples, decode_output_size);
1255
1256    if pad_scanline {
1257        let line_size = pixels_per_scanline as usize * pixel_size;
1258        let padding_size =
1259            (rsi as usize * block_size as usize - pixels_per_scanline as usize) * pixel_size;
1260        remove_padding(&mut raw_bytes, line_size, padding_size);
1261    }
1262
1263    let result = if deinterleave {
1264        let len = std::cmp::min(raw_bytes.len(), output_size);
1265        deinterleave_buffer(&raw_bytes[..len], (bits_per_pixel / 8) as usize)
1266    } else {
1267        raw_bytes.truncate(output_size);
1268        raw_bytes
1269    };
1270
1271    Ok(result)
1272}
1273
1274// ===========================================================================
1275//  Tests
1276// ===========================================================================
1277#[cfg(test)]
1278mod tests {
1279    use super::*;
1280
1281    fn roundtrip(
1282        data: &[u8],
1283        bits_per_pixel: u32,
1284        pixels_per_block: u32,
1285        pixels_per_scanline: u32,
1286        options_mask: u32,
1287    ) {
1288        let compressed = compress(
1289            data,
1290            bits_per_pixel,
1291            pixels_per_block,
1292            pixels_per_scanline,
1293            options_mask,
1294        )
1295        .expect("compress failed");
1296        let decompressed = decompress(
1297            &compressed,
1298            data.len(),
1299            bits_per_pixel,
1300            pixels_per_block,
1301            pixels_per_scanline,
1302            options_mask,
1303        )
1304        .expect("decompress failed");
1305        assert_eq!(
1306            data,
1307            &decompressed[..],
1308            "roundtrip mismatch for bpp={bits_per_pixel}"
1309        );
1310    }
1311
1312    #[test]
1313    fn test_roundtrip_u8() {
1314        let data: Vec<u8> = (0..256u16).map(|i| (i & 0xFF) as u8).collect();
1315        // NN (preprocess) + MSB
1316        roundtrip(&data, 8, 16, 256, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1317    }
1318
1319    #[test]
1320    fn test_roundtrip_u8_no_preprocess() {
1321        let data: Vec<u8> = (0..128).collect();
1322        roundtrip(&data, 8, 16, 128, SZ_MSB_OPTION_MASK);
1323    }
1324
1325    #[test]
1326    fn test_roundtrip_u16() {
1327        let mut data = Vec::new();
1328        for i in 0..128u16 {
1329            data.push((i >> 8) as u8);
1330            data.push((i & 0xFF) as u8);
1331        }
1332        roundtrip(&data, 16, 16, 128, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1333    }
1334
1335    #[test]
1336    fn test_roundtrip_u16_lsb() {
1337        let mut data = Vec::new();
1338        for i in 0..128u16 {
1339            data.push((i & 0xFF) as u8);
1340            data.push((i >> 8) as u8);
1341        }
1342        roundtrip(&data, 16, 16, 128, SZ_NN_OPTION_MASK);
1343    }
1344
1345    #[test]
1346    fn test_roundtrip_u32_interleaved() {
1347        let values: Vec<u32> = (0..64).collect();
1348        let mut data = Vec::new();
1349        for &v in &values {
1350            data.extend_from_slice(&v.to_be_bytes());
1351        }
1352        roundtrip(&data, 32, 16, 64, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1353    }
1354
1355    #[test]
1356    fn test_roundtrip_f32() {
1357        let values: Vec<f32> = (0..64).map(|i| i as f32 * 1.5).collect();
1358        let mut data = Vec::new();
1359        for &v in &values {
1360            data.extend_from_slice(&v.to_be_bytes());
1361        }
1362        roundtrip(&data, 32, 16, 64, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1363    }
1364
1365    #[test]
1366    fn test_roundtrip_f64() {
1367        let values: Vec<f64> = (0..32).map(|i| i as f64 * 2.5).collect();
1368        let mut data = Vec::new();
1369        for &v in &values {
1370            data.extend_from_slice(&v.to_be_bytes());
1371        }
1372        roundtrip(&data, 64, 16, 32, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1373    }
1374
1375    #[test]
1376    fn test_roundtrip_zeros() {
1377        let data = vec![0u8; 256];
1378        roundtrip(&data, 8, 16, 256, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1379    }
1380
1381    #[test]
1382    fn test_roundtrip_constant() {
1383        let data = vec![42u8; 128];
1384        roundtrip(&data, 8, 16, 128, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1385    }
1386
1387    #[test]
1388    fn test_roundtrip_scanline_padding() {
1389        // pixels_per_scanline not a multiple of pixels_per_block
1390        // 100 pixels, block=16 => rsi=7, padded=112
1391        let data: Vec<u8> = (0..100).collect();
1392        roundtrip(&data, 8, 16, 100, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1393    }
1394
1395    #[test]
1396    fn test_roundtrip_small_block() {
1397        let data: Vec<u8> = (0..32).collect();
1398        roundtrip(&data, 8, 8, 32, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1399    }
1400
1401    #[test]
1402    fn test_roundtrip_u8_random_like() {
1403        // Data with varying patterns to exercise different code paths
1404        let data: Vec<u8> = (0..256).map(|i| ((i * 7 + 13) % 256) as u8).collect();
1405        roundtrip(&data, 8, 16, 256, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK);
1406    }
1407
1408    #[test]
1409    fn test_compress_rejects_misaligned_32bit_input() {
1410        // 32-bit pixels: input must be a multiple of 4 bytes. A 254-byte
1411        // buffer is 2 bytes short of 64 pixels; without the alignment check
1412        // interleave_buffer would silently drop the trailing 2 bytes.
1413        let data = vec![1u8; 254];
1414        let err = compress(&data, 32, 16, 64, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK)
1415            .expect_err("misaligned 32-bit input must be rejected");
1416        assert!(
1417            err.contains("not a multiple of pixel size"),
1418            "unexpected error message: {err}"
1419        );
1420    }
1421
1422    #[test]
1423    fn test_compress_rejects_misaligned_64bit_input() {
1424        // 64-bit pixels: input must be a multiple of 8 bytes. bits_to_bytes
1425        // caps at 4, so the true pixel width (8) must be derived directly.
1426        let data = vec![1u8; 36]; // 4 full 8-byte pixels + 4 stray bytes
1427        let err = compress(&data, 64, 16, 32, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK)
1428            .expect_err("misaligned 64-bit input must be rejected");
1429        assert!(
1430            err.contains("not a multiple of pixel size"),
1431            "unexpected error message: {err}"
1432        );
1433    }
1434
1435    #[test]
1436    fn test_compress_rejects_misaligned_16bit_input() {
1437        // 16-bit non-interleaved path: input must be a multiple of 2 bytes.
1438        let data = vec![1u8; 15];
1439        let err = compress(&data, 16, 16, 128, SZ_NN_OPTION_MASK)
1440            .expect_err("misaligned 16-bit input must be rejected");
1441        assert!(
1442            err.contains("not a multiple of pixel size"),
1443            "unexpected error message: {err}"
1444        );
1445    }
1446
1447    #[test]
1448    fn test_compress_accepts_aligned_8bit_odd_length() {
1449        // 8-bit pixels have width 1: any length is aligned and must compress.
1450        let data: Vec<u8> = (0..101u32).map(|i| i as u8).collect();
1451        compress(&data, 8, 16, 101, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK)
1452            .expect("8-bit input of any length must be accepted");
1453    }
1454
1455    #[test]
1456    fn test_decompress_rejects_misaligned_output_size() {
1457        // Symmetric guard: a 32-bit output_size not a multiple of 4 would be
1458        // truncated by deinterleave_buffer; it must be rejected instead.
1459        let err = decompress(
1460            &[0u8; 8],
1461            254,
1462            32,
1463            16,
1464            64,
1465            SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK,
1466        )
1467        .expect_err("misaligned 32-bit output size must be rejected");
1468        assert!(
1469            err.contains("not a multiple of pixel size"),
1470            "unexpected error message: {err}"
1471        );
1472    }
1473
1474    #[test]
1475    fn test_compress_reduces_size() {
1476        // Highly compressible data
1477        let data = vec![0u8; 1024];
1478        let compressed = compress(&data, 8, 16, 1024, SZ_NN_OPTION_MASK | SZ_MSB_OPTION_MASK)
1479            .expect("compress failed");
1480        assert!(
1481            compressed.len() < data.len(),
1482            "compression should reduce size for zeros"
1483        );
1484    }
1485}