fn sextet(c: u8) -> Option<u8> {
match c {
b'A'..=b'Z' => Some(c - b'A'),
b'a'..=b'z' => Some(c - b'a' + 26),
b'0'..=b'9' => Some(c - b'0' + 52),
b'+' => Some(62),
b'/' => Some(63),
_ => None,
}
}
pub fn decode(s: &str) -> Result<Vec<u8>, String> {
let bytes = s.as_bytes();
let body_len = bytes.iter().take_while(|&&c| c != b'=').count();
let pad_len = bytes.len() - body_len;
if pad_len > 2 || bytes[body_len..].iter().any(|&c| c != b'=') {
return Err(format!("invalid base64 padding (pad_len={pad_len})"));
}
let core = &bytes[..body_len];
let mut out = Vec::with_capacity(core.len() / 4 * 3 + 3);
let mut acc: u32 = 0;
let mut nbits: u32 = 0;
for &c in core {
let v = sextet(c).ok_or_else(|| format!("invalid base64 char: {c:#x}"))?;
acc = (acc << 6) | v as u32;
nbits += 6;
if nbits >= 8 {
nbits -= 8;
out.push((acc >> nbits) as u8);
}
}
if nbits > 0 && (acc & ((1 << nbits) - 1)) != 0 {
return Err("invalid base64 trailing bits".to_string());
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn encode(bytes: &[u8]) -> String {
const TABLE: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
let mut out = String::with_capacity(bytes.len().div_ceil(3) * 4);
for chunk in bytes.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 n = (b0 << 16) | (b1 << 8) | b2;
out.push(TABLE[((n >> 18) & 63) as usize] as char);
out.push(TABLE[((n >> 12) & 63) as usize] as char);
out.push(if chunk.len() > 1 {
TABLE[((n >> 6) & 63) as usize] as char
} else {
'='
});
out.push(if chunk.len() > 2 {
TABLE[(n & 63) as usize] as char
} else {
'='
});
}
out
}
#[test]
fn round_trip_arbitrary_lengths() {
for len in 0..32usize {
let data: Vec<u8> = (0..len).map(|i| (i as u8).wrapping_mul(31)).collect();
let encoded = encode(&data);
assert_eq!(decode(&encoded).unwrap(), data, "round-trip len={len}");
}
}
#[test]
fn decode_known_vector() {
assert_eq!(decode("aGVsbG8=").unwrap(), b"hello");
assert_eq!(decode("Zm9vYmFy").unwrap(), b"foobar");
assert_eq!(decode("").unwrap(), b"");
}
#[test]
fn rejects_invalid_chars() {
assert!(decode("aGVsbG8*").is_err());
assert!(decode("not base64 ###").is_err());
}
}