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}