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 decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
322        // Decode is not cached — passthrough to inner.
323        self.inner.decode(token_ids, skip_special_tokens)
324    }
325}
326
327impl Tokenizer for CachedTokenizer {
328    fn vocab_size(&self) -> Option<usize> {
329        self.inner.vocab_size()
330    }
331
332    fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
333        self.inner.token_to_id(token)
334    }
335
336    fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
337        self.inner.special_token_ids()
338    }
339
340    fn num_special_tokens_added(&self) -> Result<usize> {
341        Ok(0)
342    }
343}
344
345#[cfg(test)]
346mod tests {
347    use super::*;
348    use crate::HuggingFaceTokenizer;
349    use std::sync::{Mutex, atomic::AtomicU64, atomic::Ordering};
350    use tokenizers::Tokenizer as HfTokenizer;
351
352    struct FailingTokenizer;
353
354    struct SegmentTokenizer;
355
356    impl Encoder for SegmentTokenizer {
357        fn encode(&self, input: &str) -> Result<Encoding> {
358            Ok(Encoding::Sp(vec![input.len() as u32]))
359        }
360
361        fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
362            inputs.iter().map(|input| self.encode(input)).collect()
363        }
364
365        fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
366            let ids = segments
367                .iter()
368                .flat_map(|segment| [segment.allow_special as u32, segment.text.len() as u32])
369                .collect();
370            Ok(Encoding::Sp(ids))
371        }
372    }
373
374    impl Decoder for SegmentTokenizer {
375        fn decode(
376            &self,
377            _token_ids: &[TokenIdType],
378            _skip_special_tokens: bool,
379        ) -> Result<DecodeResult> {
380            Ok(DecodeResult::Complete(String::new()))
381        }
382    }
383
384    impl Tokenizer for SegmentTokenizer {
385        fn validate_prefix_cache(&self) -> Result<()> {
386            Ok(())
387        }
388    }
389
390    impl Encoder for FailingTokenizer {
391        fn encode(&self, _input: &str) -> Result<Encoding> {
392            Err(anyhow::anyhow!("intentional encode failure"))
393        }
394
395        fn encode_batch(&self, _inputs: &[&str]) -> Result<Vec<Encoding>> {
396            Err(anyhow::anyhow!("intentional encode failure"))
397        }
398    }
399
400    impl Decoder for FailingTokenizer {
401        fn decode(
402            &self,
403            _token_ids: &[TokenIdType],
404            _skip_special_tokens: bool,
405        ) -> Result<DecodeResult> {
406            Err(anyhow::anyhow!("intentional decode failure"))
407        }
408    }
409
410    impl Tokenizer for FailingTokenizer {
411        fn validate_prefix_cache(&self) -> Result<()> {
412            Ok(())
413        }
414
415        fn vocab_size(&self) -> Option<usize> {
416            None
417        }
418    }
419
420    const TINYLLAMA_PATH: &str = concat!(
421        env!("CARGO_MANIFEST_DIR"),
422        "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
423    );
424
425    fn inner() -> Arc<dyn Tokenizer> {
426        Arc::new(HuggingFaceTokenizer::from_file(TINYLLAMA_PATH).expect("load TinyLlama"))
427    }
428
429    fn specials() -> Vec<String> {
430        vec!["<s>".into(), "</s>".into()]
431    }
432
433    fn collect_token_usage(
434        tokenizer: CachedTokenizer,
435    ) -> (CachedTokenizer, Arc<Mutex<Vec<CacheTokenUsage>>>) {
436        let events = Arc::new(Mutex::new(Vec::new()));
437        let observed = events.clone();
438        let tokenizer = tokenizer.with_token_observer(Arc::new(move |usage| {
439            observed.lock().unwrap().push(usage);
440        }));
441        (tokenizer, events)
442    }
443
444    #[test]
445    fn rejects_hf_tokenizer_that_adds_special_tokens() {
446        let tokenizer: Arc<dyn Tokenizer> = Arc::new(
447            HuggingFaceTokenizer::from_file(TINYLLAMA_PATH)
448                .expect("load TinyLlama")
449                .with_options(crate::TokenizerOptions {
450                    add_special_tokens: true,
451                }),
452        );
453
454        let result = CachedTokenizer::new(tokenizer, specials(), 4096);
455        let Err(error) = result else {
456            panic!("add_special_tokens=true must be rejected");
457        };
458        assert_eq!(
459            error.to_string(),
460            "HuggingFace tokenizers configured with add_special_tokens=true must remain uncached"
461        );
462    }
463
464    #[test]
465    fn empty_specials_passes_through_correctly() {
466        // Empty token strings carry no boundary information and must not make L1 active.
467        let tok = inner();
468        let (cached, events) = collect_token_usage(
469            CachedTokenizer::new(tok.clone(), vec![String::new()], 4096)
470                .expect("TinyLlama must support prefix caching"),
471        );
472        let s = "<s>hello world</s>";
473        let a = cached.encode(s).unwrap();
474        let b = tok.encode(s).unwrap();
475        assert_eq!(a.token_ids(), b.token_ids());
476        let stats = cached.cache_stats();
477        assert_eq!(stats.entries, 0);
478        assert_eq!(stats.misses, 0, "empty specials must not increment misses");
479        assert_eq!(stats.hits, 0);
480        assert!(
481            events.lock().unwrap().is_empty(),
482            "empty specials must not emit token usage"
483        );
484    }
485
486    #[test]
487    fn laguna_overlapping_specials_bypass_cache() {
488        const TOKENIZER_JSON: &str = r#"{
489            "version": "1.0",
490            "truncation": null,
491            "padding": null,
492            "added_tokens": [
493                {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
494                {"id": 2, "content": "〈|EOS|〉", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
495                {"id": 14, "content": "〈|", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
496                {"id": 15, "content": "|〉", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
497            ],
498            "normalizer": null,
499            "pre_tokenizer": null,
500            "post_processor": null,
501            "decoder": null,
502            "model": {
503                "type": "WordLevel",
504                "vocab": {"<unk>": 0, "〈|EOS|〉": 2, "〈|": 14, "|〉": 15, "tail": 16},
505                "unk_token": "<unk>"
506            }
507        }"#;
508
509        let hf = HfTokenizer::from_bytes(TOKENIZER_JSON).expect("load test tokenizer");
510        let tok: Arc<dyn Tokenizer> = Arc::new(HuggingFaceTokenizer::from_tokenizer(hf));
511        let overlapping = vec!["〈|EOS|〉".into(), "〈|".into(), "|〉".into()];
512        let (cached, events) = collect_token_usage(
513            CachedTokenizer::new(tok.clone(), overlapping, 4096)
514                .expect("HuggingFace tokenizer must support prefix caching"),
515        );
516
517        let expected = tok.encode("〈|EOS|〉").unwrap();
518        assert_eq!(expected.token_ids(), &[2]);
519        assert_eq!(
520            cached.encode("〈|EOS|〉").unwrap().token_ids(),
521            expected.token_ids()
522        );
523        let stats = cached.cache_stats();
524        assert_eq!(stats.entries, 0);
525        assert_eq!(
526            stats.misses, 0,
527            "overlapping specials must not increment misses"
528        );
529        assert_eq!(stats.hits, 0);
530        assert!(
531            events.lock().unwrap().is_empty(),
532            "overlapping specials must not emit token usage"
533        );
534    }
535
536    #[test]
537    fn segmented_encoding_passes_through_without_caching() {
538        let inner: Arc<dyn Tokenizer> = Arc::new(SegmentTokenizer);
539        let segments = [
540            EncodeSegment::new("<ctl>", true),
541            EncodeSegment::new("user content", false),
542        ];
543        let expected = inner.encode_segments(&segments).unwrap();
544
545        for special_tokens in [Vec::new(), vec!["<ctl>".to_string()]] {
546            let l1_enabled = !special_tokens.is_empty();
547            let (cached, events) = collect_token_usage(
548                CachedTokenizer::new(inner.clone(), special_tokens, 4096)
549                    .expect("test tokenizer supports prefix caching"),
550            );
551            let actual = cached.encode_segments(&segments).unwrap();
552
553            assert_eq!(actual.token_ids(), expected.token_ids());
554            let stats = cached.cache_stats();
555            assert_eq!(stats.entries, 0);
556            assert_eq!(stats.hits, 0);
557            assert_eq!(stats.misses, 0);
558            let events = events.lock().unwrap();
559            if l1_enabled {
560                assert_eq!(
561                    events.as_slice(),
562                    &[CacheTokenUsage {
563                        cached_tokens: 0,
564                        uncached_tokens: expected.token_ids().len(),
565                    }]
566                );
567            } else {
568                assert!(events.is_empty());
569            }
570        }
571    }
572
573    #[test]
574    fn token_observer_reports_full_miss_and_partial_hit_with_and_without_extension() {
575        for extend_on_hit in [false, true] {
576            let tok = inner();
577            let hits = Arc::new(AtomicU64::new(0));
578            let misses = Arc::new(AtomicU64::new(0));
579            let hit_counter = hits.clone();
580            let miss_counter = misses.clone();
581            let cached = CachedTokenizer::new(tok, specials(), 64 * 1024)
582                .expect("TinyLlama must support prefix caching")
583                .with_extend(extend_on_hit)
584                .with_observer(
585                    Arc::new(move || {
586                        hit_counter.fetch_add(1, Ordering::Relaxed);
587                    }),
588                    Arc::new(move || {
589                        miss_counter.fetch_add(1, Ordering::Relaxed);
590                    }),
591                );
592            let (cached, events) = collect_token_usage(cached);
593
594            let shared = "<s>system\nYou are helpful.</s><s>user\n";
595            let first = format!("{shared}First question?</s>");
596            let second = format!("{shared}Second different prompt entirely.</s>");
597
598            let first_encoding = cached.encode(&first).unwrap();
599            let second_encoding = cached.encode(&second).unwrap();
600
601            let events = events.lock().unwrap();
602            assert_eq!(events.len(), 2);
603            assert_eq!(
604                events[0],
605                CacheTokenUsage {
606                    cached_tokens: 0,
607                    uncached_tokens: first_encoding.token_ids().len(),
608                }
609            );
610            assert!(events[1].cached_tokens > 0);
611            assert!(events[1].uncached_tokens > 0);
612            assert_eq!(
613                events[1].cached_tokens + events[1].uncached_tokens,
614                second_encoding.token_ids().len()
615            );
616            assert_eq!(hits.load(Ordering::Relaxed), 1);
617            assert_eq!(misses.load(Ordering::Relaxed), 1);
618        }
619    }
620
621    #[test]
622    fn token_observer_does_not_report_failed_encodes() {
623        let tokenizer: Arc<dyn Tokenizer> = Arc::new(FailingTokenizer);
624        let (cached, events) = collect_token_usage(
625            CachedTokenizer::new(tokenizer, specials(), 4096)
626                .expect("test tokenizer explicitly supports prefix caching"),
627        );
628
629        assert!(cached.encode("<s>this fails</s>").is_err());
630        assert!(events.lock().unwrap().is_empty());
631    }
632
633    #[test]
634    fn two_turn_chat_correctness_and_hit() {
635        let tok = inner();
636        let cached = CachedTokenizer::new(tok.clone(), specials(), 64 * 1024)
637            .expect("TinyLlama must support prefix caching");
638
639        let template = "<s>system\nYou are helpful.</s><s>user\n";
640        let first = format!("{template}First question?</s>");
641        let second = format!("{template}Second different prompt entirely.</s>");
642
643        // Warm the cache.
644        let _ = cached.encode(&first).unwrap();
645
646        // Second request: shared prefix → L1 hit, suffix-only fresh encode.
647        let cached_second = cached.encode(&second).unwrap();
648        let plain_second = tok.encode(&second).unwrap();
649        assert_eq!(
650            cached_second.token_ids(),
651            plain_second.token_ids(),
652            "cached encode must equal plain encode for second turn"
653        );
654
655        let stats = cached.cache_stats();
656        assert!(stats.hits >= 1, "expected L1 hit on second request");
657    }
658
659    #[test]
660    fn decode_passes_through() {
661        let tok = inner();
662        let cached = CachedTokenizer::new(tok.clone(), specials(), 4096)
663            .expect("TinyLlama must support prefix caching");
664        let enc = cached.encode("<s>hello</s>").unwrap();
665        let direct = tok.decode(enc.token_ids(), false).unwrap();
666        let through = cached.decode(enc.token_ids(), false).unwrap();
667        assert_eq!(direct, through);
668    }
669
670    #[test]
671    fn encode_batch_uses_cache() {
672        let tok = inner();
673        let (cached, events) = collect_token_usage(
674            CachedTokenizer::new(tok.clone(), specials(), 64 * 1024)
675                .expect("TinyLlama must support prefix caching"),
676        );
677        let shared = "<s>system\nShared persona.</s><s>user\n";
678        let inputs = [
679            format!("{shared}q1</s>"),
680            format!("{shared}q2</s>"),
681            format!("{shared}q3</s>"),
682        ];
683        let refs: Vec<&str> = inputs.iter().map(String::as_str).collect();
684        let outs = cached.encode_batch(&refs).unwrap();
685        assert_eq!(outs.len(), 3);
686        let events = events.lock().unwrap();
687        assert_eq!(events.len(), outs.len());
688        for (event, output) in events.iter().zip(&outs) {
689            assert_eq!(
690                event.cached_tokens + event.uncached_tokens,
691                output.token_ids().len()
692            );
693        }
694        assert_eq!(events[0].cached_tokens, 0);
695        assert!(events[1..].iter().all(|event| event.cached_tokens > 0));
696        // First call populates, second/third hit.
697        assert!(cached.cache_stats().hits >= 2, "expected hits on q2 and q3");
698    }
699
700    #[test]
701    fn vocab_introspection_forwards_to_inner() {
702        let tok = inner();
703        let cached = CachedTokenizer::new(tok.clone(), specials(), 4096)
704            .expect("TinyLlama must support prefix caching");
705        assert_eq!(cached.vocab_size(), tok.vocab_size());
706        assert_eq!(
707            cached.token_to_id("<s>").unwrap(),
708            tok.token_to_id("<s>").unwrap()
709        );
710        assert_eq!(
711            cached.special_token_ids().unwrap(),
712            tok.special_token_ids().unwrap()
713        );
714    }
715
716    #[test]
717    fn special_token_accounting_matches_cached_encoder_behavior() {
718        let cached = CachedTokenizer::new(inner(), specials(), 4096)
719            .expect("TinyLlama must support prefix caching")
720            .with_options(crate::TokenizerOptions {
721                add_special_tokens: true,
722            });
723        let cached_ids = cached.encode("hello").unwrap();
724        let hf_ids = HuggingFaceTokenizer::from_file(TINYLLAMA_PATH)
725            .expect("load TinyLlama")
726            .with_options(crate::TokenizerOptions {
727                add_special_tokens: true,
728            })
729            .encode("hello")
730            .unwrap();
731
732        assert_eq!(cached.num_special_tokens_added().unwrap(), 0);
733        assert_eq!(hf_ids.token_ids().len(), cached_ids.token_ids().len() + 1);
734        assert_eq!(&hf_ids.token_ids()[1..], cached_ids.token_ids());
735    }
736
737    #[test]
738    fn unoverridden_introspection_methods_use_defaults() {
739        let tokenizer = SegmentTokenizer;
740        assert_eq!(tokenizer.vocab_size(), None);
741        assert!(tokenizer.token_to_id("anything").is_err());
742        assert!(tokenizer.special_token_ids().is_err());
743        assert!(tokenizer.num_special_tokens_added().is_err());
744    }
745}