use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use web_async::Lock;
const ENTRY_OVERHEAD: u64 = 256;
#[derive(Clone, Default)]
pub struct Pool {
inner: Arc<Inner>,
}
struct Inner {
used: AtomicU64,
capacity: AtomicU64,
epoch: web_async::time::Instant,
lru: Lock<Lru>,
}
impl Default for Inner {
fn default() -> Self {
Self {
used: AtomicU64::new(0),
capacity: AtomicU64::new(u64::MAX),
epoch: web_async::time::Instant::now(),
lru: Lock::default(),
}
}
}
#[derive(Default)]
struct Lru {
heap: BinaryHeap<Reverse<(u64, u64)>>,
entries: HashMap<u64, Arc<Entry>>,
next_id: u64,
}
pub(crate) struct Entry {
id: u64,
bytes: AtomicU64,
last_access: AtomicU64,
pinned: AtomicBool,
evict: Box<dyn Fn() + Send + Sync>,
epoch: web_async::time::Instant,
}
impl Entry {
fn now(&self) -> u64 {
self.epoch.elapsed().as_millis() as u64
}
pub(crate) fn touch(&self) {
self.last_access.store(self.now(), Ordering::Relaxed);
}
pub(crate) fn set_pinned(&self, pinned: bool) {
self.pinned.store(pinned, Ordering::Relaxed);
}
}
impl Pool {
pub fn new(capacity: u64) -> Self {
let pool = Self::default();
pool.inner.capacity.store(capacity, Ordering::Relaxed);
pool
}
pub fn unbounded() -> Self {
Self::default()
}
pub fn capacity(&self) -> Option<u64> {
match self.inner.capacity.load(Ordering::Relaxed) {
u64::MAX => None,
capacity => Some(capacity),
}
}
pub fn used(&self) -> u64 {
self.inner.used.load(Ordering::Relaxed)
}
pub fn resize(&self, capacity: impl Into<Option<u64>>) {
let capacity = capacity.into().unwrap_or(u64::MAX);
self.inner.capacity.store(capacity, Ordering::Relaxed);
self.evict();
}
pub fn same_pool(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.inner, &other.inner)
}
#[cfg(test)]
pub(crate) fn debug_entries(&self) -> (u64, Vec<(u64, u64, u64, bool)>) {
let now = self.inner.epoch.elapsed().as_millis() as u64;
let lru = self.inner.lru.lock();
let mut entries: Vec<_> = lru
.entries
.values()
.map(|e| {
(
e.id,
e.last_access.load(Ordering::Relaxed),
e.bytes.load(Ordering::Relaxed),
e.pinned.load(Ordering::Relaxed),
)
})
.collect();
entries.sort_unstable();
(now, entries)
}
pub(crate) fn register(&self, evict: Box<dyn Fn() + Send + Sync>) -> Charge {
let inner = self.inner.clone();
let entry = {
let mut lru = inner.lru.lock();
let id = lru.next_id;
lru.next_id += 1;
let entry = Arc::new(Entry {
id,
bytes: AtomicU64::new(ENTRY_OVERHEAD),
last_access: AtomicU64::new(0),
pinned: AtomicBool::new(false),
evict,
epoch: inner.epoch,
});
entry.touch();
lru.entries.insert(id, entry.clone());
lru.heap.push(Reverse((entry.last_access.load(Ordering::Relaxed), id)));
entry
};
inner.used.fetch_add(ENTRY_OVERHEAD, Ordering::Relaxed);
Charge {
inner: Some((inner, entry)),
}
}
pub(crate) fn evict(&self) {
let inner = &self.inner;
if inner.used.load(Ordering::Relaxed) <= inner.capacity.load(Ordering::Relaxed) {
return;
}
let mut victims = Vec::new();
{
let mut lru = inner.lru.lock();
let mut freed = 0u64;
let mut pinned = Vec::new();
let mut budget = 2 * lru.heap.len();
while budget > 0
&& inner.used.load(Ordering::Relaxed).saturating_sub(freed) > inner.capacity.load(Ordering::Relaxed)
{
budget -= 1;
let Some(Reverse((snapshot, id))) = lru.heap.pop() else {
break;
};
let Some(entry) = lru.entries.get(&id) else {
continue;
};
let access = entry.last_access.load(Ordering::Relaxed);
if access != snapshot {
lru.heap.push(Reverse((access, id)));
continue;
}
if entry.pinned.load(Ordering::Relaxed) {
pinned.push(Reverse((access, id)));
continue;
}
let entry = lru.entries.remove(&id).unwrap();
freed += entry.bytes.load(Ordering::Relaxed);
victims.push(entry);
}
for slot in pinned {
lru.heap.push(slot);
}
}
for victim in victims {
(victim.evict)();
}
}
}
impl std::fmt::Debug for Pool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Pool")
.field("used", &self.used())
.field("capacity", &self.capacity())
.finish()
}
}
#[derive(Default)]
pub(crate) struct Charge {
inner: Option<(Arc<Inner>, Arc<Entry>)>,
}
impl Charge {
pub(crate) fn add(&self, n: u64) {
if let Some((inner, entry)) = &self.inner {
entry.bytes.fetch_add(n, Ordering::Relaxed);
inner.used.fetch_add(n, Ordering::Relaxed);
entry.touch();
}
}
pub(crate) fn sub(&self, n: u64) {
if let Some((inner, entry)) = &self.inner {
entry.bytes.fetch_sub(n, Ordering::Relaxed);
inner.used.fetch_sub(n, Ordering::Relaxed);
}
}
pub(crate) fn clear(&self) {
if let Some((inner, entry)) = &self.inner {
let bytes = entry.bytes.swap(0, Ordering::Relaxed);
inner.used.fetch_sub(bytes, Ordering::Relaxed);
}
}
pub(crate) fn touch(&self) {
if let Some((_, entry)) = &self.inner {
entry.touch();
}
}
pub(crate) fn entry(&self) -> Option<Arc<Entry>> {
self.inner.as_ref().map(|(_, entry)| entry.clone())
}
}
impl Drop for Charge {
fn drop(&mut self) {
let Some((inner, entry)) = self.inner.take() else {
return;
};
let bytes = entry.bytes.swap(0, Ordering::Relaxed);
inner.used.fetch_sub(bytes, Ordering::Relaxed);
inner.lru.lock().entries.remove(&entry.id);
}
}
#[cfg(test)]
mod test {
use super::*;
use std::time::Duration;
fn flag() -> (Arc<AtomicBool>, Box<dyn Fn() + Send + Sync>) {
let evicted = Arc::new(AtomicBool::new(false));
let hook = evicted.clone();
(
evicted,
Box::new(move || {
hook.store(true, Ordering::Relaxed);
}),
)
}
#[test]
fn unbounded_never_evicts() {
let pool = Pool::unbounded();
let (evicted, hook) = flag();
let charge = pool.register(hook);
charge.add(1 << 40);
pool.evict();
assert!(!evicted.load(Ordering::Relaxed));
assert_eq!(pool.used(), (1 << 40) + ENTRY_OVERHEAD);
drop(charge);
assert_eq!(pool.used(), 0);
}
#[test]
fn detached_charge_is_noop() {
let charge = Charge::default();
charge.add(123);
charge.sub(23);
charge.clear();
charge.touch();
assert!(charge.entry().is_none());
}
#[tokio::test]
async fn evicts_least_recently_used() {
tokio::time::pause();
let pool = Pool::new(3 * ENTRY_OVERHEAD + 2500);
let (evicted_a, hook_a) = flag();
let (evicted_b, hook_b) = flag();
let (evicted_c, hook_c) = flag();
let a = pool.register(hook_a);
a.add(1000);
tokio::time::advance(Duration::from_millis(10)).await;
let b = pool.register(hook_b);
b.add(1000);
tokio::time::advance(Duration::from_millis(10)).await;
a.touch();
tokio::time::advance(Duration::from_millis(10)).await;
let c = pool.register(hook_c);
c.add(1000);
pool.evict();
assert!(!evicted_a.load(Ordering::Relaxed));
assert!(evicted_b.load(Ordering::Relaxed));
assert!(!evicted_c.load(Ordering::Relaxed));
b.clear();
assert_eq!(pool.used(), 2 * (1000 + ENTRY_OVERHEAD));
}
#[tokio::test]
async fn pinned_entries_survive() {
tokio::time::pause();
let pool = Pool::new(ENTRY_OVERHEAD);
let (evicted_a, hook_a) = flag();
let (evicted_b, hook_b) = flag();
let a = pool.register(hook_a);
a.entry().unwrap().set_pinned(true);
a.add(1000);
tokio::time::advance(Duration::from_millis(10)).await;
let b = pool.register(hook_b);
b.add(1000);
pool.evict();
assert!(!evicted_a.load(Ordering::Relaxed));
assert!(evicted_b.load(Ordering::Relaxed));
b.clear();
drop(b);
pool.evict();
assert!(!evicted_a.load(Ordering::Relaxed));
}
#[tokio::test]
async fn resize_evicts() {
tokio::time::pause();
let pool = Pool::unbounded();
let (evicted_a, hook_a) = flag();
let a = pool.register(hook_a);
a.add(1000);
pool.evict();
assert!(!evicted_a.load(Ordering::Relaxed));
pool.resize(100);
assert!(evicted_a.load(Ordering::Relaxed));
pool.resize(None);
assert_eq!(pool.capacity(), None);
}
#[tokio::test]
async fn touched_entry_is_rekeyed_not_evicted() {
tokio::time::pause();
let pool = Pool::new(2 * ENTRY_OVERHEAD + 1500);
let (evicted_a, hook_a) = flag();
let (evicted_b, hook_b) = flag();
let a = pool.register(hook_a);
a.add(1000);
tokio::time::advance(Duration::from_millis(10)).await;
let b = pool.register(hook_b);
b.add(1000);
tokio::time::advance(Duration::from_millis(10)).await;
a.touch();
tokio::time::advance(Duration::from_millis(10)).await;
pool.evict();
assert!(!evicted_a.load(Ordering::Relaxed));
assert!(evicted_b.load(Ordering::Relaxed));
}
#[test]
fn dropped_charge_leaves_stale_heap_slot() {
let pool = Pool::new(0);
let (evicted_a, hook_a) = flag();
let a = pool.register(hook_a);
drop(a);
let (evicted_b, hook_b) = flag();
let b = pool.register(hook_b);
b.add(1);
pool.evict();
assert!(!evicted_a.load(Ordering::Relaxed));
assert!(evicted_b.load(Ordering::Relaxed));
}
}