const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
pub(super) struct DataUri<'a> {
pub(super) media_type: String,
pub(super) payload: &'a str,
pub(super) base64: bool,
}
pub(super) fn parse(value: &str) -> Option<DataUri<'_>> {
let rest = strip_prefix_ignoring_case(value.trim_start(), "data:")?;
let (meta, payload) = rest.split_once(',')?;
let mut parameters = meta.split(';');
let media_type = parameters.next().unwrap_or_default().trim().to_lowercase();
let base64 = parameters.any(|p| p.trim().eq_ignore_ascii_case("base64"));
Some(DataUri {
media_type,
payload,
base64,
})
}
fn strip_prefix_ignoring_case<'a>(value: &'a str, prefix: &str) -> Option<&'a str> {
let head = value.get(..prefix.len())?;
head.eq_ignore_ascii_case(prefix)
.then(|| value.get(prefix.len()..))
.flatten()
}
pub(super) fn decode(payload: &str, budget: &mut u64) -> Option<Vec<u8>> {
let mut count = 0usize;
for byte in significant(payload) {
symbol(byte)?;
count = count.checked_add(1)?;
}
let remainder = count.checked_rem(4)?;
if remainder == 1 {
return None;
}
let len = count
.checked_div(4)?
.checked_mul(3)?
.checked_add(remainder.saturating_sub(1))?;
*budget = budget.checked_sub(u64::try_from(len).ok()?)?;
let mut out = Vec::with_capacity(len);
let mut accumulator = 0u32;
let mut bits = 0u32;
for byte in significant(payload) {
accumulator = (accumulator << 6) | u32::from(symbol(byte)?);
bits = bits.saturating_add(6);
if bits >= 8 {
bits = bits.saturating_sub(8);
out.push(u8::try_from((accumulator >> bits) & 0xFF).ok()?);
}
}
Some(out)
}
fn significant(payload: &str) -> impl Iterator<Item = u8> + '_ {
payload
.bytes()
.take_while(|byte| *byte != b'=')
.filter(|byte| !byte.is_ascii_whitespace())
}
fn symbol(byte: u8) -> Option<u8> {
match byte {
b'A'..=b'Z' => Some(byte.saturating_sub(b'A')),
b'a'..=b'z' => Some(byte.saturating_sub(b'a').saturating_add(26)),
b'0'..=b'9' => Some(byte.saturating_sub(b'0').saturating_add(52)),
b'+' => Some(62),
b'/' => Some(63),
_ => None,
}
}
pub(super) fn encode(media_type: &str, bytes: &[u8]) -> String {
let mut out = String::with_capacity(
bytes
.len()
.saturating_div(3)
.saturating_mul(4)
.saturating_add(media_type.len())
.saturating_add(16),
);
out.push_str("data:");
out.push_str(media_type);
out.push_str(";base64,");
for chunk in bytes.chunks(3) {
let (a, b, c) = (
u32::from(*chunk.first().unwrap_or(&0)),
u32::from(*chunk.get(1).unwrap_or(&0)),
u32::from(*chunk.get(2).unwrap_or(&0)),
);
let triple = (a << 16) | (b << 8) | c;
for shift in [18u32, 12, 6, 0] {
let index = usize::try_from((triple >> shift) & 0x3F).unwrap_or(0);
out.push(char::from(*ALPHABET.get(index).unwrap_or(&b'A')));
}
let pad = 3usize.saturating_sub(chunk.len());
out.truncate(out.len().saturating_sub(pad));
for _ in 0..pad {
out.push('=');
}
}
out
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used, clippy::indexing_slicing)]
use super::*;
fn round_trip(bytes: &[u8]) {
let uri = encode("image/png", bytes);
let parsed = parse(&uri).unwrap();
assert_eq!(parsed.media_type, "image/png");
assert!(parsed.base64);
let mut budget = u64::MAX;
assert_eq!(decode(parsed.payload, &mut budget).unwrap(), bytes);
}
#[test]
fn every_length_up_to_a_few_blocks_round_trips() {
for len in 0..64usize {
let bytes: Vec<u8> = (0..len)
.map(|i| u8::try_from(i % 251).unwrap_or(0))
.collect();
round_trip(&bytes);
}
round_trip(&[0x89, b'P', b'N', b'G', 0x0D, 0x0A, 0x1A, 0x0A]);
}
#[test]
fn the_encoding_matches_the_reference_vectors() {
for (input, expected) in [
(&b""[..], ""),
(b"f", "Zg=="),
(b"fo", "Zm8="),
(b"foo", "Zm9v"),
(b"foob", "Zm9vYg=="),
(b"fooba", "Zm9vYmE="),
(b"foobar", "Zm9vYmFy"),
] {
let uri = encode("text/plain", input);
assert_eq!(
uri,
format!("data:text/plain;base64,{expected}"),
"{input:?}"
);
}
}
#[test]
fn a_wrapped_payload_decodes_and_a_corrupt_one_does_not() {
let mut budget = u64::MAX;
assert_eq!(decode("Zm9v\n YmFy", &mut budget).unwrap(), b"foobar");
let mut budget = u64::MAX;
assert!(decode("Zm9v*YmFy", &mut budget).is_none());
let mut budget = u64::MAX;
assert!(decode("Zm9vY", &mut budget).is_none());
}
#[test]
fn decoding_is_charged_before_the_memory_is_committed() {
let payload = encode("image/png", &vec![0u8; 4096]);
let payload = payload.split_once(',').unwrap().1;
let mut plenty = u64::MAX;
assert!(decode(payload, &mut plenty).is_some());
let mut stingy = 100u64;
assert!(decode(payload, &mut stingy).is_none());
assert_eq!(stingy, 100, "a refused decode must not spend the budget");
}
#[test]
fn the_media_type_is_parsed_but_never_trusted_to_choose_a_handler() {
let uri = parse("data:image/jpeg;charset=utf-8;base64,Zm9v").unwrap();
assert_eq!(uri.media_type, "image/jpeg");
assert!(uri.base64, "`;base64` is a parameter, not a fixed position");
let plain = parse("data:image/png,%89PNG").unwrap();
assert!(!plain.base64);
assert!(parse("DATA:image/png;base64,Zg==").is_some());
assert!(parse("https://example.invalid/x.png").is_none());
assert!(parse("data:no-comma").is_none());
}
#[test]
fn arbitrary_text_does_not_panic_the_codec() {
for case in [
"",
"data:",
"data:,",
"data:;base64,",
",",
"data:a/b;base64",
"=",
"====",
"data:image/png;base64,=A",
"\u{feff}data:x,y",
] {
let _ = parse(case);
let _ = decode(case, &mut 1024);
}
}
}