Skip to main content

dynamo_tokenizers/cache/
l1.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3//
4// SPDX-FileCopyrightText: Copyright (c) 2024 Simo Lin, Chang Su, Keyang Ru (llm-tokenizer authors)
5//
6// Portions adapted from sgl-project/llm-tokenizer v1.3.2 (Apache-2.0).
7// Upstream: https://github.com/lightseekorg/smg
8// Modifications: removed `add_special_tokens` plumbing (Dynamo's Encoder has no such
9// flag), bound `insert_at_boundaries` on `Encoder` rather than `Tokenizer`, retargeted
10// imports onto `crate::traits`.
11
12//! L1 Cache: Special-token boundary prefix cache
13//!
14//! Caches tokenization results at ALL special token boundaries.
15//! Special tokens (like `<|im_start|>`, `<|im_end|>`) are atomic in BPE tokenizers
16//! (`special: true, normalized: false`), making them the ONLY safe split points that
17//! guarantee correctness: `tokenize(prefix) + tokenize(suffix) == tokenize(prefix + suffix)`.
18//!
19//! No fallback to whitespace/punctuation — better to not cache than risk corruption.
20//!
21//! Storage and eviction are delegated to a weighted [`moka`] `sync::Cache` (W-TinyLFU):
22//! entries are keyed by the blake3 digest of `input[0..boundary]` and weighed by their
23//! resident token-vector bytes, so the byte budget is enforced — and recency/frequency
24//! tracked — by moka rather than by hand.
25
26use std::{
27    hash::BuildHasherDefault,
28    mem::size_of_val,
29    sync::{
30        Arc,
31        atomic::{AtomicU64, Ordering},
32    },
33};
34
35use aho_corasick::AhoCorasick;
36use moka::sync::Cache;
37use rustc_hash::FxHasher;
38
39use crate::{TokenIdType, traits::Encoder};
40
41/// Hash type for cache keys
42type Blake3Hash = [u8; 32];
43
44/// Keys are blake3 digests (already uniformly distributed), so a fast non-DoS-resistant
45/// hasher suffices — no need for the default SipHash.
46type PrefixHasher = BuildHasherDefault<FxHasher>;
47
48/// Weighted W-TinyLFU cache mapping a prefix's blake3 digest to its cumulative tokens.
49type PrefixCache = Cache<Blake3Hash, Arc<[TokenIdType]>, PrefixHasher>;
50
51/// Request-local lookup result. The deepest digest can differ from the matched key.
52pub(super) struct PrefixMatch {
53    pub(super) tokens: Arc<[TokenIdType]>,
54    pub(super) prefix_len: usize,
55    deepest_boundary: usize,
56    deepest_hash: Option<Blake3Hash>,
57}
58
59/// A miss retains lookup's boundary hashes for population without another input scan.
60pub(super) enum PrefixLookup {
61    Hit(PrefixMatch),
62    Miss(Vec<(usize, Blake3Hash)>),
63}
64
65/// Hash sorted boundary prefixes incrementally.
66fn hash_prefixes<'a>(
67    input: &'a str,
68    boundaries: &'a [usize],
69) -> impl Iterator<Item = (usize, Blake3Hash)> + 'a {
70    let mut hasher = blake3::Hasher::new();
71    let mut last_pos = 0;
72    boundaries.iter().map(move |&boundary_pos| {
73        hasher.update(&input.as_bytes()[last_pos..boundary_pos]);
74        last_pos = boundary_pos;
75        (boundary_pos, *hasher.finalize().as_bytes())
76    })
77}
78
79/// Positions immediately after each special-token occurrence in `text`.
80///
81/// Callers supply token strings that the inner tokenizer treats as atomic, so a boundary
82/// immediately after a selected occurrence is a safe split point:
83/// `tokenize(prefix) + tokenize(suffix) == tokenize(prefix + suffix)`. The overlapping scan
84/// is safe only when registered special-token occurrences cannot overlap; construction
85/// screens out other sets with [`first_unsafe_overlap`]. A boundary at the end of the input
86/// is omitted because there is no suffix to encode.
87fn boundaries_with(text: &str, matcher: &AhoCorasick) -> Vec<usize> {
88    let mut boundaries: Vec<usize> = matcher
89        .find_overlapping_iter(text)
90        .map(|m| m.end())
91        .filter(|&end| end < text.len())
92        .collect();
93    boundaries.sort_unstable();
94    boundaries.dedup();
95    boundaries
96}
97
98fn has_nontrivial_self_overlap(token: &str) -> bool {
99    let bytes = token.as_bytes();
100    (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
101}
102
103fn tokens_can_overlap(a: &str, b: &str) -> bool {
104    if a.contains(b) || b.contains(a) {
105        return true;
106    }
107
108    let a = a.as_bytes();
109    let b = b.as_bytes();
110    let max_overlap = a.len().min(b.len());
111    (1..max_overlap).any(|overlap| {
112        a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
113    })
114}
115
116/// Returns the first pair of special tokens whose occurrences can overlap.
117///
118/// [`boundaries_with`] reports the end of *every* occurrence of *every* special token.
119/// That equals the tokenizer's own segmentation only when occurrences cannot overlap;
120/// otherwise a reported boundary can land strictly inside the span the tokenizer actually
121/// consumed, and splitting there breaks the module invariant
122/// `tokenize(prefix) + tokenize(suffix) == tokenize(prefix + suffix)`.
123pub(super) fn first_unsafe_overlap(special_tokens: &[String]) -> Option<(&str, &str)> {
124    for (index, token) in special_tokens.iter().enumerate() {
125        if token.is_empty() {
126            continue;
127        }
128        if has_nontrivial_self_overlap(token) {
129            return Some((token, token));
130        }
131        for other in &special_tokens[index + 1..] {
132            if !other.is_empty() && token != other && tokens_can_overlap(token, other) {
133                return Some((token, other));
134            }
135        }
136    }
137
138    None
139}
140
141/// Test-only reference: build a one-off automaton and find boundaries. Production goes
142/// through [`L1Cache::boundaries`], which reuses a process-once automaton.
143#[cfg(test)]
144fn find_special_token_boundaries(text: &str, special_tokens: &[&str]) -> Vec<usize> {
145    if special_tokens.is_empty() {
146        return Vec::new();
147    }
148    let matcher = AhoCorasick::new(special_tokens)
149        .expect("special tokens form a valid Aho-Corasick automaton");
150    boundaries_with(text, &matcher)
151}
152
153/// Optional per-event observer. `on_hit` runs after each cache hit, `on_miss`
154/// after each miss — wired by `CachedTokenizer::with_observer` to push events
155/// straight into Prometheus counters without a periodic sampling step.
156pub type CacheEventFn = Arc<dyn Fn() + Send + Sync>;
157
158/// L1 cache: prefix matching at special-token boundaries, backed by a weighted W-TinyLFU
159/// [`moka`] cache that owns storage, recency/frequency tracking, and eviction. Hit/miss
160/// counts (our notion of a *prefix* hit) are tracked separately for metrics.
161pub struct L1Cache {
162    /// Prefix entries keyed by the blake3 digest of `input[0..boundary]`.
163    cache: PrefixCache,
164    /// Aho-Corasick automaton over the special tokens, built once at construction (`None`
165    /// when there are no special tokens). Lets boundary detection be a single pass.
166    matcher: Option<AhoCorasick>,
167    hits: AtomicU64,
168    misses: AtomicU64,
169    on_hit: Option<CacheEventFn>,
170    on_miss: Option<CacheEventFn>,
171}
172
173impl L1Cache {
174    /// `special_tokens` is the atomic special-token set whose boundaries the cache splits
175    /// at; an empty set leaves L1 inert (no boundaries, no entries).
176    pub fn new(max_memory: usize, mut special_tokens: Vec<String>) -> Self {
177        special_tokens.retain(|token| !token.is_empty());
178
179        // Capacity is the byte budget; each entry weighs its resident token-vector bytes
180        // (the prefix text is hashed and discarded, never stored). moka's W-TinyLFU policy
181        // admits/evicts to keep the weighted size within budget.
182        let cache = Cache::builder()
183            .max_capacity(max_memory as u64)
184            .weigher(|_k: &Blake3Hash, tokens: &Arc<[TokenIdType]>| -> u32 {
185                size_of_val(tokens.as_ref()).min(u32::MAX as usize) as u32
186            })
187            .build_with_hasher(PrefixHasher::default());
188
189        // Build the boundary automaton once; `None` when there are no special tokens.
190        let matcher = (!special_tokens.is_empty()).then(|| {
191            AhoCorasick::new(&special_tokens)
192                .expect("special tokens form a valid Aho-Corasick automaton")
193        });
194
195        Self {
196            cache,
197            matcher,
198            hits: AtomicU64::new(0),
199            misses: AtomicU64::new(0),
200            on_hit: None,
201            on_miss: None,
202        }
203    }
204
205    /// Install hit/miss callbacks. Replaces any previously-set observers.
206    pub fn set_observer(&mut self, on_hit: CacheEventFn, on_miss: CacheEventFn) {
207        self.on_hit = Some(on_hit);
208        self.on_miss = Some(on_miss);
209    }
210
211    /// Special-token boundaries in `text` via the process-once Aho-Corasick automaton built
212    /// at construction — a single pass over the input rather than one `str::find` sweep per
213    /// token. Empty when the cache has no special tokens.
214    fn boundaries(&self, text: &str) -> Vec<usize> {
215        match &self.matcher {
216            Some(matcher) => boundaries_with(text, matcher),
217            None => Vec::new(),
218        }
219    }
220
221    /// Try to find the longest prefix match at a special-token boundary.
222    ///
223    /// Returns `(cached_tokens, byte_offset, deepest_boundary)` if found. The caller
224    /// extends the cached tokens with a fresh encode of `input[byte_offset..]`;
225    /// `deepest_boundary` is the deepest special-token boundary in `input` (end-exclusive),
226    /// handed back so [`extend_after_match`] need not rescan the input for it.
227    pub fn longest_prefix_match(&self, input: &str) -> Option<(Arc<[TokenIdType]>, usize, usize)> {
228        match self.lookup_prefix(input) {
229            PrefixLookup::Hit(matched) => {
230                Some((matched.tokens, matched.prefix_len, matched.deepest_boundary))
231            }
232            PrefixLookup::Miss(_) => None,
233        }
234    }
235
236    /// Look up the longest cached prefix, retaining hashes for extension or population.
237    /// The returned offsets and digests must be used with this same input.
238    pub(super) fn lookup_prefix(&self, input: &str) -> PrefixLookup {
239        let boundaries = self.boundaries(input);
240
241        if boundaries.is_empty() {
242            self.misses.fetch_add(1, Ordering::Relaxed);
243            if let Some(cb) = &self.on_miss {
244                cb();
245            }
246            return PrefixLookup::Miss(Vec::new());
247        }
248
249        let prefix_hashes: Vec<_> = hash_prefixes(input, &boundaries).collect();
250
251        for &(boundary_pos, hash_bytes) in prefix_hashes.iter().rev() {
252            if let Some(tokens) = self.cache.get(&hash_bytes) {
253                self.hits.fetch_add(1, Ordering::Relaxed);
254                if let Some(cb) = &self.on_hit {
255                    cb();
256                }
257                // Share cached tokens to avoid copying the prefix during lookup.
258                let &(deepest_boundary, deepest_hash) =
259                    prefix_hashes.last().expect("prefix hashes is non-empty");
260                return PrefixLookup::Hit(PrefixMatch {
261                    tokens,
262                    prefix_len: boundary_pos,
263                    deepest_boundary,
264                    deepest_hash: Some(deepest_hash),
265                });
266            }
267        }
268
269        self.misses.fetch_add(1, Ordering::Relaxed);
270        if let Some(cb) = &self.on_miss {
271            cb();
272        }
273        PrefixLookup::Miss(prefix_hashes)
274    }
275
276    /// Insert prefix entries at every special-token boundary (e.g. to pre-seed the cache).
277    ///
278    /// Uses incremental hashing and incremental tokenization (per-segment encode of the
279    /// delta text between adjacent boundaries) so populating N entries costs one full
280    /// re-tokenize total, split across the segments. The miss path uses
281    /// [`Self::populate_and_encode`] instead, which reuses this same work to *also* return
282    /// the full token vector (avoiding a redundant second tokenization).
283    pub fn insert_at_boundaries<E: Encoder + ?Sized>(
284        &self,
285        input: &str,
286        tokenizer: &E,
287    ) -> anyhow::Result<()> {
288        let boundaries = self.boundaries(input);
289        if boundaries.is_empty() {
290            return Ok(());
291        }
292        self.populate_boundaries(input, hash_prefixes(input, &boundaries), tokenizer)?;
293        Ok(())
294    }
295
296    /// Miss-path encode: tokenize `input` exactly once, caching the cumulative prefix at
297    /// every special-token boundary as we go, and return the full token-id vector. This
298    /// replaces a separate full `encode` + [`Self::insert_at_boundaries`], which together
299    /// tokenized the input ~twice (once for the result, once split across segments).
300    ///
301    /// The concatenation of the per-segment encodes equals an uncached `encode(input)`
302    /// because special tokens are atomic in BPE — the same invariant the hit path relies
303    /// on. Returns token-ids only; the caller wraps them in [`crate::Encoding::Sp`].
304    pub fn populate_and_encode<E: Encoder + ?Sized>(
305        &self,
306        input: &str,
307        tokenizer: &E,
308    ) -> anyhow::Result<Vec<TokenIdType>> {
309        let boundaries = self.boundaries(input);
310        self.populate_and_encode_with_hashes(input, hash_prefixes(input, &boundaries), tokenizer)
311    }
312
313    /// Populate a miss using sorted boundary hashes, then encode the tail.
314    /// All offsets and digests must describe this same input; public callers compute
315    /// them lazily, while CachedTokenizer supplies hashes retained from lookup.
316    pub(super) fn populate_and_encode_with_hashes<E: Encoder + ?Sized>(
317        &self,
318        input: &str,
319        prefix_hashes: impl Iterator<Item = (usize, Blake3Hash)>,
320        tokenizer: &E,
321    ) -> anyhow::Result<Vec<TokenIdType>> {
322        let (mut running, tail_start) =
323            self.populate_boundaries(input, prefix_hashes, tokenizer)?;
324        if tail_start == 0 {
325            return Ok(tokenizer.encode(input)?.token_ids().to_vec());
326        }
327
328        let tail = tokenizer.encode(&input[tail_start..])?;
329        running.extend_from_slice(tail.token_ids());
330        Ok(running)
331    }
332
333    /// Tokenize each segment and cache its cumulative prefix with the supplied digest.
334    /// Return the running tokens and last boundary. Earlier entries survive later encode failures.
335    fn populate_boundaries<E: Encoder + ?Sized>(
336        &self,
337        input: &str,
338        prefix_hashes: impl Iterator<Item = (usize, Blake3Hash)>,
339        tokenizer: &E,
340    ) -> anyhow::Result<(Vec<TokenIdType>, usize)> {
341        #[cfg(debug_assertions)]
342        let mut validation_hasher = blake3::Hasher::new();
343        let mut running_tokens: Vec<TokenIdType> = Vec::new();
344        let mut last_pos = 0;
345
346        for (boundary_pos, hash_bytes) in prefix_hashes {
347            #[cfg(debug_assertions)]
348            {
349                validation_hasher.update(&input.as_bytes()[last_pos..boundary_pos]);
350                debug_assert_eq!(hash_bytes, *validation_hasher.finalize().as_bytes());
351            }
352
353            // Incremental tokenization. Dynamo's Encoder has no `add_special_tokens`
354            // parameter — equivalent to upstream always passing `false` past the first
355            // segment (which is also what Dynamo's HF impl always does for the first).
356            let seg = tokenizer.encode(&input[last_pos..boundary_pos])?;
357            running_tokens.extend_from_slice(seg.token_ids());
358
359            let prefix_tokens: Arc<[TokenIdType]> = running_tokens.as_slice().into();
360            self.cache.insert(hash_bytes, prefix_tokens);
361
362            last_pos = boundary_pos;
363        }
364
365        Ok((running_tokens, last_pos))
366    }
367
368    /// Extend the cache on a *partial* hit so the next turn of a growing conversation
369    /// hits deeper. Given the `(prefix_tokens, prefix_len, deepest_boundary)` returned by
370    /// [`longest_prefix_match`], tokenize the remaining suffix and cache the cumulative
371    /// prefix at the suffix's **deepest** special-token boundary, then return the full
372    /// merged token vector.
373    ///
374    /// Deepest-only is intentional: in an append-only conversation the next turn always
375    /// reaches the deepest boundary, so caching it bounds per-turn work to the newest
376    /// exchange; shallow/branching coverage already comes from the miss path's
377    /// [`insert_at_boundaries`]. Splitting at special-token boundaries is correctness-safe
378    /// because special tokens are atomic in BPE
379    /// (`tokenize(a) + tokenize(b) == tokenize(a + b)`).
380    ///
381    /// Note: unlike the read-only fast path, this **writes** to the cache on a hit
382    /// (one insert + possible eviction). It relies on the same best-effort memory
383    /// accounting as [`insert_at_boundaries`].
384    pub fn extend_after_match<E: Encoder + ?Sized>(
385        &self,
386        input: &str,
387        prefix_tokens: Arc<[TokenIdType]>,
388        prefix_len: usize,
389        deepest_boundary: usize,
390        tokenizer: &E,
391    ) -> anyhow::Result<Vec<TokenIdType>> {
392        self.extend_after_match_with_hash(
393            input,
394            PrefixMatch {
395                tokens: prefix_tokens,
396                prefix_len,
397                deepest_boundary,
398                deepest_hash: None,
399            },
400            tokenizer,
401        )
402    }
403
404    /// Extend a partial hit, reusing lookup's deepest digest for the new entry.
405    /// `matched` must describe this same input. The public compatibility wrapper alone
406    /// omits the digest and computes it here when an insertion is needed.
407    pub(super) fn extend_after_match_with_hash<E: Encoder + ?Sized>(
408        &self,
409        input: &str,
410        matched: PrefixMatch,
411        tokenizer: &E,
412    ) -> anyhow::Result<Vec<TokenIdType>> {
413        let PrefixMatch {
414            tokens: prefix_tokens,
415            prefix_len,
416            deepest_boundary,
417            deepest_hash,
418        } = matched;
419        // Boundaries exclude input.len(), so the trailing segment is nonempty.
420        let deepest = (deepest_boundary > prefix_len).then_some(deepest_boundary);
421
422        let Some(deepest) = deepest else {
423            let suffix_enc = tokenizer.encode(&input[prefix_len..])?;
424            // Reserve once to avoid copying the cached prefix during vector growth.
425            let mut merged = Vec::with_capacity(prefix_tokens.len() + suffix_enc.token_ids().len());
426            merged.extend_from_slice(&prefix_tokens);
427            merged.extend_from_slice(suffix_enc.token_ids());
428            return Ok(merged);
429        };
430
431        // Encode both segments first to reserve capacity without recopying the prefix.
432        let seg_a = tokenizer.encode(&input[prefix_len..deepest])?;
433        let seg_b = tokenizer.encode(&input[deepest..])?;
434        let mut cumulative = Vec::with_capacity(
435            prefix_tokens.len() + seg_a.token_ids().len() + seg_b.token_ids().len(),
436        );
437        cumulative.extend_from_slice(&prefix_tokens);
438        cumulative.extend_from_slice(seg_a.token_ids());
439
440        let hash_bytes = deepest_hash.unwrap_or_else(|| {
441            let mut hasher = blake3::Hasher::new();
442            hasher.update(&input.as_bytes()[..deepest]);
443            *hasher.finalize().as_bytes()
444        });
445        debug_assert_eq!(
446            hash_bytes,
447            *blake3::hash(&input.as_bytes()[..deepest]).as_bytes()
448        );
449
450        // Copy only the populated prefix, excluding capacity reserved for the tail.
451        let tokens: Arc<[TokenIdType]> = cumulative.as_slice().into();
452        self.cache.insert(hash_bytes, tokens);
453
454        cumulative.extend_from_slice(seg_b.token_ids());
455        Ok(cumulative)
456    }
457
458    /// Number of live entries. Flushes moka's deferred maintenance first so the count is
459    /// exact rather than lagging behind pending inserts/evictions.
460    pub fn len(&self) -> usize {
461        self.cache.run_pending_tasks();
462        self.cache.entry_count() as usize
463    }
464
465    pub fn is_empty(&self) -> bool {
466        self.len() == 0
467    }
468
469    pub fn stats(&self) -> L1CacheStats {
470        // Flush moka's deferred maintenance so entry_count / weighted_size are accurate.
471        self.cache.run_pending_tasks();
472        let hits = self.hits.load(Ordering::Relaxed);
473        let misses = self.misses.load(Ordering::Relaxed);
474        let total_requests = hits + misses;
475
476        L1CacheStats {
477            hits,
478            misses,
479            entries: self.cache.entry_count() as usize,
480            memory_bytes: self.cache.weighted_size() as usize,
481            hit_rate: if total_requests > 0 {
482                hits as f64 / total_requests as f64
483            } else {
484                0.0
485            },
486        }
487    }
488
489    pub fn clear(&self) {
490        self.cache.invalidate_all();
491        self.cache.run_pending_tasks();
492        self.hits.store(0, Ordering::Relaxed);
493        self.misses.store(0, Ordering::Relaxed);
494    }
495}
496
497#[derive(Debug, Clone)]
498pub struct L1CacheStats {
499    pub hits: u64,
500    pub misses: u64,
501    pub entries: usize,
502    pub memory_bytes: usize,
503    pub hit_rate: f64,
504}
505
506#[cfg(test)]
507mod tests {
508    use std::sync::Arc;
509
510    use super::*;
511    use crate::{HuggingFaceTokenizer, traits::Tokenizer};
512
513    // TinyLlama: real Llama BPE with `<s>` and `</s>` as added tokens with
514    // `special: true, normalized: false` — atomic in BPE, safe boundary points.
515    const TINYLLAMA_PATH: &str = concat!(
516        env!("CARGO_MANIFEST_DIR"),
517        "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
518    );
519
520    const SPECIALS: &[&str] = &["<s>", "</s>"];
521
522    fn load_tokenizer() -> Arc<dyn Tokenizer> {
523        Arc::new(HuggingFaceTokenizer::from_file(TINYLLAMA_PATH).expect("load TinyLlama"))
524    }
525
526    /// An `L1Cache` over the TinyLlama [`SPECIALS`] with the given byte budget.
527    fn test_cache(max_memory: usize) -> L1Cache {
528        L1Cache::new(
529            max_memory,
530            SPECIALS.iter().map(|s| (*s).to_string()).collect(),
531        )
532    }
533
534    #[test]
535    fn boundaries_are_after_each_special_token_occurrence() {
536        let input = "<s>system\nHi</s><s>user\nHello</s>";
537        let bounds = find_special_token_boundaries(input, SPECIALS);
538        // Drop the trailing boundary (==text.len()), so 3 not 4 boundaries.
539        assert_eq!(bounds.len(), 3);
540        for w in bounds.windows(2) {
541            assert!(w[0] < w[1], "boundaries must be strictly increasing");
542        }
543        assert!(bounds.iter().all(|&b| b < input.len()));
544    }
545
546    #[test]
547    fn no_special_tokens_yields_no_boundaries() {
548        assert!(find_special_token_boundaries("plain text", &[]).is_empty());
549    }
550
551    #[test]
552    fn unsafe_overlap_detects_containment_crossing_and_self_overlap() {
553        let cases = [
554            (vec!["〈|", "〈|EOS|〉"], Some(("〈|", "〈|EOS|〉"))),
555            (vec!["ab", "bc"], Some(("ab", "bc"))),
556            (vec!["|◊|"], Some(("|◊|", "|◊|"))),
557            (vec!["<s>", "<s>"], None),
558        ];
559
560        for (tokens, expected) in cases {
561            let tokens: Vec<String> = tokens.into_iter().map(String::from).collect();
562            assert_eq!(first_unsafe_overlap(&tokens), expected);
563        }
564    }
565
566    #[test]
567    fn llama_numbered_special_tokens_do_not_trigger_overlap_guard() {
568        let mut llama: Vec<String> = [
569            "<|begin_of_text|>",
570            "<|end_of_text|>",
571            "<|start_header_id|>",
572            "<|end_header_id|>",
573            "<|eot_id|>",
574        ]
575        .into_iter()
576        .map(String::from)
577        .collect();
578        llama.extend((0..251).map(|id| format!("<|reserved_special_token_{id}|>")));
579
580        assert_eq!(first_unsafe_overlap(&llama), None);
581    }
582
583    #[test]
584    fn insert_then_lookup_finds_shared_prefix() {
585        let cache = test_cache(1024 * 1024);
586        let tokenizer = load_tokenizer();
587
588        let warm = "<s>system\nYou are helpful.</s><s>user\nHi</s>";
589        cache
590            .insert_at_boundaries(warm, tokenizer.as_ref())
591            .unwrap();
592        assert!(!cache.is_empty());
593
594        let target = "<s>system\nYou are helpful.</s><s>user\nDifferent question</s>";
595        let (tokens, offset, _deepest) = cache
596            .longest_prefix_match(target)
597            .expect("shared prefix should match");
598        assert!(offset > 0);
599        assert!(!tokens.is_empty());
600    }
601
602    #[test]
603    fn miss_increments_misses_counter() {
604        let cache = test_cache(1024 * 1024);
605        assert!(
606            cache
607                .longest_prefix_match("plain text no specials")
608                .is_none()
609        );
610        assert_eq!(cache.stats().misses, 1);
611    }
612
613    #[test]
614    fn hit_increments_hits_counter() {
615        let cache = test_cache(1024 * 1024);
616        let tokenizer = load_tokenizer();
617        let warm = "<s>system\nA.</s><s>user\nB</s>";
618        cache
619            .insert_at_boundaries(warm, tokenizer.as_ref())
620            .unwrap();
621        let _ = cache.longest_prefix_match(warm);
622        assert!(cache.stats().hits >= 1);
623    }
624
625    #[test]
626    fn merge_invariant_holds_against_uncached_encode() {
627        // Load-bearing correctness check: cached prefix + fresh suffix encode must
628        // equal plain encode of the full input. Relies on `<s>`/`</s>` being atomic
629        // in TinyLlama's BPE (they are).
630        let cache = test_cache(1024 * 1024);
631        let tokenizer = load_tokenizer();
632
633        let template = "<s>system\nYou are helpful.</s><s>user\n";
634        let warm = format!("{template}First.</s>");
635        cache
636            .insert_at_boundaries(&warm, tokenizer.as_ref())
637            .unwrap();
638
639        let target = format!("{template}A completely different second question.</s>");
640        let (prefix_tokens, prefix_len, _deepest) = cache
641            .longest_prefix_match(&target)
642            .expect("should find prefix");
643
644        let suffix = &target[prefix_len..];
645        let suffix_enc = tokenizer.encode(suffix).unwrap();
646        // longest_prefix_match returns the shared `Arc<[u32]>`; copy into a Vec to append the suffix.
647        let mut merged = prefix_tokens.to_vec();
648        merged.extend_from_slice(suffix_enc.token_ids());
649
650        let plain = tokenizer.encode(&target).unwrap();
651        assert_eq!(
652            merged,
653            plain.token_ids(),
654            "merged tokens must equal plain encode"
655        );
656    }
657
658    #[test]
659    fn eviction_respects_memory_budget() {
660        // 4 KB budget — tight enough to force eviction after a few inserts.
661        let cache = test_cache(4 * 1024);
662        let tokenizer = load_tokenizer();
663        for i in 0..50 {
664            let input =
665                format!("<s>system\nPersona {i} chatty.</s><s>user\nTurn {i} content here.</s>");
666            cache
667                .insert_at_boundaries(&input, tokenizer.as_ref())
668                .unwrap();
669        }
670        let stats = cache.stats();
671        assert!(
672            stats.memory_bytes <= 4 * 1024,
673            "memory_bytes={} exceeds budget",
674            stats.memory_bytes
675        );
676    }
677
678    #[test]
679    fn concurrent_inserts_and_lookups_do_not_corrupt() {
680        use std::thread;
681
682        let cache = Arc::new(test_cache(1024 * 1024));
683        let tokenizer = load_tokenizer();
684
685        let mut handles = vec![];
686        for i in 0..10 {
687            let cache_c = cache.clone();
688            let tok = tokenizer.clone();
689            handles.push(thread::spawn(move || {
690                let input = format!("<s>system\nThread {i}.</s><s>user\nThread {i} body.</s>");
691                cache_c.insert_at_boundaries(&input, tok.as_ref()).unwrap();
692                let r = cache_c.longest_prefix_match(&input);
693                assert!(r.is_some(), "thread {i} expected match after insert");
694            }));
695        }
696        for h in handles {
697            h.join().unwrap();
698        }
699        assert!(cache.stats().memory_bytes > 0);
700        assert!(cache.stats().hits >= 10);
701    }
702
703    /// Build an append-only multi-turn conversation. `turns[i]` is the full prompt at
704    /// turn `i`: the system prompt, `i + 1` completed user/assistant exchanges, and a
705    /// diverging open user turn (no trailing special, so the deepest boundary is the
706    /// `<s>` that opens it). Each `turns[i]` shares a strictly longer `</s>`-bounded
707    /// prefix with `turns[i + 1]`.
708    fn growing_chat_turns(n: usize) -> Vec<String> {
709        let mut convo = String::from("<s>system\nYou are a helpful assistant.</s>");
710        let mut turns = Vec::with_capacity(n);
711        for i in 0..n {
712            convo.push_str(&format!(
713                "<s>user\nQuestion {i} please answer it.</s><s>assistant\nDetailed answer {i} follows here.</s>"
714            ));
715            turns.push(format!("{convo}<s>user\nFollow-up {i}"));
716        }
717        turns
718    }
719
720    #[test]
721    fn extend_on_hit_advances_match_depth_each_turn() {
722        // The load-bearing behavioral proof. Without extension the match offset is
723        // pinned at turn-1 depth (hits never insert); with extension it advances every
724        // turn, so the suffix re-tokenized per turn shrinks instead of growing.
725        let tok = load_tokenizer();
726        let turns = growing_chat_turns(5);
727
728        // EXTEND OFF: seed turn 0 via the miss path, then only look up (never insert).
729        let off = test_cache(8 * 1024 * 1024);
730        off.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
731        let pinned = off.longest_prefix_match(&turns[1]).expect("hit").1;
732        for t in &turns[1..] {
733            let (_toks, offset, _deepest) = off.longest_prefix_match(t).expect("hit");
734            assert_eq!(
735                offset, pinned,
736                "extend-off offset must stay pinned at turn-1 depth"
737            );
738        }
739
740        // EXTEND ON: each hit caches the deepest boundary, so the next turn hits deeper.
741        let on = test_cache(8 * 1024 * 1024);
742        on.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
743        let mut prev = 0usize;
744        for (i, t) in turns.iter().enumerate().skip(1) {
745            let (prefix_tokens, offset, deepest) = on.longest_prefix_match(t).expect("hit");
746            assert!(
747                offset > prev,
748                "turn {i}: extend-on offset {offset} must exceed previous {prev}"
749            );
750            prev = offset;
751
752            // Extending must also preserve byte-exact correctness vs an uncached encode.
753            let merged = on
754                .extend_after_match(t, prefix_tokens, offset, deepest, tok.as_ref())
755                .unwrap();
756            let plain = tok.encode(t).unwrap();
757            assert_eq!(
758                merged,
759                plain.token_ids(),
760                "turn {i}: extend merge must equal plain encode"
761            );
762        }
763
764        assert!(
765            prev > pinned,
766            "extend-on frontier ({prev}) must reach deeper than pinned extend-off depth ({pinned})"
767        );
768    }
769
770    #[test]
771    fn extend_on_hit_respects_budget_and_stays_correct() {
772        // Tiny budget forces eviction (and over-budget skips) while extending; every
773        // turn's encode must stay correct and memory must stay within budget.
774        let tok = load_tokenizer();
775        let cache = test_cache(4 * 1024);
776        let turns = growing_chat_turns(20);
777        cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
778
779        for t in &turns[1..] {
780            let merged = match cache.longest_prefix_match(t) {
781                Some((prefix_tokens, offset, deepest)) => cache
782                    .extend_after_match(t, prefix_tokens, offset, deepest, tok.as_ref())
783                    .unwrap(),
784                None => {
785                    // Full miss under eviction pressure — mirror the miss path.
786                    let enc = tok.encode(t).unwrap();
787                    cache.insert_at_boundaries(t, tok.as_ref()).unwrap();
788                    enc.token_ids().to_vec()
789                }
790            };
791            let plain = tok.encode(t).unwrap();
792            assert_eq!(
793                merged,
794                plain.token_ids(),
795                "encode must stay correct under eviction pressure"
796            );
797            assert!(
798                cache.stats().memory_bytes <= 4 * 1024,
799                "memory_bytes={} exceeds budget",
800                cache.stats().memory_bytes
801            );
802        }
803    }
804
805    #[test]
806    fn concurrent_extend_on_hit_does_not_corrupt() {
807        use std::thread;
808
809        let tok = load_tokenizer();
810        let cache = Arc::new(test_cache(8 * 1024 * 1024));
811        let turns = growing_chat_turns(8);
812        // Seed turn 0 so every thread gets at least a partial hit.
813        cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
814
815        let mut handles = vec![];
816        for _ in 0..8 {
817            let cache_c = cache.clone();
818            let tok_c = tok.clone();
819            let turns_c = turns.clone();
820            handles.push(thread::spawn(move || {
821                for t in &turns_c[1..] {
822                    if let PrefixLookup::Hit(matched) = cache_c.lookup_prefix(t) {
823                        let merged = cache_c
824                            .extend_after_match_with_hash(t, matched, tok_c.as_ref())
825                            .unwrap();
826                        let plain = tok_c.encode(t).unwrap();
827                        assert_eq!(
828                            merged,
829                            plain.token_ids(),
830                            "concurrent extend must stay correct"
831                        );
832                    }
833                }
834            }));
835        }
836        for h in handles {
837            h.join().unwrap();
838        }
839        assert!(cache.stats().memory_bytes > 0);
840    }
841
842    #[test]
843    fn extend_after_match_persists_correct_deepest_entry() {
844        let tok = load_tokenizer();
845        for unicode in [false, true] {
846            let turns: Vec<_> = growing_chat_turns(3)
847                .into_iter()
848                .map(|t| {
849                    if unicode {
850                        t.replace("system", "system 世界 🦀")
851                    } else {
852                        t
853                    }
854                })
855                .collect();
856
857            let cache = test_cache(8 * 1024 * 1024);
858            cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
859
860            let PrefixLookup::Hit(matched) = cache.lookup_prefix(&turns[1]) else {
861                panic!("partial hit on turns[1]");
862            };
863            let prefix_len = matched.prefix_len;
864            let deepest_boundary = matched.deepest_boundary;
865            assert!(deepest_boundary > prefix_len);
866            assert_eq!(
867                matched.deepest_hash,
868                Some(*blake3::hash(&turns[1].as_bytes()[..deepest_boundary]).as_bytes())
869            );
870            assert_ne!(
871                matched.deepest_hash,
872                Some(*blake3::hash(&turns[1].as_bytes()[..prefix_len]).as_bytes())
873            );
874            let entries_before = cache.stats().entries;
875
876            let _merged = cache
877                .extend_after_match_with_hash(&turns[1], matched, tok.as_ref())
878                .unwrap();
879
880            assert_eq!(
881                cache.stats().entries,
882                entries_before + 1,
883                "extend must persist exactly one (deepest) entry"
884            );
885
886            let deepest = find_special_token_boundaries(&turns[1], SPECIALS)
887                .into_iter()
888                .rev()
889                .find(|&b| b > prefix_len)
890                .expect("a deeper boundary must exist in the appended turn");
891            assert_eq!(
892                deepest_boundary, deepest,
893                "longest_prefix_match must return the deepest boundary used by extend"
894            );
895
896            let (saved_tokens, saved_offset, _deepest) = cache
897                .longest_prefix_match(&turns[1])
898                .expect("hit after extend");
899            assert_eq!(
900                saved_offset, deepest,
901                "lookup must now hit at the just-saved deepest boundary"
902            );
903            let expected = tok.encode(&turns[1][..deepest]).unwrap();
904            assert_eq!(
905                &*saved_tokens,
906                expected.token_ids(),
907                "persisted entry tokens must equal the uncached encode of the cached prefix"
908            );
909        }
910    }
911
912    #[test]
913    fn extend_without_deeper_boundary_does_not_insert() {
914        let tok = load_tokenizer();
915        let cache = test_cache(8 * 1024 * 1024);
916        cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
917        for input in ["<s>世界", "<s>世界</s>"] {
918            let PrefixLookup::Hit(matched) = cache.lookup_prefix(input) else {
919                panic!("expected hit");
920            };
921            assert_eq!(matched.prefix_len, matched.deepest_boundary);
922            let entries = cache.len();
923            let merged = cache
924                .extend_after_match_with_hash(input, matched, tok.as_ref())
925                .unwrap();
926            assert_eq!(merged, tok.encode(input).unwrap().token_ids());
927            assert_eq!(cache.len(), entries);
928        }
929    }
930
931    struct FailAt {
932        call: std::sync::atomic::AtomicUsize,
933        fail_at: usize,
934    }
935    impl Encoder for FailAt {
936        fn encode(&self, _: &str) -> crate::Result<crate::Encoding> {
937            if self.call.fetch_add(1, Ordering::Relaxed) == self.fail_at {
938                anyhow::bail!("suffix failed");
939            }
940            Ok(crate::Encoding::Sp(vec![1]))
941        }
942        fn encode_batch(&self, inputs: &[&str]) -> crate::Result<Vec<crate::Encoding>> {
943            inputs.iter().map(|s| self.encode(s)).collect()
944        }
945    }
946
947    #[test]
948    fn hash_reuse_does_not_insert_when_either_suffix_encode_fails() {
949        let tok = load_tokenizer();
950        let cache = test_cache(8 * 1024 * 1024);
951        cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
952        let input = "<s>世界</s><s>tail";
953        for fail_at in [0, 1] {
954            let PrefixLookup::Hit(matched) = cache.lookup_prefix(input) else {
955                panic!("expected hit");
956            };
957            let entries = cache.len();
958            let error = cache
959                .extend_after_match_with_hash(
960                    input,
961                    matched,
962                    &FailAt {
963                        call: 0.into(),
964                        fail_at,
965                    },
966                )
967                .unwrap_err();
968            assert_eq!(error.to_string(), "suffix failed");
969            assert_eq!(cache.len(), entries);
970        }
971    }
972
973    #[test]
974    #[cfg(debug_assertions)]
975    fn reused_hashes_reject_mismatched_input_before_insertion() {
976        use std::panic::{AssertUnwindSafe, catch_unwind};
977
978        let tok = load_tokenizer();
979        let cache = test_cache(8 * 1024 * 1024);
980        cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
981        let PrefixLookup::Hit(matched) = cache.lookup_prefix("<s>世界</s><s>tail") else {
982            panic!("expected hit");
983        };
984        let entries = cache.len();
985        assert!(
986            catch_unwind(AssertUnwindSafe(|| {
987                cache.extend_after_match_with_hash("<s>日本</s><s>tail", matched, tok.as_ref())
988            }))
989            .is_err()
990        );
991        assert_eq!(cache.len(), entries);
992
993        let cache = test_cache(8 * 1024 * 1024);
994        let PrefixLookup::Miss(hashes) = cache.lookup_prefix("a<s>世界</s>tail") else {
995            panic!("expected miss");
996        };
997        assert!(
998            catch_unwind(AssertUnwindSafe(|| {
999                cache.populate_and_encode_with_hashes(
1000                    "b<s>世界</s>tail",
1001                    hashes.into_iter(),
1002                    tok.as_ref(),
1003                )
1004            }))
1005            .is_err()
1006        );
1007        assert!(cache.is_empty());
1008    }
1009
1010    #[test]
1011    fn boundaries_detected_for_multibyte_deepseek_tool_tokens() {
1012        // `find_special_token_boundaries` keys off byte offsets; DeepSeek's tool tokens use
1013        // multibyte code points (| = U+FF5C, ▁ = U+2581, 3 bytes each). A boundary must
1014        // land immediately after each occurrence at a valid char boundary, so the cache can
1015        // split a tool-call block at its special tokens without panicking on a slice.
1016        let specials = &["<|tool▁calls▁begin|>", "<|tool▁call▁end|>"];
1017        let text = "<|tool▁calls▁begin|>payload<|tool▁call▁end|>tail";
1018        let bounds = find_special_token_boundaries(text, specials);
1019
1020        let after_begin = "<|tool▁calls▁begin|>".len();
1021        let after_end = text.find("<|tool▁call▁end|>").unwrap() + "<|tool▁call▁end|>".len();
1022        assert_eq!(bounds, vec![after_begin, after_end]);
1023        for &b in &bounds {
1024            assert!(
1025                text.is_char_boundary(b),
1026                "boundary {b} is not a char boundary"
1027            );
1028            let _ = &text[..b]; // must not panic
1029        }
1030    }
1031
1032    fn populate_miss<E: Encoder + ?Sized>(
1033        cache: &L1Cache,
1034        input: &str,
1035        tokenizer: &E,
1036        reuse_hashes: bool,
1037    ) -> anyhow::Result<Vec<TokenIdType>> {
1038        if reuse_hashes {
1039            let PrefixLookup::Miss(hashes) = cache.lookup_prefix(input) else {
1040                panic!("expected miss");
1041            };
1042            cache.populate_and_encode_with_hashes(input, hashes.into_iter(), tokenizer)
1043        } else {
1044            cache.populate_and_encode(input, tokenizer)
1045        }
1046    }
1047
1048    #[test]
1049    fn populate_and_encode_matches_uncached_and_seeds_cache() {
1050        let tok = load_tokenizer();
1051        for input in [
1052            "<s>system\nYou are helpful.</s><s>user\nHello there, friend.</s>",
1053            "<s>system\n世界 🦀</s><s>user\nこんにちは</s>tail",
1054        ] {
1055            let plain = tok.encode(input).unwrap();
1056            let boundaries = find_special_token_boundaries(input, SPECIALS);
1057            for reuse_hashes in [false, true] {
1058                let cache = test_cache(8 * 1024 * 1024);
1059                let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1060                assert_eq!(
1061                    got,
1062                    plain.token_ids(),
1063                    "fused miss encode must equal uncached encode"
1064                );
1065
1066                let mut expected_bytes = 0;
1067                for &boundary in &boundaries {
1068                    let hash = *blake3::hash(&input.as_bytes()[..boundary]).as_bytes();
1069                    let saved = cache.cache.get(&hash).expect("every prefix is cached");
1070                    let expected = tok.encode(&input[..boundary]).unwrap();
1071                    assert_eq!(&*saved, expected.token_ids());
1072                    expected_bytes += size_of_val(expected.token_ids());
1073                }
1074                let stats = cache.stats();
1075                assert_eq!(stats.entries, boundaries.len());
1076                assert_eq!(stats.memory_bytes, expected_bytes);
1077                assert_eq!(stats.hits, 0);
1078                assert_eq!(stats.misses, u64::from(reuse_hashes));
1079                let (_t, offset, deepest) = cache
1080                    .longest_prefix_match(input)
1081                    .expect("hit after populate");
1082                assert_eq!(offset, *boundaries.last().unwrap());
1083                assert_eq!(deepest, offset);
1084            }
1085        }
1086    }
1087
1088    #[test]
1089    fn populate_and_encode_handles_inputs_without_special_tokens() {
1090        let tok = load_tokenizer();
1091        for input in ["", "plain text with no special tokens at all", "<s>"] {
1092            for reuse_hashes in [false, true] {
1093                let cache = test_cache(8 * 1024 * 1024);
1094                let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1095                let plain = tok.encode(input).unwrap();
1096                assert_eq!(got, plain.token_ids());
1097                assert!(cache.is_empty(), "nothing cacheable without boundaries");
1098                assert_eq!(cache.stats().misses, u64::from(reuse_hashes));
1099            }
1100        }
1101    }
1102
1103    #[test]
1104    fn populate_and_encode_handles_trailing_special_token() {
1105        // The boundary at input.len() is excluded, leaving the final `</s>` in the tail.
1106        let tok = load_tokenizer();
1107        let input = "<s>system\nDone.</s>";
1108        for reuse_hashes in [false, true] {
1109            let cache = test_cache(8 * 1024 * 1024);
1110            let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1111            let plain = tok.encode(input).unwrap();
1112            assert_eq!(
1113                got,
1114                plain.token_ids(),
1115                "tail-segment assembly must be exact"
1116            );
1117        }
1118    }
1119
1120    #[test]
1121    fn miss_encode_failure_retains_only_completed_prefixes() {
1122        let input = "<s>世界</s><s>tail";
1123        let boundaries = find_special_token_boundaries(input, SPECIALS);
1124        for fail_at in 0..=boundaries.len() {
1125            for reuse_hashes in [false, true] {
1126                let cache = test_cache(8 * 1024 * 1024);
1127                let encoder = FailAt {
1128                    call: 0.into(),
1129                    fail_at,
1130                };
1131                let error = populate_miss(&cache, input, &encoder, reuse_hashes).unwrap_err();
1132                assert_eq!(error.to_string(), "suffix failed");
1133                assert_eq!(cache.len(), fail_at);
1134                for (index, &boundary) in boundaries.iter().enumerate() {
1135                    let hash = *blake3::hash(&input.as_bytes()[..boundary]).as_bytes();
1136                    let saved = cache.cache.get(&hash);
1137                    if index < fail_at {
1138                        assert_eq!(&*saved.unwrap(), vec![1; index + 1]);
1139                    } else {
1140                        assert!(saved.is_none());
1141                    }
1142                }
1143            }
1144        }
1145    }
1146}