use std::cmp::Reverse;
use std::collections::BinaryHeap;
use super::tree::{align_down, match_len, NodeId, RadixTree, ROOT};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MatchResult {
pub cached_len: usize,
pub node: NodeId,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InsertResult {
pub cached_len: usize,
pub node: NodeId,
pub inserted_len: usize,
}
#[derive(Debug)]
pub struct RadixCache {
tree: RadixTree,
evictable_size: usize,
protected_size: usize,
clock: i64,
}
impl RadixCache {
pub fn new(page_size: usize) -> Self {
RadixCache {
tree: RadixTree::new(page_size),
evictable_size: 0,
protected_size: 0,
clock: 0,
}
}
pub fn page_size(&self) -> usize {
self.tree.page_size()
}
pub fn tree(&self) -> &RadixTree {
&self.tree
}
pub fn evictable_size(&self) -> usize {
self.evictable_size
}
pub fn protected_size(&self) -> usize {
self.protected_size
}
pub fn total_size(&self) -> usize {
self.evictable_size + self.protected_size
}
pub fn matched_indices(&self, node: NodeId) -> Vec<u32> {
self.tree.path_value(node)
}
fn tick(&mut self) -> i64 {
self.clock += 1;
self.clock
}
fn tree_walk(&mut self, input_ids: &[u32]) -> (NodeId, usize) {
let page_size = self.tree.page_size();
let tic = self.tick();
let mut prefix_len = 0usize;
let mut node = ROOT;
while prefix_len < input_ids.len() {
let child = match self.tree.child(node, &input_ids[prefix_len..]) {
Some(child) => child,
None => return (node, prefix_len),
};
node = child;
let matched = align_down(
match_len(&self.tree.node(node).key, &input_ids[prefix_len..]),
page_size,
);
prefix_len += matched;
if matched != self.tree.node(node).length() {
node = self.tree.split_at(node, matched);
self.tree.node_mut(node).timestamp = tic;
return (node, prefix_len);
}
self.tree.node_mut(node).timestamp = tic;
}
(node, prefix_len)
}
pub fn match_prefix(&mut self, input_ids: &[u32]) -> MatchResult {
let (node, cached_len) = self.tree_walk(input_ids);
MatchResult { cached_len, node }
}
pub fn insert_prefix(&mut self, input_ids: &[u32], indices: &[u32]) -> InsertResult {
assert_eq!(
input_ids.len(),
indices.len(),
"one page index per inserted token"
);
let insert_len = align_down(input_ids.len(), self.tree.page_size());
let input_ids = &input_ids[..insert_len];
let indices = &indices[..insert_len];
let (node, prefix_len) = self.tree_walk(input_ids);
let mut node = node;
if prefix_len != insert_len {
let tic = self.clock;
let new_node = self.tree.alloc(tic);
self.tree.set_key_value(
new_node,
input_ids[prefix_len..].to_vec(),
indices[prefix_len..].to_vec(),
);
self.tree.set_parent(new_node, node);
self.evictable_size += self.tree.node(new_node).length();
node = new_node;
}
InsertResult {
cached_len: prefix_len,
node,
inserted_len: insert_len,
}
}
pub fn lock(&mut self, node: NodeId) {
let mut cur = node;
while !self.tree.node(cur).is_root() {
if self.tree.node(cur).ref_count == 0 {
let length = self.tree.node(cur).length();
self.evictable_size -= length;
self.protected_size += length;
}
self.tree.node_mut(cur).ref_count += 1;
cur = self.tree.parent(cur).expect("non-root has a parent");
}
}
pub fn unlock(&mut self, node: NodeId) {
let mut cur = node;
while !self.tree.node(cur).is_root() {
let count = self.tree.node(cur).ref_count;
assert!(count > 0, "unlock without a matching lock");
self.tree.node_mut(cur).ref_count = count - 1;
if count - 1 == 0 {
let length = self.tree.node(cur).length();
self.evictable_size += length;
self.protected_size -= length;
}
cur = self.tree.parent(cur).expect("non-root has a parent");
}
}
pub fn evict(&mut self, size: usize) -> Vec<u32> {
if size == 0 {
return Vec::new();
}
assert!(
size <= self.evictable_size,
"cannot evict {size} tokens, only {} are evictable",
self.evictable_size
);
let mut heap: BinaryHeap<Reverse<(i64, u32)>> = self
.tree
.leaves()
.into_iter()
.filter(|id| self.tree.node(*id).ref_count == 0)
.map(|id| Reverse((self.tree.node(id).timestamp, id.0)))
.collect();
let mut evicted: Vec<u32> = Vec::new();
let mut evicted_size = 0usize;
while evicted_size < size {
let Reverse((_, raw)) = heap
.pop()
.expect("evictable_size promised more than the tree holds");
let id = NodeId(raw);
let node = self.tree.node(id);
debug_assert!(node.ref_count == 0 && node.is_leaf() && !node.is_root());
evicted_size += node.length();
evicted.extend_from_slice(&node.value);
self.evictable_size -= node.length();
let parent = self.tree.unlink(id);
if self.tree.node(parent).is_leaf()
&& self.tree.node(parent).ref_count == 0
&& !self.tree.node(parent).is_root()
{
heap.push(Reverse((self.tree.node(parent).timestamp, parent.0)));
}
}
evicted
}
pub fn check_integrity(&self) {
self.tree.check_structure();
let mut evictable = 0usize;
let mut protected = 0usize;
for id in self.tree.walk() {
let node = self.tree.node(id);
if node.ref_count == 0 {
evictable += node.length();
} else {
protected += node.length();
}
}
assert_eq!(evictable, self.evictable_size, "evictable_size drifted");
assert_eq!(protected, self.protected_size, "protected_size drifted");
}
}
#[cfg(test)]
mod tests {
use super::*;
const P: usize = 4;
fn ids(range: std::ops::Range<u32>) -> Vec<u32> {
range.collect()
}
fn pages(start: u32, len: usize) -> Vec<u32> {
(start..start + len as u32).collect()
}
#[test]
fn a_second_request_reuses_the_first_ones_pages() {
let mut cache = RadixCache::new(P);
let prompt = ids(100..108);
let slots = pages(0, 8);
let res = cache.insert_prefix(&prompt, &slots);
assert_eq!(res.cached_len, 0, "nothing was cached before");
assert_eq!(cache.evictable_size(), 8);
let m = cache.match_prefix(&ids(100..108));
assert_eq!(m.cached_len, 8);
assert_eq!(cache.matched_indices(m.node), slots);
}
#[test]
fn reuse_is_truncated_to_whole_pages() {
let mut cache = RadixCache::new(P);
cache.insert_prefix(&ids(100..108), &pages(0, 8));
let mut diverging = ids(100..106);
diverging.extend([999, 999]);
let m = cache.match_prefix(&diverging);
assert_eq!(m.cached_len, 4, "six matched tokens are one whole page");
assert_eq!(cache.matched_indices(m.node), pages(0, 4));
}
#[test]
fn a_trailing_partial_page_is_not_inserted() {
let mut cache = RadixCache::new(P);
let res = cache.insert_prefix(&ids(0..6), &pages(0, 6));
assert_eq!(res.inserted_len, 4);
assert_eq!(cache.total_size(), 4);
assert_eq!(cache.match_prefix(&ids(0..6)).cached_len, 4);
}
#[test]
fn a_sub_page_prompt_resolves_to_the_root() {
let mut cache = RadixCache::new(P);
cache.insert_prefix(&ids(0..8), &pages(0, 8));
let m = cache.match_prefix(&ids(0..3));
assert_eq!(m.cached_len, 0);
assert!(cache.matched_indices(m.node).is_empty());
}
#[test]
fn reinserting_a_cached_prefix_reports_the_duplicate_span() {
let mut cache = RadixCache::new(P);
let first = pages(0, 8);
cache.insert_prefix(&ids(100..108), &first);
let second = pages(50, 8);
let res = cache.insert_prefix(&ids(100..108), &second);
assert_eq!(res.cached_len, 8, "the caller must free all eight");
assert_eq!(
cache.matched_indices(res.node),
first,
"the tree keeps its own pages"
);
assert_eq!(cache.total_size(), 8, "and stores nothing new");
}
#[test]
fn an_inserted_node_owns_its_indices() {
let mut cache = RadixCache::new(P);
let mut slots = pages(0, 4);
let res = cache.insert_prefix(&ids(0..4), &slots);
slots.iter_mut().for_each(|s| *s = 999);
assert_eq!(cache.matched_indices(res.node), pages(0, 4));
}
#[test]
fn locking_walks_to_the_root_and_only_the_zero_to_one_edge_moves_size() {
let mut cache = RadixCache::new(P);
cache.insert_prefix(&ids(0..8), &pages(0, 8));
let head = cache.match_prefix(&ids(0..4)); let tail = cache.match_prefix(&ids(0..8));
cache.lock(tail.node);
assert_eq!(cache.protected_size(), 8);
assert_eq!(cache.evictable_size(), 0);
cache.lock(tail.node);
assert_eq!(cache.protected_size(), 8);
cache.unlock(tail.node);
assert_eq!(cache.protected_size(), 8);
cache.unlock(tail.node);
assert_eq!(cache.protected_size(), 0);
assert_eq!(cache.evictable_size(), 8);
cache.lock(head.node);
assert_eq!(cache.protected_size(), 4);
assert_eq!(cache.evictable_size(), 4);
cache.check_integrity();
}
#[test]
fn eviction_takes_the_least_recently_matched_leaf() {
let mut cache = RadixCache::new(P);
cache.insert_prefix(&ids(0..4), &pages(0, 4));
cache.insert_prefix(&ids(100..104), &pages(10, 4));
cache.match_prefix(&ids(0..4));
let freed = cache.evict(1);
assert_eq!(freed, pages(10, 4), "whole nodes only, oldest first");
assert_eq!(cache.evictable_size(), 4);
cache.check_integrity();
}
#[test]
fn a_locked_leaf_is_never_evicted() {
let mut cache = RadixCache::new(P);
let a = cache.insert_prefix(&ids(0..4), &pages(0, 4));
cache.insert_prefix(&ids(100..104), &pages(10, 4));
cache.lock(a.node);
let freed = cache.evict(4);
assert_eq!(freed, pages(10, 4));
cache.check_integrity();
}
#[test]
fn eviction_cascades_into_a_newly_childless_parent() {
let mut cache = RadixCache::new(P);
cache.insert_prefix(&ids(0..8), &pages(0, 8));
cache.match_prefix(&ids(0..4));
let mut freed = cache.evict(8);
freed.sort_unstable();
assert_eq!(freed, pages(0, 8), "both levels, one call");
assert_eq!(cache.total_size(), 0);
cache.check_integrity();
}
#[test]
#[should_panic(expected = "cannot evict")]
fn evicting_more_than_the_tree_holds_is_a_caller_bug() {
let mut cache = RadixCache::new(P);
cache.insert_prefix(&ids(0..4), &pages(0, 4));
cache.evict(8);
}
#[test]
fn evicting_nothing_is_always_safe() {
let mut cache = RadixCache::new(P);
assert!(cache.evict(0).is_empty());
cache.check_integrity();
}
#[test]
fn page_size_one_degenerates_to_a_token_trie() {
let mut cache = RadixCache::new(1);
cache.insert_prefix(&ids(0..5), &pages(0, 5));
assert_eq!(cache.match_prefix(&ids(0..3)).cached_len, 3);
let m = cache.match_prefix(&ids(0..5));
assert_eq!(m.cached_len, 5);
assert_eq!(cache.matched_indices(m.node), pages(0, 5));
cache.check_integrity();
}
#[test]
fn pages_are_conserved_across_a_mixed_workload() {
let mut cache = RadixCache::new(P);
let mut next_page = 0u32;
let mut returned: Vec<u32> = Vec::new();
for round in 0..6u32 {
let prompt: Vec<u32> = (0..8).map(|i| if i < 4 { i } else { i + round }).collect();
let slots = pages(next_page, 8);
next_page += 8;
let m = cache.match_prefix(&prompt);
cache.lock(m.node);
let res = cache.insert_prefix(&prompt, &slots);
returned.extend_from_slice(&slots[..res.cached_len]);
cache.unlock(m.node);
cache.check_integrity();
}
let mut live: Vec<u32> = cache
.tree()
.walk()
.into_iter()
.flat_map(|id| cache.tree().node(id).value.clone())
.collect();
live.extend(cache.evict(cache.evictable_size()));
live.extend(returned);
live.sort_unstable();
live.dedup();
assert_eq!(
live.len(),
next_page as usize,
"every page is accounted for exactly once"
);
}
}