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 decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
322 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 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 let _ = cached.encode(&first).unwrap();
645
646 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 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}