use std::borrow::Cow;
use std::fmt;
use std::fs::File;
use std::io::{BufWriter, Write};
use std::path::Path;
use png::{BitDepth, ColorType, Compression, Encoder, Filter};
use crate::image_io::buffer::{ColorSpace, CoverSource, ImageBuffer};
use crate::image_io::envelope;
use crate::pipeline::error::OutputError;
pub(crate) const LENGTH_HEADER_BYTES: usize = 4;
pub(crate) const LENGTH_HEADER_BITS: usize = LENGTH_HEADER_BYTES * 8;
const POSITIONS_PER_HEADER_BIT: usize = 64;
pub(crate) const HEADER_POSITIONS: usize = LENGTH_HEADER_BITS * POSITIONS_PER_HEADER_BIT;
pub(crate) const FRAME_OVERHEAD_BYTES: usize = LENGTH_HEADER_BYTES + 4;
pub(crate) const MIN_CIPHERTEXT_BYTES: usize = 16;
pub(crate) fn split_regions<T>(positions: &[T]) -> (&[T], &[T]) {
positions.split_at(HEADER_POSITIONS.min(positions.len()))
}
pub(crate) fn split_regions_mut<T>(positions: &mut [T]) -> (&mut [T], &mut [T]) {
let boundary = HEADER_POSITIONS.min(positions.len());
positions.split_at_mut(boundary)
}
fn carrier_offset(index: usize, color_space: ColorSpace) -> usize {
index * color_space.bytes_per_pixel()
}
pub(crate) fn gather_cover_symbols(image: &ImageBuffer, permutation: &[usize]) -> Vec<u8> {
let color_space = image.color_space();
let pixels = image.pixels();
permutation
.iter()
.map(|&index| {
pixels
.get(carrier_offset(index, color_space))
.copied()
.unwrap_or(0)
})
.collect()
}
pub(crate) fn apply_cover_symbols(image: &mut ImageBuffer, permutation: &[usize], symbols: &[u8]) {
let color_space = image.color_space();
let pixels = image.pixels_mut();
for (&index, &symbol) in permutation.iter().zip(symbols.iter()) {
if let Some(sample) = pixels.get_mut(carrier_offset(index, color_space)) {
*sample = symbol;
}
}
}
pub(crate) fn reorder_costs(costs: &[f32], permutation: &[usize]) -> Vec<f32> {
permutation
.iter()
.map(|&index| costs.get(index).copied().unwrap_or(0.0))
.collect()
}
pub(crate) fn encode_length_header(ciphertext_len: usize) -> Option<[u8; LENGTH_HEADER_BYTES]> {
u32::try_from(ciphertext_len).ok().map(u32::to_be_bytes)
}
pub(crate) fn decode_length_header(header: &[u8]) -> Option<usize> {
let bytes: [u8; LENGTH_HEADER_BYTES] = header.get(..LENGTH_HEADER_BYTES)?.try_into().ok()?;
Some(u32::from_be_bytes(bytes) as usize)
}
fn png_layout(color_space: ColorSpace) -> (ColorType, BitDepth) {
match color_space {
ColorSpace::Rgb8 => (ColorType::Rgb, BitDepth::Eight),
ColorSpace::Rgb16 => (ColorType::Rgb, BitDepth::Sixteen),
ColorSpace::Rgba8 => (ColorType::Rgba, BitDepth::Eight),
ColorSpace::Luma8 => (ColorType::Grayscale, BitDepth::Eight),
}
}
fn encoding_failed<E: fmt::Display>(err: E) -> OutputError {
OutputError::EncodingFailed(err.to_string())
}
fn configured_encoder<W: Write>(sink: W, image: &ImageBuffer) -> Encoder<'static, W> {
let (width, height) = image.dimensions();
let (color_type, bit_depth) = png_layout(image.color_space());
let mut encoder = Encoder::new(sink, width, height);
encoder.set_color(color_type);
encoder.set_depth(bit_depth);
encoder.set_compression(Compression::Fast);
encoder.set_filter(Filter::Adaptive);
encoder
}
fn big_endian_samples(image: &ImageBuffer) -> Cow<'_, [u8]> {
let pixels = image.pixels();
if image.color_space() != ColorSpace::Rgb16 {
return Cow::Borrowed(pixels);
}
Cow::Owned(
pixels
.chunks_exact(2)
.flat_map(|pair| match pair {
[low, high] => [*high, *low],
_ => [0, 0],
})
.collect::<Vec<u8>>(),
)
}
fn compress_samples(image: &ImageBuffer, samples: &[u8]) -> Result<Vec<u8>, OutputError> {
let mut encoded: Vec<u8> = Vec::new();
{
let encoder = configured_encoder(&mut encoded, image);
let mut writer = encoder.write_header().map_err(encoding_failed)?;
writer.write_image_data(samples).map_err(encoding_failed)?;
writer.finish().map_err(encoding_failed)?;
}
let mut stream = Vec::new();
let mut offset = envelope::SIGNATURE_LEN;
while let Some((code, data, next)) = envelope::read_chunk(&encoded, offset) {
if &code == b"IDAT" {
stream.extend_from_slice(data);
}
offset = next;
}
Ok(stream)
}
pub(crate) fn write_png(image: &ImageBuffer, path: &Path) -> Result<(), OutputError> {
let (width, height) = image.dimensions();
let color_space = image.color_space();
let expected = (width as usize)
.checked_mul(height as usize)
.and_then(|pixels| pixels.checked_mul(color_space.bytes_per_pixel()));
if expected != Some(image.pixels().len()) {
return Err(OutputError::MalformedBuffer);
}
let stream = compress_samples(image, &big_endian_samples(image))?;
let file = File::create(path).map_err(encoding_failed)?;
let envelope = image.envelope();
let mut writer = configured_encoder(BufWriter::new(file), image)
.write_header()
.map_err(encoding_failed)?;
for chunk in envelope.preserved_chunks() {
writer
.write_chunk(png::chunk::ChunkType(chunk.kind().type_code()), chunk.data())
.map_err(encoding_failed)?;
}
for piece in stream.chunks(envelope.idat_chunk_size()) {
writer
.write_chunk(png::chunk::IDAT, piece)
.map_err(encoding_failed)?;
}
writer.finish().map_err(encoding_failed)
}
#[cfg(test)]
mod tests {
#![allow(clippy::expect_used)]
#![allow(clippy::panic)]
use super::*;
use image::ImageFormat;
use tempfile::NamedTempFile;
use crate::image_io::envelope::PngEnvelope;
const WIDTH: u32 = 5;
const HEIGHT: u32 = 4;
fn container(color_space: ColorSpace) -> ImageBuffer {
let stride = color_space.bytes_per_pixel();
let pixel_count = (WIDTH * HEIGHT) as usize;
let pixels = (0..pixel_count * stride)
.map(|offset| {
if offset % stride == 0 {
(offset / stride) as u8
} else {
0xEE
}
})
.collect();
ImageBuffer::new(pixels, WIDTH, HEIGHT, color_space)
}
fn chunks(bytes: &[u8]) -> Vec<(String, usize)> {
let mut found = Vec::new();
let mut offset = 8usize;
while offset + 8 <= bytes.len() {
let length = u32::from_be_bytes(
bytes[offset..offset + 4]
.try_into()
.expect("four bytes of length"),
) as usize;
let code = String::from_utf8_lossy(&bytes[offset + 4..offset + 8]).into_owned();
found.push((code, length));
offset += 12 + length;
}
found
}
fn noisy_container(side: u32, color_space: ColorSpace) -> ImageBuffer {
let len = side as usize * side as usize * color_space.bytes_per_pixel();
let mut state = 0x2545_F491_4F6C_DD1Du64;
let pixels = (0..len)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(state >> 33) as u8
})
.collect();
ImageBuffer::new(pixels, side, side, color_space)
}
fn envelope_of(chunks: &[(&[u8; 4], Vec<u8>)]) -> PngEnvelope {
let mut bytes = vec![0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A];
for (code, data) in chunks {
bytes.extend_from_slice(&(data.len() as u32).to_be_bytes());
bytes.extend_from_slice(*code);
bytes.extend_from_slice(data);
bytes.extend_from_slice(&[0, 0, 0, 0]);
}
PngEnvelope::read(&bytes)
}
#[test]
fn the_frame_is_cut_at_a_constant_boundary() {
let long: Vec<u8> = vec![0; HEADER_POSITIONS + 17];
let (header, payload) = split_regions(&long);
assert_eq!(header.len(), HEADER_POSITIONS);
assert_eq!(payload.len(), 17);
let mut short = vec![0u8; 9];
let (header, payload) = split_regions_mut(&mut short);
assert_eq!(header.len(), 9);
assert!(payload.is_empty());
}
#[test]
fn cover_symbols_round_trip_through_the_permutation() {
for color_space in [
ColorSpace::Rgb8,
ColorSpace::Rgba8,
ColorSpace::Luma8,
ColorSpace::Rgb16,
] {
let mut image = container(color_space);
let permutation = vec![3usize, 0, 7, 1];
let gathered = gather_cover_symbols(&image, &permutation);
assert_eq!(gathered, vec![3u8, 0, 7, 1], "layout {color_space:?}");
let written: Vec<u8> = gathered.iter().map(|symbol| symbol ^ 1).collect();
apply_cover_symbols(&mut image, &permutation, &written);
assert_eq!(
gather_cover_symbols(&image, &permutation),
written,
"layout {color_space:?}"
);
let stride = color_space.bytes_per_pixel();
assert!(image
.pixels()
.iter()
.enumerate()
.filter(|(offset, _)| offset % stride != 0)
.all(|(_, &sample)| sample == 0xEE));
}
}
#[test]
fn positions_outside_the_container_are_inert() {
let mut image = container(ColorSpace::Rgb8);
let permutation = vec![0usize, 9_999];
assert_eq!(gather_cover_symbols(&image, &permutation), vec![0u8, 0]);
let before = image.pixels().to_vec();
apply_cover_symbols(&mut image, &permutation, &[0, 42]);
assert_eq!(image.pixels(), before.as_slice());
}
#[test]
fn costs_follow_the_permutation() {
let costs = [0.5f32, 1.5, 2.5, 3.5];
assert_eq!(
reorder_costs(&costs, &[2, 0, 3, 1]),
vec![2.5f32, 0.5, 3.5, 1.5]
);
assert_eq!(reorder_costs(&costs, &[9]), vec![0.0f32]);
}
#[test]
fn the_length_header_round_trips() {
let header = encode_length_header(0x0102_0304).expect("a small length must encode");
assert_eq!(header, [0x01, 0x02, 0x03, 0x04]);
assert_eq!(decode_length_header(&header), Some(0x0102_0304));
assert_eq!(decode_length_header(&[0, 0, 0, 7, 9, 9]), Some(7));
assert_eq!(decode_length_header(&[0, 0, 7]), None);
}
#[test]
fn an_unrepresentable_length_does_not_encode() {
assert!(encode_length_header(u32::MAX as usize).is_some());
assert!(encode_length_header(u32::MAX as usize + 1).is_none());
}
#[test]
fn every_layout_survives_a_write_and_a_read() {
for color_space in [
ColorSpace::Rgb8,
ColorSpace::Rgba8,
ColorSpace::Luma8,
ColorSpace::Rgb16,
] {
let image = container(color_space);
let file = NamedTempFile::new().expect("temporary stego file");
if let Err(error) = write_png(&image, file.path()) {
panic!("a {color_space:?} container must be writable: {error}");
}
let bytes = std::fs::read(file.path()).expect("the written file must be readable");
assert_eq!(&bytes[..4], &[0x89, b'P', b'N', b'G']);
let decoded = image::load_from_memory_with_format(&bytes, ImageFormat::Png)
.expect("the written png must decode");
assert_eq!(decoded.width(), WIDTH);
assert_eq!(decoded.height(), HEIGHT);
let recovered: Vec<u8> = match color_space {
ColorSpace::Rgb8 => decoded.into_rgb8().into_raw(),
ColorSpace::Rgba8 => decoded.into_rgba8().into_raw(),
ColorSpace::Luma8 => decoded.into_luma8().into_raw(),
ColorSpace::Rgb16 => decoded
.into_rgb16()
.into_raw()
.into_iter()
.flat_map(u16::to_le_bytes)
.collect(),
};
assert_eq!(
recovered,
image.pixels(),
"a {color_space:?} container must decode to the samples it was written from"
);
}
}
#[test]
fn sixteen_bit_samples_are_written_big_endian() {
let image = ImageBuffer::new(
vec![0x34, 0x12, 0x78, 0x56, 0xBC, 0x9A],
1,
1,
ColorSpace::Rgb16,
);
assert_eq!(
big_endian_samples(&image).as_ref(),
&[0x12, 0x34, 0x56, 0x78, 0x9A, 0xBC]
);
let flat = container(ColorSpace::Rgba8);
assert_eq!(big_endian_samples(&flat).as_ref(), flat.pixels());
}
#[test]
fn the_original_chunks_are_repeated_and_the_rest_are_not() {
let envelope = envelope_of(&[
(b"IHDR", vec![0; 13]),
(b"sRGB", vec![0]),
(b"gAMA", 45_455u32.to_be_bytes().to_vec()),
(b"eXIf", vec![0x45; 120]),
(b"tEXt", b"Software\0a camera".to_vec()),
(b"pHYs", vec![0, 0, 0x0B, 0x13, 0, 0, 0x0B, 0x13, 1]),
(b"IDAT", vec![0; 8192]),
(b"IEND", Vec::new()),
]);
let image = ImageBuffer::with_envelope(
container(ColorSpace::Rgb8).pixels().to_vec(),
WIDTH,
HEIGHT,
ColorSpace::Rgb8,
envelope,
);
let file = NamedTempFile::new().expect("temporary stego file");
write_png(&image, file.path()).expect("the container must be writable");
let bytes = std::fs::read(file.path()).expect("the written file must be readable");
let written: Vec<String> = chunks(&bytes)
.into_iter()
.map(|(code, _)| code)
.filter(|code| code != "IDAT")
.collect();
assert_eq!(written, vec!["IHDR", "sRGB", "gAMA", "pHYs", "IEND"]);
let physical = bytes
.windows(4)
.position(|window| window == b"pHYs")
.expect("the chunk must be in the file");
assert_eq!(
&bytes[physical + 4..physical + 13],
&[0, 0, 0x0B, 0x13, 0, 0, 0x0B, 0x13, 1]
);
}
#[test]
fn the_pixel_stream_is_cut_the_way_the_original_was() {
let split = 4096usize;
let envelope = envelope_of(&[
(b"IHDR", vec![0; 13]),
(b"IDAT", vec![0; split]),
(b"IDAT", vec![0; split]),
(b"IDAT", vec![0; 17]),
(b"IEND", Vec::new()),
]);
assert_eq!(envelope.idat_chunk_size(), split);
let noisy = noisy_container(200, ColorSpace::Rgb8);
let image = ImageBuffer::with_envelope(
noisy.pixels().to_vec(),
200,
200,
ColorSpace::Rgb8,
envelope,
);
let file = NamedTempFile::new().expect("temporary stego file");
write_png(&image, file.path()).expect("the container must be writable");
let bytes = std::fs::read(file.path()).expect("the written file must be readable");
let idat: Vec<usize> = chunks(&bytes)
.into_iter()
.filter(|(code, _)| code == "IDAT")
.map(|(_, length)| length)
.collect();
assert!(
idat.len() > 20,
"the stream must be cut into many chunks, got {}",
idat.len()
);
let (last, full) = idat.split_last().expect("at least one chunk");
assert!(full.iter().all(|&length| length == split), "{idat:?}");
assert!(*last <= split, "{idat:?}");
let decoded = image::load_from_memory_with_format(&bytes, ImageFormat::Png)
.expect("the written png must decode");
assert_eq!(decoded.into_rgb8().into_raw(), noisy.pixels());
}
#[test]
fn a_container_without_an_original_wears_the_default_profile() {
let image = noisy_container(200, ColorSpace::Rgb8);
let file = NamedTempFile::new().expect("temporary stego file");
write_png(&image, file.path()).expect("the container must be writable");
let bytes = std::fs::read(file.path()).expect("the written file must be readable");
let written = chunks(&bytes);
let names: Vec<&str> = written
.iter()
.map(|(code, _)| code.as_str())
.filter(|code| *code != "IDAT")
.collect();
assert_eq!(names, vec!["IHDR", "gAMA", "cHRM", "sRGB", "pHYs", "IEND"]);
let idat: Vec<usize> = written
.iter()
.filter(|(code, _)| code == "IDAT")
.map(|(_, length)| *length)
.collect();
assert!(idat.len() > 4, "{idat:?}");
let (last, full) = idat.split_last().expect("at least one chunk");
assert!(full.iter().all(|&length| length == 8192), "{idat:?}");
assert!(*last <= 8192, "{idat:?}");
}
#[test]
fn a_buffer_that_belies_its_geometry_is_refused() {
let image = ImageBuffer::new(vec![0u8; 5], WIDTH, HEIGHT, ColorSpace::Rgb8);
let file = NamedTempFile::new().expect("temporary stego file");
let error = write_png(&image, file.path()).expect_err("a short buffer must not be encoded");
assert!(
matches!(error, OutputError::MalformedBuffer),
"got: {error:?}"
);
}
#[test]
fn an_unwritable_path_is_reported() {
let image = container(ColorSpace::Rgb8);
let directory = tempfile::tempdir().expect("temporary directory");
let error = write_png(
&image,
&directory.path().join("no-such-dir").join("out.png"),
)
.expect_err("a path under a missing directory must fail");
assert!(
matches!(error, OutputError::EncodingFailed(_)),
"got: {error:?}"
);
}
}