const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
pub fn encode(bytes: &[u8]) -> String {
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 = *chunk.get(1).unwrap_or(&0) as u32;
let b2 = *chunk.get(2).unwrap_or(&0) as u32;
let n = (b0 << 16) | (b1 << 8) | b2;
out.push(ALPHABET[((n >> 18) & 63) as usize] as char);
out.push(ALPHABET[((n >> 12) & 63) as usize] as char);
out.push(if chunk.len() > 1 {
ALPHABET[((n >> 6) & 63) as usize] as char
} else {
'='
});
out.push(if chunk.len() > 2 {
ALPHABET[(n & 63) as usize] as char
} else {
'='
});
}
out
}
pub fn encode_f32_le(v: &[f32]) -> String {
let mut bytes = Vec::with_capacity(v.len() * 4);
for x in v {
bytes.extend_from_slice(&x.to_le_bytes());
}
encode(&bytes)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_match_rfc4648_test_vectors() {
assert_eq!(encode(b""), "");
assert_eq!(encode(b"f"), "Zg==");
assert_eq!(encode(b"fo"), "Zm8=");
assert_eq!(encode(b"foo"), "Zm9v");
assert_eq!(encode(b"foob"), "Zm9vYg==");
assert_eq!(encode(b"fooba"), "Zm9vYmE=");
assert_eq!(encode(b"foobar"), "Zm9vYmFy");
}
#[test]
fn should_round_trip_f32_le_through_a_reference_decoder() {
let v = vec![0.0f32, 1.5, -2.25, 1e9, f32::MIN_POSITIVE];
let s = encode_f32_le(&v);
let back = decode_ref(&s);
let got: Vec<f32> = back
.chunks_exact(4)
.map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.collect();
assert_eq!(got, v);
}
fn decode_ref(s: &str) -> Vec<u8> {
let val = |c: u8| -> u32 {
match c {
b'A'..=b'Z' => (c - b'A') as u32,
b'a'..=b'z' => (c - b'a' + 26) as u32,
b'0'..=b'9' => (c - b'0' + 52) as u32,
b'+' => 62,
b'/' => 63,
_ => 0,
}
};
let bytes = s.as_bytes();
let mut out = Vec::new();
for quad in bytes.chunks(4) {
let n =
(val(quad[0]) << 18) | (val(quad[1]) << 12) | (val(quad[2]) << 6) | val(quad[3]);
out.push((n >> 16) as u8);
if quad[2] != b'=' {
out.push((n >> 8) as u8);
}
if quad[3] != b'=' {
out.push(n as u8);
}
}
out
}
}
pub fn decode(s: &str) -> Result<Vec<u8>, String> {
let val = |c: u8| -> Result<u32, String> {
Ok(match c {
b'A'..=b'Z' => (c - b'A') as u32,
b'a'..=b'z' => (c - b'a') as u32 + 26,
b'0'..=b'9' => (c - b'0') as u32 + 52,
b'+' | b'-' => 62, b'/' | b'_' => 63, _ => return Err(format!("invalid base64 byte {:?}", c as char)),
})
};
let cleaned: Vec<u8> = s
.bytes()
.filter(|c| !c.is_ascii_whitespace() && *c != b'=')
.collect();
let mut out = Vec::with_capacity(cleaned.len() / 4 * 3);
for chunk in cleaned.chunks(4) {
if chunk.len() == 1 {
return Err("truncated base64 (a lone trailing character)".into());
}
let mut n = 0u32;
for (i, &c) in chunk.iter().enumerate() {
n |= val(c)? << (18 - 6 * i);
}
out.push((n >> 16) as u8);
if chunk.len() > 2 {
out.push((n >> 8) as u8);
}
if chunk.len() > 3 {
out.push(n as u8);
}
}
Ok(out)
}
#[cfg(test)]
mod decode_tests {
use super::*;
#[test]
fn round_trips_arbitrary_bytes() {
for len in 0..64usize {
let bytes: Vec<u8> = (0..len).map(|i| (i * 37 % 251) as u8).collect();
assert_eq!(decode(&encode(&bytes)).unwrap(), bytes, "len {len}");
}
}
#[test]
fn tolerates_url_safe_and_missing_padding() {
let bytes = vec![0xfb, 0xff, 0xbf, 0x00, 0x10];
let std = encode(&bytes);
assert_eq!(decode(std.trim_end_matches('=')).unwrap(), bytes);
let url = std.replace('+', "-").replace('/', "_");
assert_eq!(decode(&url).unwrap(), bytes);
}
#[test]
fn rejects_garbage_rather_than_corrupting_an_image() {
assert!(decode("not*base64").is_err());
}
}