use crate::error::DecodeError;
use crate::png::Chunk;
use crate::types::{ChunkType, FilterType};
use core::cmp::min;
use miniz_oxide::inflate::TINFLStatus;
use miniz_oxide::inflate::core::inflate_flags::{
TINFL_FLAG_COMPUTE_ADLER32, TINFL_FLAG_HAS_MORE_INPUT, TINFL_FLAG_PARSE_ZLIB_HEADER,
};
use miniz_oxide::inflate::core::{DecompressorOxide, decompress_with_limit};
#[cfg(feature = "alloc")]
extern crate alloc;
#[cfg(feature = "alloc")]
use alloc::{vec, vec::Vec};
pub struct ChunkDecompressor<'src, T> {
decompressor: DecompressorOxide,
data_chunks: &'src [u8],
next_chunk_start: Option<usize>,
current_chunk: Option<&'src [u8]>,
chunk_end: bool,
buffer: T,
data_pos: usize,
buffer_count: usize,
flags: u32,
total_decompressed: usize, }
impl<'src, 'buf> ChunkDecompressor<'src, &'buf mut [u8]> {
pub fn new_ref(data_chunks: &'src [u8], buffer: &'buf mut [u8], check_crc: bool) -> Self {
Self::new(data_chunks, buffer, check_crc)
}
}
#[cfg(feature = "alloc")]
impl<'src> ChunkDecompressor<'src, Vec<u8>> {
pub fn new_vec(data_chunks: &'src [u8], check_crc: bool) -> Self {
Self::new(data_chunks, vec![0_u8; 1024 << 5], check_crc)
}
}
impl<'src> ChunkDecompressor<'src, [u8; 1024 << 5]> {
pub fn new_static(data_chunks: &'src [u8], check_crc: bool) -> Self {
Self::new(data_chunks, [0_u8; 1024 << 5], check_crc)
}
}
impl<'src, T> ChunkDecompressor<'src, T>
where
T: AsRef<[u8]> + AsMut<[u8]>,
{
fn new(data_chunks: &'src [u8], buffer: T, check_crc: bool) -> Self {
let decompressor = DecompressorOxide::new();
let mut flags = TINFL_FLAG_PARSE_ZLIB_HEADER | TINFL_FLAG_HAS_MORE_INPUT;
if check_crc {
flags |= TINFL_FLAG_COMPUTE_ADLER32;
}
ChunkDecompressor {
decompressor,
data_chunks,
next_chunk_start: Some(0),
current_chunk: None,
chunk_end: false,
buffer,
data_pos: 0,
buffer_count: 0,
flags,
total_decompressed: 0,
}
}
fn check_chunk_data(&mut self) {
if let Some(chunk) = self.current_chunk
&& !chunk.is_empty()
{
return;
}
loop {
if let Some(next_start) = self.next_chunk_start
&& next_start < self.data_chunks.len()
{
let next_chunk = Chunk::from_bytes(self.data_chunks, next_start, false).unwrap();
if next_chunk.end < self.data_chunks.len() {
self.next_chunk_start = Some(next_chunk.end);
} else {
self.next_chunk_start = None;
}
if next_chunk.chunk_type == ChunkType::ImageData && !next_chunk.data.is_empty() {
self.current_chunk = Some(next_chunk.data);
return;
}
} else {
self.current_chunk = None;
self.chunk_end = true;
return;
}
}
}
fn get_enough_data(&mut self, size: usize) -> Result<(), DecodeError> {
debug_assert!(
size <= self.buffer.as_ref().len(),
"Decompression buffer too small (need {})",
size
);
if self.buffer_count >= size {
return Ok(());
}
let mut buffer_pos = self.buffer_count + self.data_pos;
if buffer_pos >= self.buffer.as_ref().len() {
buffer_pos -= self.buffer.as_ref().len();
}
self.check_chunk_data();
let next_data = match self.current_chunk {
None => &[], Some(x) => x,
};
let available_bytes = self.buffer.as_ref().len() - self.buffer_count;
let (status, in_count, out_count) = decompress_with_limit(
&mut self.decompressor,
next_data,
self.buffer.as_mut(),
buffer_pos,
available_bytes,
self.flags,
);
if let Some(chunk) = &mut self.current_chunk {
*chunk = &(*chunk)[in_count..];
if chunk.is_empty() && self.next_chunk_start.is_none() {
self.chunk_end = true;
}
}
self.buffer_count += out_count;
self.total_decompressed += out_count;
debug_assert!(
buffer_pos + out_count <= self.buffer.as_ref().len(),
"decompress wrapped around"
);
if (status as i32) < 0 {
return Err(DecodeError::Decompress(status));
}
match status {
TINFLStatus::Done if !self.chunk_end => {
return Err(DecodeError::InvalidChunk);
}
TINFLStatus::NeedsMoreInput if self.chunk_end => {
return Err(DecodeError::InvalidChunk);
}
_ => {}
}
self.get_enough_data(size)
}
fn remove_data(&mut self, size: usize) {
extern crate alloc;
self.data_pos += size;
if self.data_pos >= self.buffer.as_ref().len() {
self.data_pos -= self.buffer.as_ref().len();
}
self.buffer_count -= size;
}
fn filter_type(&mut self) -> Result<FilterType, DecodeError> {
let byte = self.buffer.as_ref()[self.data_pos];
FilterType::try_from(byte).map_err(|_| DecodeError::InvalidFilterType)
}
fn copy_to_slice(&self, target: &mut [u8]) {
let count = target.len();
debug_assert!(
count < self.buffer_count,
"copy_to_slice, error slice too big {} > {}",
count + 1,
self.buffer_count
);
let buffer_end = min(self.data_pos + 1 + count, self.buffer.as_ref().len());
let next_count = buffer_end - self.data_pos - 1;
target[..next_count].copy_from_slice(&self.buffer.as_ref()[self.data_pos + 1..buffer_end]);
let count = count - next_count;
if count > 0 {
let next_pos = next_count;
target[next_pos..].copy_from_slice(&self.buffer.as_ref()[..count]);
}
}
fn enumerate(&self, count: usize) -> impl Iterator<Item = (usize, u8)> {
let main_count = count + 1;
let end = min(self.data_pos + main_count, self.buffer.as_ref().len());
self.buffer.as_ref()[self.data_pos..end]
.iter()
.chain(if end == self.buffer.as_ref().len() {
self.buffer.as_ref()[0..main_count - (self.buffer.as_ref().len() - self.data_pos)]
.iter()
} else {
[].iter()
})
.skip(1)
.copied()
.enumerate()
}
pub fn decode_next_scanline(
&mut self,
last_scanline: &mut [u8],
bytes_per_pixel: usize,
) -> Result<(), DecodeError> {
self.get_enough_data(last_scanline.len() + 1)?;
let filter_type = self.filter_type()?;
match filter_type {
FilterType::None => self.copy_to_slice(last_scanline),
FilterType::Sub => {
let mut left_pixel = [0_u8; 8];
self.enumerate(last_scanline.len())
.fold(0, |byte, (i, value)| {
let left = left_pixel[byte];
last_scanline[i] = value.wrapping_add(left);
left_pixel[byte] = last_scanline[i];
(byte + 1) % bytes_per_pixel
});
}
FilterType::Up => {
for (i, value) in self.enumerate(last_scanline.len()) {
last_scanline[i] = value.wrapping_add(last_scanline[i]);
}
}
FilterType::Average => {
let mut left_pixel = [0_u8; 8];
self.enumerate(last_scanline.len())
.fold(0, |byte, (i, value)| {
let left = left_pixel[byte];
let top = last_scanline[i];
let average = (left as u16 + top as u16) / 2;
last_scanline[i] = value.wrapping_add(average as u8);
left_pixel[byte] = last_scanline[i];
(byte + 1) % bytes_per_pixel
});
}
FilterType::Paeth => {
let mut top_left_pixel = [0_u8; 8];
let mut left_pixel = [0_u8; 8];
self.enumerate(last_scanline.len())
.fold(0, |byte, (i, value)| {
let a = left_pixel[byte] as i16;
let b = last_scanline[i] as i16;
let c = top_left_pixel[byte] as i16;
let p = a + b - c; let pa = (p - a).abs(); let pb = (p - b).abs();
let pc = (p - c).abs();
let predictor = if pa <= pb && pa <= pc {
left_pixel[byte]
} else if pb <= pc {
last_scanline[i]
} else {
top_left_pixel[byte]
};
top_left_pixel[byte] = last_scanline[i];
last_scanline[i] = value.wrapping_add(predictor);
left_pixel[byte] = last_scanline[i];
(byte + 1) % bytes_per_pixel
});
}
}
self.remove_data(last_scanline.len() + 1);
Ok(())
}
pub fn reset(&mut self) {
todo!()
}
}
#[cfg(test)]
mod tests {
extern crate std;
use super::*;
use crate::ParsedPng;
use crate::colors::AlphaColor;
use std::fs;
use std::prelude::v1::*;
#[test]
fn list_chunks() {
let bytes = fs::read("sekiro.png").unwrap();
let png = ParsedPng::from_bytes(&bytes, true, AlphaColor).unwrap();
let mut decompressor = ChunkDecompressor::new_static(png.data_chunks, true);
for _ in 0..35 {
decompressor.current_chunk = None;
decompressor.check_chunk_data();
assert!(decompressor.current_chunk.is_some(), "Missing chunk");
assert!(!decompressor.chunk_end, "Decompression ended early");
assert_eq!(
decompressor.current_chunk.unwrap().len(),
32_768,
"Incorrect chunk size"
);
}
decompressor.current_chunk = None;
decompressor.check_chunk_data();
assert!(decompressor.current_chunk.is_some(), "Missing chunk");
assert!(!decompressor.chunk_end, "Decompression ended early");
assert_eq!(
decompressor.current_chunk.unwrap().len(),
7_663,
"Incorrect chunk size"
);
decompressor.current_chunk = None;
decompressor.check_chunk_data();
assert!(decompressor.chunk_end, "Decompression ended late");
}
#[test]
fn read_chunks() {
let bytes = fs::read("sekiro.png").unwrap();
let png = ParsedPng::from_bytes(&bytes, true, AlphaColor).unwrap();
let mut decompressor = ChunkDecompressor::new_static(png.data_chunks, true);
let mut scanline = vec![0_u8; 5120];
for _ in 0..720 {
let r = decompressor.get_enough_data(5121);
assert!(r.is_ok(), "Get data Error");
decompressor.copy_to_slice(&mut scanline);
assert_eq!(
decompressor.enumerate(5120).count(),
5120,
"Enumerate can't count"
);
let enumeration: Vec<u8> = decompressor.enumerate(5120).map(|(_, x)| x).collect();
assert_eq!(
enumeration, scanline,
"Enumerate misaligned with copy to slice"
);
decompressor.remove_data(5121);
}
assert_eq!(decompressor.buffer_count, 0, "Main buffer left");
assert!(decompressor.chunk_end, "Decompression left some data");
}
#[test]
fn decode() {
let bytes = fs::read("sekiro.png").unwrap();
let png = ParsedPng::from_bytes(&bytes, true, AlphaColor).unwrap();
let mut decompressor = ChunkDecompressor::new_static(png.data_chunks, true);
println!("Size: {}", size_of::<DecompressorOxide>());
let mut scanline = vec![0_u8; 5120];
for _ in 0..720 {
let r = decompressor.decode_next_scanline(&mut scanline, 4);
assert!(r.is_ok(), "Get data Error");
}
assert_eq!(decompressor.buffer_count, 0, "Main buffer left");
assert!(decompressor.chunk_end, "Decompression left some data");
}
}