use crate::fdct;
use crate::format::{ICC_CHUNK, ZIGZAG, icc_segments, marker};
use crate::huffman::HuffmanEncoder;
use crate::tables;
use otf_pixels_core::{
EncodeOptions, Encoder, ImageDescriptor, PixelFormat, PixelsError, Result, Sink,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum Subsampling {
None,
Horizontal,
#[default]
Both,
}
impl Subsampling {
const fn factors(self) -> (u32, u32) {
match self {
Self::None => (1, 1),
Self::Horizontal => (2, 1),
Self::Both => (2, 2),
}
}
}
#[derive(Debug)]
pub struct JpegEncoder {
quality: u8,
subsampling: Subsampling,
state: Option<State>,
icc: Option<Vec<u8>>,
}
#[derive(Debug)]
struct State {
descriptor: ImageDescriptor,
grayscale: bool,
factors: (u32, u32),
luma_quant: [u16; 64],
chroma_quant: [u16; 64],
luma_dc: HuffmanEncoder,
luma_ac: HuffmanEncoder,
chroma_dc: HuffmanEncoder,
chroma_ac: HuffmanEncoder,
band: Vec<u8>,
band_rows: u32,
band_height: u32,
mcus_per_line: u32,
luma: Plane,
chroma_blue: Plane,
chroma_red: Plane,
chroma_full: Vec<u8>,
predictors: [i32; 3],
writer: BitWriter,
rows_written: u32,
}
#[derive(Debug, Default)]
struct Plane {
stride: usize,
height: usize,
samples: Vec<u8>,
}
impl Plane {
fn new(stride: usize, height: usize) -> Self {
Self {
stride,
height,
samples: vec![128; stride * height],
}
}
fn block(&self, x: usize, y: usize, out: &mut [u8; 64]) {
for row in 0..8 {
let start = (y + row) * self.stride + x;
let source = self
.samples
.get(start..)
.and_then(|rest| rest.get(..8))
.unwrap_or(&[128; 8]);
if let Some(target) = out.get_mut(row * 8..row * 8 + 8) {
target.copy_from_slice(source);
}
}
}
}
fn chroma_plane(grayscale: bool, stride: usize, height: u32) -> Plane {
if grayscale {
Plane::default()
} else {
Plane::new(stride, height as usize)
}
}
#[derive(Debug, Default)]
struct BitWriter {
bytes: Vec<u8>,
accumulator: u32,
bits: u32,
}
impl BitWriter {
fn write(&mut self, code: u32, length: u32) {
if length == 0 || length > 32 {
return;
}
let mask = if length >= 32 {
u32::MAX
} else {
(1_u32 << length) - 1
};
self.accumulator = (self.accumulator << length.min(31)) | (code & mask);
self.bits += length;
while self.bits >= 8 {
let byte = ((self.accumulator >> (self.bits - 8)) & 0xFF) as u8;
self.bytes.push(byte);
if byte == 0xFF {
self.bytes.push(0x00);
}
self.bits -= 8;
}
self.accumulator &= (1_u32 << self.bits) - 1;
}
fn flush(&mut self) {
if self.bits > 0 {
let padding = 8 - self.bits;
self.write((1 << padding) - 1, padding);
}
}
}
fn magnitude(value: i32) -> (u32, u32) {
if value == 0 {
return (0, 0);
}
let size = 32 - value.unsigned_abs().leading_zeros();
let bits = if value < 0 {
(value - 1) as u32 & ((1_u32 << size) - 1)
} else {
value as u32
};
(size, bits)
}
impl JpegEncoder {
#[must_use]
pub const fn new() -> Self {
Self {
quality: EncodeOptions::DEFAULT_QUALITY,
subsampling: Subsampling::Both,
state: None,
icc: None,
}
}
pub fn with_quality(quality: u8) -> Result<Self> {
if !(1..=100).contains(&quality) {
return Err(PixelsError::invalid_argument(
"quality",
format!("must be in 1..=100, got {quality}"),
));
}
Ok(Self {
quality,
subsampling: if quality >= 90 {
Subsampling::None
} else {
Subsampling::Both
},
state: None,
icc: None,
})
}
#[must_use]
pub const fn with_subsampling(mut self, subsampling: Subsampling) -> Self {
self.subsampling = subsampling;
self
}
#[must_use]
pub fn from_options(options: &EncodeOptions) -> Self {
Self::with_quality(options.quality).unwrap_or_else(|_| Self::new())
}
#[must_use]
pub const fn subsampling(&self) -> Subsampling {
self.subsampling
}
}
impl Default for JpegEncoder {
fn default() -> Self {
Self::new()
}
}
fn channels_of(format: PixelFormat) -> Result<usize> {
match format {
PixelFormat::Gray8 => Ok(1),
PixelFormat::GrayA8 => Ok(2),
PixelFormat::Rgb8 => Ok(3),
PixelFormat::Rgba8 => Ok(4),
other => Err(PixelsError::unsupported(format!(
"JPEG encoding needs an 8-bit format; got {other}. Convert first."
))),
}
}
fn pixel_rgb(row: &[u8], index: usize, format: PixelFormat) -> [u8; 3] {
let channels = format.channels();
let at = index * channels;
let get = |offset: usize| u32::from(row.get(at + offset).copied().unwrap_or(0));
let blend = |value: u32, alpha: u32| ((value * alpha + 127) / 255) as u8;
match format {
PixelFormat::Gray8 => {
let value = get(0) as u8;
[value, value, value]
}
PixelFormat::GrayA8 => {
let value = blend(get(0), get(1));
[value, value, value]
}
PixelFormat::Rgba8 => {
let alpha = get(3);
[
blend(get(0), alpha),
blend(get(1), alpha),
blend(get(2), alpha),
]
}
_ => [get(0) as u8, get(1) as u8, get(2) as u8],
}
}
fn rgb_to_ycbcr([r, g, b]: [u8; 3]) -> [u8; 3] {
const HALF: i32 = 1 << 15;
let (r, g, b) = (i32::from(r), i32::from(g), i32::from(b));
let y = (19_595 * r + 38_470 * g + 7_471 * b + HALF) >> 16;
let cb = ((-11_056 * r - 21_712 * g + 32_768 * b + HALF) >> 16) + 128;
let cr = ((32_768 * r - 27_440 * g - 5_328 * b + HALF) >> 16) + 128;
[
y.clamp(0, 255) as u8,
cb.clamp(0, 255) as u8,
cr.clamp(0, 255) as u8,
]
}
impl State {
fn fill_planes(&mut self) {
let width = self.descriptor.width as usize;
let format = self.descriptor.pixel;
let row_bytes = self.descriptor.row_bytes();
let (h, v) = (self.factors.0 as usize, self.factors.1 as usize);
let stride = self.luma.stride;
let last = (self.band_rows.max(1) - 1) as usize;
for y in 0..self.luma.height {
let row = self
.band
.get(y.min(last) * row_bytes..)
.and_then(|rest| rest.get(..row_bytes))
.unwrap_or(&[]);
for x in 0..stride {
let rgb = pixel_rgb(row, x.min(width - 1), format);
let ycbcr = if self.grayscale {
[rgb[0], 128, 128]
} else {
rgb_to_ycbcr(rgb)
};
if let Some(slot) = self.luma.samples.get_mut(y * stride + x) {
*slot = ycbcr[0];
}
if self.grayscale {
continue;
}
if let Some(slot) = self.chroma_full.get_mut((y * stride + x) * 2..) {
if let Some(pair) = slot.get_mut(..2) {
pair.copy_from_slice(&[ycbcr[1], ycbcr[2]]);
}
}
}
}
if self.grayscale {
return;
}
let count = (h * v) as u32;
for cy in 0..self.chroma_blue.height {
for cx in 0..self.chroma_blue.stride {
let (mut blue, mut red) = (0_u32, 0_u32);
for dy in 0..v {
for dx in 0..h {
let y = (cy * v + dy).min(self.luma.height - 1);
let x = (cx * h + dx).min(stride - 1);
let at = (y * stride + x) * 2;
blue += u32::from(self.chroma_full.get(at).copied().unwrap_or(128));
red += u32::from(self.chroma_full.get(at + 1).copied().unwrap_or(128));
}
}
let at = cy * self.chroma_blue.stride + cx;
if let Some(slot) = self.chroma_blue.samples.get_mut(at) {
*slot = ((blue + count / 2) / count) as u8;
}
if let Some(slot) = self.chroma_red.samples.get_mut(at) {
*slot = ((red + count / 2) / count) as u8;
}
}
}
}
fn encode_band(&mut self) {
let (h, v) = self.factors;
let mut samples = [0_u8; 64];
let mut coefficients = [0_i64; 64];
let mut quantized = [0_i32; 64];
let [luma_predictor, blue_predictor, red_predictor] = &mut self.predictors;
for mcu in 0..self.mcus_per_line {
for block_y in 0..v {
for block_x in 0..h {
let x = ((mcu * h + block_x) * 8) as usize;
let y = (block_y * 8) as usize;
self.luma.block(x, y, &mut samples);
fdct::block(&samples, &mut coefficients);
fdct::quantize(&coefficients, &self.luma_quant, &mut quantized);
encode_block(
&mut self.writer,
&quantized,
&self.luma_dc,
&self.luma_ac,
luma_predictor,
);
}
}
if self.grayscale {
continue;
}
let x = (mcu * 8) as usize;
for (plane, predictor) in [
(&self.chroma_blue, &mut *blue_predictor),
(&self.chroma_red, &mut *red_predictor),
] {
plane.block(x, 0, &mut samples);
fdct::block(&samples, &mut coefficients);
fdct::quantize(&coefficients, &self.chroma_quant, &mut quantized);
encode_block(
&mut self.writer,
&quantized,
&self.chroma_dc,
&self.chroma_ac,
predictor,
);
}
}
}
}
fn encode_block(
writer: &mut BitWriter,
block: &[i32; 64],
dc: &HuffmanEncoder,
ac: &HuffmanEncoder,
predictor: &mut i32,
) {
let value = block.first().copied().unwrap_or(0);
let difference = value.wrapping_sub(*predictor);
*predictor = value;
let (size, bits) = magnitude(difference);
if let Some((code, length)) = dc.code(size as u8) {
writer.write(code, length);
}
writer.write(bits, size);
let mut run = 0_u32;
for index in 1..64 {
let value = ZIGZAG
.get(index)
.and_then(|&at| block.get(at))
.copied()
.unwrap_or(0);
if value == 0 {
run += 1;
continue;
}
while run >= 16 {
if let Some((code, length)) = ac.code(0xF0) {
writer.write(code, length);
}
run -= 16;
}
let (size, bits) = magnitude(value);
if let Some((code, length)) = ac.code(((run as u8) << 4) | size as u8) {
writer.write(code, length);
}
writer.write(bits, size);
run = 0;
}
if run > 0 {
if let Some((code, length)) = ac.code(0x00) {
writer.write(code, length);
}
}
}
fn write_marker(code: u8, sink: &mut dyn Sink) -> Result<()> {
sink.write_all(&[0xFF, code])
}
fn write_segment(code: u8, payload: &[u8], sink: &mut dyn Sink) -> Result<()> {
let Ok(length) = u16::try_from(payload.len() + 2) else {
return Err(PixelsError::unsupported(format!(
"a {code:#04x} segment of {} bytes does not fit a 16-bit length",
payload.len()
)));
};
sink.write_all(&[0xFF, code])?;
sink.write_all(&length.to_be_bytes())?;
sink.write_all(payload)
}
fn write_quant_table(slot: u8, steps: &[u16; 64], payload: &mut Vec<u8>) {
payload.push(slot & 0x0F);
for &position in &ZIGZAG {
payload.push(steps.get(position).copied().unwrap_or(1).clamp(1, 255) as u8);
}
}
fn write_huffman_table(
class: u8,
slot: u8,
counts: &[u8; 16],
values: &[u8],
payload: &mut Vec<u8>,
) {
payload.push(((class & 0x0F) << 4) | (slot & 0x0F));
payload.extend_from_slice(counts);
payload.extend_from_slice(values);
}
impl Encoder for JpegEncoder {
fn set_icc_profile(&mut self, profile: Option<&[u8]>) -> Result<()> {
if self.state.is_some() {
return Err(PixelsError::invalid_argument(
"profile",
"the ICC profile must be set before write_header",
));
}
if let Some(profile) = profile.filter(|p| p.len() > 255 * ICC_CHUNK) {
return Err(PixelsError::unsupported(format!(
"a {}-byte ICC profile needs more than 255 JPEG segments",
profile.len()
)));
}
self.icc = profile.map(<[u8]>::to_vec);
Ok(())
}
fn write_header(&mut self, desc: &ImageDescriptor, sink: &mut dyn Sink) -> Result<()> {
if self.state.is_some() {
return Err(PixelsError::invalid_argument(
"descriptor",
"write_header called more than once",
));
}
let channels = channels_of(desc.pixel)?;
if desc.width > u32::from(u16::MAX) || desc.height > u32::from(u16::MAX) {
return Err(PixelsError::unsupported(format!(
"JPEG dimensions are 16-bit; {}x{} does not fit",
desc.width, desc.height
)));
}
let grayscale = channels <= 2;
let factors = if grayscale {
(1, 1)
} else {
self.subsampling.factors()
};
let luma_quant = tables::scale_quant(&tables::LUMA_QUANT, self.quality);
let chroma_quant = tables::scale_quant(&tables::CHROMA_QUANT, self.quality);
let (h, v) = factors;
let band_height = v * 8;
let mcus_per_line = desc.width.div_ceil(h * 8);
let luma_stride = (mcus_per_line * h * 8) as usize;
let chroma_stride = (mcus_per_line * 8) as usize;
sink.write_all(&[0xFF, marker::SOI])?;
write_segment(
marker::APP0,
&[
b'J', b'F', b'I', b'F', 0, 1, 2, 0, 0, 1, 0, 1, 0, 0, ],
sink,
)?;
if let Some(profile) = &self.icc {
for segment in icc_segments(profile) {
write_segment(marker::APP2, &segment, sink)?;
}
}
let mut payload = Vec::new();
write_quant_table(0, &luma_quant, &mut payload);
if !grayscale {
write_quant_table(1, &chroma_quant, &mut payload);
}
write_segment(marker::DQT, &payload, sink)?;
let mut payload = vec![8];
payload.extend_from_slice(&(desc.height as u16).to_be_bytes());
payload.extend_from_slice(&(desc.width as u16).to_be_bytes());
if grayscale {
payload.push(1);
payload.extend_from_slice(&[1, 0x11, 0]);
} else {
payload.push(3);
payload.extend_from_slice(&[1, ((h as u8) << 4) | v as u8, 0]);
payload.extend_from_slice(&[2, 0x11, 1]);
payload.extend_from_slice(&[3, 0x11, 1]);
}
write_segment(marker::SOF0, &payload, sink)?;
let mut payload = Vec::new();
write_huffman_table(
0,
0,
&tables::LUMA_DC_COUNTS,
&tables::LUMA_DC_VALUES,
&mut payload,
);
write_huffman_table(
1,
0,
&tables::LUMA_AC_COUNTS,
&tables::LUMA_AC_VALUES,
&mut payload,
);
if !grayscale {
write_huffman_table(
0,
1,
&tables::CHROMA_DC_COUNTS,
&tables::CHROMA_DC_VALUES,
&mut payload,
);
write_huffman_table(
1,
1,
&tables::CHROMA_AC_COUNTS,
&tables::CHROMA_AC_VALUES,
&mut payload,
);
}
write_segment(marker::DHT, &payload, sink)?;
let mut payload = Vec::new();
if grayscale {
payload.push(1);
payload.extend_from_slice(&[1, 0x00]);
} else {
payload.push(3);
payload.extend_from_slice(&[1, 0x00, 2, 0x11, 3, 0x11]);
}
payload.extend_from_slice(&[0, 63, 0]);
write_segment(marker::SOS, &payload, sink)?;
self.state = Some(State {
descriptor: *desc,
grayscale,
factors,
luma_quant,
chroma_quant,
luma_dc: HuffmanEncoder::new(&tables::LUMA_DC_COUNTS, &tables::LUMA_DC_VALUES)?,
luma_ac: HuffmanEncoder::new(&tables::LUMA_AC_COUNTS, &tables::LUMA_AC_VALUES)?,
chroma_dc: HuffmanEncoder::new(&tables::CHROMA_DC_COUNTS, &tables::CHROMA_DC_VALUES)?,
chroma_ac: HuffmanEncoder::new(&tables::CHROMA_AC_COUNTS, &tables::CHROMA_AC_VALUES)?,
band: vec![0; desc.row_bytes() * band_height as usize],
band_rows: 0,
band_height,
mcus_per_line,
luma: Plane::new(luma_stride, band_height as usize),
chroma_full: if grayscale {
Vec::new()
} else {
vec![128; luma_stride * band_height as usize * 2]
},
chroma_blue: chroma_plane(grayscale, chroma_stride, band_height / v),
chroma_red: chroma_plane(grayscale, chroma_stride, band_height / v),
predictors: [0; 3],
writer: BitWriter::default(),
rows_written: 0,
});
Ok(())
}
fn write_row(&mut self, row: &[u8], sink: &mut dyn Sink) -> Result<()> {
let Some(state) = self.state.as_mut() else {
return Err(PixelsError::invalid_argument(
"row",
"write_row called before write_header",
));
};
let expected = state.descriptor.row_bytes();
if row.len() != expected {
return Err(PixelsError::invalid_argument(
"row",
format!("row is {} bytes, expected {expected}", row.len()),
));
}
if state.rows_written >= state.descriptor.height {
return Err(PixelsError::invalid_argument(
"row",
format!("more than {} rows written", state.descriptor.height),
));
}
let at = state.band_rows as usize * expected;
if let Some(slot) = state
.band
.get_mut(at..)
.and_then(|rest| rest.get_mut(..expected))
{
slot.copy_from_slice(row);
}
state.band_rows += 1;
state.rows_written += 1;
if state.band_rows == state.band_height {
state.fill_planes();
state.encode_band();
state.band_rows = 0;
sink.write_all(&state.writer.bytes)?;
state.writer.bytes.clear();
}
Ok(())
}
fn finish(&mut self, sink: &mut dyn Sink) -> Result<()> {
let Some(state) = self.state.as_mut() else {
return Err(PixelsError::invalid_argument(
"sink",
"finish called before write_header",
));
};
if state.rows_written < state.descriptor.height {
return Err(PixelsError::malformed(
"jpeg",
format!(
"{} of {} rows were written",
state.rows_written, state.descriptor.height
),
));
}
if state.band_rows > 0 {
state.fill_planes();
state.encode_band();
state.band_rows = 0;
}
state.writer.flush();
sink.write_all(&state.writer.bytes)?;
state.writer.bytes.clear();
write_marker(marker::EOI, sink)?;
sink.flush()
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::indexing_slicing,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
#[test]
fn magnitude_categories_mirror_the_decoder() {
for value in [-2047_i32, -255, -8, -3, -2, -1, 1, 2, 3, 8, 255, 2047] {
let (size, bits) = magnitude(value);
assert!(size <= 15, "{value}: size {size}");
let threshold = 1_i32 << (size - 1);
let raw = bits as i32;
let decoded = if raw < threshold {
raw - (1_i32 << size) + 1
} else {
raw
};
assert_eq!(decoded, value, "{value} round-tripped as {decoded}");
}
assert_eq!(magnitude(0), (0, 0));
}
#[test]
fn the_bit_writer_stuffs_ff_bytes() {
let mut writer = BitWriter::default();
writer.write(0xFF, 8);
assert_eq!(
writer.bytes,
vec![0xFF, 0x00],
"a literal FF must be stuffed"
);
let mut writer = BitWriter::default();
writer.write(0b1010, 4);
writer.write(0b0101, 4);
assert_eq!(writer.bytes, vec![0b1010_0101]);
}
#[test]
fn flushing_pads_with_one_bits() {
let mut writer = BitWriter::default();
writer.write(0b101, 3);
writer.flush();
assert_eq!(writer.bytes, vec![0b1011_1111]);
let mut writer = BitWriter::default();
writer.write(0xAB, 8);
writer.flush();
assert_eq!(writer.bytes, vec![0xAB]);
}
#[test]
fn colour_conversion_round_trips_through_the_decoder() {
for rgb in [
[0, 0, 0],
[255, 255, 255],
[128, 128, 128],
[255, 0, 0],
[0, 255, 0],
[0, 0, 255],
[37, 142, 201],
] {
let ycbcr = rgb_to_ycbcr(rgb);
if rgb[0] == rgb[1] && rgb[1] == rgb[2] {
assert_eq!(ycbcr[0], rgb[0], "grey should map to its own luma");
assert_eq!([ycbcr[1], ycbcr[2]], [128, 128], "grey has no chroma");
}
let back = crate::decoder::ycbcr_to_rgb(ycbcr);
for channel in 0..3 {
assert!(
back[channel].abs_diff(rgb[channel]) <= 2,
"{rgb:?} -> {ycbcr:?} -> {back:?}"
);
}
}
}
#[test]
fn unsupported_pixel_formats_are_refused_at_the_header() {
for format in [
PixelFormat::Gray16,
PixelFormat::Rgb16,
PixelFormat::Rgba16,
PixelFormat::RgbF32,
PixelFormat::RgbaF32,
] {
let descriptor = ImageDescriptor::new(8, 8, format).unwrap();
let mut sink = Vec::new();
let error = JpegEncoder::new()
.write_header(&descriptor, &mut sink)
.unwrap_err();
assert_eq!(
error.code(),
otf_pixels_core::ErrorCode::Unsupported,
"{format}"
);
assert!(sink.is_empty(), "{format}: bytes were written anyway");
}
}
#[test]
fn quality_selects_subsampling_but_can_be_overridden() {
assert_eq!(
JpegEncoder::with_quality(80).unwrap().subsampling(),
Subsampling::Both
);
assert_eq!(
JpegEncoder::with_quality(95).unwrap().subsampling(),
Subsampling::None
);
assert_eq!(
JpegEncoder::with_quality(95)
.unwrap()
.with_subsampling(Subsampling::Horizontal)
.subsampling(),
Subsampling::Horizontal
);
assert!(JpegEncoder::with_quality(0).is_err());
assert!(JpegEncoder::with_quality(101).is_err());
}
}