sql_dialect_fmt_encoding/
lib.rs1const UTF8_BOM: &[u8] = &[0xEF, 0xBB, 0xBF];
22const UTF16_LE_BOM: &[u8] = &[0xFF, 0xFE];
23const UTF16_BE_BOM: &[u8] = &[0xFE, 0xFF];
24
25#[derive(Clone, Copy, Debug, Eq, PartialEq)]
30#[non_exhaustive]
31pub enum TextEncoding {
32 Utf8,
34 Utf8Bom,
36 Utf16Le,
38 Utf16Be,
40 OpaqueBytes,
42}
43
44#[derive(Clone, Debug, Eq, PartialEq)]
50pub struct DecodedText {
51 kind: DecodedKind,
52}
53
54#[derive(Clone, Debug, Eq, PartialEq)]
55enum DecodedKind {
56 Text {
57 encoding: TextEncoding,
58 text: String,
59 },
60 Opaque {
61 bytes: Vec<u8>,
62 reason: OpaqueReason,
63 },
64}
65
66#[derive(Clone, Copy, Debug, Eq, PartialEq)]
71#[non_exhaustive]
72pub enum OpaqueReason {
73 InvalidUtf8,
75 OddLengthUtf16,
77 InvalidUtf16,
79}
80
81impl DecodedText {
82 pub fn decode(bytes: &[u8]) -> Self {
86 if bytes.starts_with(UTF8_BOM) {
87 return decode_utf8(&bytes[UTF8_BOM.len()..], TextEncoding::Utf8Bom, bytes);
88 }
89 if bytes.starts_with(UTF16_LE_BOM) {
90 return decode_utf16(&bytes[UTF16_LE_BOM.len()..], TextEncoding::Utf16Le, bytes);
91 }
92 if bytes.starts_with(UTF16_BE_BOM) {
93 return decode_utf16(&bytes[UTF16_BE_BOM.len()..], TextEncoding::Utf16Be, bytes);
94 }
95 decode_utf8(bytes, TextEncoding::Utf8, bytes)
96 }
97
98 pub fn encoding(&self) -> TextEncoding {
101 match &self.kind {
102 DecodedKind::Text { encoding, .. } => *encoding,
103 DecodedKind::Opaque { .. } => TextEncoding::OpaqueBytes,
104 }
105 }
106
107 pub fn as_str(&self) -> Option<&str> {
109 match &self.kind {
110 DecodedKind::Text { text, .. } => Some(text),
111 DecodedKind::Opaque { .. } => None,
112 }
113 }
114
115 pub fn opaque_reason(&self) -> Option<OpaqueReason> {
117 match &self.kind {
118 DecodedKind::Text { .. } => None,
119 DecodedKind::Opaque { reason, .. } => Some(*reason),
120 }
121 }
122
123 pub fn encode(&self) -> Vec<u8> {
126 match &self.kind {
127 DecodedKind::Text { encoding, text } => encode_text(*encoding, text),
128 DecodedKind::Opaque { bytes, .. } => bytes.clone(),
129 }
130 }
131
132 pub fn map_text(&self, edit: impl FnOnce(&str) -> String) -> Self {
135 match &self.kind {
136 DecodedKind::Text { encoding, text } => DecodedText {
137 kind: DecodedKind::Text {
138 encoding: *encoding,
139 text: edit(text),
140 },
141 },
142 DecodedKind::Opaque { .. } => self.clone(),
143 }
144 }
145}
146
147fn decode_utf8(bytes: &[u8], encoding: TextEncoding, original: &[u8]) -> DecodedText {
148 match std::str::from_utf8(bytes) {
149 Ok(text) => DecodedText {
150 kind: DecodedKind::Text {
151 encoding,
152 text: text.to_owned(),
153 },
154 },
155 Err(_) => opaque(original, OpaqueReason::InvalidUtf8),
156 }
157}
158
159fn decode_utf16(bytes: &[u8], encoding: TextEncoding, original: &[u8]) -> DecodedText {
160 if !bytes.len().is_multiple_of(2) {
161 return opaque(original, OpaqueReason::OddLengthUtf16);
162 }
163
164 let words = bytes.chunks_exact(2).map(|chunk| match encoding {
165 TextEncoding::Utf16Le => u16::from_le_bytes([chunk[0], chunk[1]]),
166 TextEncoding::Utf16Be => u16::from_be_bytes([chunk[0], chunk[1]]),
167 _ => unreachable!("decode_utf16 is only called for UTF-16 encodings"),
168 });
169
170 match String::from_utf16(&words.collect::<Vec<_>>()) {
171 Ok(text) => DecodedText {
172 kind: DecodedKind::Text { encoding, text },
173 },
174 Err(_) => opaque(original, OpaqueReason::InvalidUtf16),
175 }
176}
177
178fn encode_text(encoding: TextEncoding, text: &str) -> Vec<u8> {
179 match encoding {
180 TextEncoding::Utf8 => text.as_bytes().to_vec(),
181 TextEncoding::Utf8Bom => {
182 let mut bytes = Vec::with_capacity(UTF8_BOM.len() + text.len());
183 bytes.extend_from_slice(UTF8_BOM);
184 bytes.extend_from_slice(text.as_bytes());
185 bytes
186 }
187 TextEncoding::Utf16Le => {
188 let mut bytes = Vec::with_capacity(UTF16_LE_BOM.len() + text.len() * 2);
189 bytes.extend_from_slice(UTF16_LE_BOM);
190 for word in text.encode_utf16() {
191 bytes.extend_from_slice(&word.to_le_bytes());
192 }
193 bytes
194 }
195 TextEncoding::Utf16Be => {
196 let mut bytes = Vec::with_capacity(UTF16_BE_BOM.len() + text.len() * 2);
197 bytes.extend_from_slice(UTF16_BE_BOM);
198 for word in text.encode_utf16() {
199 bytes.extend_from_slice(&word.to_be_bytes());
200 }
201 bytes
202 }
203 TextEncoding::OpaqueBytes => unreachable!("opaque values are encoded from original bytes"),
204 }
205}
206
207fn opaque(bytes: &[u8], reason: OpaqueReason) -> DecodedText {
208 DecodedText {
209 kind: DecodedKind::Opaque {
210 bytes: bytes.to_vec(),
211 reason,
212 },
213 }
214}
215
216#[cfg(test)]
217mod tests {
218 use super::*;
219
220 #[test]
221 fn utf8_without_bom_round_trips() {
222 let bytes = "SELECT '長芋';\n".as_bytes();
223 let decoded = DecodedText::decode(bytes);
224
225 assert_eq!(decoded.encoding(), TextEncoding::Utf8);
226 assert_eq!(decoded.as_str(), Some("SELECT '長芋';\n"));
227 assert_eq!(decoded.encode(), bytes);
228 }
229
230 #[test]
231 fn utf8_bom_round_trips_and_preserves_bom() {
232 let mut bytes = UTF8_BOM.to_vec();
233 bytes.extend_from_slice("SELECT 1;\n".as_bytes());
234
235 let decoded = DecodedText::decode(&bytes);
236
237 assert_eq!(decoded.encoding(), TextEncoding::Utf8Bom);
238 assert_eq!(decoded.as_str(), Some("SELECT 1;\n"));
239 assert_eq!(decoded.encode(), bytes);
240 }
241
242 #[test]
243 fn utf16_le_round_trips_with_unicode() {
244 let text = "SELECT '長芋';\n";
245 let bytes = encode_text(TextEncoding::Utf16Le, text);
246
247 let decoded = DecodedText::decode(&bytes);
248
249 assert_eq!(decoded.encoding(), TextEncoding::Utf16Le);
250 assert_eq!(decoded.as_str(), Some(text));
251 assert_eq!(decoded.encode(), bytes);
252 }
253
254 #[test]
255 fn utf16_be_round_trips_with_unicode() {
256 let text = "SELECT '長芋';\n";
257 let bytes = encode_text(TextEncoding::Utf16Be, text);
258
259 let decoded = DecodedText::decode(&bytes);
260
261 assert_eq!(decoded.encoding(), TextEncoding::Utf16Be);
262 assert_eq!(decoded.as_str(), Some(text));
263 assert_eq!(decoded.encode(), bytes);
264 }
265
266 #[test]
267 fn opaque_invalid_utf8_preserves_original_bytes() {
268 let bytes = [0x53, 0x45, 0xFF, 0x4C];
269 let decoded = DecodedText::decode(&bytes);
270
271 assert_eq!(decoded.encoding(), TextEncoding::OpaqueBytes);
272 assert_eq!(decoded.as_str(), None);
273 assert_eq!(decoded.opaque_reason(), Some(OpaqueReason::InvalidUtf8));
274 assert_eq!(decoded.encode(), bytes);
275 }
276
277 #[test]
278 fn opaque_invalid_utf16_preserves_original_bytes() {
279 let bytes = [0xFF, 0xFE, 0x00];
280 let decoded = DecodedText::decode(&bytes);
281
282 assert_eq!(decoded.encoding(), TextEncoding::OpaqueBytes);
283 assert_eq!(decoded.opaque_reason(), Some(OpaqueReason::OddLengthUtf16));
284 assert_eq!(decoded.encode(), bytes);
285 }
286
287 #[test]
288 fn map_text_preserves_original_encoding() {
289 let source = encode_text(TextEncoding::Utf16Le, "select 1\n");
290 let decoded = DecodedText::decode(&source);
291
292 let edited = decoded.map_text(|text| text.to_uppercase());
293
294 assert_eq!(edited.encoding(), TextEncoding::Utf16Le);
295 assert_eq!(
296 DecodedText::decode(&edited.encode()).as_str(),
297 Some("SELECT 1\n")
298 );
299 }
300}