Skip to main content

dynamo_mocker/cache/
radix_cache.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Radix-tree KV cache for SGLang engine simulation.
5//!
6//! Reference: sglang/python/sglang/srt/mem_cache/radix_cache.py
7
8use dynamo_kv_router::protocols::{BlockHashOptions, LocalBlockHash, compute_block_hash_for_seq};
9use rustc_hash::{FxHashMap, FxHashSet};
10use slotmap::{SlotMap, new_key_type};
11use std::time::Instant;
12
13new_key_type! {
14    /// Stable identifier for a tree node inside the [`RadixCache`].
15    pub struct NodeId;
16}
17
18/// Physical page identifier in the simulated SGLang KV pool.
19#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
20pub struct KvPageId(usize);
21
22impl KvPageId {
23    pub(crate) fn from_token_index(index: usize, page_size: usize) -> Self {
24        Self(index / page_size)
25    }
26
27    pub(crate) fn first_token_index(self, page_size: usize) -> usize {
28        self.0 * page_size
29    }
30
31    pub(crate) fn terminal_token_index(self, page_size: usize) -> usize {
32        self.first_token_index(page_size) + page_size - 1
33    }
34}
35
36/// Manages free / allocated pages for the simulated SGLang KV cache.
37///
38/// SGLang's paged allocator owns and frees whole pages even though its public
39/// interface returns flattened token indices. This pool preserves that
40/// ownership model and only expands page IDs while a request is active.
41pub struct PagePool {
42    next_fresh: usize,
43    free: Vec<KvPageId>,
44    total_pages: usize,
45    page_size: usize,
46}
47
48impl PagePool {
49    pub fn new(total_tokens: usize, page_size: usize) -> Self {
50        assert!(page_size >= 1, "page_size must be >= 1");
51        Self {
52            next_fresh: 0,
53            free: Vec::new(),
54            total_pages: total_tokens / page_size,
55            page_size,
56        }
57    }
58
59    pub fn allocate_pages(&mut self, count: usize) -> Option<Vec<KvPageId>> {
60        if self.available_pages() < count {
61            return None;
62        }
63
64        let recycled = count.min(self.free.len());
65        let fresh = count - recycled;
66        let mut pages = Vec::with_capacity(count);
67        pages.extend(self.free.drain(self.free.len() - recycled..));
68        pages.extend((self.next_fresh..self.next_fresh + fresh).map(KvPageId));
69        self.next_fresh += fresh;
70        Some(pages)
71    }
72
73    #[cfg(test)]
74    pub fn allocate(&mut self, token_count: usize) -> Option<Vec<usize>> {
75        let mut indices = Vec::new();
76        self.allocate_indices_into(token_count, &mut indices)
77            .then_some(indices)
78    }
79
80    /// Append flattened indices for `new_tokens`, allocating whole pages only
81    /// when the request's current final page has no remaining slots.
82    pub fn allocate_indices_into(&mut self, new_tokens: usize, indices: &mut Vec<usize>) -> bool {
83        if new_tokens == 0 {
84            return true;
85        }
86
87        let available_in_last_page = indices
88            .last()
89            .map_or(0, |last| self.page_size - 1 - (last % self.page_size));
90        let tokens_requiring_pages = new_tokens.saturating_sub(available_in_last_page);
91        let required_pages = tokens_requiring_pages.div_ceil(self.page_size);
92        if self.available_pages() < required_pages {
93            return false;
94        }
95
96        indices.reserve(new_tokens);
97        let from_existing = new_tokens.min(available_in_last_page);
98        if from_existing > 0 {
99            let start = indices.last().copied().expect("last page must exist") + 1;
100            indices.extend(start..start + from_existing);
101        }
102
103        let remaining = new_tokens - from_existing;
104        let Some(pages) = self.allocate_pages(required_pages) else {
105            return false;
106        };
107        for (page_idx, page) in pages.into_iter().enumerate() {
108            let take = remaining
109                .saturating_sub(page_idx * self.page_size)
110                .min(self.page_size);
111            let start = page.first_token_index(self.page_size);
112            indices.extend(start..start + take);
113        }
114        true
115    }
116
117    pub fn expand_pages(&self, pages: &[KvPageId], token_count: usize) -> Vec<usize> {
118        assert!(
119            token_count <= pages.len() * self.page_size,
120            "cannot expand {token_count} tokens from {} pages of size {}",
121            pages.len(),
122            self.page_size
123        );
124        let mut indices = Vec::with_capacity(token_count);
125        for (page_idx, page) in pages.iter().copied().enumerate() {
126            let take = token_count
127                .saturating_sub(page_idx * self.page_size)
128                .min(self.page_size);
129            let start = page.first_token_index(self.page_size);
130            indices.extend(start..start + take);
131        }
132        indices
133    }
134
135    pub fn free_pages(&mut self, pages: &[KvPageId]) {
136        self.free.extend_from_slice(pages);
137    }
138
139    /// Free every distinct page represented by a contiguous request-index list.
140    pub fn free_indices(&mut self, indices: &[usize]) -> Vec<KvPageId> {
141        let mut pages = Vec::with_capacity(indices.len().div_ceil(self.page_size));
142        for &index in indices {
143            let page = KvPageId::from_token_index(index, self.page_size);
144            if pages.last().copied() != Some(page) {
145                pages.push(page);
146            }
147        }
148        self.free_pages(&pages);
149        pages
150    }
151
152    #[cfg(test)]
153    pub fn free(&mut self, indices: &[usize]) {
154        self.free_indices(indices);
155    }
156
157    pub fn available_pages(&self) -> usize {
158        self.free.len() + self.total_pages - self.next_fresh
159    }
160
161    pub fn available(&self) -> usize {
162        self.available_pages() * self.page_size
163    }
164
165    pub fn total(&self) -> usize {
166        self.total_pages * self.page_size
167    }
168}
169
170/// A single node in the radix tree.
171pub struct TreeNode {
172    /// Children keyed by the first complete page on the child edge.
173    pub children: FxHashMap<LocalBlockHash, NodeId>,
174    pub parent: Option<NodeId>,
175    /// One content identity per complete page stored on this compressed edge.
176    ///
177    /// The mocker intentionally uses the router's 64-bit local block hash as
178    /// page identity so completed radix state does not retain token IDs.
179    /// Consequently, as in router-side indexing, hash collisions are treated
180    /// as identical pages rather than guarded by an exact-token comparison.
181    pub key: Vec<LocalBlockHash>,
182    /// One physical page ID per key. Length = `key.len()`.
183    pub value: Vec<KvPageId>,
184    /// Walk-to-root reference count (protected when > 0).
185    pub lock_ref: usize,
186    /// Monotonic timestamp for LRU eviction.
187    pub last_access_time: Instant,
188}
189
190/// Radix tree for SGLang KV cache simulation.
191pub struct RadixCache {
192    nodes: SlotMap<NodeId, TreeNode>,
193    root: NodeId,
194    pub page_pool: PagePool,
195    page_size: usize,
196    /// Total token count in evictable nodes.
197    pub evictable_leaves: FxHashSet<NodeId>,
198    pub evictable_size: usize,
199    /// Total token count in protected (locked) nodes.
200    pub protected_size: usize,
201}
202
203impl RadixCache {
204    pub fn new(total_tokens: usize, page_size: usize) -> Self {
205        assert!(page_size >= 1, "page_size must be >= 1");
206        let mut nodes = SlotMap::with_key();
207        let root = nodes.insert(TreeNode {
208            children: FxHashMap::default(),
209            parent: None,
210            key: Vec::new(),
211            value: Vec::new(),
212            lock_ref: 0,
213            last_access_time: Instant::now(),
214        });
215        Self {
216            nodes,
217            root,
218            page_pool: PagePool::new(total_tokens, page_size),
219            page_size,
220            evictable_leaves: FxHashSet::default(),
221            evictable_size: 0,
222            protected_size: 0,
223        }
224    }
225
226    pub fn root(&self) -> NodeId {
227        self.root
228    }
229    pub fn node(&self, id: NodeId) -> &TreeNode {
230        &self.nodes[id]
231    }
232    pub fn page_size(&self) -> usize {
233        self.page_size
234    }
235    pub fn num_nodes(&self) -> usize {
236        self.nodes.len()
237    }
238
239    fn page_hashes(&self, tokens: &[u32]) -> Vec<LocalBlockHash> {
240        compute_block_hash_for_seq(tokens, self.page_size as u32, BlockHashOptions::default())
241    }
242
243    fn page_ids(&self, indices: &[usize], page_count: usize) -> Vec<KvPageId> {
244        assert!(
245            indices.len() >= page_count * self.page_size,
246            "not enough token indices for {page_count} complete pages"
247        );
248        indices
249            .chunks_exact(self.page_size)
250            .take(page_count)
251            .map(|chunk| {
252                let page = KvPageId::from_token_index(chunk[0], self.page_size);
253                let start = page.first_token_index(self.page_size);
254                assert!(
255                    chunk.iter().copied().eq(start..start + self.page_size),
256                    "SGLang cached pages must contain contiguous page-aligned indices"
257                );
258                page
259            })
260            .collect()
261    }
262
263    fn key_match(key0: &[LocalBlockHash], key1: &[LocalBlockHash]) -> usize {
264        key0.iter().zip(key1).take_while(|(a, b)| a == b).count()
265    }
266
267    pub fn match_prefix(&mut self, key: &[u32]) -> (usize, NodeId) {
268        let page_keys = self.page_hashes(key);
269        let now = Instant::now();
270        self.nodes[self.root].last_access_time = now;
271
272        let mut current = self.root;
273        let mut matched_pages: usize = 0;
274
275        while matched_pages < page_keys.len() {
276            let child_id = match self.nodes[current]
277                .children
278                .get(&page_keys[matched_pages])
279                .copied()
280            {
281                Some(id) => id,
282                None => break,
283            };
284
285            let (common_len, child_len) = {
286                let child_key = &self.nodes[child_id].key;
287                (
288                    Self::key_match(child_key, &page_keys[matched_pages..]),
289                    child_key.len(),
290                )
291            };
292
293            if common_len < child_len {
294                if common_len > 0 {
295                    let intermediate = self.split_node(child_id, common_len);
296                    current = intermediate;
297                }
298                matched_pages += common_len;
299                break;
300            }
301
302            matched_pages += common_len;
303            current = child_id;
304            self.nodes[current].last_access_time = now;
305        }
306
307        (matched_pages * self.page_size, current)
308    }
309
310    /// Read-only prefix match length (does not mutate timestamps or split nodes).
311    /// Used for LPM scheduling scoring.
312    pub fn prefix_match_len(&self, key: &[u32]) -> usize {
313        let page_keys = self.page_hashes(key);
314        let mut current = self.root;
315        let mut matched_pages: usize = 0;
316
317        while matched_pages < page_keys.len() {
318            let child_id = match self.nodes[current]
319                .children
320                .get(&page_keys[matched_pages])
321                .copied()
322            {
323                Some(id) => id,
324                None => break,
325            };
326
327            let child_key = &self.nodes[child_id].key;
328            let common_len = Self::key_match(child_key, &page_keys[matched_pages..]);
329
330            if common_len < child_key.len() {
331                matched_pages += common_len;
332                break;
333            }
334
335            matched_pages += common_len;
336            current = child_id;
337        }
338
339        matched_pages * self.page_size
340    }
341
342    /// Insert a token sequence into the tree. Key is page-aligned before insertion.
343    pub fn insert(&mut self, key: &[u32], value: &[usize]) -> NodeId {
344        self.insert_from(self.root, 0, key, value, false)
345    }
346
347    /// Insert only the suffix after a retained, page-aligned prefix.
348    ///
349    /// `prefix_node` must be the locked terminal node for `prefix_len`. Keeping
350    /// that handle lets decode growth avoid walking the full sequence from the
351    /// root on every completed page.
352    pub fn insert_from_node(
353        &mut self,
354        prefix_node: NodeId,
355        prefix_len: usize,
356        key: &[u32],
357        value: &[usize],
358    ) -> NodeId {
359        self.insert_from(prefix_node, prefix_len, key, value, true)
360    }
361
362    fn insert_from(
363        &mut self,
364        start_node: NodeId,
365        prefix_len: usize,
366        key: &[u32],
367        value: &[usize],
368        allow_locked_tail_extension: bool,
369    ) -> NodeId {
370        let aligned_len = key.len() / self.page_size * self.page_size;
371        assert_eq!(
372            prefix_len % self.page_size,
373            0,
374            "prefix length must be page-aligned"
375        );
376        assert!(
377            prefix_len <= aligned_len,
378            "prefix length {prefix_len} exceeds aligned key length {aligned_len}"
379        );
380        if aligned_len == prefix_len {
381            return start_node;
382        }
383        assert!(
384            value.len() >= aligned_len,
385            "not enough token indices: need {aligned_len}, got {}",
386            value.len()
387        );
388        // `start_node` already represents the retained prefix. Hashing and
389        // validating it again on every decode-page completion makes growth
390        // quadratic in sequence length; only the newly completed suffix is
391        // needed for insertion.
392        let page_keys = self.page_hashes(&key[prefix_len..aligned_len]);
393        let page_ids = self.page_ids(&value[prefix_len..aligned_len], page_keys.len());
394
395        let now = Instant::now();
396        self.touch_path(start_node, now);
397
398        let mut current = start_node;
399        let mut key_offset = 0;
400
401        while key_offset < page_keys.len() {
402            let can_extend_leaf = current != self.root
403                && self.nodes[current].children.is_empty()
404                && (self.nodes[current].lock_ref == 0
405                    || (allow_locked_tail_extension
406                        && current == start_node
407                        && self.nodes[current].lock_ref == 1));
408            if can_extend_leaf {
409                return self.extend_leaf(
410                    current,
411                    &page_keys[key_offset..],
412                    &page_ids[key_offset..],
413                    now,
414                );
415            }
416
417            let child_id = match self.nodes[current]
418                .children
419                .get(&page_keys[key_offset])
420                .copied()
421            {
422                Some(id) => id,
423                None => {
424                    return self.create_child(
425                        current,
426                        &page_keys[key_offset..],
427                        &page_ids[key_offset..],
428                    );
429                }
430            };
431
432            let (common_len, child_len) = {
433                let child_key = &self.nodes[child_id].key;
434                (
435                    Self::key_match(child_key, &page_keys[key_offset..]),
436                    child_key.len(),
437                )
438            };
439
440            if common_len == child_len {
441                key_offset += common_len;
442                current = child_id;
443                self.nodes[current].last_access_time = now;
444            } else {
445                if common_len > 0 {
446                    let intermediate = self.split_node(child_id, common_len);
447                    key_offset += common_len;
448                    if key_offset < page_keys.len() {
449                        return self.create_child(
450                            intermediate,
451                            &page_keys[key_offset..],
452                            &page_ids[key_offset..],
453                        );
454                    }
455                    return intermediate;
456                }
457                return current;
458            }
459        }
460
461        current
462    }
463
464    fn touch_path(&mut self, node_id: NodeId, now: Instant) {
465        let mut current = Some(node_id);
466        while let Some(id) = current {
467            self.nodes[id].last_access_time = now;
468            current = self.nodes[id].parent;
469        }
470    }
471
472    fn extend_leaf(
473        &mut self,
474        node_id: NodeId,
475        key: &[LocalBlockHash],
476        value: &[KvPageId],
477        now: Instant,
478    ) -> NodeId {
479        let node = &mut self.nodes[node_id];
480        debug_assert!(node.children.is_empty());
481        debug_assert!(node.lock_ref <= 1);
482        node.key.extend_from_slice(key);
483        node.value.extend_from_slice(value);
484        node.last_access_time = now;
485        if node.lock_ref == 0 {
486            self.evictable_size += key.len() * self.page_size;
487        } else {
488            self.protected_size += key.len() * self.page_size;
489        }
490        node_id
491    }
492
493    fn split_node(&mut self, child_id: NodeId, split_pos: usize) -> NodeId {
494        let (child_parent, original_ck, prefix_key, prefix_value, suffix_ck, lock_ref, accessed) = {
495            let child = &mut self.nodes[child_id];
496            let child_parent = child.parent;
497            let original_ck = child.key[0];
498            let suffix_key = child.key.split_off(split_pos);
499            let prefix_key = std::mem::replace(&mut child.key, suffix_key);
500            let suffix_value = child.value.split_off(split_pos);
501            let prefix_value = std::mem::replace(&mut child.value, suffix_value);
502            let suffix_ck = child.key[0];
503            (
504                child_parent,
505                original_ck,
506                prefix_key,
507                prefix_value,
508                suffix_ck,
509                child.lock_ref,
510                child.last_access_time,
511            )
512        };
513
514        let mut inter_children = FxHashMap::default();
515        inter_children.insert(suffix_ck, child_id);
516
517        let intermediate = TreeNode {
518            children: inter_children,
519            parent: child_parent,
520            key: prefix_key,
521            value: prefix_value,
522            lock_ref,
523            last_access_time: accessed,
524        };
525        let inter_id = self.nodes.insert(intermediate);
526
527        let child = &mut self.nodes[child_id];
528        child.parent = Some(inter_id);
529
530        if let Some(parent_id) = child_parent {
531            self.nodes[parent_id].children.insert(original_ck, inter_id);
532        }
533
534        // Both size totals are unchanged: the intermediate and suffix split
535        // the original edge without changing its lock state or token count.
536
537        inter_id
538    }
539
540    fn create_child(
541        &mut self,
542        parent_id: NodeId,
543        key: &[LocalBlockHash],
544        value: &[KvPageId],
545    ) -> NodeId {
546        let new_node = TreeNode {
547            children: FxHashMap::default(),
548            parent: Some(parent_id),
549            key: key.to_vec(),
550            value: value.to_vec(),
551            lock_ref: 0,
552            last_access_time: Instant::now(),
553        };
554        let ck = key[0];
555        let new_id = self.nodes.insert(new_node);
556
557        self.evictable_leaves.remove(&parent_id);
558
559        self.nodes[parent_id].children.insert(ck, new_id);
560
561        self.evictable_leaves.insert(new_id);
562        self.evictable_size += key.len() * self.page_size;
563
564        new_id
565    }
566
567    pub fn is_leaf(&self, id: NodeId) -> bool {
568        self.nodes[id].children.is_empty()
569    }
570
571    pub fn inc_lock_ref(&mut self, node_id: NodeId) {
572        let mut current = Some(node_id);
573        while let Some(id) = current {
574            if id == self.root {
575                break;
576            }
577            let node = &mut self.nodes[id];
578            let tokens = node.key.len() * self.page_size;
579            node.lock_ref += 1;
580            if node.lock_ref == 1 {
581                self.evictable_leaves.remove(&id);
582                self.evictable_size -= tokens;
583                self.protected_size += tokens;
584            }
585            current = self.nodes[id].parent;
586        }
587    }
588
589    pub fn dec_lock_ref(&mut self, node_id: NodeId) {
590        let mut current = Some(node_id);
591        while let Some(id) = current {
592            if id == self.root {
593                break;
594            }
595            let node = &mut self.nodes[id];
596            if node.lock_ref == 0 {
597                tracing::warn!("dec_lock_ref on node with lock_ref == 0, skipping");
598                break;
599            }
600            node.lock_ref -= 1;
601            if node.lock_ref == 0 {
602                let tokens = node.key.len() * self.page_size;
603                self.protected_size -= tokens;
604                self.evictable_size += tokens;
605                if self.is_leaf(id) {
606                    self.evictable_leaves.insert(id);
607                }
608            }
609            current = self.nodes[id].parent;
610        }
611    }
612
613    /// Evict tokens from the cache by LRU order, rounding partial leaves to full pages.
614    /// Returns `(num_tokens_evicted, evicted_page_ids)`.
615    pub fn evict(&mut self, num_tokens: usize) -> (usize, Vec<KvPageId>) {
616        let mut evicted = 0;
617        let mut evicted_indices =
618            Vec::with_capacity(num_tokens.min(self.evictable_size).div_ceil(self.page_size));
619        while evicted < num_tokens {
620            let victim = self
621                .evictable_leaves
622                .iter()
623                .min_by_key(|&&id| self.nodes[id].last_access_time)
624                .copied();
625
626            let Some(victim_id) = victim else {
627                break;
628            };
629
630            let victim_pages = self.nodes[victim_id].key.len();
631            let victim_tokens = victim_pages * self.page_size;
632            let remaining = num_tokens - evicted;
633            let eviction_pages = remaining.div_ceil(self.page_size).min(victim_pages);
634            let eviction_len = eviction_pages * self.page_size;
635
636            // A compressed leaf may span pages. Preserve its indexed prefix when
637            // only the newest suffix pages are needed to satisfy this eviction.
638            if eviction_len < victim_tokens {
639                let split_pos = victim_pages - eviction_pages;
640                let (nodes, page_pool) = (&mut self.nodes, &mut self.page_pool);
641                let victim_node = &mut nodes[victim_id];
642                victim_node.key.truncate(split_pos);
643                let evicted_values = &victim_node.value[split_pos..];
644                page_pool.free_pages(evicted_values);
645                evicted_indices.extend_from_slice(evicted_values);
646                victim_node.value.truncate(split_pos);
647
648                self.evictable_size -= eviction_len;
649                evicted += eviction_len;
650                continue;
651            }
652
653            let victim_node = self
654                .nodes
655                .remove(victim_id)
656                .expect("evictable leaf disappeared before removal");
657            let tokens = victim_node.key.len() * self.page_size;
658            let parent_id = victim_node.parent;
659
660            self.evictable_leaves.remove(&victim_id);
661            self.evictable_size -= tokens;
662            evicted += tokens;
663
664            evicted_indices.extend_from_slice(&victim_node.value);
665            self.page_pool.free_pages(&victim_node.value);
666
667            if let Some(pid) = parent_id {
668                self.nodes[pid].children.remove(&victim_node.key[0]);
669
670                if pid != self.root
671                    && self.nodes[pid].children.is_empty()
672                    && self.nodes[pid].lock_ref == 0
673                {
674                    self.evictable_leaves.insert(pid);
675                }
676            }
677        }
678        (evicted, evicted_indices)
679    }
680
681    pub fn available_tokens(&self) -> usize {
682        self.page_pool.available()
683    }
684
685    pub fn total_tokens(&self) -> usize {
686        self.page_pool.total()
687    }
688}
689
690#[cfg(test)]
691mod tests {
692    use super::*;
693
694    #[test]
695    fn test_page_pool_allocate_extend_and_free() {
696        let mut pool = PagePool::new(12, 4);
697        assert_eq!(pool.available(), 12);
698        assert!(pool.allocate(usize::MAX).is_none());
699        let a = pool.allocate(3).unwrap();
700        assert_eq!(a.len(), 3);
701        assert_eq!(pool.available(), 8);
702        let mut extended = a.clone();
703        assert!(pool.allocate_indices_into(1, &mut extended));
704        assert_eq!(extended, vec![0, 1, 2, 3]);
705        assert_eq!(pool.available(), 8);
706        let b = pool.allocate(5).unwrap();
707        assert_eq!(pool.available(), 0);
708        assert!(pool.allocate(1).is_none());
709        pool.free(&a);
710        assert_eq!(pool.available(), 4);
711        pool.free(&b);
712        assert_eq!(pool.available(), 12);
713    }
714
715    #[test]
716    fn test_allocate_indices_into_failure_is_atomic() {
717        let mut pool = PagePool::new(8, 4);
718        let mut destination = pool.allocate(4).unwrap();
719        let _other = pool.allocate(4).unwrap();
720        let available_before = pool.available();
721        let destination_before = destination.clone();
722
723        assert!(!pool.allocate_indices_into(1, &mut destination));
724        assert_eq!(destination, destination_before);
725        assert_eq!(pool.available(), available_before);
726    }
727
728    #[test]
729    fn test_match_prefix() {
730        let mut cache = RadixCache::new(100, 1);
731
732        // Empty tree
733        let (len, node) = cache.match_prefix(&[1, 2, 3]);
734        assert_eq!(len, 0);
735        assert_eq!(node, cache.root());
736
737        // Full match
738        cache.insert(&[1, 2, 3, 4, 5], &[10, 20, 30, 40, 50]);
739        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5]).0, 5);
740
741        // Partial match with split
742        cache.insert(&[1, 2, 3, 4, 5, 6, 7], &[10, 20, 30, 40, 50, 60, 70]);
743        let (len, node) = cache.match_prefix(&[1, 2, 3, 4, 5, 9, 9]);
744        assert_eq!(len, 5);
745        let n = cache.node(node);
746        assert_eq!(n.key, cache.page_hashes(&[1, 2, 3, 4, 5]));
747        assert_eq!(
748            n.value,
749            vec![
750                KvPageId(10),
751                KvPageId(20),
752                KvPageId(30),
753                KvPageId(40),
754                KvPageId(50)
755            ]
756        );
757        let suffix_key = cache.page_hashes(&[6])[0];
758        let &suffix_id = n.children.get(&suffix_key).unwrap();
759        assert_eq!(
760            cache.node(suffix_id).value,
761            vec![KvPageId(60), KvPageId(70)]
762        );
763    }
764
765    #[test]
766    fn test_insert() {
767        let mut cache = RadixCache::new(100, 1);
768
769        // Shared prefix splits the tree
770        cache.insert(&[1, 2, 3, 4, 5], &[10, 20, 30, 40, 50]);
771        cache.insert(&[1, 2, 3, 6, 7], &[10, 20, 30, 60, 70]);
772        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5]).0, 5);
773        assert_eq!(cache.match_prefix(&[1, 2, 3, 6, 7]).0, 5);
774        assert_eq!(cache.match_prefix(&[1, 2, 3, 9]).0, 3);
775
776        // Extend existing prefix
777        let mut cache = RadixCache::new(100, 1);
778        cache.insert(&[1, 2, 3], &[10, 20, 30]);
779        cache.insert(&[1, 2, 3, 4, 5], &[10, 20, 30, 40, 50]);
780        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5]).0, 5);
781
782        // Duplicate insert is idempotent
783        cache.insert(&[1, 2, 3], &[10, 20, 30]);
784
785        // Match then insert suffix
786        let mut cache = RadixCache::new(100, 1);
787        cache.insert(&[1, 2, 3, 4, 5], &[10, 20, 30, 40, 50]);
788        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5, 6, 7, 8]).0, 5);
789        cache.insert(&[1, 2, 3, 4, 5, 6, 7, 8], &[10, 20, 30, 40, 50, 60, 70, 80]);
790        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5, 6, 7, 8]).0, 8);
791    }
792
793    #[test]
794    fn test_retained_tail_extends_unique_leaf_in_place() {
795        let mut cache = RadixCache::new(100, 4);
796        cache.insert(&[1, 2, 3, 4], &[0, 1, 2, 3]);
797        let (_, tail) = cache.match_prefix(&[1, 2, 3, 4]);
798        cache.inc_lock_ref(tail);
799        let nodes_before = cache.num_nodes();
800
801        let extended = cache.insert_from_node(
802            tail,
803            4,
804            &[1, 2, 3, 4, 5, 6, 7, 8],
805            &[0, 1, 2, 3, 4, 5, 6, 7],
806        );
807
808        assert_eq!(extended, tail);
809        assert_eq!(cache.num_nodes(), nodes_before);
810        assert_eq!(
811            cache.node(tail).key,
812            cache.page_hashes(&[1, 2, 3, 4, 5, 6, 7, 8])
813        );
814        assert_eq!(cache.protected_size, 8);
815        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5, 6, 7, 8]).0, 8);
816    }
817
818    #[test]
819    fn test_retained_tail_does_not_extend_shared_leaf_in_place() {
820        let mut cache = RadixCache::new(100, 4);
821        cache.insert(&[1, 2, 3, 4], &[0, 1, 2, 3]);
822        let (_, tail) = cache.match_prefix(&[1, 2, 3, 4]);
823        cache.inc_lock_ref(tail);
824        cache.inc_lock_ref(tail);
825        let nodes_before = cache.num_nodes();
826
827        let extended = cache.insert_from_node(
828            tail,
829            4,
830            &[1, 2, 3, 4, 5, 6, 7, 8],
831            &[0, 1, 2, 3, 4, 5, 6, 7],
832        );
833
834        assert_ne!(extended, tail);
835        assert_eq!(cache.num_nodes(), nodes_before + 1);
836        assert_eq!(cache.node(tail).key, cache.page_hashes(&[1, 2, 3, 4]));
837        assert_eq!(cache.node(extended).key, cache.page_hashes(&[5, 6, 7, 8]));
838    }
839
840    #[test]
841    fn test_page_size() {
842        // Insert and match with page_size=4
843        let mut cache = RadixCache::new(100, 4);
844        assert_eq!(cache.page_pool.total(), 100);
845        cache.insert(&[1, 2, 3, 4, 5, 6, 7], &[0, 1, 2, 3, 4, 5, 6]);
846        assert_eq!(cache.match_prefix(&[1, 2, 3, 4]).0, 4);
847        let (_, node) = cache.match_prefix(&[1, 2, 3, 4]);
848        assert_eq!(cache.node(node).value, vec![KvPageId(0)]);
849
850        cache.insert(&[1, 2, 3, 4, 5, 6, 7, 8], &[0, 1, 2, 3, 4, 5, 6, 7]);
851        assert_eq!(cache.match_prefix(&[1, 2, 3, 4, 5, 6, 7, 8]).0, 8);
852
853        // Children disambiguated by first page_size tokens
854        let mut cache = RadixCache::new(100, 4);
855        cache.insert(&[1, 2, 3, 4], &[0, 1, 2, 3]);
856        cache.insert(&[1, 2, 3, 5], &[4, 5, 6, 7]);
857        assert_eq!(cache.match_prefix(&[1, 2, 3, 4]).0, 4);
858        assert_eq!(cache.match_prefix(&[1, 2, 3, 5]).0, 4);
859        assert_eq!(cache.match_prefix(&[1, 2, 3, 6]).0, 0);
860
861        // Split at page boundary preserves value
862        let mut cache = RadixCache::new(100, 4);
863        cache.insert(&[1, 2, 3, 4, 5, 6, 7, 8], &[0, 1, 2, 3, 4, 5, 6, 7]);
864        cache.match_prefix(&[1, 2, 3, 4, 9, 9, 9, 9]);
865        let (_, node) = cache.match_prefix(&[1, 2, 3, 4]);
866        assert_eq!(cache.node(node).value, vec![KvPageId(0)]);
867    }
868
869    #[test]
870    fn completed_edges_store_one_key_and_page_id_per_page() {
871        let mut cache = RadixCache::new(256, 64);
872        let tokens = (0..128).collect::<Vec<u32>>();
873        let indices = cache.page_pool.allocate(tokens.len()).unwrap();
874        let node = cache.insert(&tokens, &indices);
875
876        assert_eq!(cache.node(node).key.len(), 2);
877        assert_eq!(cache.node(node).value.len(), 2);
878        assert_eq!(cache.node(node).value, vec![KvPageId(0), KvPageId(1)]);
879        assert_eq!(cache.match_prefix(&tokens).0, tokens.len());
880    }
881
882    #[test]
883    fn test_lock_unlock_shared_prefix() {
884        let mut cache = RadixCache::new(100, 1);
885        cache.insert(&[1, 2, 3, 4, 5], &[0, 1, 2, 3, 4]);
886        cache.insert(&[1, 2, 3, 6, 7], &[0, 1, 2, 5, 6]);
887
888        let (_, node_a) = cache.match_prefix(&[1, 2, 3, 4, 5]);
889        let (_, node_b) = cache.match_prefix(&[1, 2, 3, 6, 7]);
890
891        cache.inc_lock_ref(node_a);
892        cache.inc_lock_ref(node_b);
893        assert_eq!(cache.protected_size, 7); // 2+2+3
894
895        cache.dec_lock_ref(node_a);
896        assert!(cache.evictable_leaves.contains(&node_a));
897        cache.dec_lock_ref(node_b);
898        assert_eq!(cache.protected_size, 0);
899    }
900
901    #[test]
902    fn test_evict() {
903        // LRU order: oldest evicted first
904        let mut cache = RadixCache::new(100, 1);
905        cache.insert(&[1, 2, 3], &[0, 1, 2]);
906        let (_, n1) = cache.match_prefix(&[1, 2, 3]);
907        cache.inc_lock_ref(n1);
908        cache.dec_lock_ref(n1);
909
910        std::thread::sleep(std::time::Duration::from_millis(1));
911        cache.insert(&[4, 5, 6], &[3, 4, 5]);
912        let (_, n2) = cache.match_prefix(&[4, 5, 6]);
913        cache.inc_lock_ref(n2);
914        cache.dec_lock_ref(n2);
915
916        let (evicted_count, evicted_indices) = cache.evict(3);
917        assert_eq!(evicted_count, 3);
918        // Evicted indices should match the pool indices originally inserted for [1,2,3]
919        let mut sorted_evicted = evicted_indices.clone();
920        sorted_evicted.sort();
921        let mut expected_indices = vec![KvPageId(0), KvPageId(1), KvPageId(2)];
922        expected_indices.sort();
923        assert_eq!(
924            sorted_evicted, expected_indices,
925            "evicted indices should match inserted indices"
926        );
927        assert_eq!(cache.match_prefix(&[1, 2, 3]).0, 0); // oldest evicted
928        assert_eq!(cache.match_prefix(&[4, 5, 6]).0, 3); // newer kept
929
930        // Locked nodes are not evicted
931        let mut cache = RadixCache::new(100, 1);
932        cache.insert(&[1, 2, 3], &[0, 1, 2]);
933        cache.insert(&[4, 5, 6], &[3, 4, 5]);
934        let (_, locked) = cache.match_prefix(&[1, 2, 3]);
935        cache.inc_lock_ref(locked);
936        let (_, unlocked) = cache.match_prefix(&[4, 5, 6]);
937        cache.inc_lock_ref(unlocked);
938        cache.dec_lock_ref(unlocked);
939        let (evicted_count, evicted_indices) = cache.evict(6);
940        assert_eq!(evicted_count, 3); // only unlocked evicted
941        let mut sorted_evicted = evicted_indices;
942        sorted_evicted.sort();
943        assert_eq!(
944            sorted_evicted,
945            vec![KvPageId(3), KvPageId(4), KvPageId(5)],
946            "should evict unlocked [4,5,6] indices"
947        );
948        assert_eq!(cache.match_prefix(&[1, 2, 3]).0, 3);
949    }
950
951    #[test]
952    fn test_evictable_size_includes_unlocked_internal_prefix() {
953        let mut cache = RadixCache::new(16, 4);
954        let first = cache.page_pool.allocate(8).unwrap();
955        cache.insert(&[1; 8], &first);
956        let mut branch = first[..4].to_vec();
957        branch.extend(cache.page_pool.allocate(4).unwrap());
958        cache.insert(&[1, 1, 1, 1, 2, 2, 2, 2], &branch);
959
960        assert_eq!(cache.evictable_size, 12);
961        assert_eq!(cache.evict(12).0, 12);
962        assert_eq!(cache.available_tokens(), 16);
963    }
964
965    #[test]
966    fn test_query_methods() {
967        let cache = RadixCache::new(100, 1);
968        assert_eq!(cache.available_tokens(), 100);
969        assert_eq!(cache.total_tokens(), 100);
970
971        let cache4 = RadixCache::new(100, 4);
972        assert_eq!(cache4.available_tokens(), 100);
973        assert_eq!(cache4.total_tokens(), 100);
974    }
975}