use std::fmt;
#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub enum HashAlgorithm {
Sha256,
Sha512,
Sha1,
}
impl HashAlgorithm {
pub const fn digest_len(self) -> usize {
match self {
HashAlgorithm::Sha256 => 32,
HashAlgorithm::Sha512 => 64,
HashAlgorithm::Sha1 => 20,
}
}
pub const fn name(self) -> &'static str {
match self {
HashAlgorithm::Sha256 => "sha256",
HashAlgorithm::Sha512 => "sha512",
HashAlgorithm::Sha1 => "sha1",
}
}
pub fn from_name(name: &str) -> Option<HashAlgorithm> {
[
HashAlgorithm::Sha256,
HashAlgorithm::Sha512,
HashAlgorithm::Sha1,
]
.into_iter()
.find(|algorithm| name.eq_ignore_ascii_case(algorithm.name()))
}
}
impl fmt::Display for HashAlgorithm {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.name())
}
}
#[derive(Clone, PartialEq, Eq, Hash, Debug)]
pub struct Digest {
pub algorithm: HashAlgorithm,
pub bytes: Box<[u8]>,
}
impl Digest {
pub fn parse_hex(algorithm: HashAlgorithm, hex_text: &str) -> Result<Digest, InvalidDigest> {
let bytes = hex::decode(hex_text).map_err(|_| InvalidDigest::Encoding {
algorithm,
form: "hexadecimal",
})?;
Digest::from_bytes(algorithm, bytes)
}
pub fn parse_sri_entry(entry: &str) -> Result<Digest, InvalidDigest> {
let entry = entry.trim();
let (name, rest) = entry
.split_once('-')
.ok_or_else(|| InvalidDigest::NotAnSriEntry(entry.to_owned()))?;
let algorithm = HashAlgorithm::from_name(name)
.ok_or_else(|| InvalidDigest::UnsupportedAlgorithm(name.to_owned()))?;
let base64_text = rest.split_once('?').map_or(rest, |(digest, _)| digest);
let bytes = decode_base64(base64_text).ok_or(InvalidDigest::Encoding {
algorithm,
form: "base64",
})?;
Digest::from_bytes(algorithm, bytes)
}
fn from_bytes(algorithm: HashAlgorithm, bytes: Vec<u8>) -> Result<Digest, InvalidDigest> {
if bytes.len() != algorithm.digest_len() {
return Err(InvalidDigest::Length {
algorithm,
expected: algorithm.digest_len(),
found: bytes.len(),
});
}
Ok(Digest {
algorithm,
bytes: bytes.into_boxed_slice(),
})
}
pub fn to_hex_lowercase(&self) -> String {
hex::encode(&self.bytes)
}
}
impl fmt::Display for Digest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}-{}", self.algorithm, self.to_hex_lowercase())
}
}
fn decode_base64(text: &str) -> Option<Vec<u8>> {
let mut out = Vec::with_capacity(text.len() / 4 * 3);
let mut accumulator: u32 = 0;
let mut bits: u32 = 0;
let mut padding = 0usize;
for &byte in text.as_bytes() {
if byte == b'=' {
padding += 1;
continue;
}
if padding > 0 {
return None;
}
let value = match byte {
b'A'..=b'Z' => byte - b'A',
b'a'..=b'z' => byte - b'a' + 26,
b'0'..=b'9' => byte - b'0' + 52,
b'+' => 62,
b'/' => 63,
_ => return None,
} as u32;
accumulator = (accumulator << 6) | value;
bits += 6;
if bits >= 8 {
bits -= 8;
out.push((accumulator >> bits) as u8);
accumulator &= (1 << bits) - 1;
}
}
if padding > 2 || bits >= 6 || accumulator != 0 {
return None;
}
Some(out)
}
#[derive(Clone, PartialEq, Eq, Debug)]
pub enum InvalidDigest {
UnsupportedAlgorithm(String),
NotAnSriEntry(String),
Length {
algorithm: HashAlgorithm,
expected: usize,
found: usize,
},
Encoding {
algorithm: HashAlgorithm,
form: &'static str,
},
}
impl fmt::Display for InvalidDigest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
InvalidDigest::UnsupportedAlgorithm(name) => {
write!(f, "unsupported hash algorithm `{name}`")
}
InvalidDigest::NotAnSriEntry(entry) => {
write!(f, "`{entry}` is not an `<algorithm>-<base64>` SRI entry")
}
InvalidDigest::Length {
algorithm,
expected,
found,
} => write!(
f,
"a {algorithm} digest is {expected} bytes, but this one decodes to {found}"
),
InvalidDigest::Encoding { algorithm, form } => {
write!(f, "the {algorithm} digest is not valid {form}")
}
}
}
}
impl std::error::Error for InvalidDigest {}