use alloc::borrow::ToOwned;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec::Vec;
use crate::SizeHandling;
const DEFAULT_BYTES: usize = 1;
fn default_bytes() -> usize {
DEFAULT_BYTES
}
const SWAP_BYTES_NAMES: &[&str] = &["swap_bytes", "swap-bytes", "swapbytes"];
const DEINTERLEAVE_NAMES: &[&str] = &["deinterleave", "de_interleave", "de-interleave", "deint"];
pub const TRANSFORM_LIST_SEPARATOR: &str = "+";
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct Transformed {
pub data: Vec<u8>,
pub used_size_handling: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub enum Transform {
#[serde(alias = "swap-bytes", alias = "swapbytes")]
SwapBytes,
#[serde(alias = "de_interleave", alias = "de-interleave", alias = "deint")]
Deinterleave {
offset: usize,
stride: usize,
#[serde(default = "default_bytes", alias = "unit")]
bytes: usize,
},
}
impl Transform {
pub fn supported_forms() -> &'static [&'static str] {
&["swap_bytes", "deinterleave:<offset>/<stride>[/<bytes>]"]
}
pub fn validate(&self) -> Result<(), TransformError> {
match self {
Transform::SwapBytes => Ok(()),
Transform::Deinterleave {
offset,
stride,
bytes,
} => {
if *bytes == 0 {
return Err(TransformError::InvalidBytes { bytes: *bytes });
}
if *stride < 2 {
return Err(TransformError::InvalidStride { stride: *stride });
}
if offset >= stride {
return Err(TransformError::InvalidOffset {
offset: *offset,
stride: *stride,
});
}
bytes
.checked_mul(*stride)
.ok_or(TransformError::GroupOverflow {
bytes: *bytes,
stride: *stride,
})?;
Ok(())
}
}
}
pub fn apply(
&self,
data: &[u8],
size_handling: &SizeHandling,
blank_byte: u8,
) -> Result<Transformed, TransformError> {
self.validate()?;
match self {
Transform::SwapBytes => {
let (data, used_size_handling) = swap_bytes(data, size_handling, blank_byte)?;
Ok(Transformed {
data,
used_size_handling,
})
}
Transform::Deinterleave {
offset,
stride,
bytes,
} => Ok(Transformed {
data: deinterleave(data, *offset, *stride, *bytes)?,
used_size_handling: false,
}),
}
}
pub fn try_from_str(s: &str) -> Result<Self, TransformError> {
let (name, params) = match s.split_once(':') {
Some((name, params)) => (name.trim(), Some(params.trim())),
None => (s.trim(), None),
};
match name {
s if SWAP_BYTES_NAMES.contains(&s) => match params {
None => Ok(Transform::SwapBytes),
Some(_) => Err(TransformError::UnexpectedParameters {
name: name.to_owned(),
}),
},
s if DEINTERLEAVE_NAMES.contains(&s) => match params {
Some(params) => parse_deinterleave(params),
None => Err(TransformError::MissingParameters {
name: name.to_owned(),
}),
},
_ => Err(TransformError::UnknownTransform {
name: name.to_owned(),
}),
}
}
}
impl core::fmt::Display for Transform {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Transform::SwapBytes => write!(f, "swap_bytes"),
Transform::Deinterleave {
offset,
stride,
bytes,
} => {
if *bytes == DEFAULT_BYTES {
write!(f, "deinterleave:{offset}/{stride}")
} else {
write!(f, "deinterleave:{offset}/{stride}/{bytes}")
}
}
}
}
}
pub fn parse_transform_list(s: &str) -> Result<Vec<Transform>, TransformError> {
if s.trim().is_empty() {
return Ok(Vec::new());
}
s.split(TRANSFORM_LIST_SEPARATOR)
.map(Transform::try_from_str)
.collect()
}
pub fn format_transform_list(transforms: &[Transform]) -> String {
transforms
.iter()
.map(|t| t.to_string())
.collect::<Vec<_>>()
.join(TRANSFORM_LIST_SEPARATOR)
}
pub fn apply_transforms(
data: &[u8],
transforms: &[Transform],
size_handling: &SizeHandling,
blank_byte: u8,
) -> Result<Transformed, TransformError> {
let mut current = Transformed {
data: data.to_vec(),
used_size_handling: false,
};
for transform in transforms {
let next = transform.apply(¤t.data, size_handling, blank_byte)?;
current = Transformed {
data: next.data,
used_size_handling: current.used_size_handling || next.used_size_handling,
};
}
Ok(current)
}
fn swap_bytes(
data: &[u8],
size_handling: &SizeHandling,
blank_byte: u8,
) -> Result<(Vec<u8>, bool), TransformError> {
let padded: Vec<u8>;
let mut used_size_handling = false;
let data: &[u8] = if data.len().is_multiple_of(2) {
data
} else {
used_size_handling = true;
match size_handling {
SizeHandling::Pad => {
let mut v = data.to_vec();
v.push(blank_byte);
padded = v;
&padded
}
SizeHandling::Truncate => &data[..data.len() - 1],
SizeHandling::None | SizeHandling::Duplicate => {
return Err(TransformError::OddLength { len: data.len() });
}
}
};
Ok((
data.chunks_exact(2).flat_map(|w| [w[1], w[0]]).collect(),
used_size_handling,
))
}
fn deinterleave(
data: &[u8],
offset: usize,
stride: usize,
bytes: usize,
) -> Result<Vec<u8>, TransformError> {
let group = bytes * stride;
if !data.len().is_multiple_of(group) {
return Err(TransformError::Ragged {
len: data.len(),
group,
});
}
let mut out = Vec::with_capacity(data.len() / stride);
let mut pos = offset * bytes;
while pos < data.len() {
out.extend_from_slice(&data[pos..pos + bytes]);
pos += group;
}
Ok(out)
}
fn parse_deinterleave(params: &str) -> Result<Transform, TransformError> {
let bad = |reason: &str| TransformError::BadParameters {
params: params.to_owned(),
reason: reason.to_string(),
};
let mut parts = params.split('/');
let offset = parts.next().unwrap_or_default().trim();
let stride = parts
.next()
.ok_or_else(|| bad("expected <offset>/<stride>"))?;
let bytes = parts.next();
if parts.next().is_some() {
return Err(bad("too many values, expected <offset>/<stride>[/<bytes>]"));
}
let number = |value: &str, what: &str| -> Result<usize, TransformError> {
value
.trim()
.parse::<usize>()
.map_err(|_| bad(&format!("{what} '{}' is not a number", value.trim())))
};
let transform = Transform::Deinterleave {
offset: number(offset, "offset")?,
stride: number(stride, "stride")?,
bytes: match bytes {
Some(bytes) => number(bytes, "bytes")?,
None => DEFAULT_BYTES,
},
};
transform.validate()?;
Ok(transform)
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[non_exhaustive]
pub enum TransformError {
UnknownTransform { name: String },
UnexpectedParameters { name: String },
MissingParameters { name: String },
BadParameters { params: String, reason: String },
InvalidStride { stride: usize },
InvalidOffset { offset: usize, stride: usize },
InvalidBytes { bytes: usize },
GroupOverflow { bytes: usize, stride: usize },
OddLength { len: usize },
Ragged { len: usize, group: usize },
}
impl core::fmt::Display for TransformError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
TransformError::UnknownTransform { name } => write!(
f,
"unknown transform '{name}'\n Supported: {}",
Transform::supported_forms().join(", ")
),
TransformError::UnexpectedParameters { name } => {
write!(f, "transform '{name}' does not take parameters")
}
TransformError::MissingParameters { name } => write!(
f,
"transform '{name}' requires parameters, e.g. 'deinterleave:1/2'"
),
TransformError::BadParameters { params, reason } => {
write!(f, "invalid transform parameters '{params}': {reason}")
}
TransformError::InvalidStride { stride } => {
write!(f, "deinterleave stride must be at least 2, got {stride}")
}
TransformError::InvalidOffset { offset, stride } => write!(
f,
"deinterleave offset must be less than the stride, got offset {offset} with stride {stride}"
),
TransformError::InvalidBytes { bytes } => {
write!(
f,
"deinterleave lane width must be at least 1 byte, got {bytes}"
)
}
TransformError::GroupOverflow { bytes, stride } => write!(
f,
"deinterleave lane width {bytes} multiplied by stride {stride} overflows"
),
TransformError::OddLength { len } => write!(
f,
"swap_bytes requires an even-length image, got {len} bytes.\n Use size_handling pad or truncate to resolve the odd byte."
),
TransformError::Ragged { len, group } => write!(
f,
"deinterleave requires the image length to be a multiple of {group} bytes (lane width x stride), got {len} bytes"
),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
const W32: [u8; 8] = [0, 1, 2, 3, 4, 5, 6, 7];
fn deint(data: &[u8], offset: usize, stride: usize, bytes: usize) -> Vec<u8> {
Transform::Deinterleave {
offset,
stride,
bytes,
}
.apply(data, &SizeHandling::None, PAD)
.unwrap()
.data
}
const PAD: u8 = 0xAA;
fn swap(data: &[u8]) -> Vec<u8> {
Transform::SwapBytes
.apply(data, &SizeHandling::None, PAD)
.unwrap()
.data
}
#[test]
fn swap_bytes_reverses_each_word() {
assert_eq!(
swap(&[0x12, 0x34, 0x56, 0x78]),
vec![0x34, 0x12, 0x78, 0x56]
);
}
#[test]
fn swap_bytes_is_its_own_inverse() {
assert_eq!(swap(&swap(&W32)), W32.to_vec());
}
#[test]
fn swap_bytes_accepts_an_empty_image() {
assert_eq!(swap(&[]), Vec::<u8>::new());
}
#[test]
fn swap_bytes_rejects_odd_length_without_size_handling() {
for handling in [SizeHandling::None, SizeHandling::Duplicate] {
assert_eq!(
Transform::SwapBytes.apply(&[1, 2, 3], &handling, PAD),
Err(TransformError::OddLength { len: 3 })
);
}
}
#[test]
fn swap_bytes_pads_odd_length_with_the_blank_byte() {
assert_eq!(
Transform::SwapBytes
.apply(&[1, 2, 3], &SizeHandling::Pad, 0xFF)
.unwrap()
.data,
vec![2, 1, 0xFF, 3]
);
assert_eq!(
Transform::SwapBytes
.apply(&[1, 2, 3], &SizeHandling::Pad, PAD)
.unwrap()
.data,
vec![2, 1, PAD, 3]
);
}
#[test]
fn swap_bytes_truncates_odd_length_when_asked() {
assert_eq!(
Transform::SwapBytes
.apply(&[1, 2, 3], &SizeHandling::Truncate, PAD)
.unwrap()
.data,
vec![2, 1]
);
}
#[test]
fn deinterleave_16_bit_source_to_bytes() {
assert_eq!(deint(&W32, 0, 2, 1), vec![0, 2, 4, 6]);
assert_eq!(deint(&W32, 1, 2, 1), vec![1, 3, 5, 7]);
}
#[test]
fn deinterleave_32_bit_source_to_bytes() {
assert_eq!(deint(&W32, 0, 4, 1), vec![0, 4]);
assert_eq!(deint(&W32, 1, 4, 1), vec![1, 5]);
assert_eq!(deint(&W32, 2, 4, 1), vec![2, 6]);
assert_eq!(deint(&W32, 3, 4, 1), vec![3, 7]);
}
#[test]
fn deinterleave_32_bit_source_to_16_bit_halves() {
assert_eq!(deint(&W32, 0, 2, 2), vec![0, 1, 4, 5]);
assert_eq!(deint(&W32, 1, 2, 2), vec![2, 3, 6, 7]);
}
#[test]
fn deinterleave_unit_selects_whole_groups_of_bytes() {
let expected: Vec<u8> = W32
.iter()
.enumerate()
.filter(|(i, _)| i % 4 < 2)
.map(|(_, b)| *b)
.collect();
assert_eq!(deint(&W32, 0, 2, 2), expected);
}
#[test]
fn deinterleave_extracts_a_16_bit_lane_from_a_64_bit_image() {
let data: Vec<u8> = (0..16).collect();
assert_eq!(deint(&data, 2, 4, 2), vec![4, 5, 12, 13]);
}
#[test]
fn deinterleave_rejects_a_ragged_image() {
assert_eq!(
Transform::Deinterleave {
offset: 0,
stride: 4,
bytes: 1
}
.apply(&[0, 1, 2, 3, 4], &SizeHandling::None, PAD),
Err(TransformError::Ragged { len: 5, group: 4 })
);
}
#[test]
fn deinterleave_rejects_invalid_parameters() {
let cases = [
(0, 1, 1, TransformError::InvalidStride { stride: 1 }),
(0, 0, 1, TransformError::InvalidStride { stride: 0 }),
(
2,
2,
1,
TransformError::InvalidOffset {
offset: 2,
stride: 2,
},
),
(0, 2, 0, TransformError::InvalidBytes { bytes: 0 }),
];
for (offset, stride, bytes, expected) in cases {
assert_eq!(
Transform::Deinterleave {
offset,
stride,
bytes
}
.apply(&W32, &SizeHandling::None, PAD),
Err(expected)
);
}
}
fn chain(spec: &str) -> Vec<u8> {
apply_transforms(&W32, &parsed(spec), &SizeHandling::None, PAD)
.unwrap()
.data
}
#[test]
fn byte_wise_deinterleave_does_not_commute_with_swap_bytes() {
let deint_then_swap = chain("deinterleave:0/2+swap_bytes");
let swap_then_deint = chain("swap_bytes+deinterleave:0/2");
assert_eq!(deint_then_swap, vec![2, 0, 6, 4]);
assert_eq!(swap_then_deint, vec![1, 3, 5, 7]);
assert_ne!(deint_then_swap, swap_then_deint);
}
#[test]
fn word_aligned_deinterleave_commutes_with_swap_bytes() {
assert_eq!(chain("deinterleave:1/2/2+swap_bytes"), vec![3, 2, 7, 6]);
assert_eq!(chain("swap_bytes+deinterleave:1/2/2"), vec![3, 2, 7, 6]);
}
#[test]
fn empty_transform_list_leaves_the_image_alone() {
assert_eq!(
apply_transforms(&W32, &[], &SizeHandling::None, PAD)
.unwrap()
.data,
W32.to_vec()
);
}
#[test]
fn later_transforms_see_the_earlier_output() {
let data: Vec<u8> = (0..15).collect();
assert_eq!(
apply_transforms(
&data,
&parsed("deinterleave:0/3+swap_bytes"),
&SizeHandling::None,
PAD
),
Err(TransformError::OddLength { len: 5 })
);
}
fn parsed(s: &str) -> Vec<Transform> {
parse_transform_list(s).unwrap()
}
#[test]
fn parses_a_single_transform() {
assert_eq!(parsed("swap_bytes"), vec![Transform::SwapBytes]);
assert_eq!(
parsed("deinterleave:1/2"),
vec![Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 1
}]
);
assert_eq!(
parsed("deinterleave:1/2/2"),
vec![Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 2
}]
);
}
#[test]
fn parses_a_list_preserving_order() {
assert_eq!(
parsed("deinterleave:1/2/2+swap_bytes"),
vec![
Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 2
},
Transform::SwapBytes,
]
);
}
#[test]
fn parses_an_empty_list() {
assert!(parsed("").is_empty());
assert!(parsed(" ").is_empty());
}
#[test]
fn text_encoding_round_trips() {
for s in [
"swap_bytes",
"deinterleave:1/2",
"deinterleave:3/4/2",
"deinterleave:1/2/2+swap_bytes",
] {
assert_eq!(format_transform_list(&parsed(s)), s);
}
}
#[test]
fn default_unit_is_omitted_when_formatting() {
assert_eq!(
Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 1
}
.to_string(),
"deinterleave:1/2"
);
}
#[test]
fn rejects_bad_text() {
assert!(matches!(
Transform::try_from_str("nonsense"),
Err(TransformError::UnknownTransform { .. })
));
assert!(matches!(
Transform::try_from_str("swap_bytes:1/2"),
Err(TransformError::UnexpectedParameters { .. })
));
assert!(matches!(
Transform::try_from_str("deinterleave"),
Err(TransformError::MissingParameters { .. })
));
assert!(matches!(
Transform::try_from_str("deinterleave:1"),
Err(TransformError::BadParameters { .. })
));
assert!(matches!(
Transform::try_from_str("deinterleave:1/2/3/4"),
Err(TransformError::BadParameters { .. })
));
assert!(matches!(
Transform::try_from_str("deinterleave:x/2"),
Err(TransformError::BadParameters { .. })
));
assert!(matches!(
Transform::try_from_str("deinterleave:2/2"),
Err(TransformError::InvalidOffset { .. })
));
}
const SWAP_SPELLINGS: &[&str] = &["swap_bytes", "swap-bytes", "swapbytes"];
const DEINTERLEAVE_SPELLINGS: &[&str] =
&["deinterleave", "de_interleave", "de-interleave", "deint"];
#[test]
fn text_accepts_every_swap_bytes_spelling() {
for spelling in SWAP_SPELLINGS {
assert_eq!(
Transform::try_from_str(spelling),
Ok(Transform::SwapBytes),
"text spelling {spelling:?} not accepted"
);
}
}
#[test]
fn text_accepts_every_deinterleave_spelling() {
for spelling in DEINTERLEAVE_SPELLINGS {
assert_eq!(
Transform::try_from_str(&alloc::format!("{spelling}:1/2")),
Ok(Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 1
}),
"text spelling {spelling:?} not accepted"
);
}
}
#[test]
fn spellings_normalise_when_formatted() {
for spelling in DEINTERLEAVE_SPELLINGS {
let parsed = Transform::try_from_str(&alloc::format!("{spelling}:1/2")).unwrap();
assert_eq!(parsed.to_string(), "deinterleave:1/2");
}
for spelling in SWAP_SPELLINGS {
let parsed = Transform::try_from_str(spelling).unwrap();
assert_eq!(parsed.to_string(), "swap_bytes");
}
}
#[test]
fn serde_accepts_every_swap_bytes_spelling() {
for spelling in SWAP_SPELLINGS {
let json = alloc::format!("\"{spelling}\"");
assert_eq!(
serde_json::from_str::<Transform>(&json).ok(),
Some(Transform::SwapBytes),
"config spelling {spelling:?} not accepted"
);
}
}
#[test]
fn serde_accepts_every_deinterleave_spelling() {
for spelling in DEINTERLEAVE_SPELLINGS {
let json = alloc::format!("{{\"{spelling}\":{{\"offset\":1,\"stride\":2}}}}");
assert_eq!(
serde_json::from_str::<Transform>(&json).ok(),
Some(Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 1
}),
"config spelling {spelling:?} not accepted"
);
}
}
#[test]
fn parameterless_transform_serialises_as_a_bare_string() {
let json = serde_json::to_string(&Transform::SwapBytes).unwrap();
assert_eq!(json, "\"swap_bytes\"");
assert_eq!(
serde_json::from_str::<Transform>(&json).unwrap(),
Transform::SwapBytes
);
}
#[test]
fn parameterised_transform_serialises_as_a_map() {
let transform = Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 2,
};
let json = serde_json::to_string(&transform).unwrap();
assert_eq!(
json,
r#"{"deinterleave":{"offset":1,"stride":2,"bytes":2}}"#
);
assert_eq!(serde_json::from_str::<Transform>(&json).unwrap(), transform);
}
#[test]
fn deserialised_unit_defaults_to_one() {
let transform: Transform =
serde_json::from_str(r#"{"deinterleave":{"offset":1,"stride":2}}"#).unwrap();
assert_eq!(
transform,
Transform::Deinterleave {
offset: 1,
stride: 2,
bytes: 1
}
);
}
}