use flamberge_crypto::aes;
use crate::Result;
pub(super) fn decrypt_member(
user_key: &[u8; 16],
wrapped: &[u8],
contents: &[u8],
) -> Result<Vec<u8>> {
let page_key = aes::ecb_decrypt(user_key, wrapped)?;
let plain = aes::ecb_decrypt(&page_key, contents)?;
Ok(strip_cms_padding(&plain).to_vec())
}
fn strip_cms_padding(data: &[u8]) -> &[u8] {
let Some(&last) = data.last() else {
return data;
};
let n = last as usize;
if n == 0 || n > data.len() {
return data;
}
if (2..16).contains(&n) && data[data.len() - n] as usize != n {
return data;
}
&data[..data.len() - n]
}
#[derive(Debug, PartialEq, Eq)]
pub(super) enum CheckResult {
Passed,
Failed,
Unchecked,
}
pub(super) fn check(path: &str, contents: &[u8]) -> CheckResult {
let lower = path.to_ascii_lowercase();
if lower.ends_with(".xhtml")
|| lower.ends_with(".html")
|| lower.ends_with(".htm")
|| lower.ends_with(".xml")
{
if check_text(contents) {
CheckResult::Passed
} else {
CheckResult::Failed
}
} else if lower.ends_with(".jpg") || lower.ends_with(".jpeg") || lower.ends_with(".jpe") {
if contents.starts_with(&[0xFF, 0xD8, 0xFF]) {
CheckResult::Passed
} else {
CheckResult::Failed
}
} else {
CheckResult::Unchecked
}
}
fn check_text(contents: &[u8]) -> bool {
let (offset, stride) = if contents.starts_with(&[0xEF, 0xBB, 0xBF]) {
(3, 1) } else if contents.starts_with(&[0xFE, 0xFF]) {
(3, 2) } else if contents.starts_with(&[0xFF, 0xFE]) {
(2, 2) } else {
(0, 1) };
(0..5).all(|i| matches!(contents.get(offset + i * stride), Some(&b) if (32..=127).contains(&b)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn strips_full_padding_block() {
let mut data = vec![b'A'; 16];
data.extend([16u8; 16]);
assert_eq!(strip_cms_padding(&data), &[b'A'; 16]);
}
#[test]
fn strips_partial_and_single_byte_padding() {
assert_eq!(strip_cms_padding(b"hello\x03\x03\x03"), b"hello");
assert_eq!(strip_cms_padding(b"hello\x01"), b"hello");
}
#[test]
fn leaves_inconsistent_padding_untouched() {
assert_eq!(strip_cms_padding(b"helloX\x03"), b"helloX\x03");
}
#[test]
fn text_check_accepts_ascii_and_boms() {
assert_eq!(check("a.xhtml", b"<?xml version"), CheckResult::Passed);
assert_eq!(check("a.xhtml", b"\xef\xbb\xbf<html"), CheckResult::Passed);
}
#[test]
fn text_check_rejects_binary() {
assert_eq!(
check("a.xhtml", b"\x00\x01\x02\x03\x04\x05"),
CheckResult::Failed
);
}
#[test]
fn jpeg_check_and_unchecked_types() {
assert_eq!(check("i.jpg", b"\xff\xd8\xff\xe0"), CheckResult::Passed);
assert_eq!(check("i.jpg", b"not a jpeg"), CheckResult::Failed);
assert_eq!(check("f.otf", b"anything"), CheckResult::Unchecked);
}
#[test]
fn two_layer_round_trip() {
let user_key = [0x2bu8; 16];
let page_key = [0x77u8; 16];
let wrapped = aes::ecb_encrypt(&user_key, &page_key).unwrap();
let text = b"<?xml version=\"1.0\"?><html>hi</html>";
let mut padded = text.to_vec();
let pad = 16 - (padded.len() % 16);
padded.extend(std::iter::repeat_n(pad as u8, pad));
let ciphertext = aes::ecb_encrypt(&page_key, &padded).unwrap();
let plain = decrypt_member(&user_key, &wrapped, &ciphertext).unwrap();
assert_eq!(plain, text);
assert_eq!(check("c.xhtml", &plain), CheckResult::Passed);
}
}