use std::borrow::Cow;
use std::path::{Path, PathBuf};
use crate::error::ConversionError;
use crate::format::InputFormat;
#[derive(Debug, Clone)]
pub struct SourceDocument {
pub name: String,
pub format: InputFormat,
pub bytes: Vec<u8>,
pub path: Option<PathBuf>,
pub base_url: Option<String>,
pub encoding: Option<String>,
}
impl SourceDocument {
pub fn from_file(path: impl AsRef<Path>) -> Result<Self, ConversionError> {
let path = path.as_ref();
let ext = path.extension().and_then(|e| e.to_str()).ok_or_else(|| {
ConversionError::UnknownFormat {
hint: format!("no extension on {}", path.display()),
}
})?;
let format =
InputFormat::from_extension(ext).ok_or_else(|| ConversionError::UnknownFormat {
hint: format!("unrecognized extension '.{ext}'"),
})?;
let bytes = std::fs::read(path)?;
let name = path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("document")
.to_string();
Ok(Self {
name,
format,
bytes,
path: Some(path.to_path_buf()),
base_url: None,
encoding: None,
})
}
pub fn from_bytes(name: impl Into<String>, format: InputFormat, bytes: Vec<u8>) -> Self {
Self {
name: name.into(),
format,
bytes,
path: None,
base_url: None,
encoding: None,
}
}
pub fn with_encoding(mut self, label: Option<String>) -> Self {
self.encoding = label;
self
}
pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
self.base_url = Some(url.into());
self
}
pub fn base_dir(&self) -> Option<&Path> {
self.path.as_deref().and_then(Path::parent)
}
pub fn text(&self) -> Result<Cow<'_, str>, ConversionError> {
match &self.encoding {
Some(label) => decode_text_as(&self.bytes, label),
None => decode_text(&self.bytes),
}
}
}
pub(crate) fn decode_text_as<'a>(
bytes: &'a [u8],
label: &str,
) -> Result<Cow<'a, str>, ConversionError> {
let encoding = lookup_encoding(label).ok_or_else(|| {
ConversionError::Parse(format!(
"unknown character encoding {label:?}: use a WHATWG encoding label such as \
utf-8, windows-1252, latin1, shift_jis, euc-jp, gbk, big5, euc-kr or koi8-r"
))
})?;
let (text, had_errors) = encoding.decode_with_bom_removal(bytes);
if had_errors {
return Err(ConversionError::Parse(format!(
"input is not valid {}: it cannot be decoded with the requested encoding {label:?}",
encoding.name()
)));
}
Ok(text)
}
fn lookup_encoding(label: &str) -> Option<&'static encoding_rs::Encoding> {
let label = label.trim();
let dashed = label.replace('_', "-").to_ascii_lowercase();
encoding_rs::Encoding::for_label(label.as_bytes())
.or_else(|| encoding_rs::Encoding::for_label(dashed.as_bytes()))
.or_else(|| {
let alias = match dashed.as_str() {
"latin-1" | "iso8859-1" | "iso-latin-1" => "latin1",
"utf-8-sig" | "utf8-sig" => "utf-8",
"cp932" | "ms932" | "sjis" => "shift_jis",
"mac-roman" | "macroman" => "macintosh",
"cp1361" | "johab" | "euc-tw" => return None,
_ => return None,
};
encoding_rs::Encoding::for_label(alias.as_bytes())
})
}
pub(crate) fn decode_text(bytes: &[u8]) -> Result<Cow<'_, str>, ConversionError> {
if let Some(rest) = bytes.strip_prefix(&[0xff, 0xfe, 0x00, 0x00]) {
return decode_utf32(rest, true);
}
if let Some(rest) = bytes.strip_prefix(&[0x00, 0x00, 0xfe, 0xff]) {
return decode_utf32(rest, false);
}
if bytes.starts_with(&[0xff, 0xfe]) || bytes.starts_with(&[0xfe, 0xff]) {
let (text, _, malformed) = encoding_rs::UTF_16LE.decode(bytes);
if malformed {
return Err(ConversionError::Parse(
"input carries a UTF-16 byte-order mark but is not valid UTF-16".into(),
));
}
return Ok(Cow::Owned(text.into_owned()));
}
match std::str::from_utf8(bytes) {
Ok(text) => Ok(Cow::Borrowed(text.strip_prefix('\u{feff}').unwrap_or(text))),
Err(utf8_error) => {
let coverage = utf8_high_byte_coverage(bytes);
if coverage >= 0.5 {
return Err(ConversionError::Parse(format!(
"input is not valid UTF-8 ({utf8_error}), but {:.0}% of its non-ASCII \
bytes still form well-formed UTF-8 sequences, so it reads as damaged \
UTF-8 rather than as another encoding; decoding it as windows-1252 \
would turn the damage into text that looks fine and is not",
coverage * 100.0
)));
}
if bytes
.iter()
.any(|b| matches!(b, 0x81 | 0x8d | 0x8f | 0x90 | 0x9d))
{
return Err(ConversionError::Parse(
"input is neither UTF-8 nor windows-1252, and its encoding is not \
declared, so it cannot be decoded reliably"
.into(),
));
}
eprintln!(
"warning: input is not UTF-8; decoded it as windows-1252 — text in any \
other single-byte or multi-byte encoding will be wrong"
);
let (text, _) = encoding_rs::WINDOWS_1252.decode_without_bom_handling(bytes);
Ok(Cow::Owned(text.into_owned()))
}
}
}
fn decode_utf32(rest: &[u8], little_endian: bool) -> Result<Cow<'static, str>, ConversionError> {
if !rest.len().is_multiple_of(4) {
return Err(ConversionError::Parse("truncated UTF-32 input".into()));
}
rest.chunks_exact(4)
.map(|c| {
let b = [c[0], c[1], c[2], c[3]];
let u = if little_endian {
u32::from_le_bytes(b)
} else {
u32::from_be_bytes(b)
};
char::from_u32(u)
.ok_or_else(|| ConversionError::Parse(format!("invalid UTF-32 code point {u:#x}")))
})
.collect::<Result<String, _>>()
.map(Cow::Owned)
}
fn utf8_high_byte_coverage(raw: &[u8]) -> f64 {
let (mut high, mut covered, mut i) = (0usize, 0usize, 0usize);
while i < raw.len() {
let lead = raw[i];
if lead < 0x80 {
i += 1;
continue;
}
let len = match lead {
0xc2..=0xdf => 2,
0xe0..=0xef => 3,
0xf0..=0xf4 => 4,
_ => 0,
};
let well_formed = len > 0
&& raw
.get(i..i + len)
.is_some_and(|chunk| std::str::from_utf8(chunk).is_ok());
if well_formed {
high += len;
covered += len;
i += len;
} else {
high += 1;
i += 1;
}
}
if high == 0 {
0.0
} else {
covered as f64 / high as f64
}
}
#[cfg(test)]
mod decode_tests {
use super::{decode_text, decode_text_as};
#[test]
fn explicit_encoding_decodes_strictly() {
let sjis = b"\x93\xfa\x96\x7b";
assert_eq!(decode_text_as(sjis, "shift_jis").unwrap(), "日本");
assert_eq!(decode_text_as(sjis, "Shift_JIS").unwrap(), "日本");
assert_eq!(decode_text_as(sjis, "cp932").unwrap(), "日本");
let koi = b"\xf0\xd2\xc9\xd7\xc5\xd4";
assert_eq!(decode_text_as(koi, "koi8-r").unwrap(), "Привет");
assert_eq!(decode_text_as(koi, "koi8_r").unwrap(), "Привет");
assert_eq!(decode_text_as(b"caf\xe9", "latin-1").unwrap(), "caf\u{e9}");
assert_eq!(decode_text_as(b"\xef\xbb\xbfa", "utf-8").unwrap(), "a");
assert!(matches!(
decode_text_as(b"plain", "utf-8").unwrap(),
std::borrow::Cow::Borrowed(_)
));
let err = decode_text_as(b"caf\xe9", "utf-8").unwrap_err().to_string();
assert!(err.contains("UTF-8"), "{err}");
let err = decode_text_as(b"x", "klingon-1").unwrap_err().to_string();
assert!(err.contains("unknown character encoding"), "{err}");
}
#[test]
fn text_documents_decode_like_docling() {
assert_eq!(decode_text("caf\u{e9}".as_bytes()).unwrap(), "caf\u{e9}");
assert!(matches!(
decode_text(b"plain").unwrap(),
std::borrow::Cow::Borrowed(_)
));
assert_eq!(decode_text(b"\xef\xbb\xbf# T").unwrap(), "# T");
assert_eq!(decode_text(b"\xff\xfeh\x00i\x00").unwrap(), "hi");
assert_eq!(decode_text(b"\xfe\xff\x00h\x00i").unwrap(), "hi");
assert_eq!(decode_text(b"\xff\xfe\x00\x00h\x00\x00\x00").unwrap(), "h");
assert_eq!(
decode_text(b"caf\xe9 \x93q\x94").unwrap(),
"caf\u{e9} \u{201c}q\u{201d}"
);
let damaged = decode_text(b"caf\xc3\xa9 na\xc3\xafve \xff");
assert!(damaged.unwrap_err().to_string().contains("damaged UTF-8"));
assert!(decode_text(b"x\x81y").is_err());
}
}