1use serde::{Deserialize, Deserializer, Serialize, Serializer};
2use sha2::{Digest, Sha256};
3use std::fmt;
4use std::hash::{Hash, Hasher};
5use std::marker::PhantomData;
6use std::str::FromStr;
7
8use crate::{
9 CanonicalizationProfile, HashError, HashKindName, Kind, HASH_ALGORITHM, HASH_PROTOCOL_LABEL,
10 HASH_PROTOCOL_VERSION,
11};
12
13pub struct HashId<K: Kind> {
14 digest: [u8; 32],
15 marker: PhantomData<fn() -> K>,
16}
17
18impl<K: Kind> HashId<K> {
19 pub const fn from_digest(digest: [u8; 32]) -> Self {
20 Self {
21 digest,
22 marker: PhantomData,
23 }
24 }
25
26 pub const fn digest(&self) -> &[u8; 32] {
27 &self.digest
28 }
29
30 pub fn digest_hex(&self) -> String {
31 hex::encode(self.digest)
32 }
33
34 pub fn into_any(self) -> AnyHashId {
35 AnyHashId {
36 kind: K::NAME,
37 digest: self.digest,
38 }
39 }
40}
41
42impl<K: Kind> Copy for HashId<K> {}
43
44impl<K: Kind> Clone for HashId<K> {
45 fn clone(&self) -> Self {
46 *self
47 }
48}
49
50impl<K: Kind> PartialEq for HashId<K> {
51 fn eq(&self, other: &Self) -> bool {
52 self.digest == other.digest
53 }
54}
55
56impl<K: Kind> Eq for HashId<K> {}
57
58impl<K: Kind> Hash for HashId<K> {
59 fn hash<H: Hasher>(&self, state: &mut H) {
60 self.digest.hash(state);
61 }
62}
63
64impl<K: Kind> fmt::Display for HashId<K> {
65 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
66 write!(
67 formatter,
68 "arete:h{}:{}:{}:{}",
69 HASH_PROTOCOL_VERSION,
70 K::NAME,
71 HASH_ALGORITHM,
72 hex::encode(self.digest)
73 )
74 }
75}
76
77impl<K: Kind> fmt::Debug for HashId<K> {
78 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
79 formatter
80 .debug_tuple("HashId")
81 .field(&self.to_string())
82 .finish()
83 }
84}
85
86impl<K: Kind> FromStr for HashId<K> {
87 type Err = HashError;
88
89 fn from_str(value: &str) -> Result<Self, Self::Err> {
90 let any = AnyHashId::from_str(value)?;
91 if any.kind != K::NAME {
92 return Err(HashError::UnexpectedKind {
93 expected: K::NAME.to_string(),
94 actual: any.kind.to_string(),
95 });
96 }
97 Ok(Self::from_digest(any.digest))
98 }
99}
100
101impl<K: Kind> Serialize for HashId<K> {
102 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
103 where
104 S: Serializer,
105 {
106 serializer.serialize_str(&self.to_string())
107 }
108}
109
110impl<'de, K: Kind> Deserialize<'de> for HashId<K> {
111 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
112 where
113 D: Deserializer<'de>,
114 {
115 let value = String::deserialize(deserializer)?;
116 value.parse().map_err(serde::de::Error::custom)
117 }
118}
119
120#[derive(Clone, Copy, PartialEq, Eq, Hash)]
121pub struct AnyHashId {
122 kind: HashKindName,
123 digest: [u8; 32],
124}
125
126impl AnyHashId {
127 pub const fn from_parts(kind: HashKindName, digest: [u8; 32]) -> Self {
128 Self { kind, digest }
129 }
130
131 pub const fn kind(self) -> HashKindName {
132 self.kind
133 }
134
135 pub const fn digest(&self) -> &[u8; 32] {
136 &self.digest
137 }
138
139 pub fn typed<K: Kind>(self) -> Result<HashId<K>, HashError> {
140 if self.kind != K::NAME {
141 return Err(HashError::UnexpectedKind {
142 expected: K::NAME.to_string(),
143 actual: self.kind.to_string(),
144 });
145 }
146 Ok(HashId::from_digest(self.digest))
147 }
148}
149
150impl fmt::Display for AnyHashId {
151 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
152 write!(
153 formatter,
154 "arete:h{}:{}:{}:{}",
155 HASH_PROTOCOL_VERSION,
156 self.kind,
157 HASH_ALGORITHM,
158 hex::encode(self.digest)
159 )
160 }
161}
162
163impl fmt::Debug for AnyHashId {
164 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
165 formatter
166 .debug_tuple("AnyHashId")
167 .field(&self.to_string())
168 .finish()
169 }
170}
171
172impl FromStr for AnyHashId {
173 type Err = HashError;
174
175 fn from_str(value: &str) -> Result<Self, Self::Err> {
176 let mut parts = value.split(':');
177 if parts.next() != Some("arete") {
178 return Err(HashError::InvalidHashId("protocol must be 'arete'"));
179 }
180 let version = parts
181 .next()
182 .ok_or(HashError::InvalidHashId("missing version"))?;
183 if version != "h1" {
184 return Err(HashError::UnknownVersion(version.to_string()));
185 }
186 let kind = parts
187 .next()
188 .ok_or(HashError::InvalidHashId("missing kind"))?
189 .parse()?;
190 let algorithm = parts
191 .next()
192 .ok_or(HashError::InvalidHashId("missing algorithm"))?;
193 if algorithm != HASH_ALGORITHM {
194 return Err(HashError::UnknownAlgorithm(algorithm.to_string()));
195 }
196 let digest = parts
197 .next()
198 .ok_or(HashError::InvalidHashId("missing digest"))?;
199 if parts.next().is_some() {
200 return Err(HashError::InvalidHashId("too many components"));
201 }
202 if digest.len() != 64 {
203 return Err(HashError::InvalidHashId(
204 "digest must contain 64 hex digits",
205 ));
206 }
207 if !digest
208 .as_bytes()
209 .iter()
210 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(byte))
211 {
212 return Err(HashError::InvalidHashId(
213 "digest must be lowercase hexadecimal",
214 ));
215 }
216 let mut decoded = [0_u8; 32];
217 hex::decode_to_slice(digest, &mut decoded)
218 .map_err(|_| HashError::InvalidHashId("invalid digest"))?;
219 Ok(Self {
220 kind,
221 digest: decoded,
222 })
223 }
224}
225
226impl Serialize for AnyHashId {
227 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
228 where
229 S: Serializer,
230 {
231 serializer.serialize_str(&self.to_string())
232 }
233}
234
235impl<'de> Deserialize<'de> for AnyHashId {
236 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
237 where
238 D: Deserializer<'de>,
239 {
240 let value = String::deserialize(deserializer)?;
241 value.parse().map_err(serde::de::Error::custom)
242 }
243}
244
245pub fn framed_preimage(
246 kind: HashKindName,
247 profile: CanonicalizationProfile,
248 payload: &[u8],
249) -> Vec<u8> {
250 let mut preimage = Vec::with_capacity(
251 8 + HASH_PROTOCOL_LABEL.len()
252 + 4
253 + 8
254 + kind.as_str().len()
255 + 8
256 + profile.as_str().len()
257 + 8
258 + payload.len(),
259 );
260 push_framed_bytes(&mut preimage, HASH_PROTOCOL_LABEL.as_bytes());
261 preimage.extend_from_slice(&HASH_PROTOCOL_VERSION.to_be_bytes());
262 push_framed_bytes(&mut preimage, kind.as_str().as_bytes());
263 push_framed_bytes(&mut preimage, profile.as_str().as_bytes());
264 push_framed_bytes(&mut preimage, payload);
265 preimage
266}
267
268pub(crate) fn hash_canonical_payload<K: Kind>(payload: &[u8]) -> HashId<K> {
269 let preimage = framed_preimage(K::NAME, K::PROFILE, payload);
270 HashId::from_digest(Sha256::digest(preimage).into())
271}
272
273pub(crate) fn require_profile<K: Kind>(actual: CanonicalizationProfile) -> Result<(), HashError> {
274 if K::PROFILE != actual {
275 return Err(HashError::ProfileMismatch {
276 kind: K::NAME.to_string(),
277 expected: K::PROFILE.to_string(),
278 actual: actual.to_string(),
279 });
280 }
281 Ok(())
282}
283
284pub(crate) fn push_framed_bytes(output: &mut Vec<u8>, bytes: &[u8]) {
285 output.extend_from_slice(&(bytes.len() as u64).to_be_bytes());
286 output.extend_from_slice(bytes);
287}