use crate::courierust_error::{Error, Result};
use alloc::string::String;
use alloc::vec::Vec;
const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
const ALPHABET_URL: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
const INVALID: u8 = 0xff;
const fn decode_table() -> [u8; 256] {
let mut t = [INVALID; 256];
let mut i = 0usize;
while i < 64 {
t[ALPHABET[i] as usize] = i as u8;
i += 1;
}
t
}
static DECODE_TABLE: [u8; 256] = decode_table();
#[inline]
pub const fn encoded_len(len: usize) -> usize {
let quanta = len / 3 + if len % 3 != 0 { 1 } else { 0 };
quanta * 4
}
#[inline]
pub const fn decoded_len(chars: usize) -> usize {
chars / 4 * 3
}
#[inline]
pub fn decoded_size(input: &[u8]) -> Result<usize> {
decode_core(input, None)
}
#[inline]
pub fn validate(input: &[u8]) -> Result<()> {
decode_core(input, None).map(|_| ())
}
pub fn encode(data: &[u8]) -> String {
let mut out = String::with_capacity(encoded_len(data.len()));
encode_into(data, &mut out);
out
}
pub fn encode_into(data: &[u8], out: &mut String) {
encode_with(data, ALPHABET, true, out)
}
pub fn encode_url_no_pad(data: &[u8]) -> String {
let mut out = String::with_capacity(encoded_len(data.len()));
encode_with(data, ALPHABET_URL, false, &mut out);
out
}
fn encode_with(data: &[u8], alphabet: &[u8; 64], pad: bool, out: &mut String) {
let mut chunks = data.chunks_exact(3);
for chunk in &mut chunks {
let n = ((chunk[0] as u32) << 16) | ((chunk[1] as u32) << 8) | chunk[2] as u32;
out.push(alphabet[(n >> 18) as usize & 0x3f] as char);
out.push(alphabet[(n >> 12) as usize & 0x3f] as char);
out.push(alphabet[(n >> 6) as usize & 0x3f] as char);
out.push(alphabet[n as usize & 0x3f] as char);
}
match chunks.remainder() {
[] => {}
[a] => {
let n = (*a as u32) << 16;
out.push(alphabet[(n >> 18) as usize & 0x3f] as char);
out.push(alphabet[(n >> 12) as usize & 0x3f] as char);
if pad {
out.push('=');
out.push('=');
}
}
[a, b] => {
let n = ((*a as u32) << 16) | ((*b as u32) << 8);
out.push(alphabet[(n >> 18) as usize & 0x3f] as char);
out.push(alphabet[(n >> 12) as usize & 0x3f] as char);
out.push(alphabet[(n >> 6) as usize & 0x3f] as char);
if pad {
out.push('=');
}
}
_ => unreachable!(),
}
}
pub fn decode(input: &[u8]) -> Result<Vec<u8>> {
let mut out = Vec::with_capacity(decoded_len(input.len()));
decode_into(input, &mut out)?;
Ok(out)
}
pub fn decode_into(input: &[u8], out: &mut Vec<u8>) -> Result<()> {
let start = out.len();
match decode_core(input, Some(out)) {
Ok(_) => Ok(()),
Err(e) => {
out.truncate(start);
Err(e)
}
}
}
fn decode_core(input: &[u8], out: Option<&mut Vec<u8>>) -> Result<usize> {
if input.len() % 4 != 0 {
return Err(Error::protocol("base64: length is not a multiple of 4"));
}
let mut out = out;
if let Some(buf) = out.as_deref_mut() {
buf.reserve(decoded_len(input.len()));
}
let mut written = 0usize;
let mut finished = false;
let mut vals = [0u8; 4];
for chunk in input.chunks_exact(4) {
if finished {
return Err(Error::protocol("base64: trailing data after padding"));
}
let mut pad = 0usize;
for (i, &c) in chunk.iter().enumerate() {
if c == b'=' {
if i < 2 {
return Err(Error::protocol("base64: padding in the first half"));
}
pad += 1;
vals[i] = 0;
continue;
}
if pad > 0 {
return Err(Error::protocol("base64: data after padding"));
}
let v = DECODE_TABLE[c as usize];
if v == INVALID {
return Err(Error::protocol("base64: invalid character"));
}
vals[i] = v;
}
let n = ((vals[0] as u32) << 18)
| ((vals[1] as u32) << 12)
| ((vals[2] as u32) << 6)
| vals[3] as u32;
let (bytes, count) = match pad {
0 => ([(n >> 16) as u8, (n >> 8) as u8, n as u8], 3usize),
1 => {
if vals[2] & 0x03 != 0 {
return Err(Error::protocol("base64: non-canonical padding bits"));
}
([(n >> 16) as u8, (n >> 8) as u8, 0], 2)
}
2 => {
if vals[1] & 0x0f != 0 {
return Err(Error::protocol("base64: non-canonical padding bits"));
}
([(n >> 16) as u8, 0, 0], 1)
}
_ => return Err(Error::protocol("base64: too much padding")),
};
if let Some(buf) = out.as_deref_mut() {
buf.extend_from_slice(&bytes[..count]);
}
written += count;
if pad > 0 {
finished = true;
}
}
Ok(written)
}
#[inline]
pub fn decode_str(input: &str) -> Result<Vec<u8>> {
decode(input.as_bytes())
}
pub fn is_canonical(input: &[u8]) -> bool {
validate(input).is_ok()
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn rfc4648_vectors() {
let cases: &[(&[u8], &str)] = &[
(b"", ""),
(b"f", "Zg=="),
(b"fo", "Zm8="),
(b"foo", "Zm9v"),
(b"foob", "Zm9vYg=="),
(b"fooba", "Zm9vYmE="),
(b"foobar", "Zm9vYmFy"),
];
for (raw, encoded) in cases {
assert_eq!(encode(raw), *encoded);
assert_eq!(decode(encoded.as_bytes()).unwrap(), raw.to_vec());
assert_eq!(encoded_len(raw.len()), encoded.len());
}
}
#[test]
fn url_alphabet_vectors_are_unpadded() {
let cases: &[(&[u8], &str)] = &[
(b"", ""),
(b"f", "Zg"),
(b"fo", "Zm8"),
(b"foo", "Zm9v"),
(b"foob", "Zm9vYg"),
(b"fooba", "Zm9vYmE"),
(b"foobar", "Zm9vYmFy"),
(&[0xfb, 0xef, 0xff], "--__"),
(&[0xff], "_w"),
];
for (raw, encoded) in cases {
assert_eq!(encode_url_no_pad(raw), *encoded, "raw={raw:?}");
}
}
#[test]
fn all_byte_values_roundtrip() {
let data: Vec<u8> = (0..=255u8).collect();
let text = encode(&data);
assert_eq!(decode(text.as_bytes()).unwrap(), data);
for n in 0..=data.len() {
let t = encode(&data[..n]);
assert_eq!(decode(t.as_bytes()).unwrap(), data[..n].to_vec());
}
}
#[test]
fn rejects_non_canonical_and_malformed() {
assert!(decode(b"Zg==").is_ok());
assert!(decode(b"Zg").is_err()); assert!(decode(b"Zg=").is_err()); assert!(decode(b"Zg===").is_err());
assert!(decode(b"Z===").is_err()); assert!(decode(b"Zg==Zg==").is_err()); assert!(decode(b"Zh==").is_err()); assert!(decode(b"Zm9=").is_err()); assert!(decode(b"Zm9v\n").is_err()); assert!(decode(b"Zm9*").is_err()); assert!(decode(b"Zm9v YmFy").is_err());
assert!(decode("Zm9é".as_bytes()).is_err());
}
#[test]
fn validation_agrees_with_decoding() {
let cases: &[&[u8]] = &[
b"",
b"Zg==",
b"Zm8=",
b"Zm9v",
b"Zg",
b"Zg=",
b"Zg===",
b"Z===",
b"Zg==Zg==",
b"Zh==",
b"Zm9=",
b"Zm9v\n",
b"Zm9*",
b"A",
b"AA",
b"AAA",
b"AAAA",
b"====",
b"AB==",
b"AAB=",
b"AAAB",
];
for case in cases {
let decoded = decode(case);
assert_eq!(validate(case).is_ok(), decoded.is_ok(), "{case:?}");
assert_eq!(is_canonical(case), decoded.is_ok(), "{case:?}");
match decoded {
Ok(v) => assert_eq!(decoded_size(case).unwrap(), v.len(), "{case:?}"),
Err(_) => assert!(decoded_size(case).is_err(), "{case:?}"),
}
}
let long = encode(&(0..=255u8).collect::<Vec<u8>>());
let bytes = long.as_bytes();
for cut in 0..bytes.len() {
let slice = &bytes[..cut];
assert_eq!(is_canonical(slice), decode(slice).is_ok(), "cut={cut}");
}
}
#[test]
fn encoded_len_is_exact() {
for len in 0..64usize {
let data: Vec<u8> = (0..len).map(|i| i as u8).collect();
assert_eq!(encoded_len(len), encode(&data).len(), "len={len}");
}
}
#[test]
fn decode_leaves_output_intact_on_error() {
let mut out = vec![1, 2, 3];
assert!(decode_into(b"AAAA!AAA", &mut out).is_err());
assert_eq!(out, vec![1, 2, 3]);
}
#[test]
fn websocket_key_shape() {
let mut raw = [0u8; 16];
for (i, b) in raw.iter_mut().enumerate() {
*b = i as u8;
}
let key = encode(&raw);
assert_eq!(key.len(), 24);
assert!(key.ends_with("=="));
assert_eq!(decode(key.as_bytes()).unwrap().len(), 16);
assert!(is_canonical(key.as_bytes()));
}
}