Skip to main content

zrip_core/bitstream/
reader_reverse.rs

1#![forbid(unsafe_code)]
2
3use crate::bitstream::primitives;
4use crate::error::DecompressError;
5
6/// Reverse bitstream reader using C zstd's bitsConsumed model.
7///
8/// Instead of tracking `bits_available` (decrement on consume, increment on refill),
9/// this tracks `bits_consumed` (increment on consume, reset on reload). Peek uses
10/// a double-shift: `(container << consumed) >> (64 - n)` = 2 ops vs the old model's
11/// 3 ops (shift + mask + subtract).
12pub struct ReverseBitReader<'a> {
13    pub data: &'a [u8],
14    pub container: u64,
15    pub bits_consumed: u32,
16    pub ptr: usize,
17    pub limit_ptr: usize,
18}
19
20impl<'a> ReverseBitReader<'a> {
21    pub fn new(data: &'a [u8]) -> Result<Self, DecompressError> {
22        if data.is_empty() {
23            return Err(DecompressError::InputExhausted);
24        }
25
26        let last_byte = *data.last().unwrap();
27        if last_byte == 0 {
28            return Err(DecompressError::CorruptSequences);
29        }
30
31        let initial_consumed = last_byte.leading_zeros() + 1;
32
33        let ptr = if data.len() >= 8 { data.len() - 8 } else { 0 };
34
35        let container = if data.len() >= 8 {
36            primitives::read_u64_le_unaligned(data, ptr)
37        } else {
38            let mut val = 0u64;
39            for (i, &b) in data.iter().enumerate() {
40                val |= (b as u64) << (i * 8);
41            }
42            val
43        };
44
45        let bits_consumed = if data.len() >= 8 {
46            initial_consumed
47        } else {
48            64 - (data.len() as u32) * 8 + initial_consumed
49        };
50
51        let limit_ptr = if data.len() >= 8 { 8 } else { 0 };
52
53        Ok(Self {
54            data,
55            container,
56            bits_consumed,
57            ptr,
58            limit_ptr,
59        })
60    }
61
62    #[inline(always)]
63    pub fn refill(&mut self) {
64        if self.bits_consumed <= 7 || self.ptr == 0 {
65            return;
66        }
67        let byte_shift = (self.bits_consumed >> 3) as usize;
68        let actual_shift = byte_shift.min(self.ptr);
69        self.ptr -= actual_shift;
70        self.bits_consumed -= (actual_shift as u32) * 8;
71        if self.ptr + 8 <= self.data.len() {
72            self.container = primitives::read_u64_le_unaligned(self.data, self.ptr);
73        } else {
74            let mut val = 0u64;
75            let avail = self.data.len() - self.ptr;
76            for i in 0..avail {
77                val |= (primitives::get_byte_unchecked(self.data, self.ptr + i) as u64) << (i * 8);
78            }
79            self.container = val;
80        }
81    }
82
83    #[inline]
84    pub fn read_bits(&mut self, n: u8) -> Result<u32, DecompressError> {
85        debug_assert!(n <= 32);
86        if n == 0 {
87            return Ok(0);
88        }
89        self.refill();
90        let avail = 64u32.saturating_sub(self.bits_consumed);
91        if (n as u32) > avail {
92            return Err(DecompressError::InputExhausted);
93        }
94        let result = ((self.container << self.bits_consumed) >> (64 - n as u32)) as u32;
95        self.bits_consumed += n as u32;
96        Ok(result)
97    }
98
99    #[inline]
100    pub fn read_bits_unchecked(&mut self, n: u8) -> u32 {
101        debug_assert!(n <= 32);
102        if n == 0 {
103            return 0;
104        }
105        self.refill();
106        debug_assert!((n as u32) <= 64u32.saturating_sub(self.bits_consumed));
107        let result = ((self.container << self.bits_consumed) >> (64 - n as u32)) as u32;
108        self.bits_consumed += n as u32;
109        result
110    }
111
112    #[inline(always)]
113    pub fn consume_bits(&mut self, n: u8) {
114        debug_assert!((n as u32) <= 64u32.saturating_sub(self.bits_consumed));
115        self.bits_consumed += n as u32;
116        self.refill();
117    }
118
119    #[inline(always)]
120    pub fn read_bits_fast(&mut self, n: u8) -> u32 {
121        debug_assert!((n as u32) <= 64u32.saturating_sub(self.bits_consumed));
122        if n == 0 {
123            return 0;
124        }
125        let result = ((self.container << self.bits_consumed) >> (64 - n as u32)) as u32;
126        self.bits_consumed += n as u32;
127        result
128    }
129
130    #[inline(always)]
131    pub fn read_bits_branchless(&mut self, n: u8) -> u32 {
132        debug_assert!(n <= 32);
133        let result = ((self.container << (self.bits_consumed & 63)) >> 1 >> (63 - n as u32)) as u32;
134        self.bits_consumed += n as u32;
135        result
136    }
137
138    #[inline(always)]
139    pub fn refill_fast(&mut self) {
140        let byte_shift = (self.bits_consumed >> 3) as usize;
141        if byte_shift > self.ptr || self.ptr - byte_shift + 8 > self.data.len() {
142            return;
143        }
144        self.ptr -= byte_shift;
145        self.bits_consumed -= (byte_shift as u32) * 8;
146        self.container = primitives::read_u64_le_unaligned(self.data, self.ptr);
147    }
148
149    #[inline(always)]
150    pub fn refill_fast_unchecked(&mut self) {
151        let byte_shift = (self.bits_consumed >> 3) as usize;
152        debug_assert!(byte_shift <= self.ptr);
153        debug_assert!(self.ptr - byte_shift + 8 <= self.data.len());
154        self.ptr -= byte_shift;
155        self.bits_consumed -= (byte_shift as u32) * 8;
156        self.container = primitives::read_u64_le_unaligned(self.data, self.ptr);
157    }
158
159    #[inline]
160    pub fn peek_bits(&self, n: u8) -> u32 {
161        debug_assert!(n <= 32);
162        debug_assert!((n as u32) <= 64u32.saturating_sub(self.bits_consumed));
163        if n == 0 {
164            return 0;
165        }
166        ((self.container << self.bits_consumed) >> (64 - n as u32)) as u32
167    }
168
169    #[inline]
170    pub fn bits_remaining(&self) -> usize {
171        64usize.saturating_sub(self.bits_consumed as usize) + self.ptr * 8
172    }
173
174    #[inline]
175    pub fn is_empty(&self) -> bool {
176        self.bits_consumed >= 64 && self.ptr == 0
177    }
178}
179
180#[cfg(test)]
181mod tests {
182    use super::*;
183
184    #[test]
185    fn empty_input() {
186        assert!(ReverseBitReader::new(&[]).is_err());
187    }
188
189    #[test]
190    fn zero_last_byte() {
191        assert!(ReverseBitReader::new(&[0x00]).is_err());
192    }
193
194    #[test]
195    fn sentinel_only_no_data() {
196        let data = [0b0000_0001];
197        let r = ReverseBitReader::new(&data).unwrap();
198        assert_eq!(r.bits_remaining(), 0);
199    }
200
201    #[test]
202    fn roundtrip_with_forward_writer() {
203        use crate::bitstream::writer::BitWriter;
204
205        let mut w = BitWriter::new();
206        w.write_bits(0b101, 3);
207        w.write_bits(0b1100_1010, 8);
208        w.write_bits(0b1, 1);
209        w.close_reverse_stream();
210        let bytes = w.into_bytes();
211
212        let mut r = ReverseBitReader::new(&bytes).unwrap();
213        assert_eq!(r.read_bits(1).unwrap(), 0b1);
214        assert_eq!(r.read_bits(8).unwrap(), 0b1100_1010);
215        assert_eq!(r.read_bits(3).unwrap(), 0b101);
216        assert_eq!(r.bits_remaining(), 0);
217    }
218
219    #[test]
220    fn single_byte_with_data() {
221        let data = [0b0000_1101];
222        let mut r = ReverseBitReader::new(&data).unwrap();
223        assert_eq!(r.read_bits(3).unwrap(), 0b101);
224        assert_eq!(r.bits_remaining(), 0);
225    }
226
227    #[test]
228    fn multi_byte_stream() {
229        use crate::bitstream::writer::BitWriter;
230
231        let mut w = BitWriter::new();
232        w.write_bits(0xFF, 8);
233        w.write_bits(0x3, 2);
234        w.close_reverse_stream();
235        let bytes = w.into_bytes();
236
237        let mut r = ReverseBitReader::new(&bytes).unwrap();
238        assert_eq!(r.read_bits(2).unwrap(), 0x3);
239        assert_eq!(r.read_bits(8).unwrap(), 0xFF);
240    }
241}
242
243#[cfg(all(test, miri, not(feature = "paranoid")))]
244mod ub_tests {
245    use super::*;
246
247    #[test]
248    fn public_refill_fast_underflows_on_short_stream() {
249        // Issue: refill_fast is a safe public method, but its requirements
250        // (enough consumed bits and at least eight readable bytes after the new
251        // pointer) are enforced only with debug_asserts. On this one-byte stream,
252        // byte_shift is 8 and ptr is 0, so release builds wrap the subtraction
253        // and miri reports the resulting out-of-bounds read_u64_le_unaligned.
254        let data = [0b0000_0001];
255        let mut reader = ReverseBitReader::new(&data).unwrap();
256        reader.refill_fast();
257    }
258}