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