Skip to main content

rich_ext/
encoding.rs

1//! Explicit text decoding for CLI and application inputs.
2//!
3//! This extension never guesses a headerless encoding. Callers retain their
4//! existing default decoding policy unless an encoding is explicitly selected.
5
6use std::io::{Error, ErrorKind, Result};
7
8/// Supported explicit text encodings. All decoding is strict.
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub enum Encoding {
11    Utf8,
12    /// Requires a UTF-16 byte-order mark.
13    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
38/// Recognize a likely UTF-16 BOM, excluding ambiguous UTF-32 signatures.
39pub 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    /// Decode bytes, consuming a matching BOM and rejecting malformed data.
46    /// `Utf16` requires a BOM; endian-specific variants also accept headerless
47    /// data, but reject a contradictory BOM. Generic `Utf16` rejects ambiguous
48    /// UTF-32 signatures; select an endian explicitly for BOM + leading NUL.
49    /// Newline policy belongs to callers.
50    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}