Skip to main content

deser_core/adapters/bytes/
encodings.rs

1//! The base64 encodings of bytes.
2use 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::Unexpected, "invalid base64 string")
11}
12
13const STANDARD: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
14const URL_SAFE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
15
16const INVALID: u8 = 0xff;
17
18/// Maps characters of both alphabets to their values.
19static 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    // SAFETY: the string only gets ASCII characters (the characters of the
40    // alphabet, `=` and the zeros from `resize`), it stays valid UTF-8 also
41    // if this panics.
42    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
68/// Decodes base64 leniently.
69///
70/// Both alphabets are accepted (also mixed) and the padding is optional.
71/// If there is padding it has to be complete.  Unused bits have to be zero.
72pub(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    // values are below 64, `INVALID` is not
86    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 with the standard alphabet and padding (RFC 4648 section 4).
138    ///
139    /// This is the default representation of bytes.  Decoding is lenient
140    /// (see [adapters documentation](crate::adapters#bytes)).
141    Base64,
142    "base64",
143    STANDARD,
144    true
145);
146
147base64!(
148    /// Base64 with the standard alphabet without padding.
149    ///
150    /// Decoding is lenient (see [adapters documentation](crate::adapters#bytes)).
151    Base64NoPad,
152    "base64-nopad",
153    STANDARD,
154    false
155);
156
157base64!(
158    /// Base64 with the URL-safe alphabet and padding (RFC 4648 section 5).
159    ///
160    /// Decoding is lenient (see [adapters documentation](crate::adapters#bytes)).
161    Base64Url,
162    "base64url",
163    URL_SAFE,
164    true
165);
166
167base64!(
168    /// Base64 with the URL-safe alphabet without padding.
169    ///
170    /// Decoding is lenient (see [adapters documentation](crate::adapters#bytes)).
171    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        // RFC 4648 section 10
190        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    /// Encodes base64 bit by bit.
213    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        // the long inputs cover the loops over whole blocks, miri is slow
237        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        // encoding appends
256        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}