pub(crate) struct DecodedImage {
pub data: Vec<u8>,
pub alpha: Option<Vec<u8>>,
pub width: u32,
pub height: u32,
pub color_space: &'static str,
pub is_jpeg: bool,
}
pub(crate) fn decode_image(data: &[u8], content_type: &str) -> Option<DecodedImage> {
if content_type.contains("jpeg") || content_type.contains("jpg") || is_jpeg(data) {
decode_jpeg(data)
} else {
decode_png_or_other(data)
}
}
fn is_jpeg(data: &[u8]) -> bool {
data.len() >= 2 && data[0] == 0xFF && data[1] == 0xD8
}
fn decode_jpeg(data: &[u8]) -> Option<DecodedImage> {
let (width, height) = jpeg_dimensions(data)?;
Some(DecodedImage {
data: data.to_vec(),
alpha: None,
width,
height,
color_space: "DeviceRGB",
is_jpeg: true,
})
}
fn jpeg_dimensions(data: &[u8]) -> Option<(u32, u32)> {
let mut i = 2; while i < data.len() {
if data[i] != 0xFF {
return None;
}
while data.get(i) == Some(&0xFF) {
i += 1;
}
let marker = *data.get(i)?;
i += 1;
if marker == 0xD9 {
return None;
}
if matches!(marker, 0x01 | 0xD0..=0xD8) {
continue;
}
let length = u16::from_be_bytes([*data.get(i)?, *data.get(i + 1)?]) as usize;
if length < 2 {
return None;
}
let segment_end = i.checked_add(length)?;
let segment = data.get(i..segment_end)?;
if matches!(marker, 0xC0..=0xC2) {
let height = u16::from_be_bytes([*segment.get(3)?, *segment.get(4)?]) as u32;
let width = u16::from_be_bytes([*segment.get(5)?, *segment.get(6)?]) as u32;
return Some((width, height));
}
i = segment_end;
}
None
}
fn decode_png_or_other(data: &[u8]) -> Option<DecodedImage> {
decode_png(data)
}
fn decode_png(data: &[u8]) -> Option<DecodedImage> {
if data.len() < 8 || &data[0..8] != b"\x89PNG\r\n\x1a\n" {
return None;
}
let mut pos = 8;
let mut width = 0u32;
let mut height = 0u32;
let mut bit_depth = 0u8;
let mut color_type = 0u8;
let mut idat_data = Vec::new();
let mut palette: Vec<u8> = Vec::new();
let mut palette_alpha: Vec<u8> = Vec::new();
while pos + 8 <= data.len() {
let chunk_len =
u32::from_be_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]) as usize;
let chunk_type = &data[pos + 4..pos + 8];
let chunk_data_start = pos + 8;
let chunk_data_end = chunk_data_start + chunk_len;
if chunk_data_end > data.len() {
break;
}
match chunk_type {
b"IHDR" => {
if chunk_len >= 13 {
let d = &data[chunk_data_start..];
width = u32::from_be_bytes([d[0], d[1], d[2], d[3]]);
height = u32::from_be_bytes([d[4], d[5], d[6], d[7]]);
bit_depth = d[8];
color_type = d[9];
}
}
b"IDAT" => {
idat_data.extend_from_slice(&data[chunk_data_start..chunk_data_end]);
}
b"PLTE" => {
palette.clear();
palette.extend_from_slice(&data[chunk_data_start..chunk_data_end]);
}
b"tRNS" => {
palette_alpha.clear();
palette_alpha.extend_from_slice(&data[chunk_data_start..chunk_data_end]);
}
b"IEND" => break,
_ => {}
}
pos = chunk_data_end + 4; }
if width == 0 || height == 0 || idat_data.is_empty() || bit_depth != 8 {
return None;
}
let decompressed = miniz_oxide::inflate::decompress_to_vec_zlib(&idat_data).ok()?;
let channels: usize = match color_type {
0 => 1, 2 => 3, 3 => 1, 4 => 2, 6 => 4, _ => return None,
};
let stride = width as usize * channels;
let expected = (stride + 1) * height as usize; if decompressed.len() < expected {
return None;
}
let mut unfiltered = vec![0u8; stride * height as usize];
let mut prev_row = vec![0u8; stride];
for y in 0..height as usize {
let row_start = y * (stride + 1);
let filter_type = decompressed[row_start];
let raw = &decompressed[row_start + 1..row_start + 1 + stride];
let out_start = y * stride;
let out = &mut unfiltered[out_start..out_start + stride];
match filter_type {
0 => {
out.copy_from_slice(raw);
}
1 => {
for i in 0..stride {
let a = if i >= channels { out[i - channels] } else { 0 };
out[i] = raw[i].wrapping_add(a);
}
}
2 => {
for i in 0..stride {
out[i] = raw[i].wrapping_add(prev_row[i]);
}
}
3 => {
for i in 0..stride {
let a = if i >= channels {
out[i - channels] as u16
} else {
0
};
let b = prev_row[i] as u16;
out[i] = raw[i].wrapping_add(((a + b) / 2) as u8);
}
}
4 => {
for i in 0..stride {
let a = if i >= channels {
out[i - channels] as i32
} else {
0
};
let b = prev_row[i] as i32;
let c = if i >= channels {
prev_row[i - channels] as i32
} else {
0
};
out[i] = raw[i].wrapping_add(paeth_predictor(a, b, c));
}
}
_ => {
out.copy_from_slice(raw);
}
}
prev_row.copy_from_slice(out);
}
match color_type {
0 => {
Some(DecodedImage {
data: unfiltered,
alpha: None,
width,
height,
color_space: "DeviceGray",
is_jpeg: false,
})
}
2 => {
Some(DecodedImage {
data: unfiltered,
alpha: None,
width,
height,
color_space: "DeviceRGB",
is_jpeg: false,
})
}
3 => {
if palette.len() < 3 {
return None;
}
let entries = palette.len() / 3;
let pixel_count = (width * height) as usize;
let mut rgb = Vec::with_capacity(pixel_count * 3);
let mut alpha = Vec::with_capacity(pixel_count);
let mut all_opaque = true;
for &index in unfiltered.iter().take(pixel_count) {
let i = index as usize;
if i < entries {
rgb.extend_from_slice(&palette[i * 3..i * 3 + 3]);
} else {
rgb.extend_from_slice(&[0, 0, 0]);
}
let a = palette_alpha.get(i).copied().unwrap_or(255);
if a != 255 {
all_opaque = false;
}
alpha.push(a);
}
Some(DecodedImage {
data: rgb,
alpha: if all_opaque { None } else { Some(alpha) },
width,
height,
color_space: "DeviceRGB",
is_jpeg: false,
})
}
4 => {
let pixel_count = (width * height) as usize;
let mut gray = Vec::with_capacity(pixel_count);
let mut alpha = Vec::with_capacity(pixel_count);
for i in 0..pixel_count {
gray.push(unfiltered[i * 2]);
alpha.push(unfiltered[i * 2 + 1]);
}
Some(DecodedImage {
data: gray,
alpha: Some(alpha),
width,
height,
color_space: "DeviceGray",
is_jpeg: false,
})
}
6 => {
let pixel_count = (width * height) as usize;
let mut rgb = Vec::with_capacity(pixel_count * 3);
let mut alpha = Vec::with_capacity(pixel_count);
let mut all_opaque = true;
for i in 0..pixel_count {
rgb.push(unfiltered[i * 4]);
rgb.push(unfiltered[i * 4 + 1]);
rgb.push(unfiltered[i * 4 + 2]);
let a = unfiltered[i * 4 + 3];
alpha.push(a);
if a != 255 {
all_opaque = false;
}
}
Some(DecodedImage {
data: rgb,
alpha: if all_opaque { None } else { Some(alpha) },
width,
height,
color_space: "DeviceRGB",
is_jpeg: false,
})
}
_ => None,
}
}
fn paeth_predictor(a: i32, b: i32, c: i32) -> u8 {
let p = a + b - c;
let pa = (p - a).abs();
let pb = (p - b).abs();
let pc = (p - c).abs();
if pa <= pb && pa <= pc {
a as u8
} else if pb <= pc {
b as u8
} else {
c as u8
}
}
#[cfg(test)]
mod tests {
use super::*;
fn jpeg_with_restart_marker_before_sof() -> Vec<u8> {
let mut jpeg = vec![0xFF, 0xD8];
jpeg.extend_from_slice(&[0xFF, 0xD0]);
jpeg.extend_from_slice(&[
0xFF, 0xC0, 0x00, 0x0B, 0x08, 0x00, 0x02, 0x00, 0x03, 0x03, 0x01, 0x11, 0x00,
]);
jpeg
}
#[test]
fn decode_jpeg_pass_through() {
let mut jpeg = vec![0xFF, 0xD8];
jpeg.extend_from_slice(&[0xFF, 0xE0, 0x00, 0x02]);
jpeg.extend_from_slice(&[
0xFF, 0xC0, 0x00, 0x0B, 0x08, 0x00, 0x02, 0x00, 0x03, 0x03, 0x01, 0x11, 0x00,
]);
jpeg.extend_from_slice(&[0xFF, 0xD9]);
let result = decode_image(&jpeg, "image/jpeg");
assert!(result.is_some());
let decoded = result.unwrap();
assert!(decoded.is_jpeg);
assert_eq!(decoded.width, 3);
assert_eq!(decoded.height, 2);
assert_eq!(decoded.color_space, "DeviceRGB");
assert!(decoded.alpha.is_none());
}
#[test]
fn decode_invalid_data_returns_none() {
let result = decode_image(b"not an image", "image/png");
assert!(result.is_none());
}
#[test]
fn jpeg_detection() {
assert!(is_jpeg(&[0xFF, 0xD8, 0xFF]));
assert!(!is_jpeg(&[0x89, 0x50, 0x4E, 0x47])); assert!(!is_jpeg(&[0xFF])); }
#[test]
fn jpeg_restart_marker_before_sof_preserves_dimensions() {
assert_eq!(
jpeg_dimensions(&jpeg_with_restart_marker_before_sof()),
Some((3, 2))
);
}
#[test]
fn jpeg_bytes_after_eoi_cannot_supply_dimensions() {
let mut jpeg = vec![0xFF, 0xD8, 0xFF, 0xD9];
jpeg.extend_from_slice(&[
0xFF, 0xC0, 0x00, 0x0B, 0x08, 0x00, 0x02, 0x00, 0x03, 0x03, 0x01, 0x11, 0x00,
]);
assert_eq!(jpeg_dimensions(&jpeg), None);
}
#[test]
fn every_truncated_jpeg_header_returns_without_panicking() {
let jpeg = jpeg_with_restart_marker_before_sof();
for length in 0..jpeg.len() {
assert_eq!(jpeg_dimensions(&jpeg[..length]), None, "length {length}");
}
assert_eq!(jpeg_dimensions(&jpeg), Some((3, 2)));
}
fn png_chunk(kind: &[u8; 4], body: &[u8]) -> Vec<u8> {
let mut c = Vec::new();
c.extend_from_slice(&(body.len() as u32).to_be_bytes());
c.extend_from_slice(kind);
c.extend_from_slice(body);
c.extend_from_slice(&[0, 0, 0, 0]);
c
}
fn indexed_png(
width: u32,
height: u32,
palette: &[u8],
trns: Option<&[u8]>,
indices: &[u8],
) -> Vec<u8> {
let mut ihdr = Vec::new();
ihdr.extend_from_slice(&width.to_be_bytes());
ihdr.extend_from_slice(&height.to_be_bytes());
ihdr.push(8); ihdr.push(3); ihdr.extend_from_slice(&[0, 0, 0]);
let mut raw = Vec::new();
for y in 0..height as usize {
raw.push(0);
raw.extend_from_slice(&indices[y * width as usize..(y + 1) * width as usize]);
}
let idat = miniz_oxide::deflate::compress_to_vec_zlib(&raw, 6);
let mut out = b"\x89PNG\r\n\x1a\n".to_vec();
out.extend_from_slice(&png_chunk(b"IHDR", &ihdr));
if !palette.is_empty() {
out.extend_from_slice(&png_chunk(b"PLTE", palette));
}
if let Some(t) = trns {
out.extend_from_slice(&png_chunk(b"tRNS", t));
}
out.extend_from_slice(&png_chunk(b"IDAT", &idat));
out.extend_from_slice(&png_chunk(b"IEND", &[]));
out
}
#[test]
fn indexed_png_expands_palette_to_rgb() {
let palette = [255, 0, 0, 0, 255, 0, 0, 0, 255];
let png = indexed_png(2, 2, &palette, None, &[0, 1, 2, 0]);
let decoded = decode_image(&png, "image/png").expect("indexed PNG must decode");
assert_eq!((decoded.width, decoded.height), (2, 2));
assert_eq!(decoded.color_space, "DeviceRGB");
assert!(!decoded.is_jpeg);
assert_eq!(
decoded.data,
vec![255, 0, 0, 0, 255, 0, 0, 0, 255, 255, 0, 0]
);
assert!(
decoded.alpha.is_none(),
"a fully opaque image should not carry an alpha channel"
);
}
#[test]
fn indexed_png_honours_trns_alpha() {
let palette = [255, 0, 0, 0, 255, 0];
let png = indexed_png(2, 1, &palette, Some(&[0]), &[0, 1]);
let decoded = decode_image(&png, "image/png").expect("indexed PNG must decode");
assert_eq!(decoded.data, vec![255, 0, 0, 0, 255, 0]);
assert_eq!(decoded.alpha, Some(vec![0, 255]));
}
#[test]
fn indexed_png_without_palette_is_rejected() {
let png = indexed_png(1, 1, &[], None, &[0]);
assert!(decode_image(&png, "image/png").is_none());
}
#[test]
fn indexed_png_out_of_range_index_falls_back_to_black() {
let palette = [255, 0, 0]; let png = indexed_png(2, 1, &palette, None, &[0, 5]);
let decoded = decode_image(&png, "image/png").expect("should still decode");
assert_eq!(decoded.data, vec![255, 0, 0, 0, 0, 0]);
}
}