use crate::error::{Error, Result};
use crate::prelude::*;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum AlphaCompression {
None,
Lossless,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum AlphaFilter {
None = 0,
Horizontal = 1,
Vertical = 2,
Gradient = 3,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub struct AlphaHeader {
pub compression: AlphaCompression,
pub filter: AlphaFilter,
pub preprocessing: u8,
}
#[must_use]
pub const fn build_header(
compression: AlphaCompression,
filter: AlphaFilter,
preprocessing: u8,
) -> u8 {
let method = match compression {
AlphaCompression::None => 0,
AlphaCompression::Lossless => 1,
};
method | ((filter as u8) << 2) | ((preprocessing & 0b11) << 4)
}
pub fn parse_header(chunk: &[u8]) -> Result<(AlphaHeader, &[u8])> {
let (&byte, data) = chunk.split_first().ok_or(Error::Truncated)?;
let compression = match byte & 0b11 {
0 => AlphaCompression::None,
1 => AlphaCompression::Lossless,
_ => return Err(Error::InvalidContainer),
};
let filter = match (byte >> 2) & 0b11 {
0 => AlphaFilter::None,
1 => AlphaFilter::Horizontal,
2 => AlphaFilter::Vertical,
_ => AlphaFilter::Gradient,
};
let preprocessing = (byte >> 4) & 0b11;
if preprocessing > 1 {
return Err(Error::InvalidContainer);
}
if byte >> 6 != 0 {
return Err(Error::InvalidContainer);
}
Ok((
AlphaHeader {
compression,
filter,
preprocessing,
},
data,
))
}
#[allow(
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
reason = "the sum is clamped into 0..=255 before the byte cast, exactly as \
libwebp GradientPredictor_C stores its clamped int into a uint8_t"
)]
fn gradient_predictor(a: u8, b: u8, c: u8) -> u8 {
(i32::from(a) + i32::from(b) - i32::from(c)).clamp(0, 255) as u8
}
fn horizontal(prev: Option<&[u8]>, row: &mut [u8]) {
let mut acc = prev.map_or(0, |p| p.first().copied().unwrap_or(0));
for value in row.iter_mut() {
acc = acc.wrapping_add(*value);
*value = acc;
}
}
fn vertical(prev: &[u8], row: &mut [u8]) {
for (value, &top) in row.iter_mut().zip(prev) {
*value = top.wrapping_add(*value);
}
}
fn gradient(prev: &[u8], row: &mut [u8]) {
let mut top_left = prev.first().copied().unwrap_or(0);
let mut left = top_left;
for (value, &top) in row.iter_mut().zip(prev) {
let grad = gradient_predictor(left, top, top_left);
left = value.wrapping_add(grad);
top_left = top;
*value = left;
}
}
pub fn unfilter_row(filter: AlphaFilter, prev: Option<&[u8]>, row: &mut [u8]) {
match (filter, prev) {
(AlphaFilter::None, _) => {},
(AlphaFilter::Horizontal, _) | (AlphaFilter::Vertical | AlphaFilter::Gradient, None) => {
horizontal(prev, row);
},
(AlphaFilter::Vertical, Some(prev)) => vertical(prev, row),
(AlphaFilter::Gradient, Some(prev)) => gradient(prev, row),
}
}
pub fn unfilter(filter: AlphaFilter, plane: &mut [u8], width: usize, height: usize) {
let Some(expected) = width.checked_mul(height) else {
return;
};
if width == 0 || height == 0 || plane.len() != expected {
return;
}
for r in 0..height {
let (done, rest) = plane.split_at_mut(r * width);
let cur = &mut rest[..width];
let prev = if r == 0 {
None
} else {
Some(&done[(r - 1) * width..r * width])
};
unfilter_row(filter, prev, cur);
}
}
fn forward_horizontal(above: Option<&[u8]>, orig: &[u8], dst: &mut [u8]) {
let mut left = above.map_or(0, |a| a.first().copied().unwrap_or(0));
for (d, &value) in dst.iter_mut().zip(orig) {
*d = value.wrapping_sub(left);
left = value;
}
}
fn forward_gradient(above: &[u8], orig: &[u8], dst: &mut [u8]) {
let mut top_left = above.first().copied().unwrap_or(0);
let mut left = top_left;
for ((d, &value), &top) in dst.iter_mut().zip(orig).zip(above) {
*d = value.wrapping_sub(gradient_predictor(left, top, top_left));
left = value;
top_left = top;
}
}
#[must_use]
pub fn filter_plane(filter: AlphaFilter, plane: &[u8], width: usize, height: usize) -> Vec<u8> {
if width == 0 || height == 0 || width.checked_mul(height) != Some(plane.len()) {
return plane.to_vec();
}
let mut out = vec![0u8; plane.len()];
for r in 0..height {
let orig = &plane[r * width..(r + 1) * width];
let above = if r == 0 {
None
} else {
Some(&plane[(r - 1) * width..r * width])
};
let dst = &mut out[r * width..(r + 1) * width];
match (filter, above) {
(AlphaFilter::None, _) => dst.copy_from_slice(orig),
(AlphaFilter::Horizontal, _)
| (AlphaFilter::Vertical | AlphaFilter::Gradient, None) => {
forward_horizontal(above, orig, dst);
},
(AlphaFilter::Vertical, Some(above)) => {
for ((d, &value), &top) in dst.iter_mut().zip(orig).zip(above) {
*d = value.wrapping_sub(top);
}
},
(AlphaFilter::Gradient, Some(above)) => forward_gradient(above, orig, dst),
}
}
out
}
#[cfg(test)]
mod tests {
use proptest::prelude::*;
use super::{
AlphaCompression, AlphaFilter, AlphaHeader, build_header, filter_plane, parse_header,
unfilter, unfilter_row,
};
use crate::error::Error;
#[test]
fn none_filter_is_identity() {
let mut plane = [10u8, 20, 30, 40];
unfilter(AlphaFilter::None, &mut plane, 2, 2);
assert_eq!(plane, [10, 20, 30, 40]);
}
#[test]
fn filter_plane_returns_input_clone_on_dimension_mismatch() {
let plane = [10u8, 20, 30];
assert_eq!(
filter_plane(AlphaFilter::Horizontal, &plane, 2, 2),
plane.to_vec()
);
}
#[test]
fn horizontal_running_sum_with_wrap_and_seed() {
let mut plane = [10u8, 5, 250, 3, 1, 2];
unfilter(AlphaFilter::Horizontal, &mut plane, 3, 2);
assert_eq!(plane, [10, 15, 9, 13, 14, 16]);
}
#[test]
fn vertical_adds_pixel_above_with_wrap() {
let mut plane = [10u8, 20, 30, 250, 5, 100];
unfilter(AlphaFilter::Vertical, &mut plane, 3, 2);
assert_eq!(plane, [10, 30, 60, 4, 35, 160]);
}
#[test]
fn gradient_col0_interior_and_saturation() {
let mut plane = [50u8, 150, 66, 90, 10, 100, 200, 5];
unfilter(AlphaFilter::Gradient, &mut plane, 4, 2);
assert_eq!(plane, [50, 200, 10, 100, 60, 54, 200, 4]);
}
#[test]
fn all_filters_agree_on_row0() {
let deltas = [10u8, 5, 250, 3];
let expected = [10u8, 15, 9, 12];
for filter in [
AlphaFilter::Horizontal,
AlphaFilter::Vertical,
AlphaFilter::Gradient,
] {
let mut plane = deltas;
unfilter(filter, &mut plane, 4, 1);
assert_eq!(plane, expected, "{filter:?} row 0");
}
}
#[test]
fn unfilter_row_row0_matches_horizontal_or_identity() {
let deltas = [10u8, 5, 250, 3];
let running_sum = [10u8, 15, 9, 12];
for filter in [
AlphaFilter::None,
AlphaFilter::Horizontal,
AlphaFilter::Vertical,
AlphaFilter::Gradient,
] {
let mut row = deltas;
unfilter_row(filter, None, &mut row);
let want = if filter == AlphaFilter::None {
deltas
} else {
running_sum
};
assert_eq!(row, want, "{filter:?}");
}
}
#[test]
fn unfilter_guards_bad_sizes() {
let mut plane = [1u8, 2, 3];
unfilter(AlphaFilter::Horizontal, &mut plane, 2, 2);
assert_eq!(plane, [1, 2, 3]);
let mut p2 = [9u8, 9];
unfilter(AlphaFilter::Gradient, &mut p2, 0, 5);
unfilter(AlphaFilter::Gradient, &mut p2, 5, 0);
assert_eq!(p2, [9, 9]);
}
#[test]
fn parse_header_empty_is_truncated() {
assert_eq!(parse_header(&[]).unwrap_err(), Error::Truncated);
}
#[test]
fn parse_header_fields_and_data_slice() {
let (header, data) = parse_header(&[0x1D, 0xAA, 0xBB]).unwrap();
assert_eq!(
header,
AlphaHeader {
compression: AlphaCompression::Lossless,
filter: AlphaFilter::Gradient,
preprocessing: 1,
}
);
assert_eq!(data, &[0xAA, 0xBB]);
}
#[test]
fn parse_header_accepts_every_filter() {
for (byte, filter) in [
(0x00u8, AlphaFilter::None),
(0x04, AlphaFilter::Horizontal),
(0x08, AlphaFilter::Vertical),
(0x0C, AlphaFilter::Gradient),
] {
let bytes = [byte];
let (header, data) = parse_header(&bytes).unwrap();
assert_eq!(header.filter, filter);
assert_eq!(header.compression, AlphaCompression::None);
assert_eq!(header.preprocessing, 0);
assert!(data.is_empty());
}
}
#[test]
fn parse_header_reads_both_methods() {
assert_eq!(
parse_header(&[0x00]).unwrap().0.compression,
AlphaCompression::None
);
assert_eq!(
parse_header(&[0x01]).unwrap().0.compression,
AlphaCompression::Lossless
);
}
#[test]
fn parse_header_rejects_reserved_bits_and_bad_fields() {
assert_eq!(parse_header(&[0x40]).unwrap_err(), Error::InvalidContainer);
assert_eq!(parse_header(&[0x80]).unwrap_err(), Error::InvalidContainer);
assert_eq!(parse_header(&[0x02]).unwrap_err(), Error::InvalidContainer);
assert_eq!(parse_header(&[0x03]).unwrap_err(), Error::InvalidContainer);
assert_eq!(parse_header(&[0x20]).unwrap_err(), Error::InvalidContainer);
assert_eq!(parse_header(&[0x30]).unwrap_err(), Error::InvalidContainer);
}
fn any_filter() -> impl Strategy<Value = AlphaFilter> {
prop_oneof![
Just(AlphaFilter::None),
Just(AlphaFilter::Horizontal),
Just(AlphaFilter::Vertical),
Just(AlphaFilter::Gradient),
]
}
#[test]
fn build_header_round_trips_through_parse() {
for compression in [AlphaCompression::None, AlphaCompression::Lossless] {
for filter in [
AlphaFilter::None,
AlphaFilter::Horizontal,
AlphaFilter::Vertical,
AlphaFilter::Gradient,
] {
for preprocessing in 0u8..=1 {
let bytes = [build_header(compression, filter, preprocessing)];
let (header, data) = parse_header(&bytes).unwrap();
assert_eq!(header.compression, compression);
assert_eq!(header.filter, filter);
assert_eq!(header.preprocessing, preprocessing);
assert!(data.is_empty());
}
}
}
assert_eq!(
build_header(AlphaCompression::Lossless, AlphaFilter::Gradient, 1) >> 6,
0
);
}
proptest! {
#[test]
fn forward_then_unfilter_round_trips(
(width, height, plane) in (1usize..=17, 1usize..=17).prop_flat_map(|(w, h)| {
prop::collection::vec(any::<u8>(), w * h).prop_map(move |plane| (w, h, plane))
}),
filter in any_filter(),
) {
let mut filtered = filter_plane(filter, &plane, width, height);
unfilter(filter, &mut filtered, width, height);
prop_assert_eq!(filtered, plane);
}
#[test]
fn unfilter_never_panics_and_is_length_guarded(
filter in any_filter(),
width in 0usize..=40,
height in 0usize..=40,
bytes in prop::collection::vec(any::<u8>(), 0..=64),
) {
let original = bytes.clone();
let mut buf = bytes;
unfilter(filter, &mut buf, width, height);
prop_assert_eq!(buf.len(), original.len());
let matches_plane =
width != 0 && height != 0 && width.checked_mul(height) == Some(original.len());
if !matches_plane {
prop_assert_eq!(buf, original);
}
}
#[test]
fn parse_header_never_panics(bytes in prop::collection::vec(any::<u8>(), 0..=8)) {
if let Ok((_, data)) = parse_header(&bytes) {
prop_assert_eq!(data, &bytes[1..]);
}
}
}
}