use std::collections::HashMap;
use super::{NodeId, RadixCache};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Handle {
pub salt: Option<u64>,
pub node: NodeId,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SaltedMatch {
pub cached_len: usize,
pub handle: Handle,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SaltedInsert {
pub cached_len: usize,
pub inserted_len: usize,
pub handle: Handle,
}
#[derive(Debug)]
pub struct SaltedRadix {
page_size: usize,
namespaces: HashMap<Option<u64>, RadixCache>,
order: Vec<Option<u64>>,
evict_cursor: usize,
}
impl SaltedRadix {
pub fn new(page_size: usize) -> Self {
SaltedRadix {
page_size,
namespaces: HashMap::new(),
order: Vec::new(),
evict_cursor: 0,
}
}
pub fn page_size(&self) -> usize {
self.page_size
}
fn tree(&mut self, salt: Option<u64>) -> &mut RadixCache {
let page_size = self.page_size;
self.namespaces.entry(salt).or_insert_with(|| {
self.order.push(salt);
RadixCache::new(page_size)
})
}
pub fn match_prefix(&mut self, salt: Option<u64>, input_ids: &[u32]) -> SaltedMatch {
let m = self.tree(salt).match_prefix(input_ids);
SaltedMatch {
cached_len: m.cached_len,
handle: Handle { salt, node: m.node },
}
}
pub fn insert_prefix(
&mut self,
salt: Option<u64>,
input_ids: &[u32],
indices: &[u32],
) -> SaltedInsert {
let r = self.tree(salt).insert_prefix(input_ids, indices);
SaltedInsert {
cached_len: r.cached_len,
inserted_len: r.inserted_len,
handle: Handle { salt, node: r.node },
}
}
pub fn matched_indices(&mut self, handle: Handle) -> Vec<u32> {
self.tree(handle.salt).matched_indices(handle.node)
}
pub fn lock(&mut self, handle: Handle) {
self.tree(handle.salt).lock(handle.node);
}
pub fn unlock(&mut self, handle: Handle) {
self.tree(handle.salt).unlock(handle.node);
}
pub fn evictable_size(&self) -> usize {
self.namespaces.values().map(|c| c.evictable_size()).sum()
}
pub fn protected_size(&self) -> usize {
self.namespaces.values().map(|c| c.protected_size()).sum()
}
pub fn total_size(&self) -> usize {
self.namespaces.values().map(|c| c.total_size()).sum()
}
pub fn evict(&mut self, size: usize) -> Vec<u32> {
let mut freed = Vec::new();
if size == 0 || self.order.is_empty() {
return freed;
}
let mut taken = 0usize;
loop {
let mut progress = false;
for _ in 0..self.order.len() {
let salt = self.order[self.evict_cursor % self.order.len()];
self.evict_cursor = self.evict_cursor.wrapping_add(1);
let cache = self
.namespaces
.get_mut(&salt)
.expect("every ordered namespace exists");
if cache.evictable_size() == 0 {
continue;
}
let got = cache.evict(1);
taken += got.len();
freed.extend(got);
progress = true;
if taken >= size {
return freed;
}
}
if !progress {
return freed;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const PAGE: usize = 4;
fn ids(n: usize) -> Vec<u32> {
(0..n as u32).collect()
}
#[test]
fn two_salts_cannot_see_each_others_prefixes() {
let mut cache = SaltedRadix::new(PAGE);
let tokens = ids(8);
let pages: Vec<u32> = vec![7; 8];
cache.insert_prefix(Some(1), &tokens, &pages);
assert_eq!(
cache.match_prefix(Some(1), &tokens).cached_len,
8,
"a caller must match its own prefix"
);
assert_eq!(
cache.match_prefix(Some(2), &tokens).cached_len,
0,
"a different salt was served the first caller's pages"
);
assert_eq!(
cache.match_prefix(None, &tokens).cached_len,
0,
"an unsalted request was served a salted caller's pages"
);
}
#[test]
fn the_shared_namespace_does_not_leak_into_a_salted_one() {
let mut cache = SaltedRadix::new(PAGE);
let tokens = ids(8);
cache.insert_prefix(None, &tokens, &[3; 8]);
assert_eq!(cache.match_prefix(None, &tokens).cached_len, 8);
assert_eq!(
cache.match_prefix(Some(9), &tokens).cached_len,
0,
"a salted caller matched the shared namespace"
);
}
#[test]
fn a_lock_protects_the_namespace_it_matched_in() {
let mut cache = SaltedRadix::new(PAGE);
let tokens = ids(8);
cache.insert_prefix(Some(1), &tokens, &[1; 8]);
cache.insert_prefix(Some(2), &tokens, &[2; 8]);
assert_eq!(cache.evictable_size(), 16, "two namespaces, eight each");
let m = cache.match_prefix(Some(1), &tokens);
cache.lock(m.handle);
assert_eq!(cache.protected_size(), 8, "one namespace's worth is held");
assert_eq!(cache.evictable_size(), 8, "the other is still evictable");
cache.unlock(m.handle);
assert_eq!(cache.protected_size(), 0);
assert_eq!(cache.evictable_size(), 16);
}
#[test]
fn eviction_reaches_across_namespaces() {
let mut cache = SaltedRadix::new(PAGE);
for salt in 0..4u64 {
cache.insert_prefix(Some(salt), &ids(8), &[salt as u32; 8]);
}
assert_eq!(cache.evictable_size(), 32);
let freed = cache.evict(24);
assert!(
freed.len() >= 24,
"eviction stopped inside one namespace: freed {}",
freed.len()
);
let touched: std::collections::BTreeSet<u32> = freed.into_iter().collect();
assert!(
touched.len() > 1,
"every freed page came from one namespace: {touched:?}"
);
}
#[test]
fn asking_for_more_than_is_held_returns_what_there_is() {
let mut cache = SaltedRadix::new(PAGE);
cache.insert_prefix(Some(1), &ids(8), &[1; 8]);
let freed = cache.evict(1_000);
assert_eq!(freed.len(), 8);
assert_eq!(cache.evictable_size(), 0);
assert!(cache.evict(1_000).is_empty(), "a second call must not spin");
}
#[test]
fn a_namespace_is_created_on_first_use() {
let mut cache = SaltedRadix::new(PAGE);
assert_eq!(cache.total_size(), 0);
assert!(cache.evict(4).is_empty(), "nothing to evict yet");
cache.insert_prefix(Some(5), &ids(4), &[0; 4]);
assert_eq!(cache.total_size(), 4);
}
}