Skip to main content

rustlavel_db/
base64.rs

1//! Base64, needed by SCRAM authentication.
2//!
3//! Thirty lines is cheaper than a dependency, and this is the only place the
4//! framework needs it.
5
6const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
7
8pub fn encode(input: &[u8]) -> String {
9    let mut out = String::with_capacity(input.len().div_ceil(3) * 4);
10
11    for chunk in input.chunks(3) {
12        let b = [chunk[0], *chunk.get(1).unwrap_or(&0), *chunk.get(2).unwrap_or(&0)];
13        let bits = (u32::from(b[0]) << 16) | (u32::from(b[1]) << 8) | u32::from(b[2]);
14
15        out.push(ALPHABET[(bits >> 18 & 0x3f) as usize] as char);
16        out.push(ALPHABET[(bits >> 12 & 0x3f) as usize] as char);
17        out.push(if chunk.len() > 1 { ALPHABET[(bits >> 6 & 0x3f) as usize] as char } else { '=' });
18        out.push(if chunk.len() > 2 { ALPHABET[(bits & 0x3f) as usize] as char } else { '=' });
19    }
20
21    out
22}
23
24pub fn decode(input: &str) -> Option<Vec<u8>> {
25    let mut out = Vec::with_capacity(input.len() / 4 * 3);
26    let mut buffer = 0u32;
27    let mut bits = 0u32;
28
29    for byte in input.bytes() {
30        let value = match byte {
31            b'A'..=b'Z' => byte - b'A',
32            b'a'..=b'z' => byte - b'a' + 26,
33            b'0'..=b'9' => byte - b'0' + 52,
34            b'+' => 62,
35            b'/' => 63,
36            b'=' | b'\n' | b'\r' => continue,
37            _ => return None,
38        };
39
40        buffer = (buffer << 6) | u32::from(value);
41        bits += 6;
42        if bits >= 8 {
43            bits -= 8;
44            out.push((buffer >> bits) as u8);
45        }
46    }
47
48    Some(out)
49}
50
51#[cfg(test)]
52mod tests {
53    use super::*;
54
55    #[test]
56    fn matches_the_rfc_4648_vectors() {
57        for (plain, encoded) in [
58            ("", ""),
59            ("f", "Zg=="),
60            ("fo", "Zm8="),
61            ("foo", "Zm9v"),
62            ("foob", "Zm9vYg=="),
63            ("fooba", "Zm9vYmE="),
64            ("foobar", "Zm9vYmFy"),
65        ] {
66            assert_eq!(encode(plain.as_bytes()), encoded, "encoding {plain:?}");
67            assert_eq!(decode(encoded).unwrap(), plain.as_bytes(), "decoding {encoded:?}");
68        }
69    }
70
71    #[test]
72    fn round_trips_binary() {
73        let bytes: Vec<u8> = (0..=255).collect();
74        assert_eq!(decode(&encode(&bytes)).unwrap(), bytes);
75    }
76
77    #[test]
78    fn rejects_characters_outside_the_alphabet() {
79        assert!(decode("not base64!").is_none());
80    }
81}