217 lines
6.9 KiB
Rust
217 lines
6.9 KiB
Rust
use std::cmp::Reverse;
|
|
use std::collections::{BinaryHeap, HashMap, HashSet};
|
|
use std::sync::Arc;
|
|
use std::time::{Duration, Instant};
|
|
|
|
use tokio::sync::{Notify, RwLock};
|
|
|
|
type Expiry = Reverse<(Instant, u64, usize)>;
|
|
|
|
#[derive(Default)]
|
|
struct BlockMapState {
|
|
locations: HashMap<u64, HashMap<usize, CacheLocation>>,
|
|
expirations: BinaryHeap<Expiry>,
|
|
}
|
|
|
|
struct CacheLocation {
|
|
deadline: Instant,
|
|
}
|
|
|
|
#[derive(Debug, Default, PartialEq, Eq)]
|
|
pub struct VacuumStats {
|
|
pub locations_removed: usize,
|
|
pub blocks_removed: usize,
|
|
}
|
|
|
|
/// Best-effort knowledge of when each backend last accessed a block prefix.
|
|
#[derive(Clone)]
|
|
pub struct BlockMap {
|
|
inner: Arc<RwLock<BlockMapState>>,
|
|
expiry_changed: Arc<Notify>,
|
|
}
|
|
|
|
impl BlockMap {
|
|
pub fn new() -> Self {
|
|
Self {
|
|
inner: Arc::new(RwLock::new(BlockMapState::default())),
|
|
expiry_changed: Arc::new(Notify::new()),
|
|
}
|
|
}
|
|
|
|
/// Return unexpired locations and remove stale associations for `hash`.
|
|
pub async fn fresh_locations(&self, hash: u64) -> HashSet<usize> {
|
|
let now = Instant::now();
|
|
let mut state = self.inner.write().await;
|
|
let mut remove_hash = false;
|
|
let locations = match state.locations.get_mut(&hash) {
|
|
Some(locations) => {
|
|
locations.retain(|_, location| location.deadline > now);
|
|
remove_hash = locations.is_empty();
|
|
locations.keys().copied().collect()
|
|
}
|
|
None => HashSet::new(),
|
|
};
|
|
if remove_hash {
|
|
state.locations.remove(&hash);
|
|
}
|
|
locations
|
|
}
|
|
|
|
/// Record or refresh a backend's access time for every supplied block.
|
|
pub async fn record_access(&self, hashes: &[u64], backend_idx: usize, ttl: Duration) {
|
|
let now = Instant::now();
|
|
let deadline = now.checked_add(ttl).expect("cache TTL is too large");
|
|
let mut state = self.inner.write().await;
|
|
let previous_earliest = state
|
|
.expirations
|
|
.peek()
|
|
.map(|Reverse((deadline, _, _))| *deadline);
|
|
for &hash in hashes {
|
|
state
|
|
.locations
|
|
.entry(hash)
|
|
.or_default()
|
|
.insert(backend_idx, CacheLocation { deadline });
|
|
state
|
|
.expirations
|
|
.push(Reverse((deadline, hash, backend_idx)));
|
|
}
|
|
drop(state);
|
|
if !hashes.is_empty() && previous_earliest.is_none_or(|earliest| deadline < earliest) {
|
|
self.expiry_changed.notify_one();
|
|
}
|
|
}
|
|
|
|
/// Remove every location whose current access timestamp has expired.
|
|
/// Stale heap entries created by refreshes are ignored.
|
|
pub async fn vacuum(&self) -> VacuumStats {
|
|
let now = Instant::now();
|
|
let mut state = self.inner.write().await;
|
|
let mut stats = VacuumStats::default();
|
|
|
|
while state
|
|
.expirations
|
|
.peek()
|
|
.is_some_and(|Reverse((deadline, _, _))| *deadline <= now)
|
|
{
|
|
let Reverse((deadline, hash, backend_idx)) = state.expirations.pop().unwrap();
|
|
let mut remove_block = false;
|
|
if let Some(locations) = state.locations.get_mut(&hash) {
|
|
let expired = locations
|
|
.get(&backend_idx)
|
|
.is_some_and(|location| location.deadline == deadline);
|
|
if expired {
|
|
locations.remove(&backend_idx);
|
|
stats.locations_removed += 1;
|
|
remove_block = locations.is_empty();
|
|
}
|
|
}
|
|
if remove_block {
|
|
state.locations.remove(&hash);
|
|
stats.blocks_removed += 1;
|
|
}
|
|
}
|
|
|
|
stats
|
|
}
|
|
|
|
/// Wait until the current heap minimum is due. Earlier insertions wake the
|
|
/// wait so the deadline can be recalculated.
|
|
pub async fn wait_until_expiry(&self) {
|
|
loop {
|
|
let notified = self.expiry_changed.notified();
|
|
let deadline = self
|
|
.inner
|
|
.read()
|
|
.await
|
|
.expirations
|
|
.peek()
|
|
.map(|Reverse((deadline, _, _))| *deadline);
|
|
|
|
match deadline {
|
|
Some(deadline) => {
|
|
tokio::select! {
|
|
_ = tokio::time::sleep_until(tokio::time::Instant::from_std(deadline)) => return,
|
|
_ = notified => {}
|
|
}
|
|
}
|
|
None => notified.await,
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
const TTL: Duration = Duration::from_millis(10);
|
|
|
|
#[tokio::test]
|
|
async fn test_unknown_hash_has_no_locations() {
|
|
let map = BlockMap::new();
|
|
assert!(map.fresh_locations(42).await.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_record_preserves_and_refreshes_all_locations() {
|
|
let map = BlockMap::new();
|
|
map.record_access(&[42], 0, TTL).await;
|
|
map.record_access(&[42], 1, TTL).await;
|
|
assert_eq!(map.fresh_locations(42).await, HashSet::from([0, 1]));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_expired_location_is_invalidated_lazily() {
|
|
let map = BlockMap::new();
|
|
map.record_access(&[42], 0, Duration::from_millis(1)).await;
|
|
tokio::time::sleep(Duration::from_millis(2)).await;
|
|
assert!(map.fresh_locations(42).await.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_vacuum_removes_expired_location_and_block() {
|
|
let map = BlockMap::new();
|
|
let ttl = Duration::from_millis(1);
|
|
map.record_access(&[42], 0, ttl).await;
|
|
tokio::time::sleep(Duration::from_millis(2)).await;
|
|
|
|
assert_eq!(
|
|
map.vacuum().await,
|
|
VacuumStats {
|
|
locations_removed: 1,
|
|
blocks_removed: 1,
|
|
}
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_vacuum_ignores_stale_deadline_after_refresh() {
|
|
let map = BlockMap::new();
|
|
let ttl = Duration::from_millis(4);
|
|
map.record_access(&[42], 0, ttl).await;
|
|
tokio::time::sleep(Duration::from_millis(3)).await;
|
|
map.record_access(&[42], 0, ttl).await;
|
|
tokio::time::sleep(Duration::from_millis(2)).await;
|
|
|
|
assert_eq!(map.vacuum().await, VacuumStats::default());
|
|
assert_eq!(map.fresh_locations(42).await, HashSet::from([0]));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_earlier_deadline_wakes_waiter() {
|
|
let map = BlockMap::new();
|
|
map.record_access(&[1], 0, Duration::from_millis(100)).await;
|
|
|
|
let waiting_map = map.clone();
|
|
let waiter = tokio::spawn(async move { waiting_map.wait_until_expiry().await });
|
|
tokio::time::sleep(Duration::from_millis(2)).await;
|
|
map.record_access(&[2], 0, Duration::from_millis(5)).await;
|
|
|
|
tokio::time::timeout(Duration::from_millis(30), waiter)
|
|
.await
|
|
.expect("waiter should reschedule to the earlier deadline")
|
|
.unwrap();
|
|
}
|
|
}
|