Skip to main content

luct_scanner/
utils.rs

1use luct_store::StringStoreValue;
2use serde::{
3    Deserialize, Serialize,
4    de::{DeserializeOwned, Visitor},
5};
6use std::{marker::PhantomData, ops::Deref};
7use web_time::{Duration, SystemTime};
8
9/// Wrapper around a type to indicate, that the contained value has been validated
10///
11/// When wrapping a `T` into [`Validated`], it means that the value has been validated and will be
12/// trusted from now on.
13#[derive(Debug, Clone, Eq, PartialOrd, Ord, Serialize)]
14pub struct Validated<T> {
15    inner: T,
16    validated_at: SystemTime,
17}
18
19impl<'de, T: Deserialize<'de>> Deserialize<'de> for Validated<T> {
20    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
21    where
22        D: serde::Deserializer<'de>,
23    {
24        struct ValidatedVisitor<T>(PhantomData<T>);
25        const FIELDS: [&str; 2] = ["inner", "validated_at"];
26
27        impl<'de, T: Deserialize<'de>> Visitor<'de> for ValidatedVisitor<T> {
28            type Value = Validated<T>;
29
30            fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
31                formatter.write_str("struct Validated")
32            }
33
34            fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
35            where
36                A: serde::de::MapAccess<'de>,
37            {
38                let mut inner = None;
39                let mut validated_at = None;
40                while let Some(key) = map.next_key::<String>()? {
41                    match key.as_str() {
42                        "inner" => {
43                            if inner.is_some() {
44                                return Err(serde::de::Error::duplicate_field("inner"));
45                            }
46                            inner = Some(map.next_value()?);
47                        }
48                        "validated_at" => {
49                            if validated_at.is_some() {
50                                return Err(serde::de::Error::duplicate_field("validated_at"));
51                            }
52                            validated_at = Some(map.next_value()?)
53                        }
54                        value => return Err(serde::de::Error::unknown_field(value, &FIELDS)),
55                    }
56                }
57
58                let inner = inner.ok_or_else(|| serde::de::Error::missing_field("inner"))?;
59                let validated_at =
60                    validated_at.ok_or_else(|| serde::de::Error::missing_field("validated_at"))?;
61
62                Ok(Validated {
63                    inner,
64                    validated_at,
65                })
66            }
67
68            fn visit_seq<V>(self, mut seq: V) -> Result<Self::Value, V::Error>
69            where
70                V: serde::de::SeqAccess<'de>,
71            {
72                let validated_at: u64 = seq
73                    .next_element()?
74                    .ok_or_else(|| serde::de::Error::invalid_length(0, &self))?;
75                let validated_at = SystemTime::UNIX_EPOCH
76                    .checked_add(Duration::from_millis(validated_at))
77                    .unwrap();
78                let inner = seq
79                    .next_element()?
80                    .ok_or_else(|| serde::de::Error::invalid_length(1, &self))?;
81                Ok(Validated {
82                    inner,
83                    validated_at,
84                })
85            }
86        }
87
88        deserializer.deserialize_any(ValidatedVisitor(PhantomData))
89    }
90}
91
92impl<T: PartialEq> PartialEq for Validated<T> {
93    fn eq(&self, other: &Self) -> bool {
94        // NOTE: The validated_at should not influence equality
95        self.inner == other.inner
96    }
97}
98
99impl<T> Validated<T> {
100    pub fn new(inner: T) -> Self {
101        Self {
102            inner,
103            validated_at: SystemTime::now(),
104        }
105    }
106
107    pub fn validated_at(&self) -> SystemTime {
108        self.validated_at
109    }
110
111    pub fn inner(&self) -> &T {
112        &self.inner
113    }
114}
115
116impl<T> Deref for Validated<T> {
117    type Target = T;
118
119    fn deref(&self) -> &Self::Target {
120        &self.inner
121    }
122}
123
124impl<T: StringStoreValue + Serialize + DeserializeOwned> StringStoreValue for Validated<T> {
125    fn serialize_value(&self) -> String {
126        serde_json::to_string(self).unwrap()
127    }
128
129    fn deserialize_value(value: &str) -> Option<Self> {
130        serde_json::from_str(value)
131            .ok()
132            .or_else(|| Self::deserialize_value_legacy(value))
133    }
134}
135
136impl<T: StringStoreValue> Validated<T> {
137    // NOTE: Version 0.1 parses STHs this way. We can drop this, once we have implemented log sth removal
138    fn deserialize_value_legacy(value: &str) -> Option<Self> {
139        let (validated_at, inner): (u64, String) = serde_json::from_str(value).ok()?;
140
141        let validated_at =
142            SystemTime::UNIX_EPOCH.checked_add(Duration::from_millis(validated_at))?;
143        let inner = T::deserialize_value(&inner)?;
144
145        Some(Self {
146            inner,
147            validated_at,
148        })
149    }
150}
151
152#[cfg(test)]
153mod tests {
154    use std::time::UNIX_EPOCH;
155
156    use super::*;
157
158    #[derive(Debug, PartialEq, Eq, Serialize, Deserialize)]
159    struct TestStruct {
160        a: u64,
161        b: String,
162    }
163
164    #[test]
165    fn validated_json_roundtrip() {
166        let test_data = Validated::new(TestStruct {
167            a: 5,
168            b: String::from("Test"),
169        });
170        let json = serde_json::to_string(&test_data).unwrap();
171        let new_test_data = serde_json::from_str(&json).unwrap();
172        assert_eq!(test_data, new_test_data)
173    }
174
175    #[test]
176    fn legacy_validated_json() {
177        let test_data = Validated::new(TestStruct {
178            a: 5,
179            b: String::from("Test"),
180        });
181        let now_str = serde_json::to_string(
182            &test_data
183                .validated_at()
184                .duration_since(UNIX_EPOCH)
185                .unwrap()
186                .as_millis(),
187        )
188        .unwrap();
189
190        let legacy_validated = format!("[{}, {{\"a\": 5, \"b\": \"Test\"}}]", now_str);
191        let new_test_data: Validated<TestStruct> = serde_json::from_str(&legacy_validated).unwrap();
192        assert_eq!(test_data, new_test_data)
193    }
194}