1use std::{
27 hash::BuildHasherDefault,
28 mem::size_of_val,
29 sync::{
30 Arc,
31 atomic::{AtomicU64, Ordering},
32 },
33};
34
35use aho_corasick::AhoCorasick;
36use moka::sync::Cache;
37use rustc_hash::FxHasher;
38
39use crate::{TokenIdType, traits::Encoder};
40
41type Blake3Hash = [u8; 32];
43
44type PrefixHasher = BuildHasherDefault<FxHasher>;
47
48type PrefixCache = Cache<Blake3Hash, CachedPrefix, PrefixHasher>;
50
51#[derive(Clone)]
52struct CachedPrefix {
53 namespace: Blake3Hash,
54 tokens: Arc<[TokenIdType]>,
55}
56
57impl CachedPrefix {
58 fn weight(&self) -> u32 {
59 size_of_val(self.tokens.as_ref()).min(u32::MAX as usize) as u32
60 }
61}
62
63#[derive(Clone)]
69pub struct SharedTokenizerCache {
70 cache: PrefixCache,
71}
72
73#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
75pub struct SharedTokenizerCacheStats {
76 pub entries: usize,
77 pub memory_bytes: usize,
78}
79
80impl SharedTokenizerCache {
81 pub fn new(max_memory_bytes: usize) -> Self {
82 Self {
83 cache: Cache::builder()
84 .max_capacity(max_memory_bytes as u64)
85 .weigher(|_key: &Blake3Hash, entry: &CachedPrefix| entry.weight())
86 .build_with_hasher(PrefixHasher::default()),
87 }
88 }
89
90 pub fn max_memory_bytes(&self) -> usize {
92 self.cache.policy().max_capacity().expect("capacity is set") as usize
93 }
94
95 pub fn stats(&self) -> SharedTokenizerCacheStats {
97 self.cache.run_pending_tasks();
98 SharedTokenizerCacheStats {
99 entries: self.cache.entry_count() as usize,
100 memory_bytes: self.cache.weighted_size() as usize,
101 }
102 }
103
104 fn namespace_stats(&self, namespace: &Blake3Hash) -> SharedTokenizerCacheStats {
105 self.cache.run_pending_tasks();
106 let mut stats = SharedTokenizerCacheStats::default();
107 for (_, entry) in &self.cache {
108 if &entry.namespace == namespace {
109 stats.entries += 1;
110 stats.memory_bytes += entry.weight() as usize;
111 }
112 }
113 stats
114 }
115}
116
117fn namespace_hasher(namespace: &[u8]) -> blake3::Hasher {
118 let mut hasher = blake3::Hasher::new();
119 hasher.update(&(namespace.len() as u64).to_le_bytes());
121 hasher.update(namespace);
122 hasher
123}
124
125pub(super) struct PrefixMatch {
127 pub(super) tokens: Arc<[TokenIdType]>,
128 pub(super) prefix_len: usize,
129 deepest_boundary: usize,
130 deepest_hash: Option<Blake3Hash>,
131}
132
133pub(super) enum PrefixLookup {
135 Hit(PrefixMatch),
136 Miss(Vec<(usize, Blake3Hash)>),
137}
138
139fn hash_prefixes<'a>(
141 mut hasher: blake3::Hasher,
142 input: &'a str,
143 boundaries: &'a [usize],
144) -> impl Iterator<Item = (usize, Blake3Hash)> + 'a {
145 let mut last_pos = 0;
146 boundaries.iter().map(move |&boundary_pos| {
147 hasher.update(&input.as_bytes()[last_pos..boundary_pos]);
148 last_pos = boundary_pos;
149 (boundary_pos, *hasher.finalize().as_bytes())
150 })
151}
152
153fn boundaries_with(text: &str, matcher: &AhoCorasick) -> Vec<usize> {
162 let mut boundaries: Vec<usize> = matcher
163 .find_overlapping_iter(text)
164 .map(|m| m.end())
165 .filter(|&end| end < text.len())
166 .collect();
167 boundaries.sort_unstable();
168 boundaries.dedup();
169 boundaries
170}
171
172fn has_nontrivial_self_overlap(token: &str) -> bool {
173 let bytes = token.as_bytes();
174 (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
175}
176
177fn tokens_can_overlap(a: &str, b: &str) -> bool {
178 if a.contains(b) || b.contains(a) {
179 return true;
180 }
181
182 let a = a.as_bytes();
183 let b = b.as_bytes();
184 let max_overlap = a.len().min(b.len());
185 (1..max_overlap).any(|overlap| {
186 a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
187 })
188}
189
190pub(super) fn first_unsafe_overlap(special_tokens: &[String]) -> Option<(&str, &str)> {
198 for (index, token) in special_tokens.iter().enumerate() {
199 if token.is_empty() {
200 continue;
201 }
202 if has_nontrivial_self_overlap(token) {
203 return Some((token, token));
204 }
205 for other in &special_tokens[index + 1..] {
206 if !other.is_empty() && token != other && tokens_can_overlap(token, other) {
207 return Some((token, other));
208 }
209 }
210 }
211
212 None
213}
214
215#[cfg(test)]
218fn find_special_token_boundaries(text: &str, special_tokens: &[&str]) -> Vec<usize> {
219 if special_tokens.is_empty() {
220 return Vec::new();
221 }
222 let matcher = AhoCorasick::new(special_tokens)
223 .expect("special tokens form a valid Aho-Corasick automaton");
224 boundaries_with(text, &matcher)
225}
226
227pub type CacheEventFn = Arc<dyn Fn() + Send + Sync>;
231
232pub struct L1Cache {
236 cache: SharedTokenizerCache,
237 shared: bool,
238 namespace: Vec<u8>,
239 namespace_hash: Blake3Hash,
240 matcher: Option<AhoCorasick>,
243 hits: AtomicU64,
244 misses: AtomicU64,
245 on_hit: Option<CacheEventFn>,
246 on_miss: Option<CacheEventFn>,
247}
248
249impl L1Cache {
250 pub fn new(max_memory: usize, special_tokens: Vec<String>) -> Self {
253 Self {
254 shared: false,
255 ..Self::new_with_cache(SharedTokenizerCache::new(max_memory), special_tokens, b"")
256 }
257 }
258
259 pub fn new_with_cache(
261 cache: SharedTokenizerCache,
262 mut special_tokens: Vec<String>,
263 namespace: &[u8],
264 ) -> Self {
265 special_tokens.retain(|token| !token.is_empty());
266
267 let matcher = (!special_tokens.is_empty()).then(|| {
269 AhoCorasick::new(&special_tokens)
270 .expect("special tokens form a valid Aho-Corasick automaton")
271 });
272
273 Self {
274 cache,
275 shared: true,
276 namespace: namespace.to_vec(),
277 namespace_hash: *namespace_hasher(namespace).finalize().as_bytes(),
278 matcher,
279 hits: AtomicU64::new(0),
280 misses: AtomicU64::new(0),
281 on_hit: None,
282 on_miss: None,
283 }
284 }
285
286 fn hasher(&self) -> blake3::Hasher {
287 namespace_hasher(&self.namespace)
288 }
289
290 fn hash_prefix(&self, prefix: &[u8]) -> Blake3Hash {
291 let mut hasher = self.hasher();
292 hasher.update(prefix);
293 *hasher.finalize().as_bytes()
294 }
295
296 fn insert(&self, hash: Blake3Hash, tokens: Arc<[TokenIdType]>) {
297 self.cache.cache.insert(
298 hash,
299 CachedPrefix {
300 namespace: self.namespace_hash,
301 tokens,
302 },
303 );
304 }
305
306 pub fn set_observer(&mut self, on_hit: CacheEventFn, on_miss: CacheEventFn) {
308 self.on_hit = Some(on_hit);
309 self.on_miss = Some(on_miss);
310 }
311
312 fn boundaries(&self, text: &str) -> Vec<usize> {
316 match &self.matcher {
317 Some(matcher) => boundaries_with(text, matcher),
318 None => Vec::new(),
319 }
320 }
321
322 pub fn longest_prefix_match(&self, input: &str) -> Option<(Arc<[TokenIdType]>, usize, usize)> {
329 match self.lookup_prefix(input) {
330 PrefixLookup::Hit(matched) => {
331 Some((matched.tokens, matched.prefix_len, matched.deepest_boundary))
332 }
333 PrefixLookup::Miss(_) => None,
334 }
335 }
336
337 pub(super) fn lookup_prefix(&self, input: &str) -> PrefixLookup {
340 let boundaries = self.boundaries(input);
341
342 if boundaries.is_empty() {
343 self.misses.fetch_add(1, Ordering::Relaxed);
344 if let Some(cb) = &self.on_miss {
345 cb();
346 }
347 return PrefixLookup::Miss(Vec::new());
348 }
349
350 let prefix_hashes: Vec<_> = hash_prefixes(self.hasher(), input, &boundaries).collect();
351
352 for &(boundary_pos, hash_bytes) in prefix_hashes.iter().rev() {
353 if let Some(entry) = self.cache.cache.get(&hash_bytes) {
354 self.hits.fetch_add(1, Ordering::Relaxed);
355 if let Some(cb) = &self.on_hit {
356 cb();
357 }
358 let &(deepest_boundary, deepest_hash) =
360 prefix_hashes.last().expect("prefix hashes is non-empty");
361 return PrefixLookup::Hit(PrefixMatch {
362 tokens: entry.tokens,
363 prefix_len: boundary_pos,
364 deepest_boundary,
365 deepest_hash: Some(deepest_hash),
366 });
367 }
368 }
369
370 self.misses.fetch_add(1, Ordering::Relaxed);
371 if let Some(cb) = &self.on_miss {
372 cb();
373 }
374 PrefixLookup::Miss(prefix_hashes)
375 }
376
377 pub fn insert_at_boundaries<E: Encoder + ?Sized>(
385 &self,
386 input: &str,
387 tokenizer: &E,
388 ) -> anyhow::Result<()> {
389 let boundaries = self.boundaries(input);
390 if boundaries.is_empty() {
391 return Ok(());
392 }
393 self.populate_boundaries(
394 input,
395 hash_prefixes(self.hasher(), input, &boundaries),
396 tokenizer,
397 )?;
398 Ok(())
399 }
400
401 pub fn populate_and_encode<E: Encoder + ?Sized>(
410 &self,
411 input: &str,
412 tokenizer: &E,
413 ) -> anyhow::Result<Vec<TokenIdType>> {
414 let boundaries = self.boundaries(input);
415 self.populate_and_encode_with_hashes(
416 input,
417 hash_prefixes(self.hasher(), input, &boundaries),
418 tokenizer,
419 )
420 }
421
422 pub(super) fn populate_and_encode_with_hashes<E: Encoder + ?Sized>(
426 &self,
427 input: &str,
428 prefix_hashes: impl Iterator<Item = (usize, Blake3Hash)>,
429 tokenizer: &E,
430 ) -> anyhow::Result<Vec<TokenIdType>> {
431 let (mut running, tail_start) =
432 self.populate_boundaries(input, prefix_hashes, tokenizer)?;
433 if tail_start == 0 {
434 return Ok(tokenizer.encode(input)?.token_ids().to_vec());
435 }
436
437 let tail = tokenizer.encode(&input[tail_start..])?;
438 running.extend_from_slice(tail.token_ids());
439 Ok(running)
440 }
441
442 fn populate_boundaries<E: Encoder + ?Sized>(
445 &self,
446 input: &str,
447 prefix_hashes: impl Iterator<Item = (usize, Blake3Hash)>,
448 tokenizer: &E,
449 ) -> anyhow::Result<(Vec<TokenIdType>, usize)> {
450 #[cfg(debug_assertions)]
451 let mut validation_hasher = self.hasher();
452 let mut running_tokens: Vec<TokenIdType> = Vec::new();
453 let mut last_pos = 0;
454
455 for (boundary_pos, hash_bytes) in prefix_hashes {
456 #[cfg(debug_assertions)]
457 {
458 validation_hasher.update(&input.as_bytes()[last_pos..boundary_pos]);
459 debug_assert_eq!(hash_bytes, *validation_hasher.finalize().as_bytes());
460 }
461
462 let seg = tokenizer.encode(&input[last_pos..boundary_pos])?;
466 running_tokens.extend_from_slice(seg.token_ids());
467
468 let prefix_tokens: Arc<[TokenIdType]> = running_tokens.as_slice().into();
469 self.insert(hash_bytes, prefix_tokens);
470
471 last_pos = boundary_pos;
472 }
473
474 Ok((running_tokens, last_pos))
475 }
476
477 pub fn extend_after_match<E: Encoder + ?Sized>(
494 &self,
495 input: &str,
496 prefix_tokens: Arc<[TokenIdType]>,
497 prefix_len: usize,
498 deepest_boundary: usize,
499 tokenizer: &E,
500 ) -> anyhow::Result<Vec<TokenIdType>> {
501 self.extend_after_match_with_hash(
502 input,
503 PrefixMatch {
504 tokens: prefix_tokens,
505 prefix_len,
506 deepest_boundary,
507 deepest_hash: None,
508 },
509 tokenizer,
510 )
511 }
512
513 pub(super) fn extend_after_match_with_hash<E: Encoder + ?Sized>(
517 &self,
518 input: &str,
519 matched: PrefixMatch,
520 tokenizer: &E,
521 ) -> anyhow::Result<Vec<TokenIdType>> {
522 let PrefixMatch {
523 tokens: prefix_tokens,
524 prefix_len,
525 deepest_boundary,
526 deepest_hash,
527 } = matched;
528 let deepest = (deepest_boundary > prefix_len).then_some(deepest_boundary);
530
531 let Some(deepest) = deepest else {
532 let suffix_enc = tokenizer.encode(&input[prefix_len..])?;
533 let mut merged = Vec::with_capacity(prefix_tokens.len() + suffix_enc.token_ids().len());
535 merged.extend_from_slice(&prefix_tokens);
536 merged.extend_from_slice(suffix_enc.token_ids());
537 return Ok(merged);
538 };
539
540 let seg_a = tokenizer.encode(&input[prefix_len..deepest])?;
542 let seg_b = tokenizer.encode(&input[deepest..])?;
543 let mut cumulative = Vec::with_capacity(
544 prefix_tokens.len() + seg_a.token_ids().len() + seg_b.token_ids().len(),
545 );
546 cumulative.extend_from_slice(&prefix_tokens);
547 cumulative.extend_from_slice(seg_a.token_ids());
548
549 let hash_bytes =
550 deepest_hash.unwrap_or_else(|| self.hash_prefix(&input.as_bytes()[..deepest]));
551 debug_assert_eq!(hash_bytes, self.hash_prefix(&input.as_bytes()[..deepest]));
552
553 let tokens: Arc<[TokenIdType]> = cumulative.as_slice().into();
555 self.insert(hash_bytes, tokens);
556
557 cumulative.extend_from_slice(seg_b.token_ids());
558 Ok(cumulative)
559 }
560
561 fn storage_stats(&self) -> SharedTokenizerCacheStats {
562 if self.shared {
563 self.cache.namespace_stats(&self.namespace_hash)
564 } else {
565 self.cache.stats()
566 }
567 }
568
569 pub fn len(&self) -> usize {
572 self.storage_stats().entries
573 }
574
575 pub fn is_empty(&self) -> bool {
576 if !self.shared {
577 return self.len() == 0;
578 }
579 self.cache.cache.run_pending_tasks();
580 !self
581 .cache
582 .cache
583 .iter()
584 .any(|(_, entry)| entry.namespace == self.namespace_hash)
585 }
586
587 pub fn stats(&self) -> L1CacheStats {
588 let storage = self.storage_stats();
589 let hits = self.hits.load(Ordering::Relaxed);
590 let misses = self.misses.load(Ordering::Relaxed);
591 let total_requests = hits + misses;
592
593 L1CacheStats {
594 hits,
595 misses,
596 entries: storage.entries,
597 memory_bytes: storage.memory_bytes,
598 hit_rate: if total_requests > 0 {
599 hits as f64 / total_requests as f64
600 } else {
601 0.0
602 },
603 }
604 }
605}
606
607#[derive(Debug, Clone, Default)]
608pub struct L1CacheStats {
609 pub hits: u64,
610 pub misses: u64,
611 pub entries: usize,
612 pub memory_bytes: usize,
613 pub hit_rate: f64,
614}
615
616#[cfg(test)]
617mod tests {
618 use std::sync::Arc;
619
620 use super::*;
621 use crate::{HuggingFaceTokenizer, traits::Tokenizer};
622
623 const TINYLLAMA_PATH: &str = concat!(
626 env!("CARGO_MANIFEST_DIR"),
627 "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
628 );
629
630 const SPECIALS: &[&str] = &["<s>", "</s>"];
631
632 fn load_tokenizer() -> Arc<dyn Tokenizer> {
633 Arc::new(HuggingFaceTokenizer::from_file(TINYLLAMA_PATH).expect("load TinyLlama"))
634 }
635
636 fn test_cache(max_memory: usize) -> L1Cache {
638 L1Cache::new(
639 max_memory,
640 SPECIALS.iter().map(|s| (*s).to_string()).collect(),
641 )
642 }
643
644 #[test]
645 fn prefix_hash_length_delimits_the_namespace() {
646 let storage = SharedTokenizerCache::new(1024);
647 let mut hashes = Vec::new();
648 for (namespace, prefix) in [("", "abc"), ("a", "bc"), ("ab", "c")] {
649 let cache = L1Cache::new_with_cache(storage.clone(), vec![], namespace.as_bytes());
650 let bytes = [
651 (namespace.len() as u64).to_le_bytes().as_slice(),
652 namespace.as_bytes(),
653 prefix.as_bytes(),
654 ]
655 .concat();
656 let expected = *blake3::hash(&bytes).as_bytes();
657 assert_eq!(cache.hash_prefix(prefix.as_bytes()), expected);
658 assert!(!hashes.contains(&expected));
659 hashes.push(expected);
660 }
661 }
662
663 #[test]
664 fn boundaries_are_after_each_special_token_occurrence() {
665 let input = "<s>system\nHi</s><s>user\nHello</s>";
666 let bounds = find_special_token_boundaries(input, SPECIALS);
667 assert_eq!(bounds.len(), 3);
669 for w in bounds.windows(2) {
670 assert!(w[0] < w[1], "boundaries must be strictly increasing");
671 }
672 assert!(bounds.iter().all(|&b| b < input.len()));
673 }
674
675 #[test]
676 fn no_special_tokens_yields_no_boundaries() {
677 assert!(find_special_token_boundaries("plain text", &[]).is_empty());
678 }
679
680 #[test]
681 fn unsafe_overlap_detects_containment_crossing_and_self_overlap() {
682 let cases = [
683 (vec!["〈|", "〈|EOS|〉"], Some(("〈|", "〈|EOS|〉"))),
684 (vec!["ab", "bc"], Some(("ab", "bc"))),
685 (vec!["|◊|"], Some(("|◊|", "|◊|"))),
686 (vec!["<s>", "<s>"], None),
687 ];
688
689 for (tokens, expected) in cases {
690 let tokens: Vec<String> = tokens.into_iter().map(String::from).collect();
691 assert_eq!(first_unsafe_overlap(&tokens), expected);
692 }
693 }
694
695 #[test]
696 fn llama_numbered_special_tokens_do_not_trigger_overlap_guard() {
697 let mut llama: Vec<String> = [
698 "<|begin_of_text|>",
699 "<|end_of_text|>",
700 "<|start_header_id|>",
701 "<|end_header_id|>",
702 "<|eot_id|>",
703 ]
704 .into_iter()
705 .map(String::from)
706 .collect();
707 llama.extend((0..251).map(|id| format!("<|reserved_special_token_{id}|>")));
708
709 assert_eq!(first_unsafe_overlap(&llama), None);
710 }
711
712 #[test]
713 fn insert_then_lookup_finds_shared_prefix() {
714 let cache = test_cache(1024 * 1024);
715 let tokenizer = load_tokenizer();
716
717 let warm = "<s>system\nYou are helpful.</s><s>user\nHi</s>";
718 cache
719 .insert_at_boundaries(warm, tokenizer.as_ref())
720 .unwrap();
721 assert!(!cache.is_empty());
722
723 let target = "<s>system\nYou are helpful.</s><s>user\nDifferent question</s>";
724 let (tokens, offset, _deepest) = cache
725 .longest_prefix_match(target)
726 .expect("shared prefix should match");
727 assert!(offset > 0);
728 assert!(!tokens.is_empty());
729 }
730
731 #[test]
732 fn miss_increments_misses_counter() {
733 let cache = test_cache(1024 * 1024);
734 assert!(
735 cache
736 .longest_prefix_match("plain text no specials")
737 .is_none()
738 );
739 assert_eq!(cache.stats().misses, 1);
740 }
741
742 #[test]
743 fn hit_increments_hits_counter() {
744 let cache = test_cache(1024 * 1024);
745 let tokenizer = load_tokenizer();
746 let warm = "<s>system\nA.</s><s>user\nB</s>";
747 cache
748 .insert_at_boundaries(warm, tokenizer.as_ref())
749 .unwrap();
750 let _ = cache.longest_prefix_match(warm);
751 assert!(cache.stats().hits >= 1);
752 }
753
754 #[test]
755 fn merge_invariant_holds_against_uncached_encode() {
756 let cache = test_cache(1024 * 1024);
760 let tokenizer = load_tokenizer();
761
762 let template = "<s>system\nYou are helpful.</s><s>user\n";
763 let warm = format!("{template}First.</s>");
764 cache
765 .insert_at_boundaries(&warm, tokenizer.as_ref())
766 .unwrap();
767
768 let target = format!("{template}A completely different second question.</s>");
769 let (prefix_tokens, prefix_len, _deepest) = cache
770 .longest_prefix_match(&target)
771 .expect("should find prefix");
772
773 let suffix = &target[prefix_len..];
774 let suffix_enc = tokenizer.encode(suffix).unwrap();
775 let mut merged = prefix_tokens.to_vec();
777 merged.extend_from_slice(suffix_enc.token_ids());
778
779 let plain = tokenizer.encode(&target).unwrap();
780 assert_eq!(
781 merged,
782 plain.token_ids(),
783 "merged tokens must equal plain encode"
784 );
785 }
786
787 #[test]
788 fn eviction_respects_memory_budget() {
789 let cache = test_cache(4 * 1024);
791 let tokenizer = load_tokenizer();
792 for i in 0..50 {
793 let input =
794 format!("<s>system\nPersona {i} chatty.</s><s>user\nTurn {i} content here.</s>");
795 cache
796 .insert_at_boundaries(&input, tokenizer.as_ref())
797 .unwrap();
798 }
799 let stats = cache.stats();
800 assert!(
801 stats.memory_bytes <= 4 * 1024,
802 "memory_bytes={} exceeds budget",
803 stats.memory_bytes
804 );
805 }
806
807 #[test]
808 fn concurrent_inserts_and_lookups_do_not_corrupt() {
809 use std::thread;
810
811 let cache = Arc::new(test_cache(1024 * 1024));
812 let tokenizer = load_tokenizer();
813
814 let mut handles = vec![];
815 for i in 0..10 {
816 let cache_c = cache.clone();
817 let tok = tokenizer.clone();
818 handles.push(thread::spawn(move || {
819 let input = format!("<s>system\nThread {i}.</s><s>user\nThread {i} body.</s>");
820 cache_c.insert_at_boundaries(&input, tok.as_ref()).unwrap();
821 let r = cache_c.longest_prefix_match(&input);
822 assert!(r.is_some(), "thread {i} expected match after insert");
823 }));
824 }
825 for h in handles {
826 h.join().unwrap();
827 }
828 assert!(cache.stats().memory_bytes > 0);
829 assert!(cache.stats().hits >= 10);
830 }
831
832 fn growing_chat_turns(n: usize) -> Vec<String> {
838 let mut convo = String::from("<s>system\nYou are a helpful assistant.</s>");
839 let mut turns = Vec::with_capacity(n);
840 for i in 0..n {
841 convo.push_str(&format!(
842 "<s>user\nQuestion {i} please answer it.</s><s>assistant\nDetailed answer {i} follows here.</s>"
843 ));
844 turns.push(format!("{convo}<s>user\nFollow-up {i}"));
845 }
846 turns
847 }
848
849 #[test]
850 fn extend_on_hit_advances_match_depth_each_turn() {
851 let tok = load_tokenizer();
855 let turns = growing_chat_turns(5);
856
857 let off = test_cache(8 * 1024 * 1024);
859 off.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
860 let pinned = off.longest_prefix_match(&turns[1]).expect("hit").1;
861 for t in &turns[1..] {
862 let (_toks, offset, _deepest) = off.longest_prefix_match(t).expect("hit");
863 assert_eq!(
864 offset, pinned,
865 "extend-off offset must stay pinned at turn-1 depth"
866 );
867 }
868
869 let on = test_cache(8 * 1024 * 1024);
871 on.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
872 let mut prev = 0usize;
873 for (i, t) in turns.iter().enumerate().skip(1) {
874 let (prefix_tokens, offset, deepest) = on.longest_prefix_match(t).expect("hit");
875 assert!(
876 offset > prev,
877 "turn {i}: extend-on offset {offset} must exceed previous {prev}"
878 );
879 prev = offset;
880
881 let merged = on
883 .extend_after_match(t, prefix_tokens, offset, deepest, tok.as_ref())
884 .unwrap();
885 let plain = tok.encode(t).unwrap();
886 assert_eq!(
887 merged,
888 plain.token_ids(),
889 "turn {i}: extend merge must equal plain encode"
890 );
891 }
892
893 assert!(
894 prev > pinned,
895 "extend-on frontier ({prev}) must reach deeper than pinned extend-off depth ({pinned})"
896 );
897 }
898
899 #[test]
900 fn extend_on_hit_respects_budget_and_stays_correct() {
901 let tok = load_tokenizer();
904 let cache = test_cache(4 * 1024);
905 let turns = growing_chat_turns(20);
906 cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
907
908 for t in &turns[1..] {
909 let merged = match cache.longest_prefix_match(t) {
910 Some((prefix_tokens, offset, deepest)) => cache
911 .extend_after_match(t, prefix_tokens, offset, deepest, tok.as_ref())
912 .unwrap(),
913 None => {
914 let enc = tok.encode(t).unwrap();
916 cache.insert_at_boundaries(t, tok.as_ref()).unwrap();
917 enc.token_ids().to_vec()
918 }
919 };
920 let plain = tok.encode(t).unwrap();
921 assert_eq!(
922 merged,
923 plain.token_ids(),
924 "encode must stay correct under eviction pressure"
925 );
926 assert!(
927 cache.stats().memory_bytes <= 4 * 1024,
928 "memory_bytes={} exceeds budget",
929 cache.stats().memory_bytes
930 );
931 }
932 }
933
934 #[test]
935 fn concurrent_extend_on_hit_does_not_corrupt() {
936 use std::thread;
937
938 let tok = load_tokenizer();
939 let cache = Arc::new(test_cache(8 * 1024 * 1024));
940 let turns = growing_chat_turns(8);
941 cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
943
944 let mut handles = vec![];
945 for _ in 0..8 {
946 let cache_c = cache.clone();
947 let tok_c = tok.clone();
948 let turns_c = turns.clone();
949 handles.push(thread::spawn(move || {
950 for t in &turns_c[1..] {
951 if let PrefixLookup::Hit(matched) = cache_c.lookup_prefix(t) {
952 let merged = cache_c
953 .extend_after_match_with_hash(t, matched, tok_c.as_ref())
954 .unwrap();
955 let plain = tok_c.encode(t).unwrap();
956 assert_eq!(
957 merged,
958 plain.token_ids(),
959 "concurrent extend must stay correct"
960 );
961 }
962 }
963 }));
964 }
965 for h in handles {
966 h.join().unwrap();
967 }
968 assert!(cache.stats().memory_bytes > 0);
969 }
970
971 #[test]
972 fn extend_after_match_persists_correct_deepest_entry() {
973 let tok = load_tokenizer();
974 for unicode in [false, true] {
975 let turns: Vec<_> = growing_chat_turns(3)
976 .into_iter()
977 .map(|t| {
978 if unicode {
979 t.replace("system", "system 世界 🦀")
980 } else {
981 t
982 }
983 })
984 .collect();
985
986 let cache = test_cache(8 * 1024 * 1024);
987 cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
988
989 let PrefixLookup::Hit(matched) = cache.lookup_prefix(&turns[1]) else {
990 panic!("partial hit on turns[1]");
991 };
992 let prefix_len = matched.prefix_len;
993 let deepest_boundary = matched.deepest_boundary;
994 assert!(deepest_boundary > prefix_len);
995 assert_eq!(
996 matched.deepest_hash,
997 Some(cache.hash_prefix(&turns[1].as_bytes()[..deepest_boundary]))
998 );
999 assert_ne!(
1000 matched.deepest_hash,
1001 Some(cache.hash_prefix(&turns[1].as_bytes()[..prefix_len]))
1002 );
1003 let entries_before = cache.stats().entries;
1004
1005 let _merged = cache
1006 .extend_after_match_with_hash(&turns[1], matched, tok.as_ref())
1007 .unwrap();
1008
1009 assert_eq!(
1010 cache.stats().entries,
1011 entries_before + 1,
1012 "extend must persist exactly one (deepest) entry"
1013 );
1014
1015 let deepest = find_special_token_boundaries(&turns[1], SPECIALS)
1016 .into_iter()
1017 .rev()
1018 .find(|&b| b > prefix_len)
1019 .expect("a deeper boundary must exist in the appended turn");
1020 assert_eq!(
1021 deepest_boundary, deepest,
1022 "longest_prefix_match must return the deepest boundary used by extend"
1023 );
1024
1025 let (saved_tokens, saved_offset, _deepest) = cache
1026 .longest_prefix_match(&turns[1])
1027 .expect("hit after extend");
1028 assert_eq!(
1029 saved_offset, deepest,
1030 "lookup must now hit at the just-saved deepest boundary"
1031 );
1032 let expected = tok.encode(&turns[1][..deepest]).unwrap();
1033 assert_eq!(
1034 &*saved_tokens,
1035 expected.token_ids(),
1036 "persisted entry tokens must equal the uncached encode of the cached prefix"
1037 );
1038 }
1039 }
1040
1041 #[test]
1042 fn extend_without_deeper_boundary_does_not_insert() {
1043 let tok = load_tokenizer();
1044 let cache = test_cache(8 * 1024 * 1024);
1045 cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
1046 for input in ["<s>世界", "<s>世界</s>"] {
1047 let PrefixLookup::Hit(matched) = cache.lookup_prefix(input) else {
1048 panic!("expected hit");
1049 };
1050 assert_eq!(matched.prefix_len, matched.deepest_boundary);
1051 let entries = cache.len();
1052 let merged = cache
1053 .extend_after_match_with_hash(input, matched, tok.as_ref())
1054 .unwrap();
1055 assert_eq!(merged, tok.encode(input).unwrap().token_ids());
1056 assert_eq!(cache.len(), entries);
1057 }
1058 }
1059
1060 struct FailAt {
1061 call: std::sync::atomic::AtomicUsize,
1062 fail_at: usize,
1063 }
1064 impl Encoder for FailAt {
1065 fn encode(&self, _: &str) -> crate::Result<crate::Encoding> {
1066 if self.call.fetch_add(1, Ordering::Relaxed) == self.fail_at {
1067 anyhow::bail!("suffix failed");
1068 }
1069 Ok(crate::Encoding::Sp(vec![1]))
1070 }
1071 fn encode_batch(&self, inputs: &[&str]) -> crate::Result<Vec<crate::Encoding>> {
1072 inputs.iter().map(|s| self.encode(s)).collect()
1073 }
1074 }
1075
1076 #[test]
1077 fn hash_reuse_does_not_insert_when_either_suffix_encode_fails() {
1078 let tok = load_tokenizer();
1079 let cache = test_cache(8 * 1024 * 1024);
1080 cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
1081 let input = "<s>世界</s><s>tail";
1082 for fail_at in [0, 1] {
1083 let PrefixLookup::Hit(matched) = cache.lookup_prefix(input) else {
1084 panic!("expected hit");
1085 };
1086 let entries = cache.len();
1087 let error = cache
1088 .extend_after_match_with_hash(
1089 input,
1090 matched,
1091 &FailAt {
1092 call: 0.into(),
1093 fail_at,
1094 },
1095 )
1096 .unwrap_err();
1097 assert_eq!(error.to_string(), "suffix failed");
1098 assert_eq!(cache.len(), entries);
1099 }
1100 }
1101
1102 #[test]
1103 #[cfg(debug_assertions)]
1104 fn reused_hashes_reject_mismatched_input_before_insertion() {
1105 use std::panic::{AssertUnwindSafe, catch_unwind};
1106
1107 let tok = load_tokenizer();
1108 let cache = test_cache(8 * 1024 * 1024);
1109 cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
1110 let PrefixLookup::Hit(matched) = cache.lookup_prefix("<s>世界</s><s>tail") else {
1111 panic!("expected hit");
1112 };
1113 let entries = cache.len();
1114 assert!(
1115 catch_unwind(AssertUnwindSafe(|| {
1116 cache.extend_after_match_with_hash("<s>日本</s><s>tail", matched, tok.as_ref())
1117 }))
1118 .is_err()
1119 );
1120 assert_eq!(cache.len(), entries);
1121
1122 let cache = test_cache(8 * 1024 * 1024);
1123 let PrefixLookup::Miss(hashes) = cache.lookup_prefix("a<s>世界</s>tail") else {
1124 panic!("expected miss");
1125 };
1126 assert!(
1127 catch_unwind(AssertUnwindSafe(|| {
1128 cache.populate_and_encode_with_hashes(
1129 "b<s>世界</s>tail",
1130 hashes.into_iter(),
1131 tok.as_ref(),
1132 )
1133 }))
1134 .is_err()
1135 );
1136 assert!(cache.is_empty());
1137 }
1138
1139 #[test]
1140 fn boundaries_detected_for_multibyte_deepseek_tool_tokens() {
1141 let specials = &["<|tool▁calls▁begin|>", "<|tool▁call▁end|>"];
1146 let text = "<|tool▁calls▁begin|>payload<|tool▁call▁end|>tail";
1147 let bounds = find_special_token_boundaries(text, specials);
1148
1149 let after_begin = "<|tool▁calls▁begin|>".len();
1150 let after_end = text.find("<|tool▁call▁end|>").unwrap() + "<|tool▁call▁end|>".len();
1151 assert_eq!(bounds, vec![after_begin, after_end]);
1152 for &b in &bounds {
1153 assert!(
1154 text.is_char_boundary(b),
1155 "boundary {b} is not a char boundary"
1156 );
1157 let _ = &text[..b]; }
1159 }
1160
1161 fn populate_miss<E: Encoder + ?Sized>(
1162 cache: &L1Cache,
1163 input: &str,
1164 tokenizer: &E,
1165 reuse_hashes: bool,
1166 ) -> anyhow::Result<Vec<TokenIdType>> {
1167 if reuse_hashes {
1168 let PrefixLookup::Miss(hashes) = cache.lookup_prefix(input) else {
1169 panic!("expected miss");
1170 };
1171 cache.populate_and_encode_with_hashes(input, hashes.into_iter(), tokenizer)
1172 } else {
1173 cache.populate_and_encode(input, tokenizer)
1174 }
1175 }
1176
1177 #[test]
1178 fn populate_and_encode_matches_uncached_and_seeds_cache() {
1179 let tok = load_tokenizer();
1180 for input in [
1181 "<s>system\nYou are helpful.</s><s>user\nHello there, friend.</s>",
1182 "<s>system\n世界 🦀</s><s>user\nこんにちは</s>tail",
1183 ] {
1184 let plain = tok.encode(input).unwrap();
1185 let boundaries = find_special_token_boundaries(input, SPECIALS);
1186 for reuse_hashes in [false, true] {
1187 let cache = test_cache(8 * 1024 * 1024);
1188 let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1189 assert_eq!(
1190 got,
1191 plain.token_ids(),
1192 "fused miss encode must equal uncached encode"
1193 );
1194
1195 let mut expected_bytes = 0;
1196 for &boundary in &boundaries {
1197 let hash = cache.hash_prefix(&input.as_bytes()[..boundary]);
1198 let saved = cache
1199 .cache
1200 .cache
1201 .get(&hash)
1202 .expect("every prefix is cached");
1203 let expected = tok.encode(&input[..boundary]).unwrap();
1204 assert_eq!(&*saved.tokens, expected.token_ids());
1205 expected_bytes += size_of_val(expected.token_ids());
1206 }
1207 let stats = cache.stats();
1208 assert_eq!(stats.entries, boundaries.len());
1209 assert_eq!(stats.memory_bytes, expected_bytes);
1210 assert_eq!(stats.hits, 0);
1211 assert_eq!(stats.misses, u64::from(reuse_hashes));
1212 let (_t, offset, deepest) = cache
1213 .longest_prefix_match(input)
1214 .expect("hit after populate");
1215 assert_eq!(offset, *boundaries.last().unwrap());
1216 assert_eq!(deepest, offset);
1217 }
1218 }
1219 }
1220
1221 #[test]
1222 fn populate_and_encode_handles_inputs_without_special_tokens() {
1223 let tok = load_tokenizer();
1224 for input in ["", "plain text with no special tokens at all", "<s>"] {
1225 for reuse_hashes in [false, true] {
1226 let cache = test_cache(8 * 1024 * 1024);
1227 let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1228 let plain = tok.encode(input).unwrap();
1229 assert_eq!(got, plain.token_ids());
1230 assert!(cache.is_empty(), "nothing cacheable without boundaries");
1231 assert_eq!(cache.stats().misses, u64::from(reuse_hashes));
1232 }
1233 }
1234 }
1235
1236 #[test]
1237 fn populate_and_encode_handles_trailing_special_token() {
1238 let tok = load_tokenizer();
1240 let input = "<s>system\nDone.</s>";
1241 for reuse_hashes in [false, true] {
1242 let cache = test_cache(8 * 1024 * 1024);
1243 let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1244 let plain = tok.encode(input).unwrap();
1245 assert_eq!(
1246 got,
1247 plain.token_ids(),
1248 "tail-segment assembly must be exact"
1249 );
1250 }
1251 }
1252
1253 #[test]
1254 fn miss_encode_failure_retains_only_completed_prefixes() {
1255 let input = "<s>世界</s><s>tail";
1256 let boundaries = find_special_token_boundaries(input, SPECIALS);
1257 for fail_at in 0..=boundaries.len() {
1258 for reuse_hashes in [false, true] {
1259 let cache = test_cache(8 * 1024 * 1024);
1260 let encoder = FailAt {
1261 call: 0.into(),
1262 fail_at,
1263 };
1264 let error = populate_miss(&cache, input, &encoder, reuse_hashes).unwrap_err();
1265 assert_eq!(error.to_string(), "suffix failed");
1266 assert_eq!(cache.len(), fail_at);
1267 for (index, &boundary) in boundaries.iter().enumerate() {
1268 let hash = cache.hash_prefix(&input.as_bytes()[..boundary]);
1269 let saved = cache.cache.cache.get(&hash);
1270 if index < fail_at {
1271 assert_eq!(&*saved.unwrap().tokens, vec![1; index + 1]);
1272 } else {
1273 assert!(saved.is_none());
1274 }
1275 }
1276 }
1277 }
1278 }
1279}