use sha2::{Digest, Sha256};
use std::{
fmt,
fs::File,
io::{self, Read},
path::Path,
str::FromStr,
};
#[cfg(test)]
mod tests;
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct Sha256Digest([u8; 32]);
impl Sha256Digest {
#[must_use]
pub const fn from_bytes(bytes: [u8; 32]) -> Self {
Self(bytes)
}
#[must_use]
pub const fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
#[must_use]
pub fn compute(bytes: &[u8]) -> Self {
Self(Sha256::digest(bytes).into())
}
}
impl fmt::Display for Sha256Digest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for byte in self.0 {
write!(f, "{byte:02x}")?;
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum DigestParseError {
Length {
actual: usize,
},
Digit {
offset: usize,
},
}
impl fmt::Display for DigestParseError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Length { actual } => {
write!(f, "SHA-256 requires 64 hex bytes, received {actual}")
}
Self::Digit { offset } => write!(f, "invalid lowercase SHA-256 digit at byte {offset}"),
}
}
}
impl std::error::Error for DigestParseError {}
impl FromStr for Sha256Digest {
type Err = DigestParseError;
fn from_str(text: &str) -> Result<Self, Self::Err> {
if text.len() != 64 {
return Err(DigestParseError::Length { actual: text.len() });
}
let mut bytes = [0; 32];
for (offset, digit) in text.bytes().enumerate() {
let nibble = match digit {
b'0'..=b'9' => digit - b'0',
b'a'..=b'f' => digit - b'a' + 10,
_ => return Err(DigestParseError::Digit { offset }),
};
bytes[offset / 2] |= nibble << if offset % 2 == 0 { 4 } else { 0 };
}
Ok(Self(bytes))
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ArtifactIdentity {
pub bytes: u64,
pub sha256: Sha256Digest,
}
#[derive(Debug)]
pub enum ArtifactError {
Io(io::Error),
NotRegularFile,
LimitExceeded {
limit: u64,
},
DigestMismatch {
expected: Sha256Digest,
actual: ArtifactIdentity,
},
Allocation(std::collections::TryReserveError),
}
impl fmt::Display for ArtifactError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io(source) => write!(f, "artifact read failed: {source}"),
Self::NotRegularFile => f.write_str("artifact path is not a regular file"),
Self::LimitExceeded { limit } => write!(f, "artifact exceeds {limit} bytes"),
Self::DigestMismatch { expected, actual } => write!(
f,
"artifact SHA-256 is {}, expected {expected}",
actual.sha256
),
Self::Allocation(source) => write!(f, "artifact allocation failed: {source}"),
}
}
}
impl std::error::Error for ArtifactError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Io(source) => Some(source),
Self::Allocation(source) => Some(source),
_ => None,
}
}
}
impl From<io::Error> for ArtifactError {
fn from(source: io::Error) -> Self {
Self::Io(source)
}
}
pub fn hash_reader(
mut reader: impl Read,
max_bytes: u64,
) -> Result<ArtifactIdentity, ArtifactError> {
let mut hasher = Sha256::new();
let bytes = visit_reader(&mut reader, max_bytes, |chunk| {
hasher.update(chunk);
Ok(())
})?;
Ok(ArtifactIdentity {
bytes,
sha256: Sha256Digest(hasher.finalize().into()),
})
}
pub fn verify_reader(
reader: impl Read,
max_bytes: u64,
expected: Sha256Digest,
) -> Result<ArtifactIdentity, ArtifactError> {
let actual = hash_reader(reader, max_bytes)?;
if actual.sha256 != expected {
return Err(ArtifactError::DigestMismatch { expected, actual });
}
Ok(actual)
}
pub fn hash_file(path: &Path, max_bytes: u64) -> Result<ArtifactIdentity, ArtifactError> {
hash_reader(open_file(path, max_bytes)?, max_bytes)
}
pub fn read_file(path: &Path, max_bytes: usize) -> Result<Vec<u8>, ArtifactError> {
read_reader(open_file(path, max_bytes as u64)?, max_bytes)
}
pub fn read_reader(mut reader: impl Read, max_bytes: usize) -> Result<Vec<u8>, ArtifactError> {
let mut bytes = Vec::new();
visit_reader(&mut reader, max_bytes as u64, |chunk| {
bytes
.try_reserve_exact(chunk.len())
.map_err(ArtifactError::Allocation)?;
bytes.extend_from_slice(chunk);
Ok(())
})?;
Ok(bytes)
}
fn open_file(path: &Path, limit: u64) -> Result<File, ArtifactError> {
check_metadata(&std::fs::metadata(path)?, limit)?;
let file = File::open(path)?;
check_metadata(&file.metadata()?, limit)?;
Ok(file)
}
fn check_metadata(metadata: &std::fs::Metadata, limit: u64) -> Result<(), ArtifactError> {
if !metadata.is_file() {
return Err(ArtifactError::NotRegularFile);
}
if metadata.len() > limit {
return Err(ArtifactError::LimitExceeded { limit });
}
Ok(())
}
fn visit_reader(
reader: &mut impl Read,
limit: u64,
mut visit: impl FnMut(&[u8]) -> Result<(), ArtifactError>,
) -> Result<u64, ArtifactError> {
let mut bytes = 0_u64;
let mut buffer = [0_u8; 16 * 1024];
loop {
let remaining = usize::try_from(limit - bytes).unwrap_or(usize::MAX);
let allowance = remaining.saturating_add(1).min(buffer.len());
let count = match reader.read(&mut buffer[..allowance]) {
Ok(count) => count,
Err(source) if source.kind() == io::ErrorKind::Interrupted => continue,
Err(source) => return Err(source.into()),
};
if count == 0 {
return Ok(bytes);
}
if count as u64 > limit - bytes {
return Err(ArtifactError::LimitExceeded { limit });
}
visit(&buffer[..count])?;
bytes += count as u64;
}
}