use core::fmt;
const STD_ENCODE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
const URL_ENCODE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
const STD_DECODE: [u8; 256] = build_decode_table(STD_ENCODE);
const URL_DECODE: [u8; 256] = build_decode_table(URL_ENCODE);
#[allow(clippy::cast_possible_truncation)]
const fn build_decode_table(alphabet: &[u8; 64]) -> [u8; 256] {
let mut table = [255u8; 256];
let mut i = 0;
while i < 64 {
table[alphabet[i] as usize] = i as u8;
i += 1;
}
table
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(clippy::enum_variant_names)]
enum Base64DecodeErrorKind {
InvalidCharacter,
InvalidPadding,
InvalidLength,
NonCanonical,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Base64DecodeError {
kind: Base64DecodeErrorKind,
}
impl Base64DecodeError {
const fn new(kind: Base64DecodeErrorKind) -> Self {
Self { kind }
}
}
impl fmt::Display for Base64DecodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.kind {
Base64DecodeErrorKind::InvalidCharacter => {
write!(f, "base64: invalid character in input")
}
Base64DecodeErrorKind::InvalidPadding => {
write!(f, "base64: invalid padding")
}
Base64DecodeErrorKind::InvalidLength => {
write!(f, "base64: invalid input length")
}
Base64DecodeErrorKind::NonCanonical => {
write!(f, "base64: non-canonical encoding (non-zero trailing bits)")
}
}
}
}
impl std::error::Error for Base64DecodeError {}
#[must_use]
pub fn base64_encode(input: &[u8]) -> String {
encode_with_alphabet(input, STD_ENCODE, true)
}
#[must_use]
pub fn base64url_encode(input: &[u8]) -> String {
encode_with_alphabet(input, URL_ENCODE, false)
}
fn encode_with_alphabet(input: &[u8], alphabet: &[u8; 64], pad: bool) -> String {
if input.is_empty() {
return String::new();
}
let full_chunks = input.len() / 3;
let remainder = input.len() % 3;
let capacity = if pad {
(full_chunks + usize::from(remainder > 0)) * 4
} else {
full_chunks * 4
+ match remainder {
1 => 2,
2 => 3,
_ => 0,
}
};
let mut out = Vec::with_capacity(capacity);
let chunks = input.chunks_exact(3);
let tail = chunks.remainder();
for chunk in chunks {
let n = (u32::from(chunk[0]) << 16) | (u32::from(chunk[1]) << 8) | u32::from(chunk[2]);
out.push(alphabet[((n >> 18) & 0x3F) as usize]);
out.push(alphabet[((n >> 12) & 0x3F) as usize]);
out.push(alphabet[((n >> 6) & 0x3F) as usize]);
out.push(alphabet[(n & 0x3F) as usize]);
}
match tail.len() {
1 => {
let n = u32::from(tail[0]) << 16;
out.push(alphabet[((n >> 18) & 0x3F) as usize]);
out.push(alphabet[((n >> 12) & 0x3F) as usize]);
if pad {
out.push(b'=');
out.push(b'=');
}
}
2 => {
let n = (u32::from(tail[0]) << 16) | (u32::from(tail[1]) << 8);
out.push(alphabet[((n >> 18) & 0x3F) as usize]);
out.push(alphabet[((n >> 12) & 0x3F) as usize]);
out.push(alphabet[((n >> 6) & 0x3F) as usize]);
if pad {
out.push(b'=');
}
}
_ => {}
}
#[allow(clippy::expect_used)]
String::from_utf8(out).expect("base64 output is always valid ASCII")
}
pub fn base64_decode(input: &str) -> Result<Vec<u8>, Base64DecodeError> {
decode_impl(input.as_bytes(), &STD_DECODE, true)
}
pub fn base64url_decode(input: &str) -> Result<Vec<u8>, Base64DecodeError> {
decode_impl(input.as_bytes(), &URL_DECODE, false)
}
#[allow(clippy::many_single_char_names, clippy::cast_possible_truncation)]
fn decode_impl(
input: &[u8],
decode_table: &[u8; 256],
require_padding: bool,
) -> Result<Vec<u8>, Base64DecodeError> {
if input.is_empty() {
return Ok(Vec::new());
}
let pad_count = input.iter().rev().take_while(|&&b| b == b'=').count();
if pad_count > 2 {
return Err(Base64DecodeError::new(
Base64DecodeErrorKind::InvalidPadding,
));
}
let data = &input[..input.len() - pad_count];
if require_padding {
if input.len() % 4 != 0 {
return Err(Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength));
}
} else {
if data.len() % 4 == 1 {
return Err(Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength));
}
}
let expected_pad = match data.len() % 4 {
0 => 0,
2 => 2,
3 => 1,
_ => return Err(Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength)),
};
let padding_ok = if require_padding {
pad_count == expected_pad
} else {
pad_count == 0 || pad_count == expected_pad
};
if !padding_ok {
return Err(Base64DecodeError::new(
Base64DecodeErrorKind::InvalidPadding,
));
}
let out_len = data.len() * 3 / 4;
let mut out = Vec::with_capacity(out_len);
let chunks = data.chunks_exact(4);
let tail = chunks.remainder();
for chunk in chunks {
let a = decode_char(chunk[0], decode_table)?;
let b = decode_char(chunk[1], decode_table)?;
let c = decode_char(chunk[2], decode_table)?;
let d = decode_char(chunk[3], decode_table)?;
let n = (u32::from(a) << 18) | (u32::from(b) << 12) | (u32::from(c) << 6) | u32::from(d);
out.push((n >> 16) as u8);
out.push((n >> 8) as u8);
out.push(n as u8);
}
match tail.len() {
2 => {
let a = decode_char(tail[0], decode_table)?;
let b = decode_char(tail[1], decode_table)?;
if b & 0b0000_1111 != 0 {
return Err(Base64DecodeError::new(Base64DecodeErrorKind::NonCanonical));
}
let n = (u32::from(a) << 18) | (u32::from(b) << 12);
out.push((n >> 16) as u8);
}
3 => {
let a = decode_char(tail[0], decode_table)?;
let b = decode_char(tail[1], decode_table)?;
let c = decode_char(tail[2], decode_table)?;
if c & 0b0000_0011 != 0 {
return Err(Base64DecodeError::new(Base64DecodeErrorKind::NonCanonical));
}
let n = (u32::from(a) << 18) | (u32::from(b) << 12) | (u32::from(c) << 6);
out.push((n >> 16) as u8);
out.push((n >> 8) as u8);
}
_ => {}
}
Ok(out)
}
#[inline]
fn decode_char(byte: u8, table: &[u8; 256]) -> Result<u8, Base64DecodeError> {
let val = table[byte as usize];
if val == 255 {
return Err(Base64DecodeError::new(
Base64DecodeErrorKind::InvalidCharacter,
));
}
Ok(val)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn encode_empty() {
assert_eq!(base64_encode(b""), "");
}
#[test]
fn encode_f() {
assert_eq!(base64_encode(b"f"), "Zg==");
}
#[test]
fn encode_fo() {
assert_eq!(base64_encode(b"fo"), "Zm8=");
}
#[test]
fn encode_foo() {
assert_eq!(base64_encode(b"foo"), "Zm9v");
}
#[test]
fn encode_foob() {
assert_eq!(base64_encode(b"foob"), "Zm9vYg==");
}
#[test]
fn encode_fooba() {
assert_eq!(base64_encode(b"fooba"), "Zm9vYmE=");
}
#[test]
fn encode_foobar() {
assert_eq!(base64_encode(b"foobar"), "Zm9vYmFy");
}
#[test]
fn decode_empty() {
assert_eq!(base64_decode("").unwrap(), b"");
}
#[test]
fn decode_f() {
assert_eq!(base64_decode("Zg==").unwrap(), b"f");
}
#[test]
fn decode_fo() {
assert_eq!(base64_decode("Zm8=").unwrap(), b"fo");
}
#[test]
fn decode_foo() {
assert_eq!(base64_decode("Zm9v").unwrap(), b"foo");
}
#[test]
fn decode_foob() {
assert_eq!(base64_decode("Zm9vYg==").unwrap(), b"foob");
}
#[test]
fn decode_fooba() {
assert_eq!(base64_decode("Zm9vYmE=").unwrap(), b"fooba");
}
#[test]
fn decode_foobar() {
assert_eq!(base64_decode("Zm9vYmFy").unwrap(), b"foobar");
}
#[test]
fn standard_round_trip_binary() {
let data: Vec<u8> = (0..=255).collect();
let encoded = base64_encode(&data);
let decoded = base64_decode(&encoded).unwrap();
assert_eq!(decoded, data);
}
#[test]
fn standard_round_trip_short_lengths() {
for len in 0..=32_u8 {
let data: Vec<u8> = (0..len).collect();
let encoded = base64_encode(&data);
let decoded = base64_decode(&encoded).unwrap();
assert_eq!(decoded, data, "round-trip failed for length {len}");
}
}
#[test]
fn decode_rejects_invalid_character() {
let err = base64_decode("Zm9!").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidCharacter)
);
}
#[test]
fn decode_rejects_invalid_length() {
let err = base64_decode("AAAAA").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength)
);
}
#[test]
fn decode_rejects_missing_padding() {
let err = base64_decode("Zg").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength)
);
}
#[test]
fn decode_rejects_wrong_padding_count() {
let err = base64_decode("Zm9v=").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength)
);
}
#[test]
fn decode_rejects_triple_padding() {
let err = base64_decode("Z===").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidPadding)
);
}
#[test]
fn decode_rejects_interior_padding() {
let err = base64_decode("Z=g=").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidCharacter)
);
}
#[test]
fn url_encode_empty() {
assert_eq!(base64url_encode(b""), "");
}
#[test]
fn url_encode_f() {
assert_eq!(base64url_encode(b"f"), "Zg");
}
#[test]
fn url_encode_fo() {
assert_eq!(base64url_encode(b"fo"), "Zm8");
}
#[test]
fn url_encode_foo() {
assert_eq!(base64url_encode(b"foo"), "Zm9v");
}
#[test]
fn url_encode_foob() {
assert_eq!(base64url_encode(b"foob"), "Zm9vYg");
}
#[test]
fn url_encode_fooba() {
assert_eq!(base64url_encode(b"fooba"), "Zm9vYmE");
}
#[test]
fn url_encode_foobar() {
assert_eq!(base64url_encode(b"foobar"), "Zm9vYmFy");
}
#[test]
fn url_encode_uses_url_safe_alphabet() {
let input: &[u8] = &[0xFB, 0xFF, 0xFE];
let standard = base64_encode(input);
let url_safe = base64url_encode(input);
assert!(standard.contains('+') || standard.contains('/'));
assert!(!url_safe.contains('+'));
assert!(!url_safe.contains('/'));
}
#[test]
fn url_decode_empty() {
assert_eq!(base64url_decode("").unwrap(), b"");
}
#[test]
fn url_decode_no_padding() {
assert_eq!(base64url_decode("Zg").unwrap(), b"f");
assert_eq!(base64url_decode("Zm8").unwrap(), b"fo");
}
#[test]
fn url_decode_with_optional_padding() {
assert_eq!(base64url_decode("Zg==").unwrap(), b"f");
assert_eq!(base64url_decode("Zm8=").unwrap(), b"fo");
}
#[test]
fn url_round_trip_binary() {
let data: Vec<u8> = (0..=255).collect();
let encoded = base64url_encode(&data);
let decoded = base64url_decode(&encoded).unwrap();
assert_eq!(decoded, data);
}
#[test]
fn url_round_trip_short_lengths() {
for len in 0..=32_u8 {
let data: Vec<u8> = (0..len).collect();
let encoded = base64url_encode(&data);
let decoded = base64url_decode(&encoded).unwrap();
assert_eq!(decoded, data, "url round-trip failed for length {len}");
}
}
#[test]
fn url_decode_rejects_invalid_character() {
let err = base64url_decode("Zm9v!!!").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidCharacter)
);
}
#[test]
fn url_decode_rejects_standard_alphabet_chars() {
assert!(base64url_decode("ab+c").is_err());
assert!(base64url_decode("ab/c").is_err());
}
#[test]
fn url_decode_rejects_invalid_length() {
let err = base64url_decode("A").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength)
);
}
#[test]
fn error_display_messages() {
let invalid_char = Base64DecodeError::new(Base64DecodeErrorKind::InvalidCharacter);
assert_eq!(
invalid_char.to_string(),
"base64: invalid character in input"
);
let invalid_pad = Base64DecodeError::new(Base64DecodeErrorKind::InvalidPadding);
assert_eq!(invalid_pad.to_string(), "base64: invalid padding");
let invalid_len = Base64DecodeError::new(Base64DecodeErrorKind::InvalidLength);
assert_eq!(invalid_len.to_string(), "base64: invalid input length");
}
#[test]
fn error_implements_std_error() {
let err: Box<dyn std::error::Error> = Box::new(Base64DecodeError::new(
Base64DecodeErrorKind::InvalidCharacter,
));
let _ = err.to_string();
}
#[test]
fn decode_rejects_non_canonical_trailing_bits() {
assert_eq!(base64_decode("Zg==").unwrap(), b"f");
let err = base64_decode("Zh==").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::NonCanonical)
);
assert_eq!(base64_decode("Zm8=").unwrap(), b"fo");
let err = base64_decode("Zm9=").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::NonCanonical)
);
assert_eq!(base64url_decode("Zg").unwrap(), b"f");
let err = base64url_decode("Zh").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::NonCanonical)
);
}
#[test]
fn encoder_output_always_round_trips() {
for len in 0..=130usize {
let bytes: Vec<u8> = (0..len)
.map(|i| u8::try_from((i * 31 + 7) % 256).unwrap())
.collect();
let std = base64_encode(&bytes);
assert_eq!(base64_decode(&std).unwrap(), bytes, "std len {len}");
let url = base64url_encode(&bytes);
assert_eq!(base64url_decode(&url).unwrap(), bytes, "url len {len}");
}
}
#[test]
fn decode_url_rejects_stray_padding() {
assert_eq!(base64url_decode("Zg").unwrap(), b"f");
let err = base64url_decode("Zg=").unwrap_err();
assert_eq!(
err,
Base64DecodeError::new(Base64DecodeErrorKind::InvalidPadding)
);
}
}