Skip to main content

zrip_core/bitstream/
reader_reverse.rs

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