Skip to main content

wasm_pkg_common/
digest.rs

1use 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/// A cryptographic digest (hash) of some content.
11#[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}