#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Encoding {
Base64,
Base64Url,
}
pub const SUPPORTED_ENCODINGS: [&str; 2] = ["base64", "base64url"];
impl Encoding {
pub fn from_name(name: &str) -> Option<Self> {
match name {
"base64" => Some(Self::Base64),
"base64url" => Some(Self::Base64Url),
_ => None,
}
}
pub fn name(self) -> &'static str {
match self {
Self::Base64 => "base64",
Self::Base64Url => "base64url",
}
}
pub fn pattern(self) -> &'static str {
match self {
Self::Base64 => concat!(
"^(?:[A-Za-z0-9+/]{4})*",
"(?:[A-Za-z0-9+/][AQgw]==|[A-Za-z0-9+/]{2}[AEIMQUYcgkosw048]=)?$"
),
Self::Base64Url => concat!(
"^(?:[A-Za-z0-9_-]{4})*",
"(?:[A-Za-z0-9_-][AQgw]|[A-Za-z0-9_-]{2}[AEIMQUYcgkosw048])?$"
),
}
}
}
pub enum EncodingClass {
Supported(Encoding),
Unsupported,
}
pub fn classify(name: &str) -> EncodingClass {
match Encoding::from_name(name) {
Some(encoding) => EncodingClass::Supported(encoding),
None => EncodingClass::Unsupported,
}
}
pub fn is_valid(encoding: Encoding, value: &str) -> bool {
regex::Regex::new(encoding.pattern())
.expect("pinned contentEncoding pattern compiles")
.is_match(value)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pinned_patterns_pass_the_pattern_gate() {
for name in SUPPORTED_ENCODINGS {
let encoding = Encoding::from_name(name).expect("supported");
crate::json_schema::pattern::gate_and_normalize(encoding.pattern()).unwrap_or_else(
|error| panic!("{name} pinned pattern rejected by gate: {error:?}"),
);
}
}
#[test]
fn classify_partitions_names() {
assert!(matches!(
classify("base64"),
EncodingClass::Supported(Encoding::Base64)
));
assert!(matches!(
classify("base64url"),
EncodingClass::Supported(Encoding::Base64Url)
));
for name in [
"base16",
"quoted-printable",
"7bit",
"8bit",
"binary",
"hex",
] {
assert!(
matches!(classify(name), EncodingClass::Unsupported),
"{name}"
);
}
}
#[test]
fn base64_accepts_canonical_padded_rejects_url_and_unpadded() {
assert!(is_valid(Encoding::Base64, "Pj4+"));
assert!(is_valid(Encoding::Base64, "aGk="));
assert!(is_valid(Encoding::Base64, ""));
assert!(!is_valid(Encoding::Base64, "Pj4-"));
assert!(!is_valid(Encoding::Base64, "a-b_"));
assert!(!is_valid(Encoding::Base64, "aGk"));
assert!(!is_valid(Encoding::Base64, "aG k="));
assert!(!is_valid(Encoding::Base64, "aGk=\n"));
assert!(!is_valid(Encoding::Base64, "aGk!"));
}
#[test]
fn base64url_accepts_canonical_unpadded_rejects_std_and_padding() {
assert!(is_valid(Encoding::Base64Url, "Pj4-"));
assert!(is_valid(Encoding::Base64Url, "aGk"));
assert!(is_valid(Encoding::Base64Url, ""));
assert!(!is_valid(Encoding::Base64Url, "Pj4+"));
assert!(!is_valid(Encoding::Base64Url, "aGk="));
assert!(!is_valid(Encoding::Base64Url, "aGk=A"));
assert!(!is_valid(Encoding::Base64Url, "a"));
}
#[test]
fn rejects_non_canonical_trailing_bits() {
assert!(is_valid(Encoding::Base64, "aGk="));
assert!(!is_valid(Encoding::Base64, "aGl="));
assert!(is_valid(Encoding::Base64, "//8="));
assert!(!is_valid(Encoding::Base64, "//9="));
assert!(is_valid(Encoding::Base64, "AA=="));
assert!(!is_valid(Encoding::Base64, "AB=="));
assert!(is_valid(Encoding::Base64, "/w=="));
assert!(!is_valid(Encoding::Base64, "/x=="));
assert!(is_valid(Encoding::Base64Url, "aGk"));
assert!(!is_valid(Encoding::Base64Url, "aGl"));
assert!(is_valid(Encoding::Base64Url, "AA"));
assert!(!is_valid(Encoding::Base64Url, "AB"));
assert!(is_valid(Encoding::Base64, "YWJjaGk="));
assert!(!is_valid(Encoding::Base64, "YWJjaGl="));
for high in 0u16..=255 {
let one = [high as u8];
assert!(is_valid(Encoding::Base64, &encode(&one, true)), "{one:?}");
assert!(
is_valid(Encoding::Base64Url, &encode(&one, false)),
"{one:?}"
);
let two = [high as u8, (high * 7 % 256) as u8];
assert!(is_valid(Encoding::Base64, &encode(&two, true)), "{two:?}");
assert!(
is_valid(Encoding::Base64Url, &encode(&two, false)),
"{two:?}"
);
}
}
fn encode(bytes: &[u8], standard: bool) -> String {
const STD: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
const URL: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
let alphabet = if standard { STD } else { URL };
let mut out = String::new();
for chunk in bytes.chunks(3) {
let mut buffer = [0u8; 3];
buffer[..chunk.len()].copy_from_slice(chunk);
let triple =
u32::from(buffer[0]) << 16 | u32::from(buffer[1]) << 8 | u32::from(buffer[2]);
let significant = chunk.len() + 1;
for index in 0..significant {
let shift = 18 - 6 * index;
out.push(alphabet[((triple >> shift) & 0x3F) as usize] as char);
}
if standard {
for _ in significant..4 {
out.push('=');
}
}
}
out
}
}