1use std::io::{Error, ErrorKind, Result};
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub enum Encoding {
11 Utf8,
12 Utf16,
14 Utf16Le,
15 Utf16Be,
16}
17
18impl std::str::FromStr for Encoding {
19 type Err = String;
20
21 fn from_str(name: &str) -> std::result::Result<Self, Self::Err> {
22 match name.to_ascii_lowercase().as_str() {
23 "utf-8" => Ok(Self::Utf8),
24 "utf-16" => Ok(Self::Utf16),
25 "utf-16le" => Ok(Self::Utf16Le),
26 "utf-16be" => Ok(Self::Utf16Be),
27 _ => Err(format!(
28 "unsupported encoding {name:?}; use utf-8, utf-16, utf-16le or utf-16be"
29 )),
30 }
31 }
32}
33
34fn invalid(message: &str) -> Error {
35 Error::new(ErrorKind::InvalidData, message)
36}
37
38pub fn has_utf16_bom(bytes: &[u8]) -> bool {
40 !bytes.starts_with(&[0xff, 0xfe, 0, 0])
41 && (bytes.starts_with(&[0xff, 0xfe]) || bytes.starts_with(&[0xfe, 0xff]))
42}
43
44impl Encoding {
45 pub fn decode(self, bytes: &[u8]) -> Result<String> {
51 if self == Self::Utf8 {
52 let bytes = bytes.strip_prefix(&[0xef, 0xbb, 0xbf]).unwrap_or(bytes);
53 return std::str::from_utf8(bytes)
54 .map(str::to_owned)
55 .map_err(|_| invalid("input is not valid UTF-8"));
56 }
57 if self == Self::Utf16
58 && (bytes.starts_with(&[0xff, 0xfe, 0, 0]) || bytes.starts_with(&[0, 0, 0xfe, 0xff]))
59 {
60 return Err(invalid(
61 "UTF-32 is not supported; convert the input to UTF-8",
62 ));
63 }
64 let bom = if bytes.starts_with(&[0xff, 0xfe]) {
65 Some(Self::Utf16Le)
66 } else if bytes.starts_with(&[0xfe, 0xff]) {
67 Some(Self::Utf16Be)
68 } else {
69 None
70 };
71 let endian = match (self, bom) {
72 (Self::Utf16, None) => return Err(invalid(
73 "UTF-16 needs a BOM; select --encoding utf-16le or utf-16be for headerless input",
74 )),
75 (Self::Utf16, Some(endian)) => endian,
76 (selected, Some(endian)) if selected != endian => {
77 return Err(invalid("UTF-16 BOM conflicts with the selected encoding"))
78 }
79 (selected, _) => selected,
80 };
81 let bytes = if bom.is_some() { &bytes[2..] } else { bytes };
82 if bytes.len() % 2 != 0 {
83 return Err(invalid("invalid UTF-16: odd byte count"));
84 }
85 let units = bytes.as_chunks::<2>().0.iter().map(|&pair| {
86 if endian == Self::Utf16Le {
87 u16::from_le_bytes(pair)
88 } else {
89 u16::from_be_bytes(pair)
90 }
91 });
92 char::decode_utf16(units)
93 .collect::<std::result::Result<String, _>>()
94 .map_err(|_| invalid("invalid UTF-16: unpaired surrogate"))
95 }
96}
97
98#[cfg(test)]
99mod tests {
100 use super::*;
101
102 #[test]
103 fn explicit_endianness_preserves_unicode() {
104 let text = "Hello 漢字 🙂\r\n";
105 for encoding in [Encoding::Utf16Le, Encoding::Utf16Be] {
106 let mut bytes: Vec<u8> = text
107 .encode_utf16()
108 .flat_map(|unit| {
109 if encoding == Encoding::Utf16Le {
110 unit.to_le_bytes()
111 } else {
112 unit.to_be_bytes()
113 }
114 })
115 .collect();
116 assert_eq!(encoding.decode(&bytes).unwrap(), text);
117 assert!(Encoding::Utf16.decode(&bytes).is_err());
118 let bom = if encoding == Encoding::Utf16Le {
119 [0xff, 0xfe]
120 } else {
121 [0xfe, 0xff]
122 };
123 bytes.splice(0..0, bom);
124 assert_eq!(Encoding::Utf16.decode(&bytes).unwrap(), text);
125 assert_eq!(encoding.decode(&bytes).unwrap(), text);
126 }
127 }
128
129 #[test]
130 fn malformed_and_conflicting_input_is_rejected() {
131 assert!(Encoding::Utf16Le.decode(&[1]).is_err());
132 assert!(Encoding::Utf16Le.decode(&[0, 0xd8]).is_err());
133 assert!(Encoding::Utf16Be.decode(&[0xff, 0xfe, 65, 0]).is_err());
134 assert!(Encoding::Utf16.decode(&[0xff, 0xfe, 0, 0]).is_err());
135 assert!(!has_utf16_bom(&[0xff, 0xfe, 0, 0]));
136 assert_eq!(Encoding::Utf16Le.decode(&[0xff, 0xfe, 0, 0]).unwrap(), "\0");
137 assert!(Encoding::Utf8.decode(&[0xff]).is_err());
138 assert_eq!(
139 Encoding::Utf8.decode(b"\xef\xbb\xbfhello").unwrap(),
140 "hello"
141 );
142 }
143}