#![allow(dead_code)]
use crate::vector_select::FutureVector;
use async_trait::async_trait;
use std::collections::HashMap;
use std::fmt::Debug;
use std::hash::Hash;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::{io, mem};
use tokio::sync::{Mutex, MutexGuard, RwLock, RwLockWriteGuard};
use tracing::{error, instrument, trace};
pub(crate) struct AsyncLruCacheEntry<V> {
value: Option<Arc<V>>,
last_used: AtomicUsize,
}
struct AsyncLruCacheInner<
Key: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync,
Value: Send + Sync,
IoBackend: AsyncLruCacheBackend<Key = Key, Value = Value>,
> {
backend: IoBackend,
map: RwLock<HashMap<Key, AsyncLruCacheEntry<Value>>>,
flush_before: Mutex<Vec<Arc<dyn FlushableCache>>>,
lru_timer: AtomicUsize,
limit: usize,
}
pub(crate) struct AsyncLruCache<
K: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync,
V: Send + Sync,
B: AsyncLruCacheBackend<Key = K, Value = V>,
>(Arc<AsyncLruCacheInner<K, V, B>>);
#[async_trait(?Send)]
trait FlushableCache: Send + Sync {
async fn flush(&self) -> io::Result<()>;
async fn check_circular(&self, other: &Arc<dyn FlushableCache>) -> bool;
}
pub(crate) trait AsyncLruCacheBackend: Send + Sync {
type Key: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync;
type Value: Send + Sync;
#[allow(async_fn_in_trait)] async fn load(&self, key: Self::Key) -> io::Result<Self::Value>;
#[allow(async_fn_in_trait)] async fn flush(&self, key: Self::Key, value: Arc<Self::Value>) -> io::Result<()>;
unsafe fn evict(&self, key: Self::Key, value: Self::Value);
}
impl<
K: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync,
V: Send + Sync,
B: AsyncLruCacheBackend<Key = K, Value = V>,
> AsyncLruCache<K, V, B>
{
pub fn new(backend: B, size: usize) -> Self {
AsyncLruCache(Arc::new(AsyncLruCacheInner {
backend,
map: Default::default(),
flush_before: Default::default(),
lru_timer: AtomicUsize::new(0),
limit: size,
}))
}
pub async fn get_or_insert(&self, key: K) -> io::Result<Arc<V>> {
self.0.get_or_insert(key).await
}
pub async fn insert(&self, key: K, value: Arc<V>) -> io::Result<()> {
self.0.insert(key, value).await
}
pub async fn flush(&self) -> io::Result<()> {
self.0.flush().await
}
pub async unsafe fn invalidate(&self) -> io::Result<()> {
unsafe { self.0.invalidate() }.await
}
}
impl<
K: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync + 'static,
V: Send + Sync + 'static,
B: AsyncLruCacheBackend<Key = K, Value = V> + 'static,
> AsyncLruCache<K, V, B>
{
#[instrument(
level = "trace",
name = "AsyncLruCache::depend_on",
skip_all,
fields(
self = Arc::as_ptr(&self.0) as usize,
other = Arc::as_ptr(&other.0) as usize,
)
)]
pub async fn depend_on<
K2: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync + 'static,
V2: Send + Sync + 'static,
B2: AsyncLruCacheBackend<Key = K2, Value = V2> + 'static,
>(
&self,
other: &AsyncLruCache<K2, V2, B2>,
) -> io::Result<()> {
let cloned: Arc<AsyncLruCacheInner<K2, V2, B2>> = Arc::clone(&other.0);
let cloned: Arc<dyn FlushableCache> = cloned;
loop {
{
let mut locked = self.0.flush_before.lock().await;
if locked.iter().any(|x| Arc::ptr_eq(x, &cloned)) {
break;
}
let self_arc: Arc<AsyncLruCacheInner<K, V, B>> = Arc::clone(&self.0);
let self_arc: Arc<dyn FlushableCache> = self_arc;
if !other.0.check_circular(&self_arc).await {
trace!("No circular dependency, entering new dependency");
locked.push(cloned);
break;
}
}
trace!("Circular dependency detected, flushing other cache first");
other.0.flush().await?;
}
Ok(())
}
}
impl<
K: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync,
V: Send + Sync,
B: AsyncLruCacheBackend<Key = K, Value = V>,
> AsyncLruCacheInner<K, V, B>
{
#[instrument(level = "trace", name = "AsyncLruCache::flush_dependencies", skip_all)]
async fn flush_dependencies(
flush_before: &mut MutexGuard<'_, Vec<Arc<dyn FlushableCache>>>,
) -> io::Result<()> {
while let Some(dep) = flush_before.pop() {
trace!("Flushing dependency {:?}", Arc::as_ptr(&dep) as *const _);
if let Err(err) = dep.flush().await {
flush_before.push(dep);
return Err(err);
}
}
Ok(())
}
#[instrument(
level = "trace",
name = "AsyncLruCache::ensure_free_entry",
skip_all,
fields(self = &self as *const _ as usize),
)]
async fn ensure_free_entry(
&self,
map: &mut RwLockWriteGuard<'_, HashMap<K, AsyncLruCacheEntry<V>>>,
) -> io::Result<()> {
while map.len() >= self.limit {
trace!("{} / {} used", map.len(), self.limit);
let now = self.lru_timer.load(Ordering::Relaxed);
let (evicted_object, key, last_used) = loop {
let oldest = map.iter().fold((0, None), |oldest, (key, entry)| {
if Arc::strong_count(entry.value()) > 1 {
return oldest;
}
let age = now.wrapping_sub(entry.last_used.load(Ordering::Relaxed));
if age >= oldest.0 {
(age, Some(*key))
} else {
oldest
}
});
let Some(oldest_key) = oldest.1 else {
error!("Cannot evict entry from cache; everything is in use");
return Err(io::Error::other(
"Cannot evict entry from cache; everything is in use",
));
};
trace!("Removing entry with key {oldest_key:?}, aged {}", oldest.0);
let mut oldest_entry = map.remove(&oldest_key).unwrap();
match Arc::try_unwrap(oldest_entry.value.take().unwrap()) {
Ok(object) => {
break (
object,
oldest_key,
oldest_entry.last_used.load(Ordering::Relaxed),
)
}
Err(arc) => {
trace!("Entry is still in use, retrying");
oldest_entry.value = Some(arc);
}
}
};
let mut dep_guard = self.flush_before.lock().await;
Self::flush_dependencies(&mut dep_guard).await?;
let obj = Arc::new(evicted_object);
trace!("Flushing {key:?}");
if let Err(err) = self.backend.flush(key, Arc::clone(&obj)).await {
map.insert(
key,
AsyncLruCacheEntry {
value: Some(obj),
last_used: last_used.into(),
},
);
return Err(err);
}
let _ = Arc::into_inner(obj).expect("flush() must not clone the object");
}
Ok(())
}
async fn get_or_insert(&self, key: K) -> io::Result<Arc<V>> {
{
let map = self.map.read().await;
if let Some(entry) = map.get(&key) {
entry.last_used.store(
self.lru_timer.fetch_add(1, Ordering::Relaxed),
Ordering::Relaxed,
);
return Ok(Arc::clone(entry.value()));
}
}
let mut map = self.map.write().await;
if let Some(entry) = map.get(&key) {
entry.last_used.store(
self.lru_timer.fetch_add(1, Ordering::Relaxed),
Ordering::Relaxed,
);
return Ok(Arc::clone(entry.value()));
}
self.ensure_free_entry(&mut map).await?;
let object = Arc::new(self.backend.load(key).await?);
let new_entry = AsyncLruCacheEntry {
value: Some(Arc::clone(&object)),
last_used: AtomicUsize::new(self.lru_timer.fetch_add(1, Ordering::Relaxed)),
};
map.insert(key, new_entry);
Ok(object)
}
async fn insert(&self, key: K, value: Arc<V>) -> io::Result<()> {
let mut map = self.map.write().await;
if let Some(entry) = map.get_mut(&key) {
entry.last_used.store(
self.lru_timer.fetch_add(1, Ordering::Relaxed),
Ordering::Relaxed,
);
let mut dep_guard = self.flush_before.lock().await;
Self::flush_dependencies(&mut dep_guard).await?;
self.backend.flush(key, Arc::clone(entry.value())).await?;
entry.value = Some(value);
} else {
self.ensure_free_entry(&mut map).await?;
let new_entry = AsyncLruCacheEntry {
value: Some(value),
last_used: AtomicUsize::new(self.lru_timer.fetch_add(1, Ordering::Relaxed)),
};
map.insert(key, new_entry);
}
Ok(())
}
#[instrument(
level = "trace",
name = "AsyncLruCache::flush",
skip_all,
fields(self = &self as *const _ as usize)
)]
async fn flush(&self) -> io::Result<()> {
let mut futs = FutureVector::new();
let mut dep_guard = self.flush_before.lock().await;
Self::flush_dependencies(&mut dep_guard).await?;
let map = self.map.read().await;
for (key, entry) in map.iter() {
let key = *key;
let object = Arc::clone(entry.value());
trace!("Flushing {key:?}");
futs.push(Box::pin(self.backend.flush(key, object)));
}
futs.discarding_join().await
}
#[instrument(
level = "trace",
name = "AsyncLruCache::invalidate",
skip_all,
fields(self = &self as *const _ as usize)
)]
async unsafe fn invalidate(&self) -> io::Result<()> {
let mut in_use = Vec::new();
let mut map = self.map.write().await;
let old_map = mem::take(&mut *map);
for (key, mut entry) in old_map {
let object = entry.value.take().unwrap();
trace!("Evicting {key:?}");
match Arc::try_unwrap(object) {
Ok(object) => {
unsafe { self.backend.evict(key, object) };
}
Err(arc) => {
trace!("Entry is still in use, retaining it");
entry.value = Some(arc);
map.insert(key, entry);
in_use.push(key);
}
}
}
if in_use.is_empty() {
self.flush_before.lock().await.clear();
Ok(())
} else {
Err(io::Error::other(format!(
"Cannot invalidate cache, entries still in use: {}",
in_use
.iter()
.map(|key| format!("{key:?}"))
.collect::<Vec<String>>()
.join(", "),
)))
}
}
}
impl<V> AsyncLruCacheEntry<V> {
fn value(&self) -> &Arc<V> {
self.value.as_ref().unwrap()
}
}
#[async_trait(?Send)]
impl<
K: Clone + Copy + Debug + PartialEq + Eq + Hash + Send + Sync,
V: Send + Sync,
B: AsyncLruCacheBackend<Key = K, Value = V>,
> FlushableCache for AsyncLruCacheInner<K, V, B>
{
async fn flush(&self) -> io::Result<()> {
AsyncLruCacheInner::<K, V, B>::flush(self).await
}
async fn check_circular(&self, other: &Arc<dyn FlushableCache>) -> bool {
let deps = self.flush_before.lock().await;
for dep in deps.iter() {
if Arc::ptr_eq(dep, other) {
return true;
}
}
false
}
}