1mod 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
86pub struct CacheTokenUsage {
87 pub cached_tokens: usize,
89 pub uncached_tokens: usize,
91}
92
93pub type CacheTokenUsageFn = Arc<dyn Fn(CacheTokenUsage) + Send + Sync>;
95
96pub struct CachedTokenizer {
101 inner: Arc<dyn Tokenizer>,
102 l1: L1Cache,
103 l1_enabled: bool,
104 extend_on_hit: bool,
105 token_observer: Option<CacheTokenUsageFn>,
107}
108
109impl CachedTokenizer {
110 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 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 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 pub fn with_extend(mut self, enabled: bool) -> Self {
202 self.extend_on_hit = enabled;
203 self
204 }
205
206 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 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 pub fn cache_stats(&self) -> L1CacheStats {
242 if self.l1_enabled {
243 self.l1.stats()
244 } else {
245 L1CacheStats::default()
246 }
247 }
248
249 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 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 if !self.l1_enabled {
299 return self.inner.encode_batch(inputs);
300 }
301
302 inputs.iter().map(|&i| self.encode(i)).collect()
306 }
307
308 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
309 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 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 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 let _ = cached.encode(&first).unwrap();
650
651 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 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}