const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
const PAD: u8 = b'=';
const INVALID: u8 = 0xFF;
const DECODE_TABLE: [u8; 256] = {
let mut table = [INVALID; 256];
let mut i = 0;
while i < 64 {
table[ALPHABET[i] as usize] = i as u8;
i += 1;
}
table
};
pub(crate) fn encode(input: &[u8]) -> String {
let mut out = String::with_capacity(input.len().div_ceil(3) * 4);
for chunk in input.chunks(3) {
let b0 = chunk[0] as u32;
let b1 = if chunk.len() > 1 { chunk[1] as u32 } else { 0 };
let b2 = if chunk.len() > 2 { chunk[2] as u32 } else { 0 };
let triple = (b0 << 16) | (b1 << 8) | b2;
out.push(ALPHABET[(triple >> 18 & 0x3F) as usize] as char);
out.push(ALPHABET[(triple >> 12 & 0x3F) as usize] as char);
out.push(if chunk.len() > 1 {
ALPHABET[(triple >> 6 & 0x3F) as usize] as char
} else {
PAD as char
});
out.push(if chunk.len() > 2 {
ALPHABET[(triple & 0x3F) as usize] as char
} else {
PAD as char
});
}
out
}
pub(crate) fn decode(input: &str) -> Result<Vec<u8>, ()> {
let bytes = input.as_bytes();
if bytes.is_empty() {
return Ok(Vec::new());
}
if !bytes.len().is_multiple_of(4) {
return Err(());
}
let chunks = bytes.len() / 4;
let mut out = Vec::with_capacity(chunks * 3);
for (index, chunk) in bytes.chunks_exact(4).enumerate() {
let is_last = index == chunks - 1;
let mut sextets = [0u8; 4];
let mut pad = 0usize;
for (i, &c) in chunk.iter().enumerate() {
if c == PAD {
if !is_last || i < 2 {
return Err(());
}
pad += 1;
} else {
if pad > 0 {
return Err(());
}
let v = DECODE_TABLE[c as usize];
if v == INVALID {
return Err(());
}
sextets[i] = v;
}
}
let triple = (sextets[0] as u32) << 18
| (sextets[1] as u32) << 12
| (sextets[2] as u32) << 6
| (sextets[3] as u32);
out.push((triple >> 16) as u8);
if pad < 2 {
out.push((triple >> 8) as u8);
}
if pad < 1 {
out.push(triple as u8);
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip() {
for case in [
&b""[..],
b"a",
b"ab",
b"abc",
b"abcd",
b"hello world",
&[0u8, 255, 1, 254, 2, 253],
] {
let encoded = encode(case);
assert_eq!(decode(&encoded).unwrap(), case, "encoded as {encoded}");
}
}
#[test]
fn known_vectors() {
assert_eq!(encode(b"abcd"), "YWJjZA==");
assert_eq!(decode("YWJjZA==").unwrap(), b"abcd");
assert_eq!(encode(b"Many hands make light work."),
"TWFueSBoYW5kcyBtYWtlIGxpZ2h0IHdvcmsu");
}
#[test]
fn rejects_garbage() {
assert!(decode("YWJj").is_ok());
assert!(decode("YWJ").is_err()); assert!(decode("YW=J").is_err()); assert!(decode("====").is_err()); assert!(decode("****").is_err()); }
}