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