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,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Frame {
pub width: u32,
pub height: u32,
pub channels: Channels,
pub depth: Depth,
pub little_endian: bool,
}
impl Frame {
pub fn byte_len(&self) -> u64 {
u64::from(self.width) * u64::from(self.height) * self.channels.count() * self.depth.bytes()
}
}
pub trait Codec: Sync {
fn id(&self) -> &'static str;
fn describe(&self) -> String {
self.id().to_string()
}
fn encode(&self, src: &[u8], frame: Frame, out: &mut dyn Write) -> Result<u64>;
fn decode(&self, data: &[u8], frame: Frame, 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
}
}
pub fn libjxl_version() -> String {
let v = unsafe { jpegxl_sys::encoder::encode::JxlEncoderVersion() };
format!("{}.{}.{}", v / 1_000_000, (v / 1_000) % 1_000, v % 1_000)
}
impl Codec for Jxl {
fn id(&self) -> &'static str {
"jxl"
}
fn describe(&self) -> String {
format!(
"libjxl {} effort {}",
libjxl_version(),
self.effort.clamp(1, 10)
)
}
fn encode(&self, src: &[u8], frame: Frame, out: &mut dyn Write) -> Result<u64> {
let Frame {
width,
height,
channels: ch,
depth,
little_endian,
} = frame;
let expected = frame.byte_len();
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 ef = jpegxl_rs::encode::EncoderFrame::new(src)
.num_channels(nch)
.endianness(endianness(little_endian));
enc.encode_frame::<u8, u8>(&ef, width, height)
.map_err(|e| Error::Backend(e.to_string()))?
.data
}
Depth::Sixteen => {
let samples = as_u16(src);
let ef = jpegxl_rs::encode::EncoderFrame::new(samples.as_ref())
.num_channels(nch)
.endianness(endianness(little_endian));
enc.encode_frame::<u16, u16>(&ef, 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], frame: Frame, out: &mut dyn Write) -> Result<u64> {
let Frame {
depth,
little_endian,
..
} = frame;
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 = frame.byte_len();
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 gray16(w: u32, h: u32, le: bool) -> Frame {
Frame {
width: w,
height: h,
channels: Channels::Gray,
depth: Depth::Sixteen,
little_endian: le,
}
}
fn round_trip(w: u32, h: u32, ch: Channels, depth: Depth, le: bool) {
let c = Jxl { effort: 1 };
let f = Frame {
width: w,
height: h,
channels: ch,
depth,
little_endian: le,
};
let src = noise(f.byte_len() as usize, 99);
let mut enc = Vec::new();
c.encode(&src, f, &mut enc).unwrap();
let mut dec = Vec::new();
c.decode(&enc, f, &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], gray16(4, 4, 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 f = gray16(16, 16, le);
let mut enc = Vec::new();
c.encode(&src, f, &mut enc).unwrap();
let mut out = Vec::new();
c.decode(&enc, f, &mut out).unwrap();
assert_eq!(src, out, "le={le}");
}
}
}
#[cfg(test)]
mod version_tests {
#[test]
fn reports_the_linked_libjxl_version() {
let v = super::libjxl_version();
assert!(
v.starts_with("0.") || v.starts_with("1."),
"odd version {v}"
);
assert_eq!(
v.matches('.').count(),
2,
"expected major.minor.patch, got {v}"
);
use super::Codec;
assert!(super::Jxl::default().describe().contains(&v));
}
}