deser_core/adapters/bytes/
encodings.rs1use crate::adapters::bytes::BytesEncoding;
3use crate::error::{Error, ErrorKind};
4use alloc::string::String;
5use alloc::vec;
6use alloc::vec::Vec;
7
8#[cold]
9fn invalid() -> Error {
10 Error::new(ErrorKind::InvalidValue, "invalid base64 string")
11}
12
13const STANDARD: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
14const URL_SAFE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
15
16const INVALID: u8 = 0xff;
17
18static BASE64_VALUES: [u8; 256] = {
20 let mut table = [INVALID; 256];
21 let mut idx = 0;
22 while idx < 64 {
23 table[STANDARD[idx] as usize] = idx as u8;
24 table[URL_SAFE[idx] as usize] = idx as u8;
25 idx += 1;
26 }
27 table
28};
29
30fn encode_base64(bytes: &[u8], alphabet: &[u8; 64], pad: bool, out: &mut String) {
31 let char_at = |value: u32| alphabet[(value & 0x3f) as usize];
32 let (chunks, remainder) = bytes.as_chunks::<3>();
33 let tail_len = match remainder.len() {
34 0 => 0,
35 _ if pad => 4,
36 len => len + 1,
37 };
38 let start = out.len();
39 let buf = unsafe { out.as_mut_vec() };
43 buf.resize(start + chunks.len() * 4 + tail_len, 0);
44 let (body, tail) = buf[start..].split_at_mut(chunks.len() * 4);
45 for (&[a, b, c], dst) in chunks.iter().zip(body.as_chunks_mut::<4>().0) {
46 let n = (a as u32) << 16 | (b as u32) << 8 | c as u32;
47 *dst = [
48 char_at(n >> 18),
49 char_at(n >> 12),
50 char_at(n >> 6),
51 char_at(n),
52 ];
53 }
54 let quad = match *remainder {
55 [a] => {
56 let n = (a as u32) << 16;
57 [char_at(n >> 18), char_at(n >> 12), b'=', b'=']
58 }
59 [a, b] => {
60 let n = (a as u32) << 16 | (b as u32) << 8;
61 [char_at(n >> 18), char_at(n >> 12), char_at(n >> 6), b'=']
62 }
63 _ => return,
64 };
65 tail.copy_from_slice(&quad[..tail_len]);
66}
67
68pub(crate) fn decode_base64(s: &str) -> Result<Vec<u8>, Error> {
73 let input = s.as_bytes();
74 let padding = input
75 .iter()
76 .rev()
77 .take(2)
78 .take_while(|&&b| b == b'=')
79 .count();
80 let (chunks, remainder) = input[..input.len() - padding].as_chunks::<4>();
81 if remainder.len() == 1 || (padding > 0 && remainder.len() + padding != 4) {
82 return Err(invalid());
83 }
84
85 let value = |b: u8| BASE64_VALUES[b as usize];
87 let mut out = vec![0; chunks.len() * 3 + remainder.len().saturating_sub(1)];
88 let (body, tail) = out.split_at_mut(chunks.len() * 3);
89 for (&[a, b, c, d], dst) in chunks.iter().zip(body.as_chunks_mut::<3>().0) {
90 let (a, b, c, d) = (value(a), value(b), value(c), value(d));
91 if a | b | c | d >= 64 {
92 return Err(invalid());
93 }
94 let n = (a as u32) << 18 | (b as u32) << 12 | (c as u32) << 6 | d as u32;
95 *dst = [(n >> 16) as u8, (n >> 8) as u8, n as u8];
96 }
97 match *remainder {
98 [a, b] => {
99 let (a, b) = (value(a), value(b));
100 if a | b >= 64 || b & 0xf != 0 {
101 return Err(invalid());
102 }
103 tail[0] = a << 2 | b >> 4;
104 }
105 [a, b, c] => {
106 let (a, b, c) = (value(a), value(b), value(c));
107 if a | b | c >= 64 || c & 0x3 != 0 {
108 return Err(invalid());
109 }
110 tail.copy_from_slice(&[a << 2 | b >> 4, b << 4 | c >> 2]);
111 }
112 _ => {}
113 }
114 Ok(out)
115}
116
117macro_rules! base64 {
118 ($(#[$meta:meta])* $ty:ident, $name:expr, $alphabet:expr, $pad:expr) => {
119 $(#[$meta])*
120 pub struct $ty;
121
122 impl BytesEncoding for $ty {
123 const NAME: &'static str = $name;
124
125 fn encode(bytes: &[u8], out: &mut String) {
126 encode_base64(bytes, $alphabet, $pad, out)
127 }
128
129 fn decode(s: &str) -> Result<Vec<u8>, Error> {
130 decode_base64(s)
131 }
132 }
133 };
134}
135
136base64!(
137 Base64,
142 "base64",
143 STANDARD,
144 true
145);
146
147base64!(
148 Base64NoPad,
152 "base64-nopad",
153 STANDARD,
154 false
155);
156
157base64!(
158 Base64Url,
162 "base64url",
163 URL_SAFE,
164 true
165);
166
167base64!(
168 Base64UrlNoPad,
172 "base64url-nopad",
173 URL_SAFE,
174 false
175);
176
177#[cfg(test)]
178mod tests {
179 use super::*;
180
181 fn encode<E: BytesEncoding>(bytes: &[u8]) -> String {
182 let mut rv = String::new();
183 E::encode(bytes, &mut rv);
184 rv
185 }
186
187 #[test]
188 fn test_base64_encode() {
189 for (bytes, expected) in [
191 (&b""[..], ""),
192 (b"f", "Zg=="),
193 (b"fo", "Zm8="),
194 (b"foo", "Zm9v"),
195 (b"foob", "Zm9vYg=="),
196 (b"fooba", "Zm9vYmE="),
197 (b"foobar", "Zm9vYmFy"),
198 ] {
199 assert_eq!(encode::<Base64>(bytes), expected);
200 assert_eq!(encode::<Base64NoPad>(bytes), expected.trim_end_matches('='));
201 assert_eq!(decode_base64(expected).unwrap(), bytes);
202 assert_eq!(
203 decode_base64(expected.trim_end_matches('=')).unwrap(),
204 bytes
205 );
206 }
207 assert_eq!(encode::<Base64>(b"\xfb\xff"), "+/8=");
208 assert_eq!(encode::<Base64Url>(b"\xfb\xff"), "-_8=");
209 assert_eq!(encode::<Base64UrlNoPad>(b"\xfb\xff"), "-_8");
210 }
211
212 fn reference_base64(bytes: &[u8], alphabet: &[u8; 64], pad: bool) -> String {
214 let bits: Vec<bool> = bytes
215 .iter()
216 .flat_map(|byte| (0..8).rev().map(move |idx| byte >> idx & 1 == 1))
217 .collect();
218 let mut rv: String = bits
219 .chunks(6)
220 .map(|chunk| {
221 let value = (0..6).fold(0, |acc, idx| {
222 acc << 1 | chunk.get(idx).copied().unwrap_or(false) as usize
223 });
224 alphabet[value] as char
225 })
226 .collect();
227 while pad && !rv.len().is_multiple_of(4) {
228 rv.push('=');
229 }
230 rv
231 }
232
233 #[test]
234 fn test_base64_roundtrip() {
235 let data: Vec<u8> = (0..=255).rev().chain(0..=255).collect();
236 let long: &[usize] = if cfg!(miri) {
238 &[63, 64]
239 } else {
240 &[254, 255, 256, 511, 512]
241 };
242 for len in (0..10).chain(long.iter().copied()) {
243 let bytes = &data[..len];
244 for (encoded, alphabet, pad) in [
245 (encode::<Base64>(bytes), STANDARD, true),
246 (encode::<Base64NoPad>(bytes), STANDARD, false),
247 (encode::<Base64Url>(bytes), URL_SAFE, true),
248 (encode::<Base64UrlNoPad>(bytes), URL_SAFE, false),
249 ] {
250 assert_eq!(encoded, reference_base64(bytes, alphabet, pad));
251 assert_eq!(decode_base64(&encoded).unwrap(), bytes);
252 }
253 }
254
255 let mut out = String::from("x");
257 Base64::encode(b"f", &mut out);
258 Base64UrlNoPad::encode(b"\xfb\xff", &mut out);
259 assert_eq!(out, "xZg==-_8");
260 }
261
262 #[test]
263 fn test_base64_invalid_chars() {
264 let valid = encode::<Base64>(&[0xa5; 8]);
265 for idx in 0..valid.len() - 1 {
266 for invalid in [b' ', b'.', b'\n', b'=', 0xc3] {
267 let mut bytes = valid.clone().into_bytes();
268 bytes[idx] = invalid;
269 let s = String::from_utf8_lossy(&bytes);
270 assert!(decode_base64(&s).is_err(), "{s:?}");
271 }
272 }
273 }
274
275 #[test]
276 fn test_base64_decode() {
277 assert_eq!(decode_base64("+/8=").unwrap(), b"\xfb\xff");
278 assert_eq!(decode_base64("-_8").unwrap(), b"\xfb\xff");
279 assert_eq!(decode_base64("+_8").unwrap(), b"\xfb\xff");
280 for invalid in [
281 "Z", "Zg=", "Zg===", "Zm9v=", "Zm9v==", "Z===", "Zh==", "Zm9=", "Zm 9v", "Zm9v\n", "=",
282 "==", "Zg==Zg==",
283 ] {
284 assert!(decode_base64(invalid).is_err(), "{:?}", invalid);
285 }
286 }
287}