Skip to main content

cached_store_gcs/
lib.rs

1use async_trait::async_trait;
2use cached::IOCachedAsync;
3use chrono::{DateTime, Utc};
4use google_cloud_storage::client::{Client, ClientConfig};
5use google_cloud_storage::http::objects::delete::DeleteObjectRequest;
6use google_cloud_storage::http::objects::download::Range;
7use google_cloud_storage::http::objects::get::GetObjectRequest;
8use google_cloud_storage::http::objects::upload::{Media, UploadObjectRequest, UploadType};
9use serde::de::DeserializeOwned;
10use serde::Serialize;
11use std::fmt::Display;
12use std::marker::PhantomData;
13use std::time::Duration;
14
15const ENV_BUCKET_KEY: &str = "CACHED_GCS_BUCKET";
16
17use thiserror::Error;
18
19#[derive(Error, Debug)]
20pub enum GcsCacheBuildError {
21    #[error("gcs client creation error")]
22    ClientBuild {
23        error: google_cloud_storage::client::google_cloud_auth::error::Error,
24    },
25    #[error("Bucket name not specified or invalid in env var {env_key:?}: {error:?}")]
26    MissingBucket {
27        env_key: String,
28        error: std::env::VarError,
29    },
30}
31
32pub struct GcsCache<K, V> {
33    pub ttl: Duration,
34    pub client: Client,
35    pub bucket: String,
36    pub prefix: String,
37    _phantom: PhantomData<(K, V)>,
38}
39
40impl<K, V> GcsCache<K, V>
41where
42    K: Display,
43    V: Serialize + DeserializeOwned,
44{
45    fn generate_key(&self, key: &K) -> String {
46        format!("{}{}", self.prefix, key)
47    }
48
49    pub async fn new(ttl: Duration, prefix: &str) -> Result<GcsCache<K, V>, GcsCacheBuildError> {
50        Ok(GcsCache {
51            ttl,
52            client: Client::new(
53                ClientConfig::default()
54                    .with_auth()
55                    .await
56                    .map_err(|e| GcsCacheBuildError::ClientBuild { error: e })?,
57            ),
58            bucket: std::env::var(ENV_BUCKET_KEY).map_err(|e| {
59                GcsCacheBuildError::MissingBucket {
60                    env_key: ENV_BUCKET_KEY.to_string(),
61                    error: e,
62                }
63            })?,
64            prefix: prefix.to_string(),
65            _phantom: PhantomData,
66        })
67    }
68}
69
70#[derive(Error, Debug)]
71pub enum GcsCacheError {
72    #[error("gcs error")]
73    CloudStorageError(#[from] google_cloud_storage::http::Error),
74    #[error("Error deserializing cached value: {cached_value:?}: {error:?}")]
75    CacheDeserializationError {
76        cached_value: Vec<u8>,
77        error: serde_json::Error,
78    },
79    #[error("Error serializing cached value: {error:?}")]
80    CacheSerializationError { error: serde_json::Error },
81}
82
83#[derive(serde::Serialize, serde::Deserialize)]
84struct CachedGcsValue<V> {
85    pub(crate) value: V,
86    pub(crate) expires_at: DateTime<Utc>,
87}
88
89#[async_trait]
90impl<K, V> IOCachedAsync<K, V> for GcsCache<K, V>
91where
92    K: Display + Send + Sync,
93    V: Serialize + DeserializeOwned + Send + Sync,
94{
95    type Error = GcsCacheError;
96
97    async fn cache_get(&self, key: &K) -> Result<Option<V>, Self::Error> {
98        let object_key = self.generate_key(key);
99        let result = self
100            .client
101            .download_object(
102                &GetObjectRequest {
103                    bucket: self.bucket.clone(),
104                    object: object_key,
105                    ..Default::default()
106                },
107                &Range::default(),
108            )
109            .await;
110        match result {
111            Ok(data) => {
112                let val = serde_json::from_slice::<CachedGcsValue<V>>(&data).map_err(|e| {
113                    GcsCacheError::CacheDeserializationError {
114                        cached_value: data,
115                        error: e,
116                    }
117                })?;
118                if Utc::now() <= val.expires_at {
119                    Ok(Some(val.value))
120                } else {
121                    Ok(None)
122                }
123            }
124            Err(google_cloud_storage::http::Error::HttpClient(error))
125                if error.status() == Some(http::StatusCode::NOT_FOUND) =>
126            {
127                Ok(None)
128            }
129            Err(e) => Err(GcsCacheError::CloudStorageError(e)),
130        }
131    }
132
133    async fn cache_set(&self, key: K, value: V) -> Result<Option<V>, Self::Error> {
134        let object_key = self.generate_key(&key);
135        let val = CachedGcsValue {
136            value,
137            expires_at: Utc::now() + self.ttl,
138        };
139        let data = serde_json::to_vec(&val)
140            .map_err(|e| GcsCacheError::CacheSerializationError { error: e })?;
141
142        let old = self.cache_get(&key).await?;
143
144        self.client
145            .upload_object(
146                &UploadObjectRequest {
147                    bucket: self.bucket.clone(),
148                    ..Default::default()
149                },
150                data,
151                &UploadType::Simple(Media::new(object_key)),
152            )
153            .await
154            .map_err(|e| GcsCacheError::CloudStorageError(e))?;
155
156        Ok(old)
157    }
158
159    async fn cache_remove(&self, key: &K) -> Result<Option<V>, Self::Error> {
160        let object_key = self.generate_key(key);
161
162        let old = self.cache_get(&key).await?;
163
164        self.client
165            .delete_object(&DeleteObjectRequest {
166                bucket: self.bucket.clone(),
167                object: object_key,
168                ..Default::default()
169            })
170            .await
171            .map_err(|e| GcsCacheError::CloudStorageError(e))?;
172
173        Ok(old)
174    }
175
176    fn cache_set_refresh(&mut self, refresh: bool) -> bool {
177        panic!("refresh is not yet supported on GcsCache");
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184    use std::thread::sleep;
185    use std::time::Duration;
186
187    fn now_millis() -> u128 {
188        std::time::SystemTime::now()
189            .duration_since(std::time::UNIX_EPOCH)
190            .unwrap()
191            .as_millis()
192    }
193
194    #[tokio::test]
195    async fn test_gcs_cache() {
196        let c: GcsCache<u32, u32> = GcsCache::new(
197            Duration::from_secs(5),
198            &format!("{}:gcs-cache-test/", now_millis()),
199        )
200        .await
201        .unwrap();
202
203        assert!(c.cache_get(&1).await.unwrap().is_none());
204
205        assert!(c.cache_set(1, 100).await.unwrap().is_none());
206        assert!(c.cache_get(&1).await.unwrap().is_some());
207
208        sleep(Duration::new(5, 500_000));
209        assert!(c.cache_get(&1).await.unwrap().is_none());
210    }
211}