use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use lru::LruCache;
use parking_lot::Mutex;
use crate::lexical::index::inverted::core::posting::DecodedPostingList;
fn estimated_bytes(list: &DecodedPostingList) -> usize {
let mut bytes = std::mem::size_of::<DecodedPostingList>();
bytes += list.term.len();
bytes += list.doc_ids.len() * std::mem::size_of::<u32>();
bytes += list.frequencies.len() * std::mem::size_of::<u32>();
bytes += list.weights.len() * std::mem::size_of::<f32>();
for level in &list.skip_levels {
bytes += level.len() * std::mem::size_of::<u32>();
}
if let Some(positions) = &list.positions {
for p in positions.iter().flatten() {
bytes += p.len() * std::mem::size_of::<u32>();
}
}
bytes
}
#[derive(Debug)]
struct PostingCacheInner {
lru: LruCache<String, Arc<DecodedPostingList>>,
cur_bytes: usize,
max_bytes: usize,
}
#[derive(Debug)]
pub struct PostingCache {
inner: Option<Mutex<PostingCacheInner>>,
hits: AtomicU64,
misses: AtomicU64,
}
impl PostingCache {
pub fn new(max_bytes: usize) -> Self {
let inner = (max_bytes > 0).then(|| {
Mutex::new(PostingCacheInner {
lru: LruCache::unbounded(),
cur_bytes: 0,
max_bytes,
})
});
PostingCache {
inner,
hits: AtomicU64::new(0),
misses: AtomicU64::new(0),
}
}
pub fn is_enabled(&self) -> bool {
self.inner.is_some()
}
pub fn get(&self, key: &str) -> Option<Arc<DecodedPostingList>> {
let hit = self
.inner
.as_ref()
.and_then(|inner| inner.lock().lru.get(key).cloned());
if hit.is_some() {
self.hits.fetch_add(1, Ordering::Relaxed);
} else {
self.misses.fetch_add(1, Ordering::Relaxed);
}
hit
}
pub fn put(&self, key: String, list: Arc<DecodedPostingList>) {
let Some(inner) = self.inner.as_ref() else {
return;
};
let size = estimated_bytes(&list);
let mut guard = inner.lock();
if size > guard.max_bytes {
return;
}
if let Some(old) = guard.lru.put(key, list) {
guard.cur_bytes = guard.cur_bytes.saturating_sub(estimated_bytes(&old));
}
guard.cur_bytes += size;
while guard.cur_bytes > guard.max_bytes {
match guard.lru.pop_lru() {
Some((_, evicted)) => {
guard.cur_bytes = guard.cur_bytes.saturating_sub(estimated_bytes(&evicted));
}
None => break,
}
}
}
pub fn len(&self) -> usize {
self.inner
.as_ref()
.map_or(0, |inner| inner.lock().lru.len())
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn stats(&self) -> PostingCacheStats {
PostingCacheStats {
hits: self.hits.load(Ordering::Relaxed),
misses: self.misses.load(Ordering::Relaxed),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PostingCacheStats {
pub hits: u64,
pub misses: u64,
}
#[cfg(test)]
mod tests {
use super::*;
fn list(term: &str, n: u32) -> Arc<DecodedPostingList> {
Arc::new(DecodedPostingList {
term: term.to_string(),
doc_ids: (0..n).collect(),
frequencies: vec![1; n as usize],
weights: Vec::new(),
positions: None,
skip_levels: Vec::new(),
total_frequency: n as u64,
doc_frequency: n as u64,
})
}
#[test]
fn put_then_get_returns_cached_list() {
let cache = PostingCache::new(1 << 20);
cache.put("body\u{1}rust".to_string(), list("rust", 4));
let got = cache.get("body\u{1}rust").expect("present");
assert_eq!(got.doc_ids, vec![0, 1, 2, 3]);
assert_eq!(cache.stats().hits, 1);
assert_eq!(cache.stats().misses, 0);
}
#[test]
fn miss_increments_miss_counter() {
let cache = PostingCache::new(1 << 20);
assert!(cache.get("absent").is_none());
assert_eq!(cache.stats().misses, 1);
}
#[test]
fn capacity_zero_disables_cache() {
let cache = PostingCache::new(0);
assert!(!cache.is_enabled());
cache.put("k".to_string(), list("k", 4));
assert!(cache.get("k").is_none());
assert!(cache.is_empty());
}
#[test]
fn byte_budget_evicts_least_recently_used() {
let one = estimated_bytes(&list("x", 64));
let cache = PostingCache::new(one * 2 + one / 2);
cache.put("a".to_string(), list("a", 64));
cache.put("b".to_string(), list("b", 64));
assert!(cache.get("a").is_some());
cache.put("c".to_string(), list("c", 64));
assert!(cache.get("a").is_some(), "recently used 'a' survives");
assert!(cache.get("c").is_some(), "just-inserted 'c' survives");
assert!(cache.get("b").is_none(), "LRU 'b' must be evicted");
}
#[test]
fn oversized_entry_is_not_cached() {
let big = list("big", 1024);
let cache = PostingCache::new(estimated_bytes(&big) / 2);
cache.put("big".to_string(), big);
assert!(
cache.get("big").is_none(),
"entry larger than budget is skipped"
);
assert!(cache.is_empty());
}
}