1use sha2::{Digest, Sha256};
8use std::{
9 fmt,
10 fs::File,
11 io::{self, Read},
12 path::Path,
13 str::FromStr,
14};
15
16#[cfg(unix)]
17mod no_follow;
18#[cfg(test)]
19mod tests;
20
21#[cfg(unix)]
22pub use no_follow::read_file_no_follow;
23
24#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
26pub struct Sha256Digest([u8; 32]);
27
28impl Sha256Digest {
29 #[must_use]
31 pub const fn from_bytes(bytes: [u8; 32]) -> Self {
32 Self(bytes)
33 }
34
35 #[must_use]
37 pub const fn as_bytes(&self) -> &[u8; 32] {
38 &self.0
39 }
40
41 #[must_use]
43 pub fn compute(bytes: &[u8]) -> Self {
44 Self(Sha256::digest(bytes).into())
45 }
46}
47
48impl fmt::Display for Sha256Digest {
49 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
50 for byte in self.0 {
51 write!(f, "{byte:02x}")?;
52 }
53 Ok(())
54 }
55}
56
57#[derive(Clone, Copy, Debug, Eq, PartialEq)]
59pub enum DigestParseError {
60 Length {
62 actual: usize,
64 },
65 Digit {
67 offset: usize,
69 },
70}
71
72impl fmt::Display for DigestParseError {
73 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
74 match self {
75 Self::Length { actual } => {
76 write!(f, "SHA-256 requires 64 hex bytes, received {actual}")
77 }
78 Self::Digit { offset } => write!(f, "invalid lowercase SHA-256 digit at byte {offset}"),
79 }
80 }
81}
82impl std::error::Error for DigestParseError {}
83
84impl FromStr for Sha256Digest {
85 type Err = DigestParseError;
86
87 fn from_str(text: &str) -> Result<Self, Self::Err> {
88 if text.len() != 64 {
89 return Err(DigestParseError::Length { actual: text.len() });
90 }
91 let mut bytes = [0; 32];
92 for (offset, digit) in text.bytes().enumerate() {
93 let nibble = match digit {
94 b'0'..=b'9' => digit - b'0',
95 b'a'..=b'f' => digit - b'a' + 10,
96 _ => return Err(DigestParseError::Digit { offset }),
97 };
98 bytes[offset / 2] |= nibble << if offset % 2 == 0 { 4 } else { 0 };
99 }
100 Ok(Self(bytes))
101 }
102}
103
104#[derive(Clone, Copy, Debug, Eq, PartialEq)]
106pub struct ArtifactIdentity {
107 pub bytes: u64,
109 pub sha256: Sha256Digest,
111}
112
113#[derive(Debug)]
115pub enum ArtifactError {
116 Io(io::Error),
118 NotRegularFile,
120 LimitExceeded {
122 limit: u64,
124 },
125 DigestMismatch {
127 expected: Sha256Digest,
129 actual: ArtifactIdentity,
131 },
132 Allocation(std::collections::TryReserveError),
134}
135
136impl fmt::Display for ArtifactError {
137 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
138 match self {
139 Self::Io(source) => write!(f, "artifact read failed: {source}"),
140 Self::NotRegularFile => f.write_str("artifact path is not a regular file"),
141 Self::LimitExceeded { limit } => write!(f, "artifact exceeds {limit} bytes"),
142 Self::DigestMismatch { expected, actual } => write!(
143 f,
144 "artifact SHA-256 is {}, expected {expected}",
145 actual.sha256
146 ),
147 Self::Allocation(source) => write!(f, "artifact allocation failed: {source}"),
148 }
149 }
150}
151impl std::error::Error for ArtifactError {
152 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
153 match self {
154 Self::Io(source) => Some(source),
155 Self::Allocation(source) => Some(source),
156 _ => None,
157 }
158 }
159}
160impl From<io::Error> for ArtifactError {
161 fn from(source: io::Error) -> Self {
162 Self::Io(source)
163 }
164}
165
166pub fn hash_reader(
175 mut reader: impl Read,
176 max_bytes: u64,
177) -> Result<ArtifactIdentity, ArtifactError> {
178 let mut hasher = Sha256::new();
179 let bytes = visit_reader(&mut reader, max_bytes, |chunk| {
180 hasher.update(chunk);
181 Ok(())
182 })?;
183 Ok(ArtifactIdentity {
184 bytes,
185 sha256: Sha256Digest(hasher.finalize().into()),
186 })
187}
188
189pub fn verify_reader(
194 reader: impl Read,
195 max_bytes: u64,
196 expected: Sha256Digest,
197) -> Result<ArtifactIdentity, ArtifactError> {
198 let actual = hash_reader(reader, max_bytes)?;
199 if actual.sha256 != expected {
200 return Err(ArtifactError::DigestMismatch { expected, actual });
201 }
202 Ok(actual)
203}
204
205pub fn hash_file(path: &Path, max_bytes: u64) -> Result<ArtifactIdentity, ArtifactError> {
213 hash_reader(open_file(path, max_bytes)?, max_bytes)
214}
215
216pub fn read_file(path: &Path, max_bytes: usize) -> Result<Vec<u8>, ArtifactError> {
224 read_reader(open_file(path, max_bytes as u64)?, max_bytes)
225}
226
227pub fn read_opened_file(file: File, max_bytes: usize) -> Result<Vec<u8>, ArtifactError> {
243 check_metadata(&file.metadata()?, max_bytes as u64)?;
244 read_reader(file, max_bytes)
245}
246
247pub fn read_reader(mut reader: impl Read, max_bytes: usize) -> Result<Vec<u8>, ArtifactError> {
255 let mut bytes = Vec::new();
256 visit_reader(&mut reader, max_bytes as u64, |chunk| {
257 bytes
258 .try_reserve_exact(chunk.len())
259 .map_err(ArtifactError::Allocation)?;
260 bytes.extend_from_slice(chunk);
261 Ok(())
262 })?;
263 Ok(bytes)
264}
265
266fn open_file(path: &Path, limit: u64) -> Result<File, ArtifactError> {
267 check_metadata(&std::fs::metadata(path)?, limit)?;
270 let file = File::open(path)?;
271 check_metadata(&file.metadata()?, limit)?;
272 Ok(file)
273}
274
275fn check_metadata(metadata: &std::fs::Metadata, limit: u64) -> Result<(), ArtifactError> {
276 if !metadata.is_file() {
277 return Err(ArtifactError::NotRegularFile);
278 }
279 if metadata.len() > limit {
280 return Err(ArtifactError::LimitExceeded { limit });
281 }
282 Ok(())
283}
284
285fn visit_reader(
286 reader: &mut impl Read,
287 limit: u64,
288 mut visit: impl FnMut(&[u8]) -> Result<(), ArtifactError>,
289) -> Result<u64, ArtifactError> {
290 let mut bytes = 0_u64;
291 let mut buffer = [0_u8; 16 * 1024];
292 loop {
293 let remaining = usize::try_from(limit - bytes).unwrap_or(usize::MAX);
294 let allowance = remaining.saturating_add(1).min(buffer.len());
295 let count = match reader.read(&mut buffer[..allowance]) {
296 Ok(count) => count,
297 Err(source) if source.kind() == io::ErrorKind::Interrupted => continue,
298 Err(source) => return Err(source.into()),
299 };
300 if count == 0 {
301 return Ok(bytes);
302 }
303 if count as u64 > limit - bytes {
304 return Err(ArtifactError::LimitExceeded { limit });
305 }
306 visit(&buffer[..count])?;
307 bytes += count as u64;
308 }
309}