Skip to main content

otf_pixels_codec_webp/
vp8l.rs

1//! VP8L, the WebP lossless bitstream (RFC 9649 §3).
2//!
3//! An image stream is an optional chain of reversible transforms (predictor,
4//! colour, subtract-green, colour indexing) over a pixel array coded with
5//! canonical prefix codes, LZ77 back references into the pixels already
6//! decoded, and a small hash cache of recent colours. The transforms' own data
7//! are smaller images coded the same way, so decoding recurses.
8//!
9//! Pixels are ARGB packed in a `u32` as the specification describes them:
10//! alpha in the top byte, then red, green, blue.
11
12#![allow(
13    clippy::indexing_slicing,
14    reason = "every index is bounded by construction: transform sub-images are \
15              sized by the same div_round_up their lookups use, the colour table is \
16              padded to 256, distance codes are range-checked before the map, back \
17              references are checked against the pixels decoded so far, and code \
18              lengths are below 16 by the token alphabet. The truncation and \
19              corruption tests hold the decoder to that"
20)]
21
22use otf_pixels_core::{PixelsError, Result};
23
24/// The most code-length bits any prefix code may use.
25const MAX_CODE_LENGTH: usize = 15;
26/// Bits resolved by one lookup in a prefix code's fast table.
27const FAST_BITS: u32 = 8;
28/// The literal-length code lengths' own code order (§3.7.2.1.2).
29const CODE_LENGTH_ORDER: [usize; 19] = [
30    17, 18, 0, 1, 2, 3, 4, 5, 16, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
31];
32/// Back-reference length prefix codes (§3.6.2.2).
33const LENGTH_CODES: usize = 24;
34/// Distance prefix codes.
35const DISTANCE_CODES: usize = 40;
36/// The 120 short distance codes' `(dx, dy)` neighbour offsets (§3.6.2.2.1).
37pub(crate) const DISTANCE_MAP: [(i8, i8); 120] = [
38    (0, 1),
39    (1, 0),
40    (1, 1),
41    (-1, 1),
42    (0, 2),
43    (2, 0),
44    (1, 2),
45    (-1, 2),
46    (2, 1),
47    (-2, 1),
48    (2, 2),
49    (-2, 2),
50    (0, 3),
51    (3, 0),
52    (1, 3),
53    (-1, 3),
54    (3, 1),
55    (-3, 1),
56    (2, 3),
57    (-2, 3),
58    (3, 2),
59    (-3, 2),
60    (0, 4),
61    (4, 0),
62    (1, 4),
63    (-1, 4),
64    (4, 1),
65    (-4, 1),
66    (3, 3),
67    (-3, 3),
68    (2, 4),
69    (-2, 4),
70    (4, 2),
71    (-4, 2),
72    (0, 5),
73    (3, 4),
74    (-3, 4),
75    (4, 3),
76    (-4, 3),
77    (5, 0),
78    (1, 5),
79    (-1, 5),
80    (5, 1),
81    (-5, 1),
82    (2, 5),
83    (-2, 5),
84    (5, 2),
85    (-5, 2),
86    (4, 4),
87    (-4, 4),
88    (3, 5),
89    (-3, 5),
90    (5, 3),
91    (-5, 3),
92    (0, 6),
93    (6, 0),
94    (1, 6),
95    (-1, 6),
96    (6, 1),
97    (-6, 1),
98    (2, 6),
99    (-2, 6),
100    (6, 2),
101    (-6, 2),
102    (4, 5),
103    (-4, 5),
104    (5, 4),
105    (-5, 4),
106    (3, 6),
107    (-3, 6),
108    (6, 3),
109    (-6, 3),
110    (0, 7),
111    (7, 0),
112    (1, 7),
113    (-1, 7),
114    (5, 5),
115    (-5, 5),
116    (7, 1),
117    (-7, 1),
118    (4, 6),
119    (-4, 6),
120    (6, 4),
121    (-6, 4),
122    (2, 7),
123    (-2, 7),
124    (7, 2),
125    (-7, 2),
126    (3, 7),
127    (-3, 7),
128    (7, 3),
129    (-7, 3),
130    (5, 6),
131    (-5, 6),
132    (6, 5),
133    (-6, 5),
134    (8, 0),
135    (4, 7),
136    (-4, 7),
137    (7, 4),
138    (-7, 4),
139    (8, 1),
140    (8, 2),
141    (6, 6),
142    (-6, 6),
143    (8, 3),
144    (5, 7),
145    (-5, 7),
146    (7, 5),
147    (-7, 5),
148    (8, 4),
149    (6, 7),
150    (-6, 7),
151    (7, 6),
152    (-7, 6),
153    (8, 5),
154    (7, 7),
155    (-7, 7),
156    (8, 6),
157    (8, 7),
158];
159
160fn malformed(detail: impl Into<String>) -> PixelsError {
161    PixelsError::malformed("webp", detail.into())
162}
163
164/// Reads bits least-significant first, as VP8L packs them.
165pub struct BitReader<'a> {
166    data: &'a [u8],
167    /// Bit position from the start of `data`.
168    position: usize,
169}
170
171impl<'a> BitReader<'a> {
172    /// Read from the start of `data`.
173    #[must_use]
174    pub const fn new(data: &'a [u8]) -> Self {
175        Self { data, position: 0 }
176    }
177
178    /// The next `n` (at most 32) bits without consuming them; bits past the
179    /// end read as zero, which [`BitReader::consume`] then refuses.
180    fn peek(&self, n: u32) -> u32 {
181        let byte = self.position >> 3;
182        let mut word = 0_u64;
183        for (i, &b) in self.data.iter().skip(byte).take(5).enumerate() {
184            word |= u64::from(b) << (8 * i);
185        }
186        let shifted = word >> (self.position & 7);
187        (shifted & ((1_u64 << n) - 1)) as u32
188    }
189
190    fn consume(&mut self, n: u32) -> Result<()> {
191        self.position += n as usize;
192        if self.position > self.data.len() * 8 {
193            return Err(malformed("the lossless bitstream ends early"));
194        }
195        Ok(())
196    }
197
198    /// Read `n` bits, at most 32.
199    ///
200    /// # Errors
201    ///
202    /// Returns [`PixelsError::Malformed`] past the end of the data.
203    pub fn read(&mut self, n: u32) -> Result<u32> {
204        let value = self.peek(n);
205        self.consume(n)?;
206        Ok(value)
207    }
208
209    fn flag(&mut self) -> Result<bool> {
210        Ok(self.read(1)? == 1)
211    }
212}
213
214/// One canonical prefix code.
215enum PrefixCode {
216    /// A single used symbol, which costs no bits (§3.7.2.1).
217    Single(u16),
218    /// Two or more symbols.
219    Tree {
220        /// `(symbol << 4) | length` for every `FAST_BITS`-bit window whose
221        /// code is that short, 0 where a longer code continues.
222        fast: Box<[u16; 1 << FAST_BITS]>,
223        /// Codes of each length, for the canonical walk past the fast table.
224        counts: [u16; MAX_CODE_LENGTH + 1],
225        /// Symbols in canonical order.
226        symbols: Vec<u16>,
227    },
228}
229
230impl PrefixCode {
231    /// Build a code from per-symbol lengths, which must describe a complete
232    /// tree unless exactly one symbol is used.
233    fn new(lengths: &[u8]) -> Result<Self> {
234        let mut counts = [0_u16; MAX_CODE_LENGTH + 1];
235        for &length in lengths {
236            counts[usize::from(length)] += 1;
237        }
238        counts[0] = 0;
239        let used: usize = counts.iter().map(|&c| usize::from(c)).sum();
240        if used == 0 {
241            return Err(malformed("a prefix code has no symbols"));
242        }
243        if used == 1 {
244            let symbol = lengths.iter().position(|&l| l != 0).unwrap_or(0);
245            return Ok(Self::Single(symbol as u16));
246        }
247        // Kraft: the lengths must fill the code space exactly.
248        let mut space = 1_i64 << MAX_CODE_LENGTH;
249        for (length, &count) in counts.iter().enumerate().skip(1) {
250            space -= i64::from(count) << (MAX_CODE_LENGTH - length);
251        }
252        if space != 0 {
253            return Err(malformed(
254                "a prefix code's lengths do not form a complete tree",
255            ));
256        }
257
258        let mut offsets = [0_usize; MAX_CODE_LENGTH + 2];
259        for length in 1..=MAX_CODE_LENGTH {
260            offsets[length + 1] = offsets[length] + usize::from(counts[length]);
261        }
262        let mut symbols = vec![0_u16; used];
263        let mut next = offsets;
264        for (symbol, &length) in lengths.iter().enumerate() {
265            if length != 0 {
266                let slot = &mut next[usize::from(length)];
267                symbols[*slot] = symbol as u16;
268                *slot += 1;
269            }
270        }
271
272        // Fast table: walk the canonical codes in order, placing each short
273        // one at every window whose low bits are its code bit-reversed.
274        let mut fast = Box::new([0_u16; 1 << FAST_BITS]);
275        let mut code = 0_u32;
276        let mut index = 0;
277        for length in 1..=MAX_CODE_LENGTH as u32 {
278            for _ in 0..counts[length as usize] {
279                if length <= FAST_BITS {
280                    let reversed = code.reverse_bits() >> (32 - length);
281                    let entry = (symbols[index] << 4) | length as u16;
282                    let mut slot = reversed as usize;
283                    while slot < 1 << FAST_BITS {
284                        fast[slot] = entry;
285                        slot += 1 << length;
286                    }
287                }
288                code += 1;
289                index += 1;
290            }
291            code <<= 1;
292        }
293        Ok(Self::Tree {
294            fast,
295            counts,
296            symbols,
297        })
298    }
299
300    fn read(&self, bits: &mut BitReader<'_>) -> Result<u16> {
301        match self {
302            Self::Single(symbol) => Ok(*symbol),
303            Self::Tree {
304                fast,
305                counts,
306                symbols,
307            } => {
308                let entry = fast[bits.peek(FAST_BITS) as usize];
309                if entry != 0 {
310                    bits.consume(u32::from(entry & 15))?;
311                    return Ok(entry >> 4);
312                }
313                // The canonical walk (as in zlib's `puff`), one bit at a time.
314                let (mut code, mut first, mut index) = (0_i32, 0_i32, 0_i32);
315                for &count in counts.iter().skip(1) {
316                    code |= bits.read(1)? as i32;
317                    let count = i32::from(count);
318                    if code - count < first {
319                        return symbols
320                            .get((index + code - first) as usize)
321                            .copied()
322                            .ok_or_else(|| malformed("a prefix code walked off its symbols"));
323                    }
324                    index += count;
325                    first = (first + count) << 1;
326                    code <<= 1;
327                }
328                Err(malformed("a prefix code was read past its longest length"))
329            }
330        }
331    }
332}
333
334/// Read one prefix code's lengths over an `alphabet` and build it
335/// (§3.7.2.1). `build` false parses and discards, for groups nothing uses.
336fn read_prefix_code(
337    bits: &mut BitReader<'_>,
338    alphabet: usize,
339    build: bool,
340) -> Result<Option<PrefixCode>> {
341    let mut lengths = vec![0_u8; alphabet];
342    if bits.flag()? {
343        // Simple code: one or two symbols of length 1.
344        let two = bits.flag()?;
345        let first_bits = if bits.flag()? { 8 } else { 1 };
346        let symbol = bits.read(first_bits)? as usize;
347        *lengths
348            .get_mut(symbol)
349            .ok_or_else(|| malformed("a simple code's symbol is out of range"))? = 1;
350        if two {
351            let symbol = bits.read(8)? as usize;
352            *lengths
353                .get_mut(symbol)
354                .ok_or_else(|| malformed("a simple code's symbol is out of range"))? = 1;
355        }
356    } else {
357        let mut length_lengths = [0_u8; 19];
358        let count = 4 + bits.read(4)? as usize;
359        for &position in CODE_LENGTH_ORDER.iter().take(count) {
360            length_lengths[position] = bits.read(3)? as u8;
361        }
362        let length_code = PrefixCode::new(&length_lengths)?;
363        let mut max_tokens = if bits.flag()? {
364            let width = 2 + 2 * bits.read(3)?;
365            let max = 2 + bits.read(width)? as usize;
366            if max > alphabet {
367                return Err(malformed(
368                    "a prefix code declares more symbols than its alphabet",
369                ));
370            }
371            max
372        } else {
373            alphabet
374        };
375        let mut symbol = 0;
376        let mut previous = 8_u8;
377        while symbol < alphabet {
378            if max_tokens == 0 {
379                break;
380            }
381            max_tokens -= 1;
382            let token = length_code.read(bits)?;
383            if token < 16 {
384                lengths[symbol] = token as u8;
385                symbol += 1;
386                if token != 0 {
387                    previous = token as u8;
388                }
389                continue;
390            }
391            let (repeat, value) = match token {
392                16 => (3 + bits.read(2)? as usize, previous),
393                17 => (3 + bits.read(3)? as usize, 0),
394                _ => (11 + bits.read(7)? as usize, 0),
395            };
396            let run = lengths
397                .get_mut(symbol..symbol + repeat)
398                .ok_or_else(|| malformed("a code-length run overruns the alphabet"))?;
399            run.fill(value);
400            symbol += repeat;
401        }
402    }
403    if build {
404        PrefixCode::new(&lengths).map(Some)
405    } else {
406        // Still validated: a broken unused code is still a broken stream.
407        PrefixCode::new(&lengths).map(|_| None)
408    }
409}
410
411/// The five codes for one block: green-length-cache, red, blue, alpha, distance.
412struct Group {
413    codes: [PrefixCode; 5],
414}
415
416/// A transform and the width it was read at.
417enum Transform {
418    Predictor { bits: u32, modes: Vec<u32> },
419    Color { bits: u32, elements: Vec<u32> },
420    SubtractGreen,
421    ColorIndexing { bits: u32, table: Vec<u32> },
422}
423
424pub(crate) const fn div_round_up(value: usize, bits: u32) -> usize {
425    (value + (1 << bits) - 1) >> bits
426}
427
428/// Decode a VP8L file's image stream after its 5-byte header: the image's
429/// `width * height` ARGB pixels.
430///
431/// # Errors
432///
433/// Returns [`PixelsError::Malformed`] for any stream that breaks the format.
434pub fn decode(data: &[u8], width: usize, height: usize) -> Result<Vec<u32>> {
435    let mut bits = BitReader::new(
436        data.get(5..)
437            .ok_or_else(|| malformed("VP8L header cut short"))?,
438    );
439    decode_image_stream(&mut bits, width, height, true)
440}
441
442/// Decode an image stream of the given size (§3.8). Only the top-level
443/// (`level0`) image carries transforms and meta prefix codes; an `ALPH`
444/// chunk's stream is one, with its size implicit.
445///
446/// # Errors
447///
448/// As [`decode`].
449pub fn decode_image_stream(
450    bits: &mut BitReader<'_>,
451    width: usize,
452    height: usize,
453    level0: bool,
454) -> Result<Vec<u32>> {
455    let mut transforms: Vec<(Transform, usize)> = Vec::new();
456    let mut coded_width = width;
457    if level0 {
458        let mut seen = [false; 4];
459        while bits.flag()? {
460            let kind = bits.read(2)? as usize;
461            if std::mem::replace(&mut seen[kind], true) {
462                return Err(malformed("a transform appears twice"));
463            }
464            let transform = match kind {
465                0 | 1 => {
466                    let block_bits = bits.read(3)? + 2;
467                    let sub = decode_image_stream(
468                        bits,
469                        div_round_up(coded_width, block_bits),
470                        div_round_up(height, block_bits),
471                        false,
472                    )?;
473                    if kind == 0 {
474                        Transform::Predictor {
475                            bits: block_bits,
476                            modes: sub,
477                        }
478                    } else {
479                        Transform::Color {
480                            bits: block_bits,
481                            elements: sub,
482                        }
483                    }
484                }
485                2 => Transform::SubtractGreen,
486                _ => {
487                    let size = bits.read(8)? as usize + 1;
488                    let mut table = decode_image_stream(bits, size, 1, false)?;
489                    // Stored as deltas, each channel wrapping independently.
490                    for i in 1..table.len() {
491                        table[i] = add_pixels(table[i], table[i - 1]);
492                    }
493                    let bundle = match size {
494                        0..=2 => 3,
495                        3..=4 => 2,
496                        5..=16 => 1,
497                        _ => 0,
498                    };
499                    // Unused indices decode to transparent black.
500                    table.resize(256, 0);
501                    Transform::ColorIndexing {
502                        bits: bundle,
503                        table,
504                    }
505                }
506            };
507            let read_width = coded_width;
508            if let Transform::ColorIndexing { bits: bundle, .. } = transform {
509                coded_width = div_round_up(coded_width, bundle);
510            }
511            transforms.push((transform, read_width));
512        }
513    }
514
515    let cache_bits = if bits.flag()? {
516        let cache_bits = bits.read(4)?;
517        if !(1..=11).contains(&cache_bits) {
518            return Err(malformed(format!(
519                "color cache bits {cache_bits} outside 1..=11"
520            )));
521        }
522        cache_bits
523    } else {
524        0
525    };
526
527    // Meta prefix codes: an entropy image naming each block's group.
528    let mut entropy: Option<(u32, usize, Vec<u32>)> = None;
529    let mut group_count = 1;
530    if level0 && bits.flag()? {
531        let block_bits = bits.read(3)? + 2;
532        let entropy_width = div_round_up(coded_width, block_bits);
533        let image =
534            decode_image_stream(bits, entropy_width, div_round_up(height, block_bits), false)?;
535        group_count = image
536            .iter()
537            .map(|&p| ((p >> 8) & 0xffff) as usize)
538            .max()
539            .unwrap_or(0)
540            + 1;
541        entropy = Some((block_bits, entropy_width, image));
542    }
543    let mut used = vec![entropy.is_none(); group_count];
544    if let Some((_, _, image)) = &entropy {
545        for &p in image {
546            used[((p >> 8) & 0xffff) as usize] = true;
547        }
548    }
549    let cache_size = if cache_bits > 0 {
550        1_usize << cache_bits
551    } else {
552        0
553    };
554    let alphabets = [
555        256 + LENGTH_CODES + cache_size,
556        256,
557        256,
558        256,
559        DISTANCE_CODES,
560    ];
561    let mut groups: Vec<Option<Group>> = Vec::with_capacity(group_count.min(4096));
562    for &needed in &used {
563        let mut codes = Vec::with_capacity(5);
564        for &alphabet in &alphabets {
565            codes.push(read_prefix_code(bits, alphabet, needed)?);
566        }
567        groups.push(if needed {
568            let codes: Vec<PrefixCode> = codes.into_iter().flatten().collect();
569            let codes: [PrefixCode; 5] = codes
570                .try_into()
571                .map_err(|_| malformed("a prefix code group is incomplete"))?;
572            Some(Group { codes })
573        } else {
574            None
575        });
576    }
577
578    let mut pixels = decode_pixels(
579        bits,
580        coded_width,
581        height,
582        cache_bits,
583        &groups,
584        entropy.as_ref(),
585    )?;
586
587    for (transform, read_width) in transforms.iter().rev() {
588        pixels = inverse(transform, pixels, *read_width, height);
589    }
590    Ok(pixels)
591}
592
593/// Per-channel addition, wrapping each byte.
594const fn add_pixels(a: u32, b: u32) -> u32 {
595    let alpha_green = (a & 0xff00_ff00).wrapping_add(b & 0xff00_ff00);
596    let red_blue = (a & 0x00ff_00ff).wrapping_add(b & 0x00ff_00ff);
597    (alpha_green & 0xff00_ff00) | (red_blue & 0x00ff_00ff)
598}
599
600/// A length or distance from its prefix code and extra bits (§3.6.2.2).
601fn prefix_value(bits: &mut BitReader<'_>, code: u16) -> Result<usize> {
602    let code = u32::from(code);
603    if code < 4 {
604        return Ok(code as usize + 1);
605    }
606    let extra = (code - 2) >> 1;
607    let offset = (2 + (code & 1)) << extra;
608    Ok((offset + bits.read(extra)? + 1) as usize)
609}
610
611fn decode_pixels(
612    bits: &mut BitReader<'_>,
613    width: usize,
614    height: usize,
615    cache_bits: u32,
616    groups: &[Option<Group>],
617    entropy: Option<&(u32, usize, Vec<u32>)>,
618) -> Result<Vec<u32>> {
619    let total = width
620        .checked_mul(height)
621        .ok_or_else(|| malformed("the image size overflows"))?;
622    let mut pixels: Vec<u32> = Vec::with_capacity(total);
623    let mut cache = vec![0_u32; if cache_bits > 0 { 1 << cache_bits } else { 0 }];
624    let mut cached = 0;
625    let group_at = |position: usize| -> Result<&Group> {
626        let index = match entropy {
627            None => 0,
628            Some((block_bits, entropy_width, image)) => {
629                let (x, y) = (position % width, position / width);
630                let at = (y >> block_bits) * entropy_width + (x >> block_bits);
631                ((image.get(at).copied().unwrap_or(0) >> 8) & 0xffff) as usize
632            }
633        };
634        groups
635            .get(index)
636            .and_then(Option::as_ref)
637            .ok_or_else(|| malformed("a block names a prefix code group that is missing"))
638    };
639
640    while pixels.len() < total {
641        let group = group_at(pixels.len())?;
642        let symbol = group.codes[0].read(bits)?;
643        if symbol < 256 {
644            let red = group.codes[1].read(bits)?;
645            let blue = group.codes[2].read(bits)?;
646            let alpha = group.codes[3].read(bits)?;
647            pixels.push(
648                (u32::from(alpha) << 24)
649                    | (u32::from(red) << 16)
650                    | (u32::from(symbol) << 8)
651                    | u32::from(blue),
652            );
653        } else if usize::from(symbol) < 256 + LENGTH_CODES {
654            let length = prefix_value(bits, symbol - 256)?;
655            let distance_symbol = group.codes[4].read(bits)?;
656            let code = prefix_value(bits, distance_symbol)?;
657            let distance = if code > 120 {
658                code - 120
659            } else {
660                let (dx, dy) = DISTANCE_MAP[code - 1];
661                (i64::from(dx) + i64::from(dy) * width as i64).max(1) as usize
662            };
663            let start = pixels
664                .len()
665                .checked_sub(distance)
666                .ok_or_else(|| malformed("a back reference points before the image"))?;
667            if pixels.len() + length > total {
668                return Err(malformed("a back reference runs past the image"));
669            }
670            for i in 0..length {
671                let pixel = pixels[start + i];
672                pixels.push(pixel);
673            }
674        } else {
675            let index = usize::from(symbol) - 256 - LENGTH_CODES;
676            let pixel = *cache
677                .get(index)
678                .ok_or_else(|| malformed("a color cache index is out of range"))?;
679            pixels.push(pixel);
680        }
681        if cache_bits > 0 {
682            for &pixel in &pixels[cached..] {
683                let key = pixel.wrapping_mul(0x1e35_a7bd) >> (32 - cache_bits);
684                cache[key as usize] = pixel;
685            }
686            cached = pixels.len();
687        }
688    }
689    Ok(pixels)
690}
691
692fn channel(pixel: u32, shift: u32) -> i32 {
693    ((pixel >> shift) & 0xff) as i32
694}
695
696fn per_channel(f: impl Fn(i32, i32, i32) -> i32, a: u32, b: u32, c: u32) -> u32 {
697    [24, 16, 8, 0].iter().fold(0, |out, &shift| {
698        let value = f(channel(a, shift), channel(b, shift), channel(c, shift));
699        out | ((value.clamp(0, 255) as u32) << shift)
700    })
701}
702
703fn average2(a: u32, b: u32) -> u32 {
704    per_channel(|a, b, _| (a + b) / 2, a, b, 0)
705}
706
707fn select(left: u32, top: u32, top_left: u32) -> u32 {
708    let distance = |pixel: u32| -> i32 {
709        [24, 16, 8, 0]
710            .iter()
711            .map(|&s| {
712                let estimate = channel(left, s) + channel(top, s) - channel(top_left, s);
713                (estimate - channel(pixel, s)).abs()
714            })
715            .sum()
716    };
717    if distance(left) < distance(top) {
718        left
719    } else {
720        top
721    }
722}
723
724/// Predictor `mode` from left, top, top-right and top-left (§3.5.1).
725pub(crate) fn predict(mode: u32, l: u32, t: u32, tr: u32, tl: u32) -> u32 {
726    match mode {
727        1 => l,
728        2 => t,
729        3 => tr,
730        4 => tl,
731        5 => average2(average2(l, tr), t),
732        6 => average2(l, tl),
733        7 => average2(l, t),
734        8 => average2(tl, t),
735        9 => average2(t, tr),
736        10 => average2(average2(l, tl), average2(t, tr)),
737        11 => select(l, t, tl),
738        12 => per_channel(|l, t, tl| l + t - tl, l, t, tl),
739        13 => per_channel(|a, tl, _| a + (a - tl) / 2, average2(l, t), tl, 0),
740        // 0, and the two values a 4-bit mode can hold beyond 13, which
741        // libwebp also treats as black.
742        _ => 0xff00_0000,
743    }
744}
745
746fn color_delta(t: u32, c: u32) -> i32 {
747    (i32::from(t as u8 as i8) * i32::from(c as u8 as i8)) >> 5
748}
749
750/// Undo one transform: `pixels` at the width it produced, back to `width`.
751fn inverse(transform: &Transform, mut pixels: Vec<u32>, width: usize, height: usize) -> Vec<u32> {
752    match transform {
753        Transform::SubtractGreen => {
754            for pixel in &mut pixels {
755                let green = (*pixel >> 8) & 0xff;
756                *pixel = add_pixels(*pixel, (green << 16) | green);
757            }
758            pixels
759        }
760        Transform::Color { bits, elements } => {
761            let blocks_wide = div_round_up(width, *bits);
762            for (i, pixel) in pixels.iter_mut().enumerate() {
763                let (x, y) = (i % width, i / width);
764                let element = elements[(y >> bits) * blocks_wide + (x >> bits)];
765                let (green, red, blue) =
766                    ((*pixel >> 8) & 0xff, (*pixel >> 16) & 0xff, *pixel & 0xff);
767                let new_red = (red as i32 + color_delta(element, green)) & 0xff;
768                let new_blue = (blue as i32
769                    + color_delta(element >> 8, green)
770                    + color_delta(element >> 16, new_red as u32))
771                    & 0xff;
772                *pixel = (*pixel & 0xff00_ff00) | ((new_red as u32) << 16) | new_blue as u32;
773            }
774            pixels
775        }
776        Transform::Predictor { bits, modes } => {
777            let blocks_wide = div_round_up(width, *bits);
778            for i in 0..pixels.len() {
779                let (x, y) = (i % width, i / width);
780                let prediction = if i == 0 {
781                    0xff00_0000
782                } else if y == 0 {
783                    pixels[i - 1]
784                } else if x == 0 {
785                    pixels[i - width]
786                } else {
787                    let mode = (modes[(y >> bits) * blocks_wide + (x >> bits)] >> 8) & 0xf;
788                    // The top-right of the last column is, in scan order,
789                    // the first pixel of the current row (§3.5.1).
790                    predict(
791                        mode,
792                        pixels[i - 1],
793                        pixels[i - width],
794                        pixels[i - width + 1],
795                        pixels[i - width - 1],
796                    )
797                };
798                pixels[i] = add_pixels(pixels[i], prediction);
799            }
800            pixels
801        }
802        Transform::ColorIndexing { bits, table } => {
803            let packed_width = div_round_up(width, *bits);
804            let per_byte = 1 << bits;
805            let index_bits = 8 >> bits;
806            let mask = (1_u32 << index_bits) - 1;
807            let mut out = Vec::with_capacity(width * height);
808            for y in 0..height {
809                for x in 0..width {
810                    let packed = pixels[y * packed_width + (x >> bits)];
811                    let shift = (x & (per_byte - 1)) as u32 * index_bits;
812                    let index = ((packed >> 8) >> shift) & mask;
813                    out.push(table[index as usize]);
814                }
815            }
816            out
817        }
818    }
819}
820
821#[cfg(test)]
822#[allow(
823    clippy::unwrap_used,
824    clippy::indexing_slicing,
825    reason = "tests operate on known-good values and assert shapes directly"
826)]
827mod tests {
828    use super::*;
829
830    #[test]
831    fn bits_are_read_least_significant_first() {
832        let mut bits = BitReader::new(&[0b1011_0100, 0xff]);
833        assert_eq!(bits.read(2).unwrap(), 0b00);
834        assert_eq!(bits.read(3).unwrap(), 0b101);
835        assert_eq!(bits.read(5).unwrap(), 0b11_101);
836        assert!(bits.read(7).is_err(), "only six bits remain");
837    }
838
839    #[test]
840    fn a_prefix_code_decodes_canonically() {
841        // Lengths 1, 2, 3, 3 give codes 0, 10, 110, 111 — sent MSB first,
842        // so in LSB-first bytes symbol 2 (110) arrives as bits 1, 1, 0.
843        let code = PrefixCode::new(&[1, 2, 3, 3]).unwrap();
844        // Bits in order: 0 | 1 0 | 1 1 0 | 1 1 1.
845        let mut bits = BitReader::new(&[0b1101_1010, 0b1]);
846        let decoded: Vec<u16> = (0..4).map(|_| code.read(&mut bits).unwrap()).collect();
847        assert_eq!(decoded, [0, 1, 2, 3]);
848    }
849
850    #[test]
851    fn long_codes_take_the_canonical_walk() {
852        // 2 symbols of length 1..=10, one of each up to length 11 and two
853        // of 12: lengths past FAST_BITS must still decode.
854        let mut lengths = vec![0_u8; 16];
855        for (i, l) in (1..=11).enumerate() {
856            lengths[i] = l;
857        }
858        lengths[11] = 12;
859        lengths[12] = 12;
860        let code = PrefixCode::new(&lengths).unwrap();
861        // Symbol 12 is the last code: twelve 1 bits.
862        let mut bits = BitReader::new(&[0xff, 0x0f]);
863        assert_eq!(code.read(&mut bits).unwrap(), 12);
864    }
865
866    #[test]
867    fn incomplete_and_empty_codes_are_rejected_and_one_symbol_costs_nothing() {
868        assert!(PrefixCode::new(&[1, 2]).is_err(), "incomplete");
869        assert!(PrefixCode::new(&[1, 1, 1]).is_err(), "over-subscribed");
870        assert!(PrefixCode::new(&[0, 0]).is_err(), "empty");
871        let single = PrefixCode::new(&[0, 0, 7]).unwrap();
872        let mut bits = BitReader::new(&[]);
873        assert_eq!(single.read(&mut bits).unwrap(), 2);
874    }
875
876    #[test]
877    fn prefix_values_follow_the_table() {
878        let mut bits = BitReader::new(&[0b1, 0, 0, 0]);
879        assert_eq!(prefix_value(&mut bits, 3).unwrap(), 4);
880        // Code 4: one extra bit, offset 2 -> 5..6.
881        assert_eq!(prefix_value(&mut bits, 4).unwrap(), 6);
882        // Code 39: 18 extra bits, offset 3 << 18 -> 786433.. (all zero bits).
883        let mut zeros = BitReader::new(&[0; 4]);
884        assert_eq!(prefix_value(&mut zeros, 39).unwrap(), 786_433);
885    }
886
887    #[test]
888    fn predictors_follow_their_definitions() {
889        let (l, t, tr, tl) = (0x10_20_30_40, 0x30_40_50_60, 0xff_00_ff_00, 0x00_00_00_00);
890        assert_eq!(predict(0, l, t, tr, tl), 0xff00_0000);
891        assert_eq!(predict(7, l, t, tr, tl), 0x20_30_40_50);
892        // Mode 12 clamps per channel: 0x10 + 0x30 - 0 etc.
893        assert_eq!(predict(12, l, t, tr, tl), 0x40_60_80_a0);
894        // Mode 13 halves (a - tl) truncating toward zero, as C does:
895        // 2 + (2 - 5) / 2 is 2 + -1, not 2 + -2.
896        assert_eq!(
897            predict(13, 0x00_00_00_02, 0x00_00_00_02, 0, 0x00_00_00_05),
898            1
899        );
900        assert_eq!(predict(14, l, t, tr, tl), 0xff00_0000);
901    }
902
903    #[test]
904    fn color_indexing_unpacks_bundled_pixels() {
905        // Two colours: one bit per pixel, eight per packed pixel.
906        let mut table = vec![0xff00_0000, 0xffff_ffff];
907        table.resize(256, 0);
908        let transform = Transform::ColorIndexing { bits: 3, table };
909        let packed = vec![0b1010_0101 << 8];
910        let out = inverse(&transform, packed, 8, 1);
911        let ones: Vec<bool> = out.iter().map(|&p| p == 0xffff_ffff).collect();
912        assert_eq!(ones, [true, false, true, false, false, true, false, true]);
913    }
914
915    #[test]
916    fn truncated_streams_are_malformed_not_panics() {
917        for len in 0..12 {
918            let data = vec![0xa5_u8; len];
919            let mut bits = BitReader::new(&data);
920            assert!(
921                decode_image_stream(&mut bits, 4, 4, true).is_err(),
922                "{len} bytes"
923            );
924        }
925    }
926}