use tokio::sync::{RwLock, Mutex};
use std::sync::Arc;
use std::time::{Instant, Duration};
use std::future::Future;
pub type AsyncCacheResult<T> = Result<Arc<T>, Box<dyn std::error::Error + Send + Sync>>;
pub type AsyncRefreshResult<T> = Result<T, Box<dyn std::error::Error + Send + Sync>>;
struct BasicCache<T, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
{
data: Option<Arc<T>>,
ttl: Duration,
age: std::time::Instant,
refresher: Box<dyn Fn() -> F + Send + Sync>
}
struct StatefulCache<T, S, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
S: Send + Sync
{
data: Option<Arc<T>>,
ttl: Duration,
state: Arc<Mutex<S>>,
age: std::time::Instant,
refresher: Box<dyn Fn(Arc<Mutex<S>>) -> F + Send + Sync >
}
enum DataCache<T, S, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
S: Send + Sync
{
Basic(BasicCache<T, F>),
Stateful(StatefulCache<T, S, F>),
}
impl<T, S, F> StatefulCache<T, S, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
S: Send + Sync
{
fn get_reference(&self) -> Result<Arc<T>, ()> {
if self.age.elapsed() > self.ttl {
return Err(());
}
Ok(Arc::clone(self.data.as_ref().unwrap()))
}
async fn refresh(&mut self) -> AsyncCacheResult<T> {
if self.age.elapsed() < self.ttl {
return Ok(Arc::clone(self.data.as_ref().unwrap()));
}
let data = (self.refresher)(Arc::clone(&self.state)).await;
if data.is_ok() {
self.data = Some(Arc::new(data.unwrap()));
self.age = Instant::now();
return Ok(Arc::clone(self.data.as_ref().unwrap()))
}
self.age = Instant::now();
Err(data.err().unwrap())
}
}
impl<T, F> BasicCache<T, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
{
fn get_reference(&self) -> Result<Arc<T>, ()> {
if self.age.elapsed() > self.ttl {
return Err(());
}
Ok(Arc::clone(self.data.as_ref().unwrap()))
}
async fn refresh(&mut self) -> AsyncCacheResult<T> {
if self.age.elapsed() < self.ttl {
return Ok(Arc::clone(self.data.as_ref().unwrap()));
}
let data = (self.refresher)().await;
if data.is_ok() {
self.data = Some(Arc::new(data.unwrap()));
self.age = Instant::now();
return Ok(Arc::clone(self.data.as_ref().unwrap()))
}
self.age = Instant::now();
Err(data.err().unwrap())
}
}
pub struct AsyncCache<T, S, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
S: Send + Sync
{
data: RwLock<DataCache<T, S, F>>,
}
impl< T, S, F> AsyncCache<T, S, F>
where F: Future<Output = AsyncRefreshResult<T>> + Send + Sync,
S: Send + Sync
{
pub fn new(ttl: Duration, refresher: Box<dyn Fn() -> F + Send + Sync>)
-> Self
{
let datacache = DataCache::Basic(BasicCache {
data: None,
age: Instant::now() - ttl - ttl,
ttl,
refresher
});
AsyncCache {
data: RwLock::new(datacache)
}
}
pub fn new_with_state(
state: S,
ttl: Duration,
refresher: Box<dyn Fn(Arc<Mutex<S>>) -> F + Send + Sync>
) -> Self
{
let cache = DataCache::Stateful(StatefulCache {
data: None,
age: Instant::now() - ttl - ttl,
state: Arc::new(Mutex::new(state)),
ttl,
refresher
});
AsyncCache {
data: RwLock::new(cache)
}
}
pub async fn get_data(&self) -> AsyncCacheResult<T> {
{
let cache = self.data.read().await;
match &(*cache) {
DataCache::Basic(cache) => {
let data = cache.get_reference();
if let Ok(data) = data { return Ok(data); }
},
DataCache::Stateful(cache) => {
let data = cache.get_reference();
if let Ok(data) = data { return Ok(data); }
}
}
}
let mut cache = self.data.write().await;
match &mut (*cache) {
DataCache::Basic(cache) => match cache.refresh().await {
Ok(data) => Ok(data),
Err(err) => Err(err)
},
DataCache::Stateful(cache) => match cache.refresh().await {
Ok(data) => Ok(data),
Err(err) => Err(err)
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
use tokio::time::sleep;
#[tokio::test(flavor = "multi_thread")]
async fn cache_test() {
struct State {
i: u32
}
async fn alter_state(s: Arc<Mutex<State>>) -> Result<u32, Box<dyn std::error::Error + Send + Sync >> {
let mut data = s.lock().await;
data.i = data.i + 1;
Ok(data.i)
}
let state = State { i: 0 };
let cache = Arc::new(AsyncCache::new_with_state(state, Duration::from_millis(50), Box::new(alter_state)));
let cache1 = Arc::clone(&cache);
let t1 = tokio::spawn( async move {
let mut data = cache1.get_data().await;
assert_eq!(*data.unwrap(), 1);
sleep(Duration::from_millis(10)).await;
data = cache1.get_data().await;
assert_eq!(*data.unwrap(), 1);
sleep(Duration::from_millis(50)).await;
data = cache1.get_data().await;
assert_eq!(*data.unwrap(), 2);
sleep(Duration::from_millis(20)).await;
data = cache1.get_data().await;
assert_eq!(*data.unwrap(), 2);
sleep(Duration::from_millis(50)).await;
data = cache1.get_data().await;
assert_eq!(*data.unwrap(), 3);
});
let cache2 = Arc::clone(&cache);
let t2 = tokio::spawn( async move {
let mut data = cache2.get_data().await;
assert_eq!(*data.unwrap(), 1);
sleep(Duration::from_millis(20)).await;
data = cache2.get_data().await;
assert_eq!(*data.unwrap(), 1);
sleep(Duration::from_millis(40)).await;
data = cache2.get_data().await;
assert_eq!(*data.unwrap(), 2);
sleep(Duration::from_millis(10)).await;
data = cache2.get_data().await;
assert_eq!(*data.unwrap(), 2);
sleep(Duration::from_millis(50)).await;
data = cache2.get_data().await;
assert_eq!(*data.unwrap(), 3);
});
t1.await.unwrap();
t2.await.unwrap();
}
}