Skip to main content

rskit_cache/
typed_store.rs

1use std::marker::PhantomData;
2use std::sync::Arc;
3use std::time::Duration;
4
5use serde::Serialize;
6use serde::de::DeserializeOwned;
7
8use rskit_errors::{AppError, AppResult, ErrorCode};
9
10use crate::registry::CacheStore;
11
12/// A generic, JSON-serialised store backed by a [`CacheStore`].
13///
14/// Keys are automatically prefixed with the store's `prefix` so that
15/// multiple `TypedStore` instances can coexist on the same cache store without
16/// key collisions.
17pub struct TypedStore<T> {
18    client: Arc<dyn CacheStore>,
19    prefix: String,
20    _marker: PhantomData<T>,
21}
22
23impl<T: Serialize + DeserializeOwned + Send + Sync> TypedStore<T> {
24    /// Create a new typed store that prefixes all keys with `prefix`.
25    pub fn new(client: Arc<dyn CacheStore>, prefix: impl Into<String>) -> Self {
26        Self {
27            client,
28            prefix: prefix.into(),
29            _marker: PhantomData,
30        }
31    }
32
33    /// Build the cache-store key used for storage.
34    fn full_key(&self, key: &str) -> String {
35        format!("{}:{}", self.prefix, key)
36    }
37
38    /// Retrieve a value by key, deserialising from JSON.
39    pub async fn get(&self, key: &str) -> AppResult<Option<T>> {
40        let raw = self.client.get(&self.full_key(key)).await?;
41        match raw {
42            Some(json) => {
43                let val = serde_json::from_str(&json).map_err(|e| {
44                    AppError::new(ErrorCode::Internal, format!("json deserialise error: {e}"))
45                        .with_cause(e)
46                })?;
47                Ok(Some(val))
48            }
49            None => Ok(None),
50        }
51    }
52
53    /// Store a value by key, serialising to JSON. An optional TTL may be set.
54    pub async fn set(&self, key: &str, val: &T, ttl: Option<Duration>) -> AppResult<()> {
55        let json = serde_json::to_string(val).map_err(|e| {
56            AppError::new(ErrorCode::Internal, format!("json serialise error: {e}")).with_cause(e)
57        })?;
58        self.client.set(&self.full_key(key), &json, ttl).await
59    }
60
61    /// Delete a key. Returns `true` if the key existed.
62    pub async fn delete(&self, key: &str) -> AppResult<bool> {
63        self.client.delete(&self.full_key(key)).await
64    }
65
66    /// Check whether a key exists.
67    pub async fn exists(&self, key: &str) -> AppResult<bool> {
68        self.client.exists(&self.full_key(key)).await
69    }
70}
71
72#[cfg(test)]
73mod tests {
74    use super::*;
75    use parking_lot::Mutex;
76    use serde::Serializer;
77    use std::collections::BTreeMap;
78
79    #[derive(Default)]
80    struct MemoryStore {
81        values: Mutex<BTreeMap<String, String>>,
82    }
83
84    #[async_trait::async_trait]
85    impl CacheStore for MemoryStore {
86        async fn get(&self, key: &str) -> AppResult<Option<String>> {
87            Ok(self.values.lock().get(key).cloned())
88        }
89
90        async fn set(&self, key: &str, val: &str, _ttl: Option<Duration>) -> AppResult<()> {
91            self.values.lock().insert(key.to_string(), val.to_string());
92            Ok(())
93        }
94
95        async fn delete(&self, key: &str) -> AppResult<bool> {
96            Ok(self.values.lock().remove(key).is_some())
97        }
98
99        async fn exists(&self, key: &str) -> AppResult<bool> {
100            Ok(self.values.lock().contains_key(key))
101        }
102    }
103
104    struct FailingSerialize;
105
106    impl Serialize for FailingSerialize {
107        fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
108        where
109            S: Serializer,
110        {
111            Err(serde::ser::Error::custom("boom"))
112        }
113    }
114
115    impl<'de> serde::Deserialize<'de> for FailingSerialize {
116        fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
117        where
118            D: serde::Deserializer<'de>,
119        {
120            serde::de::IgnoredAny::deserialize(deserializer)?;
121            Ok(Self)
122        }
123    }
124
125    #[tokio::test]
126    async fn typed_store_round_trips_prefixed_json() {
127        let store = Arc::new(MemoryStore::default());
128        let typed = TypedStore::<u32>::new(store.clone(), "numbers");
129
130        typed
131            .set("answer", &42, None)
132            .await
133            .expect("set should serialise");
134
135        assert!(store.exists("numbers:answer").await.expect("exists works"));
136        assert_eq!(
137            typed.get("answer").await.expect("get should deserialise"),
138            Some(42)
139        );
140        assert!(typed.delete("answer").await.expect("delete should succeed"));
141        assert_eq!(
142            typed.get("answer").await.expect("missing key succeeds"),
143            None
144        );
145    }
146
147    #[tokio::test]
148    async fn get_rejects_invalid_json() {
149        let store = Arc::new(MemoryStore::default());
150        store
151            .set("numbers:bad", "not-json", None)
152            .await
153            .expect("fixture write succeeds");
154        let typed = TypedStore::<u32>::new(store, "numbers");
155
156        let err = typed
157            .get("bad")
158            .await
159            .expect_err("invalid json should fail");
160        assert_eq!(err.code(), ErrorCode::Internal);
161    }
162
163    #[tokio::test]
164    async fn set_rejects_serialisation_errors() {
165        let store = Arc::new(MemoryStore::default());
166        let typed = TypedStore::<FailingSerialize>::new(store.clone(), "bad");
167
168        let err = typed
169            .set("value", &FailingSerialize, None)
170            .await
171            .expect_err("serialisation errors should surface");
172        assert_eq!(err.code(), ErrorCode::Internal);
173
174        store
175            .set("bad:value", "null", None)
176            .await
177            .expect("fixture write succeeds");
178        assert!(
179            typed
180                .get("value")
181                .await
182                .expect("deserialise should succeed")
183                .is_some()
184        );
185    }
186}