use crate::base::{Error, Result};
use crate::gfx::bitmap::Bitmap;
use crate::gfx::{jpeg, png};
const PNG_MAGIC: [u8; 8] = [0x89, b'P', b'N', b'G', b'\r', b'\n', 0x1A, b'\n'];
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum ImageFormat {
Png,
Jpeg,
}
pub fn sniff_format(bytes: &[u8]) -> Option<ImageFormat> {
if bytes.starts_with(&PNG_MAGIC) {
Some(ImageFormat::Png)
} else if bytes.starts_with(&[0xFF, 0xD8]) {
Some(ImageFormat::Jpeg)
} else {
None
}
}
pub fn decode_image(bytes: &[u8]) -> Result<Bitmap> {
match sniff_format(bytes) {
Some(ImageFormat::Png) => png::decode(bytes),
Some(ImageFormat::Jpeg) => jpeg::decode(bytes),
None => Err(Error::Parse(format!(
"image: unrecognized format (magic {:02X?}); PNG and baseline JPEG decode, \
GIF/WebP/AVIF/TIFF do not",
&bytes[..bytes.len().min(4)]
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::gfx::png_test_encoder::encode_rgba;
#[test]
fn routes_png_by_magic() {
let bmp = Bitmap::from_fn(3, 2, |x, y| {
crate::base::Rgba::rgb((x * 80) as u8, (y * 100) as u8, 7)
});
let png = encode_rgba(&bmp);
let out = decode_image(&png).unwrap();
assert_eq!((out.width(), out.height()), (3, 2));
assert_eq!(out.get(1, 1), bmp.get(1, 1));
}
#[test]
fn routes_jpeg_by_magic() {
let jpg = crate::gfx::jpeg_fixtures::GRAD444;
assert_eq!(sniff_format(jpg), Some(ImageFormat::Jpeg));
let out = decode_image(jpg).unwrap();
assert!(out.width() > 0 && out.height() > 0);
}
#[test]
fn unknown_magic_rejects_by_name() {
let err = decode_image(b"GIF89a....").unwrap_err();
let msg = err.to_string();
assert!(msg.contains("unrecognized format"), "{msg}");
assert!(msg.contains("PNG"), "must name what DOES decode: {msg}");
let err = decode_image(b"").unwrap_err();
assert!(err.to_string().contains("unrecognized format"));
}
#[test]
fn truncated_after_magic_fails_in_the_decoder_not_the_sniffer() {
let err = decode_image(&[0xFF, 0xD8, 0xFF]).unwrap_err();
assert!(!err.to_string().contains("unrecognized"), "{err}");
}
#[test]
fn decode_image_survives_truncation_and_marker_soup() {
let png = encode_rgba(&Bitmap::from_fn(9, 7, |x, y| {
crate::base::Rgba::rgb((x * 29) as u8, (y * 37) as u8, 128)
}));
let jpg = crate::gfx::jpeg_fixtures::GRAD420;
for src in [&png[..], jpg] {
for cut in 0..src.len().min(96) {
let _ = decode_image(&src[..cut]);
}
let mut cut = 96;
while cut < src.len() {
let _ = decode_image(&src[..cut]);
cut += 7;
}
}
let mut state = 0xDEADBEEFCAFEF00Du64;
let mut rng = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for case in 0..300 {
let len = 8 + (rng() % 300) as usize;
let mut bytes: Vec<u8> = (0..len).map(|_| rng() as u8).collect();
match case % 3 {
0 => bytes[..8].copy_from_slice(&PNG_MAGIC),
1 => {
bytes[0] = 0xFF;
bytes[1] = 0xD8;
}
_ => {}
}
let _ = decode_image(&bytes);
}
for (i, src) in [&png[..], jpg].into_iter().enumerate() {
for k in 0..200 {
let mut mutated = src.to_vec();
let pos = (rng() as usize) % mutated.len();
mutated[pos] ^= 1 << (rng() % 8);
let _ = decode_image(&mutated);
let _ = (i, k);
}
}
}
}