wasm_pkg_common/
digest.rs1use std::str::FromStr;
2
3use bytes::Bytes;
4use futures_util::{Stream, StreamExt, TryStream, TryStreamExt, future::ready, stream::once};
5use serde::{Deserialize, Serialize};
6use sha2::{Digest, Sha256};
7
8use crate::Error;
9
10#[derive(Clone, Debug, PartialEq, Eq)]
12pub enum ContentDigest {
13 Sha256 { hex: String },
14}
15
16impl ContentDigest {
17 pub fn validating_stream<S>(
18 &self,
19 stream: S,
20 ) -> impl Stream<Item = Result<Bytes, Error>> + use<S>
21 where
22 S: TryStream<Ok = Bytes, Error = Error>,
23 {
24 let want = self.clone();
25 stream.map_ok(Some).chain(once(async { Ok(None) })).scan(
26 Sha256::new(),
27 move |hasher, res| {
28 ready(match res {
29 Ok(Some(bytes)) => {
30 hasher.update(&bytes);
31 Some(Ok(bytes))
32 }
33 Ok(None) => {
34 let got: Self = std::mem::take(hasher).into();
35 if got == want {
36 None
37 } else {
38 Some(Err(Error::InvalidContent(format!(
39 "expected digest {want}, got {got}"
40 ))))
41 }
42 }
43 Err(err) => Some(Err(err)),
44 })
45 },
46 )
47 }
48}
49
50impl std::fmt::Display for ContentDigest {
51 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
52 match self {
53 ContentDigest::Sha256 { hex } => write!(f, "sha256:{hex}"),
54 }
55 }
56}
57
58impl From<Sha256> for ContentDigest {
59 fn from(hasher: Sha256) -> Self {
60 Self::Sha256 {
61 hex: base16ct::lower::encode_string(&hasher.finalize()),
62 }
63 }
64}
65
66impl<'a> TryFrom<&'a str> for ContentDigest {
67 type Error = Error;
68
69 fn try_from(value: &'a str) -> Result<Self, Self::Error> {
70 let Some(hex) = value.strip_prefix("sha256:") else {
71 return Err(Error::InvalidContentDigest(
72 "must start with 'sha256:'".into(),
73 ));
74 };
75 let hex = hex.to_lowercase();
76 if hex.len() != 64 {
77 return Err(Error::InvalidContentDigest(format!(
78 "must be 64 hex digits; got {} chars",
79 hex.len()
80 )));
81 }
82 if let Some(invalid) = hex.chars().find(|c| !c.is_ascii_hexdigit()) {
83 return Err(Error::InvalidContentDigest(format!(
84 "must be hex; got {invalid:?}"
85 )));
86 }
87 Ok(Self::Sha256 { hex })
88 }
89}
90
91impl FromStr for ContentDigest {
92 type Err = Error;
93
94 fn from_str(s: &str) -> Result<Self, Self::Err> {
95 s.try_into()
96 }
97}
98
99impl Serialize for ContentDigest {
100 fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
101 serializer.serialize_str(&self.to_string())
102 }
103}
104
105impl<'de> Deserialize<'de> for ContentDigest {
106 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
107 where
108 D: serde::Deserializer<'de>,
109 {
110 Self::from_str(&String::deserialize(deserializer)?).map_err(serde::de::Error::custom)
111 }
112}
113
114#[cfg(test)]
115mod tests {
116 use bytes::BytesMut;
117 use futures_util::stream;
118
119 use super::*;
120
121 #[tokio::test]
122 async fn test_validating_stream() {
123 let input = b"input";
124 let digest = ContentDigest::from(Sha256::new_with_prefix(input));
125 let stream = stream::iter(input.chunks(2));
126 let validating = digest.validating_stream(stream.map(|bytes| Ok(bytes.into())));
127 assert_eq!(
128 validating.try_collect::<BytesMut>().await.unwrap(),
129 &input[..]
130 );
131 }
132
133 #[tokio::test]
134 async fn test_invalidating_stream() {
135 let input = b"input";
136 let digest = ContentDigest::Sha256 {
137 hex: "doesn't match anything!".to_string(),
138 };
139 let stream = stream::iter(input.chunks(2));
140 let validating = digest.validating_stream(stream.map(|bytes| Ok(bytes.into())));
141 assert!(matches!(
142 validating.try_collect::<BytesMut>().await,
143 Err(Error::InvalidContent(_)),
144 ));
145 }
146}