use encoding_rs::{Encoding, UTF_8, UTF_16BE, UTF_16LE, WINDOWS_1252};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DocFamily {
Xml,
Json,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Repair {
ReencodedToUtf8,
StrippedControlChars,
EscapedNakedAmpersands,
}
impl Repair {
pub fn note(&self) -> &'static str {
match self {
Repair::ReencodedToUtf8 => "sanitation: reencoded-to-utf8",
Repair::StrippedControlChars => "sanitation: stripped-control-chars",
Repair::EscapedNakedAmpersands => "sanitation: escaped-naked-ampersands",
}
}
}
const UTF8_BOM: &[u8] = &[0xEF, 0xBB, 0xBF];
pub(super) fn find_sub(haystack: &[u8], needle: &[u8]) -> Option<usize> {
if needle.is_empty() || haystack.len() < needle.len() {
return None;
}
haystack.windows(needle.len()).position(|w| w == needle)
}
fn strip_utf8_bom(bytes: &[u8]) -> (&[u8], bool) {
match bytes.strip_prefix(UTF8_BOM) {
Some(rest) => (rest, true),
None => (bytes, false),
}
}
fn utf16_bom(bytes: &[u8]) -> Option<(&'static Encoding, &[u8])> {
if let Some(rest) = bytes.strip_prefix(&[0xFF, 0xFE][..]) {
return Some((UTF_16LE, rest));
}
if let Some(rest) = bytes.strip_prefix(&[0xFE, 0xFF][..]) {
return Some((UTF_16BE, rest));
}
None
}
pub fn detect_family(bytes: &[u8]) -> DocFamily {
let body = match utf16_bom(bytes) {
Some((_, rest)) => rest,
None => strip_utf8_bom(bytes).0,
};
for &b in body {
if b.is_ascii_whitespace() {
continue;
}
return if b == b'{' {
DocFamily::Json
} else {
DocFamily::Xml
};
}
DocFamily::Xml
}
pub fn refuse_internal_dtd(bytes: &[u8]) -> Result<(), String> {
const REFUSAL: &str = "internal DTD subset refused (entity-expansion guard)";
let body = strip_utf8_bom(bytes).0;
let mut i = 0;
while i < body.len() {
let rest = &body[i..];
if rest[0].is_ascii_whitespace() {
i += 1;
continue;
}
if rest[0] != b'<' {
i += 1;
continue;
}
if let Some(skip) = skip_delimited(rest, b"<!--", b"-->") {
i += skip;
continue;
}
if rest.starts_with(b"<!--") {
return Ok(()); }
if let Some(skip) = skip_delimited(rest, b"<?", b"?>") {
i += skip;
continue;
}
if rest.starts_with(b"<?") {
return Ok(());
}
if starts_with_ignore_ascii_case(rest, b"<!DOCTYPE") {
let mut j = "<!DOCTYPE".len();
let mut quote: Option<u8> = None;
while j < rest.len() {
let c = rest[j];
match quote {
Some(q) if c == q => quote = None,
Some(_) => {}
None => match c {
b'"' | b'\'' => quote = Some(c),
b'[' => return Err(REFUSAL.to_string()),
b'>' => break,
_ => {}
},
}
j += 1;
}
i += j + 1;
continue;
}
return Ok(());
}
Ok(())
}
fn starts_with_ignore_ascii_case(haystack: &[u8], needle: &[u8]) -> bool {
haystack.len() >= needle.len() && haystack[..needle.len()].eq_ignore_ascii_case(needle)
}
fn skip_delimited(rest: &[u8], open: &[u8], close: &[u8]) -> Option<usize> {
if !rest.starts_with(open) {
return None;
}
find_sub(&rest[open.len()..], close).map(|p| open.len() + p + close.len())
}
pub fn rung_reencode_utf8(input: &[u8]) -> (Vec<u8>, bool) {
if let Some((enc, rest)) = utf16_bom(input) {
let (text, _, _) = enc.decode(rest);
return (rewrite_decl_encoding(&text), true);
}
let (body, had_bom) = strip_utf8_bom(input);
if let Ok(text) = std::str::from_utf8(body) {
let declaration_misleads = xml_decl_encoding_label(body)
.and_then(|label| Encoding::for_label(&label))
.is_some_and(|enc| enc != UTF_8 && enc.decode_without_bom_handling(body).0 != text);
if declaration_misleads {
return (rewrite_decl_encoding(text), true);
}
return (body.to_vec(), had_bom);
}
let declared = xml_decl_encoding_label(body)
.and_then(|label| Encoding::for_label(&label))
.filter(|enc| *enc != UTF_8);
let enc = declared.unwrap_or(WINDOWS_1252);
let (text, _, _) = enc.decode(body);
(rewrite_decl_encoding(&text), true)
}
fn xml_decl_encoding_label(body: &[u8]) -> Option<Vec<u8>> {
if !body.starts_with(b"<?xml") {
return None;
}
let decl = &body[..find_sub(body, b"?>")?];
let mut j = find_sub(decl, b"encoding")? + "encoding".len();
while j < decl.len() && decl[j].is_ascii_whitespace() {
j += 1;
}
if decl.get(j) != Some(&b'=') {
return None;
}
j += 1;
while j < decl.len() && decl[j].is_ascii_whitespace() {
j += 1;
}
let quote = *decl.get(j)?;
if quote != b'"' && quote != b'\'' {
return None;
}
j += 1;
let start = j;
while j < decl.len() && decl[j] != quote {
j += 1;
}
if j >= decl.len() {
return None;
}
Some(decl[start..j].to_vec())
}
fn rewrite_decl_encoding(text: &str) -> Vec<u8> {
let unchanged = || text.as_bytes().to_vec();
if !text.as_bytes().starts_with(b"<?xml") {
return unchanged();
}
let Some(decl_end) = find_sub(text.as_bytes(), b"?>") else {
return unchanged();
};
let decl = &text[..decl_end];
let Some(pos) = decl.find("encoding") else {
return unchanged();
};
let after = &decl[pos + "encoding".len()..];
let eq = match after.char_indices().find(|(_, c)| !c.is_whitespace()) {
Some((i, '=')) => i,
_ => return unchanged(),
};
let rest = &after[eq + 1..];
let (q_at, quote) = match rest.char_indices().find(|(_, c)| !c.is_whitespace()) {
Some((i, c @ ('"' | '\''))) => (i, c),
_ => return unchanged(),
};
let val_start = pos + "encoding".len() + eq + 1 + q_at + 1;
let Some(len) = text[val_start..].find(quote) else {
return unchanged();
};
let mut out = String::with_capacity(text.len() + "UTF-8".len());
out.push_str(&text[..val_start]);
out.push_str("UTF-8");
out.push_str(&text[val_start + len..]);
out.into_bytes()
}
pub fn rung_strip_control_chars(input: &[u8]) -> (Vec<u8>, bool) {
let mut out = Vec::with_capacity(input.len());
let mut changed = false;
let mut i = 0;
while i < input.len() {
let b = input[i];
if b < 0x20 && !matches!(b, b'\t' | b'\n' | b'\r') {
changed = true;
i += 1;
continue;
}
if b == 0xEF
&& input.get(i + 1) == Some(&0xBF)
&& matches!(input.get(i + 2), Some(&0xBE) | Some(&0xBF))
{
changed = true;
i += 3;
continue;
}
out.push(b);
i += 1;
}
(out, changed)
}
pub fn rung_escape_naked_ampersands(input: &[u8]) -> (Vec<u8>, bool) {
let mut out = Vec::with_capacity(input.len());
let mut changed = false;
let mut i = 0;
'scan: while i < input.len() {
let rest = &input[i..];
for (open, close) in [
(&b"<![CDATA["[..], &b"]]>"[..]),
(&b"<!--"[..], &b"-->"[..]),
(&b"<?"[..], &b"?>"[..]),
] {
if rest.starts_with(open) {
let end = skip_delimited(rest, open, close).unwrap_or(rest.len());
out.extend_from_slice(&rest[..end]);
i += end;
continue 'scan;
}
}
if rest[0] == b'&' {
match valid_reference_len(rest) {
Some(n) => {
out.extend_from_slice(&rest[..n]);
i += n;
}
None => {
out.extend_from_slice(b"&");
i += 1;
changed = true;
}
}
continue;
}
out.push(rest[0]);
i += 1;
}
(out, changed)
}
const MAX_LEADING_ZEROS: usize = 16;
fn valid_reference_len(rest: &[u8]) -> Option<usize> {
debug_assert_eq!(rest[0], b'&');
let tail = &rest[1..];
for name in [&b"amp;"[..], b"lt;", b"gt;", b"apos;", b"quot;"] {
if tail.starts_with(name) {
return Some(1 + name.len());
}
}
let digits = tail.strip_prefix(b"#")?;
let (body, max, is_digit): (&[u8], usize, fn(&u8) -> bool) = match digits.strip_prefix(b"x") {
Some(hex) => (hex, 6, u8::is_ascii_hexdigit),
None => (digits, 7, u8::is_ascii_digit),
};
let zeros = body
.iter()
.take(MAX_LEADING_ZEROS + 1)
.take_while(|&&b| b == b'0')
.count();
if zeros > MAX_LEADING_ZEROS {
return None;
}
let n = body[zeros..]
.iter()
.take(max)
.take_while(|b| is_digit(b))
.count();
if zeros + n == 0 || body.get(zeros + n) != Some(&b';') {
return None;
}
Some(rest.len() - body.len() + zeros + n + 1)
}
#[cfg(test)]
type RungFn = fn(&[u8]) -> (Vec<u8>, bool);
#[cfg(test)]
pub(crate) const RUNGS_FOR_TEST: [(&str, RungFn); 3] = [
("reencode_utf8", rung_reencode_utf8),
("strip_control_chars", rung_strip_control_chars),
("escape_naked_ampersands", rung_escape_naked_ampersands),
];
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn every_rung_is_a_byte_level_noop_on_wellformed_input() {
let wellformed: &[&str] = &[
r#"<?xml version="1.0" encoding="UTF-8"?><rss version="2.0"><channel><title>t & u</title></channel></rss>"#,
r#"<rss version="2.0"><channel><description><![CDATA[a & b && c]]></description></channel></rss>"#,
r#"<feed xmlns="http://www.w3.org/2005/Atom"><title>© — <ok></title></feed>"#,
r#"<feed xmlns="http://www.w3.org/2005/Atom"><title>©</title></feed>"#,
r#"<feed xmlns="http://www.w3.org/2005/Atom"><title>—</title></feed>"#,
"<!-- a & naked amp in a comment --><rss version=\"2.0\"/>",
"<?pi with & inside?><rss version=\"2.0\"/>",
r#"<x a=''"'>' "</x>"#,
r#"<?xml version="1.0" encoding="us-ascii"?><rss version="2.0"><channel><title>t</title></channel></rss>"#,
r#"<?xml version="1.0" encoding="ISO-8859-1"?><rss version="2.0"><channel><title>t</title></channel></rss>"#,
r#"<?xml version="1.0" encoding="windows-1252"?><rss version="2.0"><channel><title>t</title></channel></rss>"#,
r#"<?xml version="1.0" encoding="ISO-2022-JP"?><rss version="2.0"><channel><title>t</title></channel></rss>"#,
];
for doc in wellformed {
let b = doc.as_bytes();
for (name, rung) in RUNGS_FOR_TEST {
let (out, changed) = rung(b);
assert!(!changed, "rung {name} changed well-formed doc: {doc}");
assert_eq!(out, b, "rung {name} output differs on: {doc}");
}
}
}
#[test]
fn naked_and_undefined_ampersands_are_escaped_defined_ones_kept() {
let input = br#"<x a="M N">Fish & Chips & more ©</x>"#;
let expect = br#"<x a="M &nbsp; N">Fish & Chips & more ©</x>"#;
let (out, changed) = rung_escape_naked_ampersands(input);
assert!(changed);
assert_eq!(out, expect);
}
#[test]
fn references_with_more_leading_zeros_than_the_cap_are_escaped() {
let input = format!("<x>&#{}169;</x>", "0".repeat(MAX_LEADING_ZEROS + 1));
let (out, changed) = rung_escape_naked_ampersands(input.as_bytes());
assert!(changed);
assert!(String::from_utf8(out).unwrap().contains("&#"));
}
#[test]
fn trailing_ampersand_at_end_of_input_is_escaped_without_panicking() {
let input = b"<x>a &";
let (out, changed) = rung_escape_naked_ampersands(input);
assert!(changed);
assert_eq!(out, b"<x>a &");
}
#[test]
fn latin1_bytes_reencode_to_utf8() {
let mut doc = br#"<?xml version="1.0" encoding="iso-8859-1"?><x>caf"#.to_vec();
doc.push(0xE9);
doc.extend_from_slice(b"</x>");
let (out, changed) = rung_reencode_utf8(&doc);
assert!(changed);
let s = std::str::from_utf8(&out).unwrap();
assert!(s.contains("café"));
assert!(
!s.contains("iso-8859-1"),
"decl encoding token rewritten: {s}"
);
}
#[test]
fn lying_non_utf8_decl_over_valid_utf8_bytes_is_corrected() {
let doc = "<?xml version=\"1.0\" encoding=\"iso-8859-1\"?><x>café</x>".as_bytes();
let (out, changed) = rung_reencode_utf8(doc);
assert!(changed);
let s = std::str::from_utf8(&out).unwrap();
assert!(s.contains("café"), "{s}");
assert!(
!s.contains("iso-8859-1"),
"decl encoding token rewritten: {s}"
);
}
#[test]
fn legacy_decl_over_ascii_only_content_is_a_noop() {
let ascii = r#"<?xml version="1.0" encoding="ISO-8859-1"?><x>plain ascii</x>"#.as_bytes();
let (out, changed) = rung_reencode_utf8(ascii);
assert!(!changed, "ASCII under a legacy label needs no repair");
assert_eq!(out, ascii);
}
#[test]
fn ascii_under_a_utf16_decl_is_still_repaired() {
let doc = r#"<?xml version="1.0" encoding="utf-16"?><x>plain ascii</x>"#.as_bytes();
let (out, changed) = rung_reencode_utf8(doc);
assert!(changed);
assert_eq!(
out,
r#"<?xml version="1.0" encoding="UTF-8"?><x>plain ascii</x>"#.as_bytes()
);
}
#[test]
fn iso_2022_jp_carrying_real_escape_sequences_is_still_repaired() {
let mut doc = br#"<?xml version="1.0" encoding="ISO-2022-JP"?><x>"#.to_vec();
doc.extend_from_slice(b"\x1b$B$3$s$K$A$O\x1b(B");
doc.extend_from_slice(b"</x>");
let (out, changed) = rung_reencode_utf8(&doc);
assert!(changed, "escape-sequence body decodes to different text");
let s = std::str::from_utf8(&out).unwrap();
assert!(
!s.contains("ISO-2022-JP") && s.contains("UTF-8"),
"decl encoding token rewritten: {s}"
);
assert!(
s.contains('\x1b'),
"body relabeled, not transcoded, so ESC survives: {s:?}"
);
assert!(
!s.contains("こんにちは"),
"body is not transcoded here: {s:?}"
);
}
#[test]
fn utf8_decl_over_utf8_bytes_stays_a_noop() {
let doc = "<?xml version=\"1.0\" encoding=\"UTF-8\"?><x>café</x>".as_bytes();
let (out, changed) = rung_reencode_utf8(doc);
assert!(!changed);
assert_eq!(out, doc);
let no_decl = "<x>café</x>".as_bytes();
let (out, changed) = rung_reencode_utf8(no_decl);
assert!(!changed);
assert_eq!(out, no_decl);
}
#[test]
fn bom_is_stripped() {
let mut doc = vec![0xEF, 0xBB, 0xBF];
doc.extend_from_slice(br#"<rss version="2.0"/>"#);
let (out, changed) = rung_reencode_utf8(&doc);
assert!(changed, "a leading BOM is a repair");
assert_eq!(out, br#"<rss version="2.0"/>"#);
}
#[test]
fn lying_utf8_decl_over_latin1_bytes_is_sniffed() {
let mut doc = br#"<?xml version="1.0" encoding="utf-8"?><x>caf"#.to_vec();
doc.push(0xE9);
doc.extend_from_slice(b"</x>");
assert!(
std::str::from_utf8(&doc).is_err(),
"fixture must be invalid UTF-8"
);
let (out, changed) = rung_reencode_utf8(&doc);
assert!(changed);
let s = std::str::from_utf8(&out).expect("output is valid UTF-8");
assert!(
s.contains("café"),
"sniffed transcode recovered the text: {s}"
);
}
#[test]
fn control_chars_stripped_tab_lf_cr_kept() {
let mut doc = b"<x>a".to_vec();
doc.push(0x08); doc.extend_from_slice(b"b\t c\n d\r e</x>");
let (out, changed) = rung_strip_control_chars(&doc);
assert!(changed);
assert_eq!(out, b"<x>ab\t c\n d\r e</x>");
assert!(!out.contains(&0x08), "0x08 removed");
let legal = b"<x>a\tb\nc\rd</x>";
let (out, changed) = rung_strip_control_chars(legal);
assert!(!changed, "tab/LF/CR are legal XML 1.0 characters");
assert_eq!(out, legal);
}
#[test]
fn u_fffe_and_u_ffff_are_stripped() {
let mut doc = b"<x>a".to_vec();
doc.extend_from_slice("\u{FFFE}".as_bytes());
doc.extend_from_slice(b"b");
doc.extend_from_slice("\u{FFFF}".as_bytes());
doc.extend_from_slice(b"c</x>");
let (out, changed) = rung_strip_control_chars(&doc);
assert!(changed);
assert_eq!(out, b"<x>abc</x>");
}
#[test]
fn internal_dtd_subset_refused() {
let doc = br#"<?xml version="1.0"?><!DOCTYPE lolz [ <!ENTITY lol "lol"> <!ENTITY lol2 "&lol;&lol;"> ]><lolz>&lol2;</lolz>"#;
let err = refuse_internal_dtd(doc).expect_err("billion-laughs prolog must be refused");
assert!(
err.contains("internal DTD subset refused"),
"error names the guard, got: {err}"
);
assert!(
err.contains("entity-expansion guard"),
"error names the class, got: {err}"
);
}
#[test]
fn plain_doctype_without_subset_not_refused() {
refuse_internal_dtd(b"<!DOCTYPE opml><opml version=\"2.0\"/>").unwrap();
}
#[test]
fn json_family_detected() {
assert_eq!(
detect_family(br#"{"version": "https://jsonfeed.org/version/1.1"}"#),
DocFamily::Json
);
let mut doc = vec![0xEF, 0xBB, 0xBF];
doc.extend_from_slice(b" \r\n\t {\"version\": \"1.1\"}");
assert_eq!(detect_family(&doc), DocFamily::Json);
assert_eq!(detect_family(br#"<rss version="2.0"/>"#), DocFamily::Xml);
let mut xml = vec![0xEF, 0xBB, 0xBF];
xml.extend_from_slice(b"\n <?xml version=\"1.0\"?><feed/>");
assert_eq!(detect_family(&xml), DocFamily::Xml);
}
#[test]
fn cdata_and_comment_regions_pass_untouched_even_with_naked_amps() {
for doc in [
&br#"<x><![CDATA[Tom & Jerry && co]]></x>"#[..],
&b"<x><!-- Tom & Jerry --></x>"[..],
&b"<x><?php echo $a & $b; ?></x>"[..],
] {
let (out, changed) = rung_escape_naked_ampersands(doc);
assert!(
!changed,
"CDATA/comment/PI region must pass untouched: {}",
String::from_utf8_lossy(doc)
);
assert_eq!(out, doc);
}
}
#[test]
fn repair_notes_are_the_contract_strings() {
assert_eq!(
Repair::ReencodedToUtf8.note(),
"sanitation: reencoded-to-utf8"
);
assert_eq!(
Repair::StrippedControlChars.note(),
"sanitation: stripped-control-chars"
);
assert_eq!(
Repair::EscapedNakedAmpersands.note(),
"sanitation: escaped-naked-ampersands"
);
}
}