use std::io::Write;
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("io: {0}")]
Io(#[from] std::io::Error),
#[error("libjxl: {0}")]
Backend(String),
#[error("byte count mismatch: got {actual}, expected {expected}")]
SizeMismatch { expected: u64, actual: u64 },
}
pub type Result<T> = std::result::Result<T, Error>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Channels {
Gray,
Rgb,
}
impl Channels {
pub fn count(self) -> u64 {
match self {
Channels::Gray => 1,
Channels::Rgb => 3,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Depth {
Eight,
Sixteen,
}
impl Depth {
pub fn from_bits(bits: u16) -> Option<Self> {
match bits {
8 => Some(Depth::Eight),
16 => Some(Depth::Sixteen),
_ => None,
}
}
pub fn bytes(self) -> u64 {
match self {
Depth::Eight => 1,
Depth::Sixteen => 2,
}
}
}
pub trait Codec: Sync {
fn id(&self) -> &'static str;
fn encode(
&self,
src: &[u8],
width: u32,
height: u32,
ch: Channels,
depth: Depth,
little_endian: bool,
out: &mut dyn Write,
) -> Result<u64>;
fn decode(
&self,
data: &[u8],
width: u32,
height: u32,
ch: Channels,
depth: Depth,
little_endian: bool,
out: &mut dyn Write,
) -> Result<u64>;
}
#[derive(Debug, Clone)]
pub struct Jxl {
pub effort: u8,
}
impl Default for Jxl {
fn default() -> Self {
Self { effort: 4 }
}
}
impl Jxl {
fn speed(&self) -> jpegxl_rs::encode::EncoderSpeed {
use jpegxl_rs::encode::EncoderSpeed::*;
match self.effort.clamp(1, 10) {
1 => Lightning,
2 => Thunder,
3 => Falcon,
4 => Cheetah,
5 => Hare,
6 => Wombat,
7 => Squirrel,
8 => Kitten,
9 => Tortoise,
_ => Glacier,
}
}
}
fn as_u16(src: &[u8]) -> std::borrow::Cow<'_, [u16]> {
match bytemuck::try_cast_slice::<u8, u16>(src) {
Ok(s) => std::borrow::Cow::Borrowed(s),
Err(_) => std::borrow::Cow::Owned(
src.chunks_exact(2)
.map(|c| u16::from_ne_bytes([c[0], c[1]]))
.collect(),
),
}
}
fn endianness(little_endian: bool) -> jpegxl_rs::Endianness {
if little_endian {
jpegxl_rs::Endianness::Little
} else {
jpegxl_rs::Endianness::Big
}
}
impl Codec for Jxl {
fn id(&self) -> &'static str {
"jxl"
}
fn encode(
&self,
src: &[u8],
width: u32,
height: u32,
ch: Channels,
depth: Depth,
little_endian: bool,
out: &mut dyn Write,
) -> Result<u64> {
let expected = u64::from(width) * u64::from(height) * ch.count() * depth.bytes();
if src.len() as u64 != expected {
return Err(Error::SizeMismatch {
expected,
actual: src.len() as u64,
});
}
let color = match ch {
Channels::Gray => jpegxl_rs::encode::ColorEncoding::SrgbLuma,
Channels::Rgb => jpegxl_rs::encode::ColorEncoding::Srgb,
};
let runner = jpegxl_rs::ThreadsRunner::default();
let mut enc = jpegxl_rs::encoder_builder()
.parallel_runner(&runner)
.lossless(true)
.uses_original_profile(true)
.speed(self.speed())
.color_encoding(color)
.has_alpha(false)
.build()
.map_err(|e| Error::Backend(e.to_string()))?;
let nch = ch.count() as u32;
let encoded: Vec<u8> = match depth {
Depth::Eight => {
let frame = jpegxl_rs::encode::EncoderFrame::new(src)
.num_channels(nch)
.endianness(endianness(little_endian));
enc.encode_frame::<u8, u8>(&frame, width, height)
.map_err(|e| Error::Backend(e.to_string()))?
.data
}
Depth::Sixteen => {
let samples = as_u16(src);
let frame = jpegxl_rs::encode::EncoderFrame::new(samples.as_ref())
.num_channels(nch)
.endianness(endianness(little_endian));
enc.encode_frame::<u16, u16>(&frame, width, height)
.map_err(|e| Error::Backend(e.to_string()))?
.data
}
};
out.write_all(&encoded)?;
Ok(encoded.len() as u64)
}
fn decode(
&self,
data: &[u8],
width: u32,
height: u32,
ch: Channels,
depth: Depth,
little_endian: bool,
out: &mut dyn Write,
) -> Result<u64> {
let runner = jpegxl_rs::ThreadsRunner::default();
let dec = jpegxl_rs::decoder_builder()
.parallel_runner(&runner)
.build()
.map_err(|e| Error::Backend(e.to_string()))?;
let expected = u64::from(width) * u64::from(height) * ch.count() * depth.bytes();
let bytes: Vec<u8> = match depth {
Depth::Eight => {
dec.decode_with::<u8>(data)
.map_err(|e| Error::Backend(e.to_string()))?
.1
}
Depth::Sixteen => {
let (_, px) = dec
.decode_with::<u16>(data)
.map_err(|e| Error::Backend(e.to_string()))?;
let mut v = Vec::with_capacity(px.len() * 2);
if little_endian {
for s in &px {
v.extend_from_slice(&s.to_le_bytes());
}
} else {
for s in &px {
v.extend_from_slice(&s.to_be_bytes());
}
}
v
}
};
if bytes.len() as u64 != expected {
return Err(Error::SizeMismatch {
expected,
actual: bytes.len() as u64,
});
}
out.write_all(&bytes)?;
Ok(bytes.len() as u64)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn noise(n: usize, seed: u64) -> Vec<u8> {
let mut s = seed | 1;
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
(s >> 24) as u8
})
.collect()
}
fn round_trip(w: u32, h: u32, ch: Channels, depth: Depth, le: bool) {
let c = Jxl { effort: 1 };
let n = (u64::from(w) * u64::from(h) * ch.count() * depth.bytes()) as usize;
let src = noise(n, 99);
let mut enc = Vec::new();
c.encode(&src, w, h, ch, depth, le, &mut enc).unwrap();
let mut dec = Vec::new();
c.decode(&enc, w, h, ch, depth, le, &mut dec).unwrap();
assert_eq!(src, dec, "{w}x{h} {ch:?} {depth:?} le={le}");
}
#[test]
fn round_trips_losslessly() {
round_trip(64, 48, Channels::Gray, Depth::Sixteen, true);
round_trip(64, 48, Channels::Gray, Depth::Sixteen, false);
round_trip(40, 24, Channels::Rgb, Depth::Eight, true);
round_trip(33, 17, Channels::Rgb, Depth::Sixteen, true);
round_trip(33, 17, Channels::Rgb, Depth::Sixteen, false);
}
#[test]
fn handles_odd_dimensions() {
round_trip(1, 1, Channels::Gray, Depth::Sixteen, true);
round_trip(7, 3, Channels::Rgb, Depth::Eight, false);
}
#[test]
fn encode_rejects_wrong_byte_count() {
let c = Jxl::default();
assert!(matches!(
c.encode(
&[0u8; 10],
4,
4,
Channels::Gray,
Depth::Sixteen,
false,
&mut Vec::new()
),
Err(Error::SizeMismatch { .. })
));
}
#[test]
fn byte_order_is_preserved() {
let c = Jxl { effort: 1 };
let src = noise(16 * 16 * 2, 3);
for le in [true, false] {
let mut enc = Vec::new();
c.encode(&src, 16, 16, Channels::Gray, Depth::Sixteen, le, &mut enc)
.unwrap();
let mut out = Vec::new();
c.decode(&enc, 16, 16, Channels::Gray, Depth::Sixteen, le, &mut out)
.unwrap();
assert_eq!(src, out, "le={le}");
}
}
}