1use std::fmt;
10use std::fs;
11use std::io;
12use std::ops;
13use std::path::Path;
14use std::str::FromStr;
15
16use hex::FromHex;
17use hex::FromHexError;
18use serde::de;
19use serde::de::Deserialize;
20use serde::de::Deserializer;
21use serde::de::Visitor;
22use serde::ser::Serialize;
23use serde::ser::Serializer;
24use sha2::Digest as _;
25use sha2::Sha256;
26
27const DIGEST_LENGTH: usize = 32;
28
29#[derive(Copy, Clone, Ord, PartialOrd, Eq, PartialEq, Hash, Default)]
31pub struct Digest([u8; DIGEST_LENGTH]);
32
33impl Digest {
34 pub fn new(data: &[u8]) -> Self {
36 Digest(Sha256::digest(data).into())
37 }
38
39 pub fn digest_reader<R: io::Read>(mut reader: R) -> io::Result<Self> {
41 let mut hasher = Sha256::default();
42 let mut buf = [0u8; 16384];
43
44 loop {
45 let n = reader.read(&mut buf)?;
46 if n == 0 {
47 break;
48 }
49 hasher.update(&buf[..n]);
50 }
51
52 Ok(Digest(hasher.finalize().into()))
53 }
54
55 pub fn digest_path<P: AsRef<Path>>(path: P) -> io::Result<Self> {
57 Self::digest_reader(fs::File::open(path)?)
58 }
59}
60
61impl fmt::LowerHex for Digest {
62 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
63 for byte in &self.0 {
64 write!(f, "{byte:02x}")?;
65 }
66 Ok(())
67 }
68}
69
70impl fmt::UpperHex for Digest {
71 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
72 for byte in &self.0 {
73 write!(f, "{byte:02X}")?;
74 }
75 Ok(())
76 }
77}
78
79impl fmt::Display for Digest {
80 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
81 <Self as fmt::LowerHex>::fmt(self, f)
82 }
83}
84
85impl fmt::Debug for Digest {
86 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
87 <Self as fmt::Display>::fmt(self, f)
88 }
89}
90
91impl ops::Deref for Digest {
92 type Target = [u8; DIGEST_LENGTH];
93
94 fn deref(&self) -> &Self::Target {
95 &self.0
96 }
97}
98
99impl Serialize for Digest {
100 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
101 where
102 S: Serializer,
103 {
104 if serializer.is_human_readable() {
105 serializer.serialize_str(&self.to_string())
106 } else {
107 serializer.serialize_bytes(self.0.as_ref())
108 }
109 }
110}
111
112impl<'de> Deserialize<'de> for Digest {
113 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
114 where
115 D: Deserializer<'de>,
116 {
117 struct DigestVisitor;
118
119 impl<'de> Visitor<'de> for DigestVisitor {
120 type Value = Digest;
121
122 fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
123 write!(f, "hex string or 32 bytes")
124 }
125
126 fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
127 where
128 E: de::Error,
129 {
130 let v = <[u8; DIGEST_LENGTH]>::from_hex(v).map_err(|e| match e {
131 FromHexError::InvalidHexCharacter { c, .. } => E::invalid_value(
132 de::Unexpected::Char(c),
133 &"string with only hexadecimal characters",
134 ),
135 FromHexError::InvalidStringLength => {
136 E::invalid_length(v.len(), &"hex string with a valid length")
137 }
138 FromHexError::OddLength => {
139 E::invalid_length(v.len(), &"hex string with an even length")
140 }
141 })?;
142 Ok(Digest(v))
143 }
144
145 fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
146 where
147 E: de::Error,
148 {
149 if v.len() != DIGEST_LENGTH {
150 return Err(E::invalid_length(v.len(), &"32 bytes"));
151 }
152 let mut inner = [0; DIGEST_LENGTH];
153 inner.copy_from_slice(v);
154 Ok(Digest(inner))
155 }
156 }
157
158 if deserializer.is_human_readable() {
159 deserializer.deserialize_str(DigestVisitor)
160 } else {
161 deserializer.deserialize_bytes(DigestVisitor)
162 }
163 }
164}
165
166impl FromHex for Digest {
167 type Error = FromHexError;
168
169 fn from_hex<T: AsRef<[u8]>>(hex: T) -> Result<Self, Self::Error> {
170 <[u8; DIGEST_LENGTH]>::from_hex(hex).map(Digest)
171 }
172}
173
174impl AsRef<[u8]> for Digest {
175 fn as_ref(&self) -> &[u8] {
176 &self.0
177 }
178}
179
180impl FromStr for Digest {
181 type Err = FromHexError;
182
183 fn from_str(s: &str) -> Result<Self, Self::Err> {
184 Digest::from_hex(s)
185 }
186}
187
188#[cfg(test)]
189mod test {
190 use super::*;
191
192 #[test]
193 fn test_from_hex() {
194 let s = "b1fbeefc23e6a149a6f7d0c2fb635bfc78f7ddc2da963ea9c6a63eb324260e6d";
195 assert_eq!(Digest::from_str(s).unwrap().to_string(), s);
196 }
197
198 #[test]
211 fn digest_is_independent_of_optimisation_level() {
212 assert_eq!(
213 Digest::new(b"").to_string(),
214 "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
215 );
216 assert_eq!(
217 Digest::new(b"abc").to_string(),
218 "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad"
219 );
220 let long = vec![0x61u8; 1_000_000];
224 assert_eq!(
225 Digest::new(&long).to_string(),
226 "cdc76e5c9914fb9281a1c7e284d73e67f1809a48a497200e046d39ccc7112cd0"
227 );
228 }
229}