1mod l1;
61
62use std::sync::Arc;
63
64pub use l1::{CacheEventFn, L1Cache, L1CacheStats};
65
66use crate::{
67 EncodeSegment, Encoding, Result, TokenIdType,
68 traits::{DecodeResult, Decoder, Encoder, Tokenizer},
69};
70
71#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76pub struct CacheTokenUsage {
77 pub cached_tokens: usize,
79 pub uncached_tokens: usize,
81}
82
83pub type CacheTokenUsageFn = Arc<dyn Fn(CacheTokenUsage) + Send + Sync>;
85
86pub struct CachedTokenizer {
91 inner: Arc<dyn Tokenizer>,
92 l1: L1Cache,
93 l1_enabled: bool,
94 extend_on_hit: bool,
97 token_observer: Option<CacheTokenUsageFn>,
99}
100
101impl CachedTokenizer {
102 pub fn new(
118 inner: Arc<dyn Tokenizer>,
119 mut special_tokens: Vec<String>,
120 max_memory_bytes: usize,
121 ) -> Result<Self> {
122 inner.validate_prefix_cache()?;
123 special_tokens.retain(|token| !token.is_empty());
124
125 let overlapping_specials = match l1::first_unsafe_overlap(&special_tokens) {
128 Some((first, second)) => {
129 tracing::warn!(
130 target: "tokenizer",
131 first_token = first,
132 second_token = second,
133 special_token_count = special_tokens.len(),
134 "special tokens can overlap; tokenizer prefix cache disabled"
135 );
136 true
137 }
138 None => false,
139 };
140
141 let l1_enabled = !special_tokens.is_empty() && !overlapping_specials;
142 let cache_tokens = if l1_enabled {
143 special_tokens
144 } else {
145 Vec::new()
146 };
147 Ok(Self {
148 inner,
149 l1: L1Cache::new(max_memory_bytes, cache_tokens),
150 l1_enabled,
151 extend_on_hit: false,
152 token_observer: None,
153 })
154 }
155
156 pub fn with_extend(mut self, enabled: bool) -> Self {
161 self.extend_on_hit = enabled;
162 self
163 }
164
165 pub fn with_observer(mut self, on_hit: CacheEventFn, on_miss: CacheEventFn) -> Self {
169 self.l1.set_observer(on_hit, on_miss);
170 self
171 }
172
173 pub fn with_token_observer(mut self, observer: CacheTokenUsageFn) -> Self {
181 self.token_observer = Some(observer);
182 self
183 }
184
185 fn observe_token_usage(&self, cached_tokens: usize, total_tokens: usize) {
186 if let Some(observer) = &self.token_observer {
187 let uncached_tokens = total_tokens
188 .checked_sub(cached_tokens)
189 .expect("cached token count cannot exceed total token count");
190 observer(CacheTokenUsage {
191 cached_tokens,
192 uncached_tokens,
193 });
194 }
195 }
196
197 pub fn cache_stats(&self) -> L1CacheStats {
199 self.l1.stats()
200 }
201
202 pub fn clear_cache(&self) {
204 self.l1.clear();
205 }
206
207 pub fn inner(&self) -> &Arc<dyn Tokenizer> {
209 &self.inner
210 }
211}
212
213impl Encoder for CachedTokenizer {
214 fn encode(&self, input: &str) -> Result<Encoding> {
215 if !self.l1_enabled {
220 return self.inner.encode(input);
221 }
222
223 if let Some((prefix_tokens, prefix_len, deepest_boundary)) =
224 self.l1.longest_prefix_match(input)
225 {
226 let cached_tokens = prefix_tokens.len();
227 let suffix = &input[prefix_len..];
228 let encoding = if suffix.is_empty() {
229 Encoding::Sp(prefix_tokens.to_vec())
230 } else if self.extend_on_hit {
231 Encoding::Sp(self.l1.extend_after_match(
235 input,
236 prefix_tokens,
237 prefix_len,
238 deepest_boundary,
239 self.inner.as_ref(),
240 )?)
241 } else {
242 let suffix_enc = self.inner.encode(suffix)?;
243 let mut merged: Vec<TokenIdType> =
246 Vec::with_capacity(prefix_tokens.len() + suffix_enc.token_ids().len());
247 merged.extend_from_slice(&prefix_tokens);
248 merged.extend_from_slice(suffix_enc.token_ids());
249 Encoding::Sp(merged)
250 };
251 self.observe_token_usage(cached_tokens, encoding.token_ids().len());
252 return Ok(encoding);
253 }
254
255 let encoding = Encoding::Sp(self.l1.populate_and_encode(input, self.inner.as_ref())?);
261 self.observe_token_usage(0, encoding.token_ids().len());
262 Ok(encoding)
263 }
264
265 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
266 if !self.l1_enabled {
270 return self.inner.encode_batch(inputs);
271 }
272
273 inputs.iter().map(|&i| self.encode(i)).collect()
277 }
278
279 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
280 let encoding = self.inner.encode_segments(segments)?;
284 if self.l1_enabled {
285 self.observe_token_usage(0, encoding.token_ids().len());
286 }
287 Ok(encoding)
288 }
289}
290
291impl Decoder for CachedTokenizer {
292 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
293 self.inner.decode(token_ids, skip_special_tokens)
295 }
296}
297
298impl Tokenizer for CachedTokenizer {}
299
300#[cfg(test)]
301mod tests {
302 use super::*;
303 use crate::HuggingFaceTokenizer;
304 use std::sync::{Mutex, atomic::AtomicU64, atomic::Ordering};
305 use tokenizers::Tokenizer as HfTokenizer;
306
307 struct FailingTokenizer;
308
309 struct SegmentTokenizer;
310
311 impl Encoder for SegmentTokenizer {
312 fn encode(&self, input: &str) -> Result<Encoding> {
313 Ok(Encoding::Sp(vec![input.len() as u32]))
314 }
315
316 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
317 inputs.iter().map(|input| self.encode(input)).collect()
318 }
319
320 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
321 let ids = segments
322 .iter()
323 .flat_map(|segment| [segment.allow_special as u32, segment.text.len() as u32])
324 .collect();
325 Ok(Encoding::Sp(ids))
326 }
327 }
328
329 impl Decoder for SegmentTokenizer {
330 fn decode(
331 &self,
332 _token_ids: &[TokenIdType],
333 _skip_special_tokens: bool,
334 ) -> Result<DecodeResult> {
335 Ok(DecodeResult::Complete(String::new()))
336 }
337 }
338
339 impl Tokenizer for SegmentTokenizer {
340 fn validate_prefix_cache(&self) -> Result<()> {
341 Ok(())
342 }
343 }
344
345 impl Encoder for FailingTokenizer {
346 fn encode(&self, _input: &str) -> Result<Encoding> {
347 Err(anyhow::anyhow!("intentional encode failure"))
348 }
349
350 fn encode_batch(&self, _inputs: &[&str]) -> Result<Vec<Encoding>> {
351 Err(anyhow::anyhow!("intentional encode failure"))
352 }
353 }
354
355 impl Decoder for FailingTokenizer {
356 fn decode(
357 &self,
358 _token_ids: &[TokenIdType],
359 _skip_special_tokens: bool,
360 ) -> Result<DecodeResult> {
361 Err(anyhow::anyhow!("intentional decode failure"))
362 }
363 }
364
365 impl Tokenizer for FailingTokenizer {
366 fn validate_prefix_cache(&self) -> Result<()> {
367 Ok(())
368 }
369 }
370
371 const TINYLLAMA_PATH: &str = concat!(
372 env!("CARGO_MANIFEST_DIR"),
373 "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
374 );
375
376 fn inner() -> Arc<dyn Tokenizer> {
377 Arc::new(HuggingFaceTokenizer::from_file(TINYLLAMA_PATH).expect("load TinyLlama"))
378 }
379
380 fn specials() -> Vec<String> {
381 vec!["<s>".into(), "</s>".into()]
382 }
383
384 fn collect_token_usage(
385 tokenizer: CachedTokenizer,
386 ) -> (CachedTokenizer, Arc<Mutex<Vec<CacheTokenUsage>>>) {
387 let events = Arc::new(Mutex::new(Vec::new()));
388 let observed = events.clone();
389 let tokenizer = tokenizer.with_token_observer(Arc::new(move |usage| {
390 observed.lock().unwrap().push(usage);
391 }));
392 (tokenizer, events)
393 }
394
395 #[test]
396 fn rejects_hf_tokenizer_that_adds_special_tokens() {
397 let tokenizer: Arc<dyn Tokenizer> = Arc::new(
398 HuggingFaceTokenizer::from_file(TINYLLAMA_PATH)
399 .expect("load TinyLlama")
400 .with_options(crate::TokenizerOptions {
401 add_special_tokens: true,
402 }),
403 );
404
405 let result = CachedTokenizer::new(tokenizer, specials(), 4096);
406 let Err(error) = result else {
407 panic!("add_special_tokens=true must be rejected");
408 };
409 assert_eq!(
410 error.to_string(),
411 "HuggingFace tokenizers configured with add_special_tokens=true must remain uncached"
412 );
413 }
414
415 #[test]
416 fn empty_specials_passes_through_correctly() {
417 let tok = inner();
419 let (cached, events) = collect_token_usage(
420 CachedTokenizer::new(tok.clone(), vec![String::new()], 4096)
421 .expect("TinyLlama must support prefix caching"),
422 );
423 let s = "<s>hello world</s>";
424 let a = cached.encode(s).unwrap();
425 let b = tok.encode(s).unwrap();
426 assert_eq!(a.token_ids(), b.token_ids());
427 let stats = cached.cache_stats();
428 assert_eq!(stats.entries, 0);
429 assert_eq!(stats.misses, 0, "empty specials must not increment misses");
430 assert_eq!(stats.hits, 0);
431 assert!(
432 events.lock().unwrap().is_empty(),
433 "empty specials must not emit token usage"
434 );
435 }
436
437 #[test]
438 fn laguna_overlapping_specials_bypass_cache() {
439 const TOKENIZER_JSON: &str = r#"{
440 "version": "1.0",
441 "truncation": null,
442 "padding": null,
443 "added_tokens": [
444 {"id": 0, "content": "<unk>", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
445 {"id": 2, "content": "〈|EOS|〉", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
446 {"id": 14, "content": "〈|", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false},
447 {"id": 15, "content": "|〉", "special": true, "single_word": false, "lstrip": false, "rstrip": false, "normalized": false}
448 ],
449 "normalizer": null,
450 "pre_tokenizer": null,
451 "post_processor": null,
452 "decoder": null,
453 "model": {
454 "type": "WordLevel",
455 "vocab": {"<unk>": 0, "〈|EOS|〉": 2, "〈|": 14, "|〉": 15, "tail": 16},
456 "unk_token": "<unk>"
457 }
458 }"#;
459
460 let hf = HfTokenizer::from_bytes(TOKENIZER_JSON).expect("load test tokenizer");
461 let tok: Arc<dyn Tokenizer> = Arc::new(HuggingFaceTokenizer::from_tokenizer(hf));
462 let overlapping = vec!["〈|EOS|〉".into(), "〈|".into(), "|〉".into()];
463 let (cached, events) = collect_token_usage(
464 CachedTokenizer::new(tok.clone(), overlapping, 4096)
465 .expect("HuggingFace tokenizer must support prefix caching"),
466 );
467
468 let expected = tok.encode("〈|EOS|〉").unwrap();
469 assert_eq!(expected.token_ids(), &[2]);
470 assert_eq!(
471 cached.encode("〈|EOS|〉").unwrap().token_ids(),
472 expected.token_ids()
473 );
474 let stats = cached.cache_stats();
475 assert_eq!(stats.entries, 0);
476 assert_eq!(
477 stats.misses, 0,
478 "overlapping specials must not increment misses"
479 );
480 assert_eq!(stats.hits, 0);
481 assert!(
482 events.lock().unwrap().is_empty(),
483 "overlapping specials must not emit token usage"
484 );
485 }
486
487 #[test]
488 fn segmented_encoding_passes_through_without_caching() {
489 let inner: Arc<dyn Tokenizer> = Arc::new(SegmentTokenizer);
490 let segments = [
491 EncodeSegment::new("<ctl>", true),
492 EncodeSegment::new("user content", false),
493 ];
494 let expected = inner.encode_segments(&segments).unwrap();
495
496 for special_tokens in [Vec::new(), vec!["<ctl>".to_string()]] {
497 let l1_enabled = !special_tokens.is_empty();
498 let (cached, events) = collect_token_usage(
499 CachedTokenizer::new(inner.clone(), special_tokens, 4096)
500 .expect("test tokenizer supports prefix caching"),
501 );
502 let actual = cached.encode_segments(&segments).unwrap();
503
504 assert_eq!(actual.token_ids(), expected.token_ids());
505 let stats = cached.cache_stats();
506 assert_eq!(stats.entries, 0);
507 assert_eq!(stats.hits, 0);
508 assert_eq!(stats.misses, 0);
509 let events = events.lock().unwrap();
510 if l1_enabled {
511 assert_eq!(
512 events.as_slice(),
513 &[CacheTokenUsage {
514 cached_tokens: 0,
515 uncached_tokens: expected.token_ids().len(),
516 }]
517 );
518 } else {
519 assert!(events.is_empty());
520 }
521 }
522 }
523
524 #[test]
525 fn token_observer_reports_full_miss_and_partial_hit_with_and_without_extension() {
526 for extend_on_hit in [false, true] {
527 let tok = inner();
528 let hits = Arc::new(AtomicU64::new(0));
529 let misses = Arc::new(AtomicU64::new(0));
530 let hit_counter = hits.clone();
531 let miss_counter = misses.clone();
532 let cached = CachedTokenizer::new(tok, specials(), 64 * 1024)
533 .expect("TinyLlama must support prefix caching")
534 .with_extend(extend_on_hit)
535 .with_observer(
536 Arc::new(move || {
537 hit_counter.fetch_add(1, Ordering::Relaxed);
538 }),
539 Arc::new(move || {
540 miss_counter.fetch_add(1, Ordering::Relaxed);
541 }),
542 );
543 let (cached, events) = collect_token_usage(cached);
544
545 let shared = "<s>system\nYou are helpful.</s><s>user\n";
546 let first = format!("{shared}First question?</s>");
547 let second = format!("{shared}Second different prompt entirely.</s>");
548
549 let first_encoding = cached.encode(&first).unwrap();
550 let second_encoding = cached.encode(&second).unwrap();
551
552 let events = events.lock().unwrap();
553 assert_eq!(events.len(), 2);
554 assert_eq!(
555 events[0],
556 CacheTokenUsage {
557 cached_tokens: 0,
558 uncached_tokens: first_encoding.token_ids().len(),
559 }
560 );
561 assert!(events[1].cached_tokens > 0);
562 assert!(events[1].uncached_tokens > 0);
563 assert_eq!(
564 events[1].cached_tokens + events[1].uncached_tokens,
565 second_encoding.token_ids().len()
566 );
567 assert_eq!(hits.load(Ordering::Relaxed), 1);
568 assert_eq!(misses.load(Ordering::Relaxed), 1);
569 }
570 }
571
572 #[test]
573 fn token_observer_does_not_report_failed_encodes() {
574 let tokenizer: Arc<dyn Tokenizer> = Arc::new(FailingTokenizer);
575 let (cached, events) = collect_token_usage(
576 CachedTokenizer::new(tokenizer, specials(), 4096)
577 .expect("test tokenizer explicitly supports prefix caching"),
578 );
579
580 assert!(cached.encode("<s>this fails</s>").is_err());
581 assert!(events.lock().unwrap().is_empty());
582 }
583
584 #[test]
585 fn two_turn_chat_correctness_and_hit() {
586 let tok = inner();
587 let cached = CachedTokenizer::new(tok.clone(), specials(), 64 * 1024)
588 .expect("TinyLlama must support prefix caching");
589
590 let template = "<s>system\nYou are helpful.</s><s>user\n";
591 let first = format!("{template}First question?</s>");
592 let second = format!("{template}Second different prompt entirely.</s>");
593
594 let _ = cached.encode(&first).unwrap();
596
597 let cached_second = cached.encode(&second).unwrap();
599 let plain_second = tok.encode(&second).unwrap();
600 assert_eq!(
601 cached_second.token_ids(),
602 plain_second.token_ids(),
603 "cached encode must equal plain encode for second turn"
604 );
605
606 let stats = cached.cache_stats();
607 assert!(stats.hits >= 1, "expected L1 hit on second request");
608 }
609
610 #[test]
611 fn decode_passes_through() {
612 let tok = inner();
613 let cached = CachedTokenizer::new(tok.clone(), specials(), 4096)
614 .expect("TinyLlama must support prefix caching");
615 let enc = cached.encode("<s>hello</s>").unwrap();
616 let direct = tok.decode(enc.token_ids(), false).unwrap();
617 let through = cached.decode(enc.token_ids(), false).unwrap();
618 assert_eq!(direct, through);
619 }
620
621 #[test]
622 fn encode_batch_uses_cache() {
623 let tok = inner();
624 let (cached, events) = collect_token_usage(
625 CachedTokenizer::new(tok.clone(), specials(), 64 * 1024)
626 .expect("TinyLlama must support prefix caching"),
627 );
628 let shared = "<s>system\nShared persona.</s><s>user\n";
629 let inputs = [
630 format!("{shared}q1</s>"),
631 format!("{shared}q2</s>"),
632 format!("{shared}q3</s>"),
633 ];
634 let refs: Vec<&str> = inputs.iter().map(String::as_str).collect();
635 let outs = cached.encode_batch(&refs).unwrap();
636 assert_eq!(outs.len(), 3);
637 let events = events.lock().unwrap();
638 assert_eq!(events.len(), outs.len());
639 for (event, output) in events.iter().zip(&outs) {
640 assert_eq!(
641 event.cached_tokens + event.uncached_tokens,
642 output.token_ids().len()
643 );
644 }
645 assert_eq!(events[0].cached_tokens, 0);
646 assert!(events[1..].iter().all(|event| event.cached_tokens > 0));
647 assert!(cached.cache_stats().hits >= 2, "expected hits on q2 and q3");
649 }
650}