Skip to main content

otf_pixels_codec_avif/av1/
bits.rs

1//! The AV1 bitstream reader.
2//!
3//! AV1 is read most-significant-bit first (AV1 spec §4.10.2). This is a
4//! crate-local reader by the same convention the rest of the workspace follows:
5//! JPEG and inflate each carry their own, because a bit order and a set of
6//! variable-length primitives are format decisions, not shared infrastructure.
7//!
8//! Every primitive here is fallible. The spec is written as if the stream never
9//! runs out — a conformant one does not — but this decoder reads
10//! attacker-controlled bytes under `unsafe_code = "forbid"` and a ban on
11//! `panic!`, so reading past the end is a returned error, never a trap.
12
13use otf_pixels_core::{PixelsError, Result};
14
15/// A most-significant-bit-first reader over an AV1 byte slice.
16///
17/// Position is tracked in bits so `byte_alignment` and the byte-granular
18/// primitives (`leb128`, `le`) can assert and act on alignment.
19pub struct BitReader<'a> {
20    data: &'a [u8],
21    /// The next bit to read, counted from the front of `data`. Bit 0 is the
22    /// most significant bit of byte 0.
23    pos: usize,
24}
25
26impl<'a> BitReader<'a> {
27    /// Wrap a byte slice. Reading starts at its first bit.
28    #[must_use]
29    pub fn new(data: &'a [u8]) -> Self {
30        Self { data, pos: 0 }
31    }
32
33    /// The current position, in bits from the start.
34    #[must_use]
35    pub fn bit_position(&self) -> usize {
36        self.pos
37    }
38
39    /// The total length of the underlying slice, in bits.
40    #[must_use]
41    pub fn bit_len(&self) -> usize {
42        self.data.len().saturating_mul(8)
43    }
44
45    /// Bits remaining before the end of the slice.
46    #[must_use]
47    pub fn bits_left(&self) -> usize {
48        self.bit_len().saturating_sub(self.pos)
49    }
50
51    /// Whether the reader sits exactly on a byte boundary.
52    #[must_use]
53    pub fn is_byte_aligned(&self) -> bool {
54        self.pos % 8 == 0
55    }
56
57    /// The current position in whole bytes — only meaningful when byte-aligned.
58    #[must_use]
59    pub fn byte_position(&self) -> usize {
60        self.pos / 8
61    }
62
63    /// Read a single bit.
64    fn read_bit(&mut self) -> Result<u32> {
65        let byte = self.pos / 8;
66        let Some(&value) = self.data.get(byte) else {
67            return Err(PixelsError::malformed(
68                "avif",
69                "the AV1 bitstream ended in the middle of a value",
70            ));
71        };
72        // Bit 0 of a byte is its most significant bit (spec §4.10.2).
73        let shift = 7 - (self.pos % 8);
74        self.pos += 1;
75        Ok(u32::from((value >> shift) & 1))
76    }
77
78    /// `f(n)` — read `n` bits as an unsigned integer, MSB first (§4.10.2).
79    ///
80    /// `n` is at most 32; the AV1 syntax never reads a wider `f(n)` in one call.
81    pub fn f(&mut self, n: u32) -> Result<u32> {
82        if n == 0 {
83            return Ok(0);
84        }
85        if n > 32 {
86            return Err(PixelsError::malformed(
87                "avif",
88                "an AV1 fixed-width read wider than 32 bits is a decoder bug",
89            ));
90        }
91        let mut value: u32 = 0;
92        for _ in 0..n {
93            // Shifting a u32 left by up to 31 and or-ing one bit never
94            // overflows; the width guard above keeps the loop within 32 steps.
95            value = (value << 1) | self.read_bit()?;
96        }
97        Ok(value)
98    }
99
100    /// `f(n)` for values that may need the full 64 bits (`le` uses it).
101    fn f64(&mut self, n: u32) -> Result<u64> {
102        if n > 64 {
103            return Err(PixelsError::malformed(
104                "avif",
105                "an AV1 fixed-width read wider than 64 bits is a decoder bug",
106            ));
107        }
108        let mut value: u64 = 0;
109        for _ in 0..n {
110            value = (value << 1) | u64::from(self.read_bit()?);
111        }
112        Ok(value)
113    }
114
115    /// Read a boolean flag — `f(1)` reported as `bool`.
116    pub fn flag(&mut self) -> Result<bool> {
117        Ok(self.f(1)? != 0)
118    }
119
120    /// `uvlc()` — unsigned variable-length code (§4.10.3).
121    ///
122    /// A run of zero bits terminated by a one, then that many value bits. A run
123    /// of 32 or more leading zeros is the spec's saturation case and yields
124    /// `u32::MAX`.
125    pub fn uvlc(&mut self) -> Result<u32> {
126        let mut leading_zeros: u32 = 0;
127        loop {
128            if self.flag()? {
129                break;
130            }
131            leading_zeros += 1;
132            if leading_zeros >= 32 {
133                return Ok(u32::MAX);
134            }
135        }
136        let value = self.f(leading_zeros)?;
137        // value + 2^leading_zeros - 1, computed without overflow: leading_zeros
138        // is < 32 here, and value < 2^leading_zeros, so the sum fits in u32.
139        Ok(value + ((1_u32 << leading_zeros) - 1))
140    }
141
142    /// `le(n)` — an `n`-byte little-endian unsigned integer (§4.10.4).
143    ///
144    /// Must be byte-aligned, which the syntax guarantees at every call site.
145    pub fn le(&mut self, n: u32) -> Result<u64> {
146        if !self.is_byte_aligned() {
147            return Err(PixelsError::malformed(
148                "avif",
149                "an AV1 le() read was not byte-aligned",
150            ));
151        }
152        let mut value: u64 = 0;
153        for i in 0..n {
154            let byte = self.f64(8)?;
155            value |= byte << (i * 8);
156        }
157        Ok(value)
158    }
159
160    /// `leb128()` — a little-endian base-128 unsigned integer (§4.10.5).
161    ///
162    /// At most eight bytes; a ninth continuation bit is malformed. Returns the
163    /// value and the number of bytes consumed so callers can bound a payload.
164    pub fn leb128(&mut self) -> Result<u64> {
165        if !self.is_byte_aligned() {
166            return Err(PixelsError::malformed(
167                "avif",
168                "an AV1 leb128() read was not byte-aligned",
169            ));
170        }
171        let mut value: u64 = 0;
172        for i in 0..8 {
173            let byte = self.f(8)?;
174            // Seven payload bits per byte, low group first.
175            value |= u64::from(byte & 0x7f) << (i * 7);
176            if byte & 0x80 == 0 {
177                return Ok(value);
178            }
179        }
180        Err(PixelsError::malformed(
181            "avif",
182            "an AV1 leb128 value ran past its eight-byte maximum",
183        ))
184    }
185
186    /// `su(n)` — a signed integer in `n+1` bits, sign last (§4.10.6).
187    pub fn su(&mut self, n: u32) -> Result<i32> {
188        let value = self.f(n + 1)? as i32;
189        let sign_mask = 1_i32 << n;
190        if value & sign_mask != 0 {
191            Ok(value - 2 * sign_mask)
192        } else {
193            Ok(value)
194        }
195    }
196
197    /// `ns(n)` — a non-symmetric unsigned integer over `[0, n)` (§4.10.7).
198    ///
199    /// Uses one fewer bit for the smaller half of the range, so it is not a
200    /// plain `f`. `n == 0` reads nothing and yields 0.
201    pub fn ns(&mut self, n: u32) -> Result<u32> {
202        if n <= 1 {
203            return Ok(0);
204        }
205        let w = floor_log2(n) + 1;
206        let m = (1_u32 << w) - n;
207        let v = self.f(w - 1)?;
208        if v < m {
209            return Ok(v);
210        }
211        let extra_bit = self.f(1)?;
212        Ok((v << 1) - m + extra_bit)
213    }
214
215    /// Advance to the next byte boundary, requiring the skipped bits be zero
216    /// (`byte_alignment()`, §5.3.5). AV1 mandates the padding be zero.
217    pub fn byte_alignment(&mut self) -> Result<()> {
218        while !self.is_byte_aligned() {
219            if self.f(1)? != 0 {
220                return Err(PixelsError::malformed(
221                    "avif",
222                    "an AV1 byte-alignment pad bit was not zero",
223                ));
224            }
225        }
226        Ok(())
227    }
228
229    /// Skip `n` bits without interpreting them.
230    pub fn skip_bits(&mut self, n: usize) -> Result<()> {
231        let end = self.pos.checked_add(n).filter(|&e| e <= self.bit_len());
232        let Some(end) = end else {
233            return Err(PixelsError::malformed(
234                "avif",
235                "an AV1 skip ran past the end of the bitstream",
236            ));
237        };
238        self.pos = end;
239        Ok(())
240    }
241}
242
243/// `FloorLog2(x)` (§4.7): the index of the most significant set bit. `x` must
244/// be non-zero, which every AV1 call site guarantees.
245#[must_use]
246pub fn floor_log2(x: u32) -> u32 {
247    // 31 - leading_zeros is the MSB index; for x >= 1 it is well-defined.
248    31 - x.leading_zeros()
249}
250
251#[cfg(test)]
252#[allow(
253    clippy::unwrap_used,
254    clippy::indexing_slicing,
255    clippy::panic,
256    clippy::unusual_byte_groupings,
257    reason = "tests operate on known-good values and assert shapes directly"
258)]
259mod tests {
260    use super::*;
261    use otf_pixels_core::ErrorCode;
262
263    #[test]
264    fn f_reads_most_significant_bit_first() {
265        // 0b1011_0010, 0b0100_0000
266        let data = [0xB2, 0x40];
267        let mut r = BitReader::new(&data);
268        assert_eq!(r.f(3).unwrap(), 0b101);
269        assert_eq!(r.f(5).unwrap(), 0b10010);
270        assert_eq!(r.f(2).unwrap(), 0b01);
271        assert_eq!(r.bit_position(), 10);
272    }
273
274    #[test]
275    fn f_of_zero_reads_nothing() {
276        let data = [0xFF];
277        let mut r = BitReader::new(&data);
278        assert_eq!(r.f(0).unwrap(), 0);
279        assert_eq!(r.bit_position(), 0);
280    }
281
282    #[test]
283    fn reading_past_the_end_is_an_error_not_a_panic() {
284        let data = [0xFF];
285        let mut r = BitReader::new(&data);
286        assert_eq!(r.f(8).unwrap(), 0xFF);
287        let err = r.f(1).unwrap_err();
288        assert_eq!(err.code(), ErrorCode::Malformed);
289    }
290
291    #[test]
292    fn uvlc_decodes_the_exponential_golomb_shape() {
293        // 1                -> 0
294        // 010              -> 1
295        // 011              -> 2
296        // 00100            -> 3
297        // Pack: 1 010 011 00100 = 1010_0110_0100_...
298        let data = [0b1010_0110, 0b0100_0000];
299        let mut r = BitReader::new(&data);
300        assert_eq!(r.uvlc().unwrap(), 0);
301        assert_eq!(r.uvlc().unwrap(), 1);
302        assert_eq!(r.uvlc().unwrap(), 2);
303        assert_eq!(r.uvlc().unwrap(), 3);
304    }
305
306    #[test]
307    fn uvlc_saturates_at_thirty_two_leading_zeros() {
308        // 32 zero bits with no terminating one: four zero bytes, then more.
309        let data = [0x00, 0x00, 0x00, 0x00, 0x80];
310        let mut r = BitReader::new(&data);
311        assert_eq!(r.uvlc().unwrap(), u32::MAX);
312    }
313
314    #[test]
315    fn leb128_reads_little_endian_base_128() {
316        // 0xE5 0x8E 0x26 -> 624485, the canonical LEB128 example.
317        let data = [0xE5, 0x8E, 0x26];
318        let mut r = BitReader::new(&data);
319        assert_eq!(r.leb128().unwrap(), 624_485);
320        assert_eq!(r.byte_position(), 3);
321    }
322
323    #[test]
324    fn leb128_rejects_a_ninth_continuation_byte() {
325        let data = [0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80];
326        let mut r = BitReader::new(&data);
327        let err = r.leb128().unwrap_err();
328        assert_eq!(err.code(), ErrorCode::Malformed);
329    }
330
331    #[test]
332    fn le_reads_little_endian_bytes() {
333        let data = [0x34, 0x12];
334        let mut r = BitReader::new(&data);
335        assert_eq!(r.le(2).unwrap(), 0x1234);
336    }
337
338    #[test]
339    fn su_recovers_negative_values() {
340        // su(3) reads 4 bits. 0b1111 -> -1; 0b0111 -> 7; 0b1000 -> -8.
341        let data = [0b1111_0111, 0b1000_0000];
342        let mut r = BitReader::new(&data);
343        assert_eq!(r.su(3).unwrap(), -1);
344        assert_eq!(r.su(3).unwrap(), 7);
345        assert_eq!(r.su(3).unwrap(), -8);
346    }
347
348    #[test]
349    fn ns_uses_one_fewer_bit_for_the_low_half() {
350        // n = 3: w = 2, m = 1. v = f(1); v<1 -> value v; else read one more.
351        // Bits 0 -> 0. Bits 10 -> (1<<1)-1+0 = 1. Bits 11 -> (1<<1)-1+1 = 2.
352        let data = [0b0_10_11_000];
353        let mut r = BitReader::new(&data);
354        assert_eq!(r.ns(3).unwrap(), 0);
355        assert_eq!(r.ns(3).unwrap(), 1);
356        assert_eq!(r.ns(3).unwrap(), 2);
357    }
358
359    #[test]
360    fn ns_of_a_power_of_two_is_plain_fixed_width() {
361        // n = 4: w = 2, m = 0, so every value is read in 2 bits.
362        let data = [0b00_01_10_11];
363        let mut r = BitReader::new(&data);
364        assert_eq!(r.ns(4).unwrap(), 0);
365        assert_eq!(r.ns(4).unwrap(), 1);
366        assert_eq!(r.ns(4).unwrap(), 2);
367        assert_eq!(r.ns(4).unwrap(), 3);
368    }
369
370    #[test]
371    fn byte_alignment_requires_zero_padding() {
372        let mut r = BitReader::new(&[0b101_00000]);
373        assert_eq!(r.f(3).unwrap(), 0b101);
374        r.byte_alignment().unwrap();
375        assert!(r.is_byte_aligned());
376        assert_eq!(r.byte_position(), 1);
377
378        let mut bad = BitReader::new(&[0b101_00001]);
379        assert_eq!(bad.f(3).unwrap(), 0b101);
380        assert_eq!(
381            bad.byte_alignment().unwrap_err().code(),
382            ErrorCode::Malformed
383        );
384    }
385
386    #[test]
387    fn floor_log2_is_the_top_set_bit() {
388        assert_eq!(floor_log2(1), 0);
389        assert_eq!(floor_log2(2), 1);
390        assert_eq!(floor_log2(3), 1);
391        assert_eq!(floor_log2(255), 7);
392        assert_eq!(floor_log2(256), 8);
393    }
394}