use crate::store::Store;
use axess_clock::{Clock, SystemClock};
use chrono::{DateTime, Duration as ChronoDuration, Utc};
use dashmap::DashMap;
use std::convert::Infallible;
use std::future::Future;
use std::hash::Hash;
use std::sync::Arc;
use std::time::Duration;
pub struct MemoryStore<K, V>
where
K: Eq + Hash + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
inner: Arc<DashMap<K, Entry<V>>>,
clock: Arc<dyn Clock>,
}
#[derive(Debug, Clone)]
struct Entry<V> {
value: V,
expires_at: DateTime<Utc>,
}
impl<K, V> std::fmt::Debug for MemoryStore<K, V>
where
K: Eq + Hash + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MemoryStore")
.field("entries", &self.inner.len())
.finish()
}
}
impl<K, V> Clone for MemoryStore<K, V>
where
K: Eq + Hash + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
clock: self.clock.clone(),
}
}
}
impl<K, V> Default for MemoryStore<K, V>
where
K: Eq + Hash + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
fn default() -> Self {
Self::new()
}
}
impl<K, V> MemoryStore<K, V>
where
K: Eq + Hash + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
pub fn new() -> Self {
Self {
inner: Arc::new(DashMap::new()),
clock: Arc::new(SystemClock),
}
}
pub fn with_clock(mut self, clock: Arc<dyn Clock>) -> Self {
self.clock = clock;
self
}
pub fn clock(&self) -> Arc<dyn Clock> {
self.clock.clone()
}
pub fn snapshot(&self) -> Vec<(K, V)> {
let now = self.clock.now();
self.inner
.iter()
.filter(|e| e.value().expires_at > now)
.map(|e| (e.key().clone(), e.value().value.clone()))
.collect()
}
pub fn update<F>(&self, key: &K, mut f: F) -> bool
where
F: FnMut(&mut V),
{
let now = self.clock.now();
match self.inner.get_mut(key) {
Some(mut entry) if entry.expires_at > now => {
f(&mut entry.value);
true
}
_ => false,
}
}
pub fn prune_expired_sync(&self) -> u64 {
let now = self.clock.now();
let before = self.inner.len();
self.inner.retain(|_, entry| entry.expires_at > now);
let after = self.inner.len();
(before - after) as u64
}
pub fn len(&self) -> usize {
self.inner.len()
}
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
pub fn physically_contains_key(&self, key: &K) -> bool {
self.inner.contains_key(key)
}
}
impl<K, V> Store<K, V> for MemoryStore<K, V>
where
K: Eq + Hash + Clone + Send + Sync + 'static,
V: Clone + Send + Sync + 'static,
{
type Error = Infallible;
fn get(&self, key: &K) -> impl Future<Output = Result<Option<V>, Self::Error>> + Send {
let now = self.clock.now();
let result = self
.inner
.get(key)
.filter(|e| e.expires_at > now)
.map(|e| e.value.clone());
async move { Ok(result) }
}
fn put(
&self,
key: &K,
value: &V,
ttl: Duration,
) -> impl Future<Output = Result<(), Self::Error>> + Send {
let now = self.clock.now();
let expires_at = ChronoDuration::from_std(ttl)
.ok()
.and_then(|d| now.checked_add_signed(d))
.unwrap_or(DateTime::<Utc>::MAX_UTC);
self.inner.insert(
key.clone(),
Entry {
value: value.clone(),
expires_at,
},
);
async { Ok(()) }
}
fn delete(&self, key: &K) -> impl Future<Output = Result<(), Self::Error>> + Send {
self.inner.remove(key);
async { Ok(()) }
}
fn prune_expired(&self) -> impl Future<Output = Result<u64, Self::Error>> + Send {
let removed = self.prune_expired_sync();
async move { Ok(removed) }
}
}
#[cfg(test)]
mod tests;