use super::{Cache, CacheError};
use async_trait::async_trait;
use scc::{HashMap, hash_map::Entry};
pub struct MemoryCache<T> {
map: HashMap<String, T>,
}
impl<T> MemoryCache<T> {
pub fn new(capacity: usize) -> Self {
Self {
map: HashMap::with_capacity(capacity),
}
}
}
#[async_trait]
impl<T: Clone + Send + Sync> Cache<T> for MemoryCache<T> {
async fn get(&self, key: &str) -> Result<Option<T>, CacheError> {
Ok(self.map.read_async(key, |_, value| value.clone()).await)
}
async fn get_with<'a>(
&self,
key: String,
future: Box<dyn Future<Output = T> + Send + 'a>,
) -> Result<T, CacheError> {
Ok(match self.map.entry_async(key).await {
Entry::Occupied(entry) => entry.get().clone(),
Entry::Vacant(entry) => {
let value = Box::into_pin(future).await;
entry.insert_entry(value.clone());
value
}
})
}
async fn set(&self, key: String, value: T) -> Result<(), CacheError> {
self.map.upsert_async(key, value).await;
Ok(())
}
async fn remove(&self, key: &str) -> Result<(), CacheError> {
self.map.remove_async(key).await;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn get_or_set() {
let cache = MemoryCache::new(1 << 10);
assert_eq!(
cache
.get_with("key".into(), Box::new(async { 42 }))
.await
.unwrap(),
42,
);
assert_eq!(
cache
.get_with("key".into(), Box::new(async { 0 }))
.await
.unwrap(),
42,
);
}
#[tokio::test]
async fn get() {
let cache = MemoryCache::new(1 << 10);
assert_eq!(cache.get("key").await.unwrap(), None);
cache
.get_with("key".into(), Box::new(async { 42 }))
.await
.unwrap();
assert_eq!(cache.get("key").await.unwrap(), Some(42));
}
#[tokio::test]
async fn set() {
let cache = MemoryCache::new(1 << 10);
cache.set("key".into(), 42).await.unwrap();
assert_eq!(cache.get("key").await.unwrap(), Some(42));
cache.set("key".into(), 2).await.unwrap();
assert_eq!(cache.get("key").await.unwrap(), Some(2));
}
}