Skip to main content

dynamo_tokenizers/cache/
mod.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 L0 layer, removed `add_special_tokens` plumbing (Dynamo's
9// `Encoder::encode` has no such flag), dropped fingerprinting, retargeted onto
10// `crate::traits::Tokenizer`.
11
12//! Tokenizer caching layer (L1: prefix matching at special-token boundaries).
13//!
14//! Wraps a cache-compatible [`Tokenizer`] in a cache that records prefix
15//! tokenizations at every special-token boundary. On a hit, the cached prefix
16//! tokens are merged with a fresh encode of the trailing suffix only — turning
17//! O(N) tokenization work into O(suffix_len) when prompts share a system prefix.
18//!
19//! # Correctness
20//!
21//! Boundaries are taken **only** at positions immediately following a registered
22//! special token (e.g. `<|im_start|>`, `<|im_end|>`, `<s>`, `</s>`). Special tokens
23//! are atomic in BPE (`special: true, normalized: false`), so splitting there
24//! preserves the invariant `tokenize(prefix) + tokenize(suffix) == tokenize(prefix + suffix)`.
25//! No fallback to whitespace or punctuation — better to miss than to corrupt.
26//! Callers must supply actual atomic special tokens recognized by the inner tokenizer;
27//! adding arbitrary strings to the boundary list is unsafe.
28//!
29//! Atomicity alone is insufficient when registered special-token strings can overlap.
30//! [`CachedTokenizer::new`] disables L1 for such sets because the boundary scanner could
31//! otherwise split inside the token selected by the underlying tokenizer.
32//!
33//! # Storage normalization
34//!
35//! When L1 is enabled, **every** `encode` returns [`Encoding::Sp`] (token-ids only) —
36//! hits merge cached prefix ids with a fresh suffix encode, and misses assemble the ids
37//! from the per-boundary segment encodes (see [`L1Cache::populate_and_encode`]) — even
38//! when the inner tokenizer would have produced [`Encoding::Hf`] (rich offsets/attention/
39//! etc). All current downstream consumers in Dynamo only call [`Encoding::token_ids`], so
40//! this lossy normalization is safe; revisit if a caller starts reading offsets or
41//! attention masks from encodings produced through the cache.
42//!
43//! # Configuration
44//!
45//! - `special_tokens: Vec<String>` — must be supplied at construction (the
46//!   [`Tokenizer`] trait is intentionally minimal and does not expose them).
47//!   An empty list disables L1: `encode`/`encode_batch` short-circuit straight
48//!   to the inner tokenizer with no lookup, no miss-counter bump, and no
49//!   insert attempt. A list whose members can overlap disables L1 identically.
50//! - `encode_segments` always passes through to the inner tokenizer without
51//!   caching. Flattening segments for L1 would discard their special-token
52//!   trust boundaries.
53//! - `max_memory_bytes` — token-ID payload byte budget, excluding keys and metadata.
54//!   Moka shares admission and eviction via W-TinyLFU. Deferred maintenance makes the
55//!   capacity approximate; it is not a limit on total process memory.
56//! - [`CachedTokenizer::new`] owns a private cache. [`CachedTokenizer::new_with_cache`]
57//!   shares storage across wrappers; entries survive a wrapper being dropped while
58//!   shared storage remains alive. Equal namespaces must identify identical tokenizer
59//!   behavior, including tokenizer files, backend, and options that affect token IDs.
60//!
61//! # Provenance
62//!
63//! Adapted from `llm-tokenizer` v1.3.2 (`cache/l1.rs`, `cache/mod.rs`). L0 and
64//! upstream fingerprinting were dropped; L1 covers the multi-turn-chat workload.
65//! Shared caches use caller-supplied namespaces to separate tokenizer identities.
66
67mod l1;
68
69use std::sync::Arc;
70
71use l1::PrefixLookup;
72pub use l1::{
73    CacheEventFn, L1Cache, L1CacheStats, SharedTokenizerCache, SharedTokenizerCacheStats,
74};
75
76use crate::{
77    EncodeSegment, Encoding, Result, TokenIdType,
78    traits::{DecodeResult, Decoder, Encoder, Tokenizer},
79};
80
81/// Token-level cache usage for one successful encode.
82///
83/// A partial cache hit reports both cached prefix tokens and uncached suffix tokens.
84/// Their sum always equals the number of tokens returned by the encode operation.
85#[derive(Debug, Clone, Copy, PartialEq, Eq)]
86pub struct CacheTokenUsage {
87    /// Tokens returned from the cached prefix.
88    pub cached_tokens: usize,
89    /// Tokens freshly encoded from the uncached suffix.
90    pub uncached_tokens: usize,
91}
92
93/// Optional observer for token-level cache usage.
94pub type CacheTokenUsageFn = Arc<dyn Fn(CacheTokenUsage) + Send + Sync>;
95
96/// Caching wrapper around an inner tokenizer.
97///
98/// Implements [`Encoder`], [`Decoder`], and [`Tokenizer`]; decode calls pass
99/// through to the inner tokenizer (decoding is fast and rarely repeated).
100pub struct CachedTokenizer {
101    inner: Arc<dyn Tokenizer>,
102    l1: L1Cache,
103    l1_enabled: bool,
104    extend_on_hit: bool,
105    /// Called once after every successful encode while L1 is active.
106    token_observer: Option<CacheTokenUsageFn>,
107}
108
109impl CachedTokenizer {
110    /// Construct a cached tokenizer.
111    ///
112    /// `special_tokens` is the list of atomic special-token strings the inner
113    /// tokenizer recognizes (typically extracted via the HuggingFace tokenizer's
114    /// `get_added_tokens_decoder()` filtering by `special == true`). An empty list
115    /// disables L1 — `encode`/`encode_batch` short-circuit to the inner tokenizer
116    /// without touching the cache or its counters. An overlapping token set also disables
117    /// L1, with a warning, because its boundaries are ambiguous.
118    ///
119    /// `max_memory_bytes` is the private token-ID payload byte budget. Moka defers
120    /// eviction, so the budget is approximate and excludes keys and metadata.
121    ///
122    /// # Errors
123    ///
124    /// Returns the inner tokenizer's compatibility error when it cannot be
125    /// safely wrapped in the prefix cache.
126    pub fn new(
127        inner: Arc<dyn Tokenizer>,
128        special_tokens: Vec<String>,
129        max_memory_bytes: usize,
130    ) -> Result<Self> {
131        Self::build(inner, special_tokens, |tokens| {
132            L1Cache::new(max_memory_bytes, tokens)
133        })
134    }
135
136    /// Construct a tokenizer using shared storage and a caller-supplied namespace.
137    ///
138    /// Equal namespaces share entries and must describe identical tokenizer behavior,
139    /// including tokenizer files, backend, and encoding options. Different namespaces
140    /// compete for the same byte budget but cannot reuse each other's token IDs.
141    /// Entries survive this wrapper being dropped while the shared cache remains alive.
142    /// The eligibility checks are the same as [`Self::new`].
143    ///
144    /// # Errors
145    ///
146    /// Returns the inner tokenizer's compatibility error if prefix caching is unsafe.
147    pub fn new_with_cache(
148        inner: Arc<dyn Tokenizer>,
149        special_tokens: Vec<String>,
150        shared_cache: SharedTokenizerCache,
151        namespace: &[u8],
152    ) -> Result<Self> {
153        Self::build(inner, special_tokens, |tokens| {
154            L1Cache::new_with_cache(shared_cache, tokens, namespace)
155        })
156    }
157
158    fn build(
159        inner: Arc<dyn Tokenizer>,
160        mut special_tokens: Vec<String>,
161        make_cache: impl FnOnce(Vec<String>) -> L1Cache,
162    ) -> Result<Self> {
163        inner.validate_prefix_cache()?;
164        special_tokens.retain(|token| !token.is_empty());
165
166        // Overlapping matches can create a cache boundary inside a token selected by the
167        // inner tokenizer. Preserve correctness by bypassing this optional optimization.
168        let overlapping_specials = match l1::first_unsafe_overlap(&special_tokens) {
169            Some((first, second)) => {
170                tracing::warn!(
171                    target: "tokenizer",
172                    first_token = first,
173                    second_token = second,
174                    special_token_count = special_tokens.len(),
175                    "special tokens can overlap; tokenizer prefix cache disabled"
176                );
177                true
178            }
179            None => false,
180        };
181
182        let l1_enabled = !special_tokens.is_empty() && !overlapping_specials;
183        let cache_tokens = if l1_enabled {
184            special_tokens
185        } else {
186            Vec::new()
187        };
188        Ok(Self {
189            inner,
190            l1: make_cache(cache_tokens),
191            l1_enabled,
192            extend_on_hit: false,
193            token_observer: None,
194        })
195    }
196
197    /// Enable partial-hit extension. When on, a partial cache hit also caches the
198    /// freshly-tokenized suffix at its deepest special-token boundary, so each turn of
199    /// a growing multi-turn conversation hits deeper than the last and per-turn
200    /// tokenization cost stops growing with conversation length. Default off.
201    pub fn with_extend(mut self, enabled: bool) -> Self {
202        self.extend_on_hit = enabled;
203        self
204    }
205
206    /// Install hit/miss callbacks so each L1 lookup pushes an event into the
207    /// supplied closures (e.g. `Prometheus::Counter::inc`). Replaces any
208    /// previously-set observer.
209    pub fn with_observer(mut self, on_hit: CacheEventFn, on_miss: CacheEventFn) -> Self {
210        self.l1.set_observer(on_hit, on_miss);
211        self
212    }
213
214    /// Install a callback that receives exact cached and uncached token counts after each
215    /// successful encode while L1 is active. A partial hit reports both categories, which
216    /// lets consumers maintain token-level cache totals and derive a reuse ratio. Replaces
217    /// any previously-set token observer.
218    ///
219    /// This observer is not called when the special-token set is empty (and L1 is therefore
220    /// disabled) or when encoding returns an error.
221    pub fn with_token_observer(mut self, observer: CacheTokenUsageFn) -> Self {
222        self.token_observer = Some(observer);
223        self
224    }
225
226    fn observe_token_usage(&self, cached_tokens: usize, total_tokens: usize) {
227        if let Some(observer) = &self.token_observer {
228            let uncached_tokens = total_tokens
229                .checked_sub(cached_tokens)
230                .expect("cached token count cannot exceed total token count");
231            observer(CacheTokenUsage {
232                cached_tokens,
233                uncached_tokens,
234            });
235        }
236    }
237
238    /// Wrapper-local hits/misses and namespace-wide entries/token bytes.
239    /// Shared storage statistics scan the namespace; private caches use Moka's totals.
240    /// Results can change under concurrent writes. Disabled wrappers report zeroes.
241    pub fn cache_stats(&self) -> L1CacheStats {
242        if self.l1_enabled {
243            self.l1.stats()
244        } else {
245            L1CacheStats::default()
246        }
247    }
248
249    /// Access the underlying tokenizer (e.g. for downcasting to a concrete type).
250    pub fn inner(&self) -> &Arc<dyn Tokenizer> {
251        &self.inner
252    }
253}
254
255impl Encoder for CachedTokenizer {
256    fn encode(&self, input: &str) -> Result<Encoding> {
257        if !self.l1_enabled {
258            return self.inner.encode(input);
259        }
260
261        let matched = match self.l1.lookup_prefix(input) {
262            PrefixLookup::Hit(matched) => matched,
263            PrefixLookup::Miss(prefix_hashes) => {
264                let encoding = Encoding::Sp(self.l1.populate_and_encode_with_hashes(
265                    input,
266                    prefix_hashes.into_iter(),
267                    self.inner.as_ref(),
268                )?);
269                self.observe_token_usage(0, encoding.token_ids().len());
270                return Ok(encoding);
271            }
272        };
273
274        let cached_tokens = matched.tokens.len();
275        let encoding = if self.extend_on_hit {
276            Encoding::Sp(self.l1.extend_after_match_with_hash(
277                input,
278                matched,
279                self.inner.as_ref(),
280            )?)
281        } else {
282            let suffix_enc = self.inner.encode(&input[matched.prefix_len..])?;
283            // Reserve once to avoid copying the cached prefix during vector growth.
284            let mut merged: Vec<TokenIdType> =
285                Vec::with_capacity(matched.tokens.len() + suffix_enc.token_ids().len());
286            merged.extend_from_slice(&matched.tokens);
287            merged.extend_from_slice(suffix_enc.token_ids());
288            Encoding::Sp(merged)
289        };
290        self.observe_token_usage(cached_tokens, encoding.token_ids().len());
291        Ok(encoding)
292    }
293
294    fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
295        // True passthrough when L1 is disabled — delegate to the inner's native
296        // batch path (which may be rayon-parallel for HF) instead of falling
297        // through per-item.
298        if !self.l1_enabled {
299            return self.inner.encode_batch(inputs);
300        }
301
302        // Per-item cache lookup — do NOT delegate to inner.encode_batch, which would
303        // bypass the cache. Sequential iteration is fine; if rayon is added later it
304        // belongs here, not inside `encode`.
305        inputs.iter().map(|&i| self.encode(i)).collect()
306    }
307
308    fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
309        // L1 indexes flattened string offsets and cannot preserve each
310        // segment's allow_special boundary. Keep the operation correct by
311        // delegating without populating or consulting the cache.
312        let encoding = self.inner.encode_segments(segments)?;
313        if self.l1_enabled {
314            self.observe_token_usage(0, encoding.token_ids().len());
315        }
316        Ok(encoding)
317    }
318}
319
320impl Decoder for CachedTokenizer {
321    fn has_unstable_suffix(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> bool {
322        self.inner
323            .has_unstable_suffix(token_ids, skip_special_tokens)
324    }
325
326    fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
327        // Decode is not cached — passthrough to inner.
328        self.inner.decode(token_ids, skip_special_tokens)
329    }
330}
331
332impl Tokenizer for CachedTokenizer {
333    fn vocab_size(&self) -> Option<usize> {
334        self.inner.vocab_size()
335    }
336
337    fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
338        self.inner.token_to_id(token)
339    }
340
341    fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
342        self.inner.special_token_ids()
343    }
344
345    fn num_special_tokens_added(&self) -> Result<usize> {
346        Ok(0)
347    }
348}
349
350#[cfg(test)]
351mod tests {
352    use super::*;
353    use crate::HuggingFaceTokenizer;
354    use std::sync::{Mutex, atomic::AtomicU64, atomic::Ordering};
355    use tokenizers::Tokenizer as HfTokenizer;
356
357    struct FailingTokenizer;
358
359    struct SegmentTokenizer;
360
361    impl Encoder for SegmentTokenizer {
362        fn encode(&self, input: &str) -> Result<Encoding> {
363            Ok(Encoding::Sp(vec![input.len() as u32]))
364        }
365
366        fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
367            inputs.iter().map(|input| self.encode(input)).collect()
368        }
369
370        fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
371            let ids = segments
372                .iter()
373                .flat_map(|segment| [segment.allow_special as u32, segment.text.len() as u32])
374                .collect();
375            Ok(Encoding::Sp(ids))
376        }
377    }
378
379    impl Decoder for SegmentTokenizer {
380        fn decode(
381            &self,
382            _token_ids: &[TokenIdType],
383            _skip_special_tokens: bool,
384        ) -> Result<DecodeResult> {
385            Ok(DecodeResult::Complete(String::new()))
386        }
387    }
388
389    impl Tokenizer for SegmentTokenizer {
390        fn validate_prefix_cache(&self) -> Result<()> {
391            Ok(())
392        }
393    }
394
395    impl Encoder for FailingTokenizer {
396        fn encode(&self, _input: &str) -> Result<Encoding> {
397            Err(anyhow::anyhow!("intentional encode failure"))
398        }
399
400        fn encode_batch(&self, _inputs: &[&str]) -> Result<Vec<Encoding>> {
401            Err(anyhow::anyhow!("intentional encode failure"))
402        }
403    }
404
405    impl Decoder for FailingTokenizer {
406        fn decode(
407            &self,
408            _token_ids: &[TokenIdType],
409            _skip_special_tokens: bool,
410        ) -> Result<DecodeResult> {
411            Err(anyhow::anyhow!("intentional decode failure"))
412        }
413    }
414
415    impl Tokenizer for FailingTokenizer {
416        fn validate_prefix_cache(&self) -> Result<()> {
417            Ok(())
418        }
419
420        fn vocab_size(&self) -> Option<usize> {
421            None
422        }
423    }
424
425    const TINYLLAMA_PATH: &str = concat!(
426        env!("CARGO_MANIFEST_DIR"),
427        "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
428    );
429
430    fn inner() -> Arc<dyn Tokenizer> {
431        Arc::new(HuggingFaceTokenizer::from_file(TINYLLAMA_PATH).expect("load TinyLlama"))
432    }
433
434    fn specials() -> Vec<String> {
435        vec!["<s>".into(), "</s>".into()]
436    }
437
438    fn collect_token_usage(
439        tokenizer: CachedTokenizer,
440    ) -> (CachedTokenizer, Arc<Mutex<Vec<CacheTokenUsage>>>) {
441        let events = Arc::new(Mutex::new(Vec::new()));
442        let observed = events.clone();
443        let tokenizer = tokenizer.with_token_observer(Arc::new(move |usage| {
444            observed.lock().unwrap().push(usage);
445        }));
446        (tokenizer, events)
447    }
448
449    #[test]
450    fn rejects_hf_tokenizer_that_adds_special_tokens() {
451        let tokenizer: Arc<dyn Tokenizer> = Arc::new(
452            HuggingFaceTokenizer::from_file(TINYLLAMA_PATH)
453                .expect("load TinyLlama")
454                .with_options(crate::TokenizerOptions {
455                    add_special_tokens: true,
456                }),
457        );
458
459        let result = CachedTokenizer::new(tokenizer, specials(), 4096);
460        let Err(error) = result else {
461            panic!("add_special_tokens=true must be rejected");
462        };
463        assert_eq!(
464            error.to_string(),
465            "HuggingFace tokenizers configured with add_special_tokens=true must remain uncached"
466        );
467    }
468
469    #[test]
470    fn empty_specials_passes_through_correctly() {
471        // Empty token strings carry no boundary information and must not make L1 active.
472        let tok = inner();
473        let (cached, events) = collect_token_usage(
474            CachedTokenizer::new(tok.clone(), vec![String::new()], 4096)
475                .expect("TinyLlama must support prefix caching"),
476        );
477        let s = "<s>hello world</s>";
478        let a = cached.encode(s).unwrap();
479        let b = tok.encode(s).unwrap();
480        assert_eq!(a.token_ids(), b.token_ids());
481        let stats = cached.cache_stats();
482        assert_eq!(stats.entries, 0);
483        assert_eq!(stats.misses, 0, "empty specials must not increment misses");
484        assert_eq!(stats.hits, 0);
485        assert!(
486            events.lock().unwrap().is_empty(),
487            "empty specials must not emit token usage"
488        );
489    }
490
491    #[test]
492    fn laguna_overlapping_specials_bypass_cache() {
493        const TOKENIZER_JSON: &str = r#"{
494            "version": "1.0",
495            "truncation": null,
496            "padding": null,
497            "added_tokens": [
498                {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
499                {"id": 2, "content": "〈|EOS|〉", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
500                {"id": 14, "content": "〈|", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
501                {"id": 15, "content": "|〉", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
502            ],
503            "normalizer": null,
504            "pre_tokenizer": null,
505            "post_processor": null,
506            "decoder": null,
507            "model": {
508                "type": "WordLevel",
509                "vocab": {"<unk>": 0, "〈|EOS|〉": 2, "〈|": 14, "|〉": 15, "tail": 16},
510                "unk_token": "<unk>"
511            }
512        }"#;
513
514        let hf = HfTokenizer::from_bytes(TOKENIZER_JSON).expect("load test tokenizer");
515        let tok: Arc<dyn Tokenizer> = Arc::new(HuggingFaceTokenizer::from_tokenizer(hf));
516        let overlapping = vec!["〈|EOS|〉".into(), "〈|".into(), "|〉".into()];
517        let (cached, events) = collect_token_usage(
518            CachedTokenizer::new(tok.clone(), overlapping, 4096)
519                .expect("HuggingFace tokenizer must support prefix caching"),
520        );
521
522        let expected = tok.encode("〈|EOS|〉").unwrap();
523        assert_eq!(expected.token_ids(), &[2]);
524        assert_eq!(
525            cached.encode("〈|EOS|〉").unwrap().token_ids(),
526            expected.token_ids()
527        );
528        let stats = cached.cache_stats();
529        assert_eq!(stats.entries, 0);
530        assert_eq!(
531            stats.misses, 0,
532            "overlapping specials must not increment misses"
533        );
534        assert_eq!(stats.hits, 0);
535        assert!(
536            events.lock().unwrap().is_empty(),
537            "overlapping specials must not emit token usage"
538        );
539    }
540
541    #[test]
542    fn segmented_encoding_passes_through_without_caching() {
543        let inner: Arc<dyn Tokenizer> = Arc::new(SegmentTokenizer);
544        let segments = [
545            EncodeSegment::new("<ctl>", true),
546            EncodeSegment::new("user content", false),
547        ];
548        let expected = inner.encode_segments(&segments).unwrap();
549
550        for special_tokens in [Vec::new(), vec!["<ctl>".to_string()]] {
551            let l1_enabled = !special_tokens.is_empty();
552            let (cached, events) = collect_token_usage(
553                CachedTokenizer::new(inner.clone(), special_tokens, 4096)
554                    .expect("test tokenizer supports prefix caching"),
555            );
556            let actual = cached.encode_segments(&segments).unwrap();
557
558            assert_eq!(actual.token_ids(), expected.token_ids());
559            let stats = cached.cache_stats();
560            assert_eq!(stats.entries, 0);
561            assert_eq!(stats.hits, 0);
562            assert_eq!(stats.misses, 0);
563            let events = events.lock().unwrap();
564            if l1_enabled {
565                assert_eq!(
566                    events.as_slice(),
567                    &[CacheTokenUsage {
568                        cached_tokens: 0,
569                        uncached_tokens: expected.token_ids().len(),
570                    }]
571                );
572            } else {
573                assert!(events.is_empty());
574            }
575        }
576    }
577
578    #[test]
579    fn token_observer_reports_full_miss_and_partial_hit_with_and_without_extension() {
580        for extend_on_hit in [false, true] {
581            let tok = inner();
582            let hits = Arc::new(AtomicU64::new(0));
583            let misses = Arc::new(AtomicU64::new(0));
584            let hit_counter = hits.clone();
585            let miss_counter = misses.clone();
586            let cached = CachedTokenizer::new(tok, specials(), 64 * 1024)
587                .expect("TinyLlama must support prefix caching")
588                .with_extend(extend_on_hit)
589                .with_observer(
590                    Arc::new(move || {
591                        hit_counter.fetch_add(1, Ordering::Relaxed);
592                    }),
593                    Arc::new(move || {
594                        miss_counter.fetch_add(1, Ordering::Relaxed);
595                    }),
596                );
597            let (cached, events) = collect_token_usage(cached);
598
599            let shared = "<s>system\nYou are helpful.</s><s>user\n";
600            let first = format!("{shared}First question?</s>");
601            let second = format!("{shared}Second different prompt entirely.</s>");
602
603            let first_encoding = cached.encode(&first).unwrap();
604            let second_encoding = cached.encode(&second).unwrap();
605
606            let events = events.lock().unwrap();
607            assert_eq!(events.len(), 2);
608            assert_eq!(
609                events[0],
610                CacheTokenUsage {
611                    cached_tokens: 0,
612                    uncached_tokens: first_encoding.token_ids().len(),
613                }
614            );
615            assert!(events[1].cached_tokens > 0);
616            assert!(events[1].uncached_tokens > 0);
617            assert_eq!(
618                events[1].cached_tokens + events[1].uncached_tokens,
619                second_encoding.token_ids().len()
620            );
621            assert_eq!(hits.load(Ordering::Relaxed), 1);
622            assert_eq!(misses.load(Ordering::Relaxed), 1);
623        }
624    }
625
626    #[test]
627    fn token_observer_does_not_report_failed_encodes() {
628        let tokenizer: Arc<dyn Tokenizer> = Arc::new(FailingTokenizer);
629        let (cached, events) = collect_token_usage(
630            CachedTokenizer::new(tokenizer, specials(), 4096)
631                .expect("test tokenizer explicitly supports prefix caching"),
632        );
633
634        assert!(cached.encode("<s>this fails</s>").is_err());
635        assert!(events.lock().unwrap().is_empty());
636    }
637
638    #[test]
639    fn two_turn_chat_correctness_and_hit() {
640        let tok = inner();
641        let cached = CachedTokenizer::new(tok.clone(), specials(), 64 * 1024)
642            .expect("TinyLlama must support prefix caching");
643
644        let template = "<s>system\nYou are helpful.</s><s>user\n";
645        let first = format!("{template}First question?</s>");
646        let second = format!("{template}Second different prompt entirely.</s>");
647
648        // Warm the cache.
649        let _ = cached.encode(&first).unwrap();
650
651        // Second request: shared prefix → L1 hit, suffix-only fresh encode.
652        let cached_second = cached.encode(&second).unwrap();
653        let plain_second = tok.encode(&second).unwrap();
654        assert_eq!(
655            cached_second.token_ids(),
656            plain_second.token_ids(),
657            "cached encode must equal plain encode for second turn"
658        );
659
660        let stats = cached.cache_stats();
661        assert!(stats.hits >= 1, "expected L1 hit on second request");
662    }
663
664    #[test]
665    fn decode_passes_through() {
666        let tok = inner();
667        let cached = CachedTokenizer::new(tok.clone(), specials(), 4096)
668            .expect("TinyLlama must support prefix caching");
669        let enc = cached.encode("<s>hello</s>").unwrap();
670        let direct = tok.decode(enc.token_ids(), false).unwrap();
671        let through = cached.decode(enc.token_ids(), false).unwrap();
672        assert_eq!(direct, through);
673    }
674
675    #[test]
676    fn encode_batch_uses_cache() {
677        let tok = inner();
678        let (cached, events) = collect_token_usage(
679            CachedTokenizer::new(tok.clone(), specials(), 64 * 1024)
680                .expect("TinyLlama must support prefix caching"),
681        );
682        let shared = "<s>system\nShared persona.</s><s>user\n";
683        let inputs = [
684            format!("{shared}q1</s>"),
685            format!("{shared}q2</s>"),
686            format!("{shared}q3</s>"),
687        ];
688        let refs: Vec<&str> = inputs.iter().map(String::as_str).collect();
689        let outs = cached.encode_batch(&refs).unwrap();
690        assert_eq!(outs.len(), 3);
691        let events = events.lock().unwrap();
692        assert_eq!(events.len(), outs.len());
693        for (event, output) in events.iter().zip(&outs) {
694            assert_eq!(
695                event.cached_tokens + event.uncached_tokens,
696                output.token_ids().len()
697            );
698        }
699        assert_eq!(events[0].cached_tokens, 0);
700        assert!(events[1..].iter().all(|event| event.cached_tokens > 0));
701        // First call populates, second/third hit.
702        assert!(cached.cache_stats().hits >= 2, "expected hits on q2 and q3");
703    }
704
705    #[test]
706    fn vocab_introspection_forwards_to_inner() {
707        let tok = inner();
708        let cached = CachedTokenizer::new(tok.clone(), specials(), 4096)
709            .expect("TinyLlama must support prefix caching");
710        assert_eq!(cached.vocab_size(), tok.vocab_size());
711        assert_eq!(
712            cached.token_to_id("<s>").unwrap(),
713            tok.token_to_id("<s>").unwrap()
714        );
715        assert_eq!(
716            cached.special_token_ids().unwrap(),
717            tok.special_token_ids().unwrap()
718        );
719    }
720
721    #[test]
722    fn special_token_accounting_matches_cached_encoder_behavior() {
723        let cached = CachedTokenizer::new(inner(), specials(), 4096)
724            .expect("TinyLlama must support prefix caching")
725            .with_options(crate::TokenizerOptions {
726                add_special_tokens: true,
727            });
728        let cached_ids = cached.encode("hello").unwrap();
729        let hf_ids = HuggingFaceTokenizer::from_file(TINYLLAMA_PATH)
730            .expect("load TinyLlama")
731            .with_options(crate::TokenizerOptions {
732                add_special_tokens: true,
733            })
734            .encode("hello")
735            .unwrap();
736
737        assert_eq!(cached.num_special_tokens_added().unwrap(), 0);
738        assert_eq!(hf_ids.token_ids().len(), cached_ids.token_ids().len() + 1);
739        assert_eq!(&hf_ids.token_ids()[1..], cached_ids.token_ids());
740    }
741
742    #[test]
743    fn unoverridden_introspection_methods_use_defaults() {
744        let tokenizer = SegmentTokenizer;
745        assert_eq!(tokenizer.vocab_size(), None);
746        assert!(tokenizer.token_to_id("anything").is_err());
747        assert!(tokenizer.special_token_ids().is_err());
748        assert!(tokenizer.num_special_tokens_added().is_err());
749    }
750}