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, Arc<[TokenIdType]>, PrefixHasher>;
50
51pub(super) struct PrefixMatch {
53 pub(super) tokens: Arc<[TokenIdType]>,
54 pub(super) prefix_len: usize,
55 deepest_boundary: usize,
56 deepest_hash: Option<Blake3Hash>,
57}
58
59pub(super) enum PrefixLookup {
61 Hit(PrefixMatch),
62 Miss(Vec<(usize, Blake3Hash)>),
63}
64
65fn hash_prefixes<'a>(
67 input: &'a str,
68 boundaries: &'a [usize],
69) -> impl Iterator<Item = (usize, Blake3Hash)> + 'a {
70 let mut hasher = blake3::Hasher::new();
71 let mut last_pos = 0;
72 boundaries.iter().map(move |&boundary_pos| {
73 hasher.update(&input.as_bytes()[last_pos..boundary_pos]);
74 last_pos = boundary_pos;
75 (boundary_pos, *hasher.finalize().as_bytes())
76 })
77}
78
79fn boundaries_with(text: &str, matcher: &AhoCorasick) -> Vec<usize> {
88 let mut boundaries: Vec<usize> = matcher
89 .find_overlapping_iter(text)
90 .map(|m| m.end())
91 .filter(|&end| end < text.len())
92 .collect();
93 boundaries.sort_unstable();
94 boundaries.dedup();
95 boundaries
96}
97
98fn has_nontrivial_self_overlap(token: &str) -> bool {
99 let bytes = token.as_bytes();
100 (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
101}
102
103fn tokens_can_overlap(a: &str, b: &str) -> bool {
104 if a.contains(b) || b.contains(a) {
105 return true;
106 }
107
108 let a = a.as_bytes();
109 let b = b.as_bytes();
110 let max_overlap = a.len().min(b.len());
111 (1..max_overlap).any(|overlap| {
112 a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
113 })
114}
115
116pub(super) fn first_unsafe_overlap(special_tokens: &[String]) -> Option<(&str, &str)> {
124 for (index, token) in special_tokens.iter().enumerate() {
125 if token.is_empty() {
126 continue;
127 }
128 if has_nontrivial_self_overlap(token) {
129 return Some((token, token));
130 }
131 for other in &special_tokens[index + 1..] {
132 if !other.is_empty() && token != other && tokens_can_overlap(token, other) {
133 return Some((token, other));
134 }
135 }
136 }
137
138 None
139}
140
141#[cfg(test)]
144fn find_special_token_boundaries(text: &str, special_tokens: &[&str]) -> Vec<usize> {
145 if special_tokens.is_empty() {
146 return Vec::new();
147 }
148 let matcher = AhoCorasick::new(special_tokens)
149 .expect("special tokens form a valid Aho-Corasick automaton");
150 boundaries_with(text, &matcher)
151}
152
153pub type CacheEventFn = Arc<dyn Fn() + Send + Sync>;
157
158pub struct L1Cache {
162 cache: PrefixCache,
164 matcher: Option<AhoCorasick>,
167 hits: AtomicU64,
168 misses: AtomicU64,
169 on_hit: Option<CacheEventFn>,
170 on_miss: Option<CacheEventFn>,
171}
172
173impl L1Cache {
174 pub fn new(max_memory: usize, mut special_tokens: Vec<String>) -> Self {
177 special_tokens.retain(|token| !token.is_empty());
178
179 let cache = Cache::builder()
183 .max_capacity(max_memory as u64)
184 .weigher(|_k: &Blake3Hash, tokens: &Arc<[TokenIdType]>| -> u32 {
185 size_of_val(tokens.as_ref()).min(u32::MAX as usize) as u32
186 })
187 .build_with_hasher(PrefixHasher::default());
188
189 let matcher = (!special_tokens.is_empty()).then(|| {
191 AhoCorasick::new(&special_tokens)
192 .expect("special tokens form a valid Aho-Corasick automaton")
193 });
194
195 Self {
196 cache,
197 matcher,
198 hits: AtomicU64::new(0),
199 misses: AtomicU64::new(0),
200 on_hit: None,
201 on_miss: None,
202 }
203 }
204
205 pub fn set_observer(&mut self, on_hit: CacheEventFn, on_miss: CacheEventFn) {
207 self.on_hit = Some(on_hit);
208 self.on_miss = Some(on_miss);
209 }
210
211 fn boundaries(&self, text: &str) -> Vec<usize> {
215 match &self.matcher {
216 Some(matcher) => boundaries_with(text, matcher),
217 None => Vec::new(),
218 }
219 }
220
221 pub fn longest_prefix_match(&self, input: &str) -> Option<(Arc<[TokenIdType]>, usize, usize)> {
228 match self.lookup_prefix(input) {
229 PrefixLookup::Hit(matched) => {
230 Some((matched.tokens, matched.prefix_len, matched.deepest_boundary))
231 }
232 PrefixLookup::Miss(_) => None,
233 }
234 }
235
236 pub(super) fn lookup_prefix(&self, input: &str) -> PrefixLookup {
239 let boundaries = self.boundaries(input);
240
241 if boundaries.is_empty() {
242 self.misses.fetch_add(1, Ordering::Relaxed);
243 if let Some(cb) = &self.on_miss {
244 cb();
245 }
246 return PrefixLookup::Miss(Vec::new());
247 }
248
249 let prefix_hashes: Vec<_> = hash_prefixes(input, &boundaries).collect();
250
251 for &(boundary_pos, hash_bytes) in prefix_hashes.iter().rev() {
252 if let Some(tokens) = self.cache.get(&hash_bytes) {
253 self.hits.fetch_add(1, Ordering::Relaxed);
254 if let Some(cb) = &self.on_hit {
255 cb();
256 }
257 let &(deepest_boundary, deepest_hash) =
259 prefix_hashes.last().expect("prefix hashes is non-empty");
260 return PrefixLookup::Hit(PrefixMatch {
261 tokens,
262 prefix_len: boundary_pos,
263 deepest_boundary,
264 deepest_hash: Some(deepest_hash),
265 });
266 }
267 }
268
269 self.misses.fetch_add(1, Ordering::Relaxed);
270 if let Some(cb) = &self.on_miss {
271 cb();
272 }
273 PrefixLookup::Miss(prefix_hashes)
274 }
275
276 pub fn insert_at_boundaries<E: Encoder + ?Sized>(
284 &self,
285 input: &str,
286 tokenizer: &E,
287 ) -> anyhow::Result<()> {
288 let boundaries = self.boundaries(input);
289 if boundaries.is_empty() {
290 return Ok(());
291 }
292 self.populate_boundaries(input, hash_prefixes(input, &boundaries), tokenizer)?;
293 Ok(())
294 }
295
296 pub fn populate_and_encode<E: Encoder + ?Sized>(
305 &self,
306 input: &str,
307 tokenizer: &E,
308 ) -> anyhow::Result<Vec<TokenIdType>> {
309 let boundaries = self.boundaries(input);
310 self.populate_and_encode_with_hashes(input, hash_prefixes(input, &boundaries), tokenizer)
311 }
312
313 pub(super) fn populate_and_encode_with_hashes<E: Encoder + ?Sized>(
317 &self,
318 input: &str,
319 prefix_hashes: impl Iterator<Item = (usize, Blake3Hash)>,
320 tokenizer: &E,
321 ) -> anyhow::Result<Vec<TokenIdType>> {
322 let (mut running, tail_start) =
323 self.populate_boundaries(input, prefix_hashes, tokenizer)?;
324 if tail_start == 0 {
325 return Ok(tokenizer.encode(input)?.token_ids().to_vec());
326 }
327
328 let tail = tokenizer.encode(&input[tail_start..])?;
329 running.extend_from_slice(tail.token_ids());
330 Ok(running)
331 }
332
333 fn populate_boundaries<E: Encoder + ?Sized>(
336 &self,
337 input: &str,
338 prefix_hashes: impl Iterator<Item = (usize, Blake3Hash)>,
339 tokenizer: &E,
340 ) -> anyhow::Result<(Vec<TokenIdType>, usize)> {
341 #[cfg(debug_assertions)]
342 let mut validation_hasher = blake3::Hasher::new();
343 let mut running_tokens: Vec<TokenIdType> = Vec::new();
344 let mut last_pos = 0;
345
346 for (boundary_pos, hash_bytes) in prefix_hashes {
347 #[cfg(debug_assertions)]
348 {
349 validation_hasher.update(&input.as_bytes()[last_pos..boundary_pos]);
350 debug_assert_eq!(hash_bytes, *validation_hasher.finalize().as_bytes());
351 }
352
353 let seg = tokenizer.encode(&input[last_pos..boundary_pos])?;
357 running_tokens.extend_from_slice(seg.token_ids());
358
359 let prefix_tokens: Arc<[TokenIdType]> = running_tokens.as_slice().into();
360 self.cache.insert(hash_bytes, prefix_tokens);
361
362 last_pos = boundary_pos;
363 }
364
365 Ok((running_tokens, last_pos))
366 }
367
368 pub fn extend_after_match<E: Encoder + ?Sized>(
385 &self,
386 input: &str,
387 prefix_tokens: Arc<[TokenIdType]>,
388 prefix_len: usize,
389 deepest_boundary: usize,
390 tokenizer: &E,
391 ) -> anyhow::Result<Vec<TokenIdType>> {
392 self.extend_after_match_with_hash(
393 input,
394 PrefixMatch {
395 tokens: prefix_tokens,
396 prefix_len,
397 deepest_boundary,
398 deepest_hash: None,
399 },
400 tokenizer,
401 )
402 }
403
404 pub(super) fn extend_after_match_with_hash<E: Encoder + ?Sized>(
408 &self,
409 input: &str,
410 matched: PrefixMatch,
411 tokenizer: &E,
412 ) -> anyhow::Result<Vec<TokenIdType>> {
413 let PrefixMatch {
414 tokens: prefix_tokens,
415 prefix_len,
416 deepest_boundary,
417 deepest_hash,
418 } = matched;
419 let deepest = (deepest_boundary > prefix_len).then_some(deepest_boundary);
421
422 let Some(deepest) = deepest else {
423 let suffix_enc = tokenizer.encode(&input[prefix_len..])?;
424 let mut merged = Vec::with_capacity(prefix_tokens.len() + suffix_enc.token_ids().len());
426 merged.extend_from_slice(&prefix_tokens);
427 merged.extend_from_slice(suffix_enc.token_ids());
428 return Ok(merged);
429 };
430
431 let seg_a = tokenizer.encode(&input[prefix_len..deepest])?;
433 let seg_b = tokenizer.encode(&input[deepest..])?;
434 let mut cumulative = Vec::with_capacity(
435 prefix_tokens.len() + seg_a.token_ids().len() + seg_b.token_ids().len(),
436 );
437 cumulative.extend_from_slice(&prefix_tokens);
438 cumulative.extend_from_slice(seg_a.token_ids());
439
440 let hash_bytes = deepest_hash.unwrap_or_else(|| {
441 let mut hasher = blake3::Hasher::new();
442 hasher.update(&input.as_bytes()[..deepest]);
443 *hasher.finalize().as_bytes()
444 });
445 debug_assert_eq!(
446 hash_bytes,
447 *blake3::hash(&input.as_bytes()[..deepest]).as_bytes()
448 );
449
450 let tokens: Arc<[TokenIdType]> = cumulative.as_slice().into();
452 self.cache.insert(hash_bytes, tokens);
453
454 cumulative.extend_from_slice(seg_b.token_ids());
455 Ok(cumulative)
456 }
457
458 pub fn len(&self) -> usize {
461 self.cache.run_pending_tasks();
462 self.cache.entry_count() as usize
463 }
464
465 pub fn is_empty(&self) -> bool {
466 self.len() == 0
467 }
468
469 pub fn stats(&self) -> L1CacheStats {
470 self.cache.run_pending_tasks();
472 let hits = self.hits.load(Ordering::Relaxed);
473 let misses = self.misses.load(Ordering::Relaxed);
474 let total_requests = hits + misses;
475
476 L1CacheStats {
477 hits,
478 misses,
479 entries: self.cache.entry_count() as usize,
480 memory_bytes: self.cache.weighted_size() as usize,
481 hit_rate: if total_requests > 0 {
482 hits as f64 / total_requests as f64
483 } else {
484 0.0
485 },
486 }
487 }
488
489 pub fn clear(&self) {
490 self.cache.invalidate_all();
491 self.cache.run_pending_tasks();
492 self.hits.store(0, Ordering::Relaxed);
493 self.misses.store(0, Ordering::Relaxed);
494 }
495}
496
497#[derive(Debug, Clone)]
498pub struct L1CacheStats {
499 pub hits: u64,
500 pub misses: u64,
501 pub entries: usize,
502 pub memory_bytes: usize,
503 pub hit_rate: f64,
504}
505
506#[cfg(test)]
507mod tests {
508 use std::sync::Arc;
509
510 use super::*;
511 use crate::{HuggingFaceTokenizer, traits::Tokenizer};
512
513 const TINYLLAMA_PATH: &str = concat!(
516 env!("CARGO_MANIFEST_DIR"),
517 "/tests/data/sample-models/TinyLlama_v1.1/tokenizer.json"
518 );
519
520 const SPECIALS: &[&str] = &["<s>", "</s>"];
521
522 fn load_tokenizer() -> Arc<dyn Tokenizer> {
523 Arc::new(HuggingFaceTokenizer::from_file(TINYLLAMA_PATH).expect("load TinyLlama"))
524 }
525
526 fn test_cache(max_memory: usize) -> L1Cache {
528 L1Cache::new(
529 max_memory,
530 SPECIALS.iter().map(|s| (*s).to_string()).collect(),
531 )
532 }
533
534 #[test]
535 fn boundaries_are_after_each_special_token_occurrence() {
536 let input = "<s>system\nHi</s><s>user\nHello</s>";
537 let bounds = find_special_token_boundaries(input, SPECIALS);
538 assert_eq!(bounds.len(), 3);
540 for w in bounds.windows(2) {
541 assert!(w[0] < w[1], "boundaries must be strictly increasing");
542 }
543 assert!(bounds.iter().all(|&b| b < input.len()));
544 }
545
546 #[test]
547 fn no_special_tokens_yields_no_boundaries() {
548 assert!(find_special_token_boundaries("plain text", &[]).is_empty());
549 }
550
551 #[test]
552 fn unsafe_overlap_detects_containment_crossing_and_self_overlap() {
553 let cases = [
554 (vec!["〈|", "〈|EOS|〉"], Some(("〈|", "〈|EOS|〉"))),
555 (vec!["ab", "bc"], Some(("ab", "bc"))),
556 (vec!["|◊|"], Some(("|◊|", "|◊|"))),
557 (vec!["<s>", "<s>"], None),
558 ];
559
560 for (tokens, expected) in cases {
561 let tokens: Vec<String> = tokens.into_iter().map(String::from).collect();
562 assert_eq!(first_unsafe_overlap(&tokens), expected);
563 }
564 }
565
566 #[test]
567 fn llama_numbered_special_tokens_do_not_trigger_overlap_guard() {
568 let mut llama: Vec<String> = [
569 "<|begin_of_text|>",
570 "<|end_of_text|>",
571 "<|start_header_id|>",
572 "<|end_header_id|>",
573 "<|eot_id|>",
574 ]
575 .into_iter()
576 .map(String::from)
577 .collect();
578 llama.extend((0..251).map(|id| format!("<|reserved_special_token_{id}|>")));
579
580 assert_eq!(first_unsafe_overlap(&llama), None);
581 }
582
583 #[test]
584 fn insert_then_lookup_finds_shared_prefix() {
585 let cache = test_cache(1024 * 1024);
586 let tokenizer = load_tokenizer();
587
588 let warm = "<s>system\nYou are helpful.</s><s>user\nHi</s>";
589 cache
590 .insert_at_boundaries(warm, tokenizer.as_ref())
591 .unwrap();
592 assert!(!cache.is_empty());
593
594 let target = "<s>system\nYou are helpful.</s><s>user\nDifferent question</s>";
595 let (tokens, offset, _deepest) = cache
596 .longest_prefix_match(target)
597 .expect("shared prefix should match");
598 assert!(offset > 0);
599 assert!(!tokens.is_empty());
600 }
601
602 #[test]
603 fn miss_increments_misses_counter() {
604 let cache = test_cache(1024 * 1024);
605 assert!(
606 cache
607 .longest_prefix_match("plain text no specials")
608 .is_none()
609 );
610 assert_eq!(cache.stats().misses, 1);
611 }
612
613 #[test]
614 fn hit_increments_hits_counter() {
615 let cache = test_cache(1024 * 1024);
616 let tokenizer = load_tokenizer();
617 let warm = "<s>system\nA.</s><s>user\nB</s>";
618 cache
619 .insert_at_boundaries(warm, tokenizer.as_ref())
620 .unwrap();
621 let _ = cache.longest_prefix_match(warm);
622 assert!(cache.stats().hits >= 1);
623 }
624
625 #[test]
626 fn merge_invariant_holds_against_uncached_encode() {
627 let cache = test_cache(1024 * 1024);
631 let tokenizer = load_tokenizer();
632
633 let template = "<s>system\nYou are helpful.</s><s>user\n";
634 let warm = format!("{template}First.</s>");
635 cache
636 .insert_at_boundaries(&warm, tokenizer.as_ref())
637 .unwrap();
638
639 let target = format!("{template}A completely different second question.</s>");
640 let (prefix_tokens, prefix_len, _deepest) = cache
641 .longest_prefix_match(&target)
642 .expect("should find prefix");
643
644 let suffix = &target[prefix_len..];
645 let suffix_enc = tokenizer.encode(suffix).unwrap();
646 let mut merged = prefix_tokens.to_vec();
648 merged.extend_from_slice(suffix_enc.token_ids());
649
650 let plain = tokenizer.encode(&target).unwrap();
651 assert_eq!(
652 merged,
653 plain.token_ids(),
654 "merged tokens must equal plain encode"
655 );
656 }
657
658 #[test]
659 fn eviction_respects_memory_budget() {
660 let cache = test_cache(4 * 1024);
662 let tokenizer = load_tokenizer();
663 for i in 0..50 {
664 let input =
665 format!("<s>system\nPersona {i} chatty.</s><s>user\nTurn {i} content here.</s>");
666 cache
667 .insert_at_boundaries(&input, tokenizer.as_ref())
668 .unwrap();
669 }
670 let stats = cache.stats();
671 assert!(
672 stats.memory_bytes <= 4 * 1024,
673 "memory_bytes={} exceeds budget",
674 stats.memory_bytes
675 );
676 }
677
678 #[test]
679 fn concurrent_inserts_and_lookups_do_not_corrupt() {
680 use std::thread;
681
682 let cache = Arc::new(test_cache(1024 * 1024));
683 let tokenizer = load_tokenizer();
684
685 let mut handles = vec![];
686 for i in 0..10 {
687 let cache_c = cache.clone();
688 let tok = tokenizer.clone();
689 handles.push(thread::spawn(move || {
690 let input = format!("<s>system\nThread {i}.</s><s>user\nThread {i} body.</s>");
691 cache_c.insert_at_boundaries(&input, tok.as_ref()).unwrap();
692 let r = cache_c.longest_prefix_match(&input);
693 assert!(r.is_some(), "thread {i} expected match after insert");
694 }));
695 }
696 for h in handles {
697 h.join().unwrap();
698 }
699 assert!(cache.stats().memory_bytes > 0);
700 assert!(cache.stats().hits >= 10);
701 }
702
703 fn growing_chat_turns(n: usize) -> Vec<String> {
709 let mut convo = String::from("<s>system\nYou are a helpful assistant.</s>");
710 let mut turns = Vec::with_capacity(n);
711 for i in 0..n {
712 convo.push_str(&format!(
713 "<s>user\nQuestion {i} please answer it.</s><s>assistant\nDetailed answer {i} follows here.</s>"
714 ));
715 turns.push(format!("{convo}<s>user\nFollow-up {i}"));
716 }
717 turns
718 }
719
720 #[test]
721 fn extend_on_hit_advances_match_depth_each_turn() {
722 let tok = load_tokenizer();
726 let turns = growing_chat_turns(5);
727
728 let off = test_cache(8 * 1024 * 1024);
730 off.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
731 let pinned = off.longest_prefix_match(&turns[1]).expect("hit").1;
732 for t in &turns[1..] {
733 let (_toks, offset, _deepest) = off.longest_prefix_match(t).expect("hit");
734 assert_eq!(
735 offset, pinned,
736 "extend-off offset must stay pinned at turn-1 depth"
737 );
738 }
739
740 let on = test_cache(8 * 1024 * 1024);
742 on.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
743 let mut prev = 0usize;
744 for (i, t) in turns.iter().enumerate().skip(1) {
745 let (prefix_tokens, offset, deepest) = on.longest_prefix_match(t).expect("hit");
746 assert!(
747 offset > prev,
748 "turn {i}: extend-on offset {offset} must exceed previous {prev}"
749 );
750 prev = offset;
751
752 let merged = on
754 .extend_after_match(t, prefix_tokens, offset, deepest, tok.as_ref())
755 .unwrap();
756 let plain = tok.encode(t).unwrap();
757 assert_eq!(
758 merged,
759 plain.token_ids(),
760 "turn {i}: extend merge must equal plain encode"
761 );
762 }
763
764 assert!(
765 prev > pinned,
766 "extend-on frontier ({prev}) must reach deeper than pinned extend-off depth ({pinned})"
767 );
768 }
769
770 #[test]
771 fn extend_on_hit_respects_budget_and_stays_correct() {
772 let tok = load_tokenizer();
775 let cache = test_cache(4 * 1024);
776 let turns = growing_chat_turns(20);
777 cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
778
779 for t in &turns[1..] {
780 let merged = match cache.longest_prefix_match(t) {
781 Some((prefix_tokens, offset, deepest)) => cache
782 .extend_after_match(t, prefix_tokens, offset, deepest, tok.as_ref())
783 .unwrap(),
784 None => {
785 let enc = tok.encode(t).unwrap();
787 cache.insert_at_boundaries(t, tok.as_ref()).unwrap();
788 enc.token_ids().to_vec()
789 }
790 };
791 let plain = tok.encode(t).unwrap();
792 assert_eq!(
793 merged,
794 plain.token_ids(),
795 "encode must stay correct under eviction pressure"
796 );
797 assert!(
798 cache.stats().memory_bytes <= 4 * 1024,
799 "memory_bytes={} exceeds budget",
800 cache.stats().memory_bytes
801 );
802 }
803 }
804
805 #[test]
806 fn concurrent_extend_on_hit_does_not_corrupt() {
807 use std::thread;
808
809 let tok = load_tokenizer();
810 let cache = Arc::new(test_cache(8 * 1024 * 1024));
811 let turns = growing_chat_turns(8);
812 cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
814
815 let mut handles = vec![];
816 for _ in 0..8 {
817 let cache_c = cache.clone();
818 let tok_c = tok.clone();
819 let turns_c = turns.clone();
820 handles.push(thread::spawn(move || {
821 for t in &turns_c[1..] {
822 if let PrefixLookup::Hit(matched) = cache_c.lookup_prefix(t) {
823 let merged = cache_c
824 .extend_after_match_with_hash(t, matched, tok_c.as_ref())
825 .unwrap();
826 let plain = tok_c.encode(t).unwrap();
827 assert_eq!(
828 merged,
829 plain.token_ids(),
830 "concurrent extend must stay correct"
831 );
832 }
833 }
834 }));
835 }
836 for h in handles {
837 h.join().unwrap();
838 }
839 assert!(cache.stats().memory_bytes > 0);
840 }
841
842 #[test]
843 fn extend_after_match_persists_correct_deepest_entry() {
844 let tok = load_tokenizer();
845 for unicode in [false, true] {
846 let turns: Vec<_> = growing_chat_turns(3)
847 .into_iter()
848 .map(|t| {
849 if unicode {
850 t.replace("system", "system 世界 🦀")
851 } else {
852 t
853 }
854 })
855 .collect();
856
857 let cache = test_cache(8 * 1024 * 1024);
858 cache.insert_at_boundaries(&turns[0], tok.as_ref()).unwrap();
859
860 let PrefixLookup::Hit(matched) = cache.lookup_prefix(&turns[1]) else {
861 panic!("partial hit on turns[1]");
862 };
863 let prefix_len = matched.prefix_len;
864 let deepest_boundary = matched.deepest_boundary;
865 assert!(deepest_boundary > prefix_len);
866 assert_eq!(
867 matched.deepest_hash,
868 Some(*blake3::hash(&turns[1].as_bytes()[..deepest_boundary]).as_bytes())
869 );
870 assert_ne!(
871 matched.deepest_hash,
872 Some(*blake3::hash(&turns[1].as_bytes()[..prefix_len]).as_bytes())
873 );
874 let entries_before = cache.stats().entries;
875
876 let _merged = cache
877 .extend_after_match_with_hash(&turns[1], matched, tok.as_ref())
878 .unwrap();
879
880 assert_eq!(
881 cache.stats().entries,
882 entries_before + 1,
883 "extend must persist exactly one (deepest) entry"
884 );
885
886 let deepest = find_special_token_boundaries(&turns[1], SPECIALS)
887 .into_iter()
888 .rev()
889 .find(|&b| b > prefix_len)
890 .expect("a deeper boundary must exist in the appended turn");
891 assert_eq!(
892 deepest_boundary, deepest,
893 "longest_prefix_match must return the deepest boundary used by extend"
894 );
895
896 let (saved_tokens, saved_offset, _deepest) = cache
897 .longest_prefix_match(&turns[1])
898 .expect("hit after extend");
899 assert_eq!(
900 saved_offset, deepest,
901 "lookup must now hit at the just-saved deepest boundary"
902 );
903 let expected = tok.encode(&turns[1][..deepest]).unwrap();
904 assert_eq!(
905 &*saved_tokens,
906 expected.token_ids(),
907 "persisted entry tokens must equal the uncached encode of the cached prefix"
908 );
909 }
910 }
911
912 #[test]
913 fn extend_without_deeper_boundary_does_not_insert() {
914 let tok = load_tokenizer();
915 let cache = test_cache(8 * 1024 * 1024);
916 cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
917 for input in ["<s>世界", "<s>世界</s>"] {
918 let PrefixLookup::Hit(matched) = cache.lookup_prefix(input) else {
919 panic!("expected hit");
920 };
921 assert_eq!(matched.prefix_len, matched.deepest_boundary);
922 let entries = cache.len();
923 let merged = cache
924 .extend_after_match_with_hash(input, matched, tok.as_ref())
925 .unwrap();
926 assert_eq!(merged, tok.encode(input).unwrap().token_ids());
927 assert_eq!(cache.len(), entries);
928 }
929 }
930
931 struct FailAt {
932 call: std::sync::atomic::AtomicUsize,
933 fail_at: usize,
934 }
935 impl Encoder for FailAt {
936 fn encode(&self, _: &str) -> crate::Result<crate::Encoding> {
937 if self.call.fetch_add(1, Ordering::Relaxed) == self.fail_at {
938 anyhow::bail!("suffix failed");
939 }
940 Ok(crate::Encoding::Sp(vec![1]))
941 }
942 fn encode_batch(&self, inputs: &[&str]) -> crate::Result<Vec<crate::Encoding>> {
943 inputs.iter().map(|s| self.encode(s)).collect()
944 }
945 }
946
947 #[test]
948 fn hash_reuse_does_not_insert_when_either_suffix_encode_fails() {
949 let tok = load_tokenizer();
950 let cache = test_cache(8 * 1024 * 1024);
951 cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
952 let input = "<s>世界</s><s>tail";
953 for fail_at in [0, 1] {
954 let PrefixLookup::Hit(matched) = cache.lookup_prefix(input) else {
955 panic!("expected hit");
956 };
957 let entries = cache.len();
958 let error = cache
959 .extend_after_match_with_hash(
960 input,
961 matched,
962 &FailAt {
963 call: 0.into(),
964 fail_at,
965 },
966 )
967 .unwrap_err();
968 assert_eq!(error.to_string(), "suffix failed");
969 assert_eq!(cache.len(), entries);
970 }
971 }
972
973 #[test]
974 #[cfg(debug_assertions)]
975 fn reused_hashes_reject_mismatched_input_before_insertion() {
976 use std::panic::{AssertUnwindSafe, catch_unwind};
977
978 let tok = load_tokenizer();
979 let cache = test_cache(8 * 1024 * 1024);
980 cache.insert_at_boundaries("<s>seed", tok.as_ref()).unwrap();
981 let PrefixLookup::Hit(matched) = cache.lookup_prefix("<s>世界</s><s>tail") else {
982 panic!("expected hit");
983 };
984 let entries = cache.len();
985 assert!(
986 catch_unwind(AssertUnwindSafe(|| {
987 cache.extend_after_match_with_hash("<s>日本</s><s>tail", matched, tok.as_ref())
988 }))
989 .is_err()
990 );
991 assert_eq!(cache.len(), entries);
992
993 let cache = test_cache(8 * 1024 * 1024);
994 let PrefixLookup::Miss(hashes) = cache.lookup_prefix("a<s>世界</s>tail") else {
995 panic!("expected miss");
996 };
997 assert!(
998 catch_unwind(AssertUnwindSafe(|| {
999 cache.populate_and_encode_with_hashes(
1000 "b<s>世界</s>tail",
1001 hashes.into_iter(),
1002 tok.as_ref(),
1003 )
1004 }))
1005 .is_err()
1006 );
1007 assert!(cache.is_empty());
1008 }
1009
1010 #[test]
1011 fn boundaries_detected_for_multibyte_deepseek_tool_tokens() {
1012 let specials = &["<|tool▁calls▁begin|>", "<|tool▁call▁end|>"];
1017 let text = "<|tool▁calls▁begin|>payload<|tool▁call▁end|>tail";
1018 let bounds = find_special_token_boundaries(text, specials);
1019
1020 let after_begin = "<|tool▁calls▁begin|>".len();
1021 let after_end = text.find("<|tool▁call▁end|>").unwrap() + "<|tool▁call▁end|>".len();
1022 assert_eq!(bounds, vec![after_begin, after_end]);
1023 for &b in &bounds {
1024 assert!(
1025 text.is_char_boundary(b),
1026 "boundary {b} is not a char boundary"
1027 );
1028 let _ = &text[..b]; }
1030 }
1031
1032 fn populate_miss<E: Encoder + ?Sized>(
1033 cache: &L1Cache,
1034 input: &str,
1035 tokenizer: &E,
1036 reuse_hashes: bool,
1037 ) -> anyhow::Result<Vec<TokenIdType>> {
1038 if reuse_hashes {
1039 let PrefixLookup::Miss(hashes) = cache.lookup_prefix(input) else {
1040 panic!("expected miss");
1041 };
1042 cache.populate_and_encode_with_hashes(input, hashes.into_iter(), tokenizer)
1043 } else {
1044 cache.populate_and_encode(input, tokenizer)
1045 }
1046 }
1047
1048 #[test]
1049 fn populate_and_encode_matches_uncached_and_seeds_cache() {
1050 let tok = load_tokenizer();
1051 for input in [
1052 "<s>system\nYou are helpful.</s><s>user\nHello there, friend.</s>",
1053 "<s>system\n世界 🦀</s><s>user\nこんにちは</s>tail",
1054 ] {
1055 let plain = tok.encode(input).unwrap();
1056 let boundaries = find_special_token_boundaries(input, SPECIALS);
1057 for reuse_hashes in [false, true] {
1058 let cache = test_cache(8 * 1024 * 1024);
1059 let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1060 assert_eq!(
1061 got,
1062 plain.token_ids(),
1063 "fused miss encode must equal uncached encode"
1064 );
1065
1066 let mut expected_bytes = 0;
1067 for &boundary in &boundaries {
1068 let hash = *blake3::hash(&input.as_bytes()[..boundary]).as_bytes();
1069 let saved = cache.cache.get(&hash).expect("every prefix is cached");
1070 let expected = tok.encode(&input[..boundary]).unwrap();
1071 assert_eq!(&*saved, expected.token_ids());
1072 expected_bytes += size_of_val(expected.token_ids());
1073 }
1074 let stats = cache.stats();
1075 assert_eq!(stats.entries, boundaries.len());
1076 assert_eq!(stats.memory_bytes, expected_bytes);
1077 assert_eq!(stats.hits, 0);
1078 assert_eq!(stats.misses, u64::from(reuse_hashes));
1079 let (_t, offset, deepest) = cache
1080 .longest_prefix_match(input)
1081 .expect("hit after populate");
1082 assert_eq!(offset, *boundaries.last().unwrap());
1083 assert_eq!(deepest, offset);
1084 }
1085 }
1086 }
1087
1088 #[test]
1089 fn populate_and_encode_handles_inputs_without_special_tokens() {
1090 let tok = load_tokenizer();
1091 for input in ["", "plain text with no special tokens at all", "<s>"] {
1092 for reuse_hashes in [false, true] {
1093 let cache = test_cache(8 * 1024 * 1024);
1094 let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1095 let plain = tok.encode(input).unwrap();
1096 assert_eq!(got, plain.token_ids());
1097 assert!(cache.is_empty(), "nothing cacheable without boundaries");
1098 assert_eq!(cache.stats().misses, u64::from(reuse_hashes));
1099 }
1100 }
1101 }
1102
1103 #[test]
1104 fn populate_and_encode_handles_trailing_special_token() {
1105 let tok = load_tokenizer();
1107 let input = "<s>system\nDone.</s>";
1108 for reuse_hashes in [false, true] {
1109 let cache = test_cache(8 * 1024 * 1024);
1110 let got = populate_miss(&cache, input, tok.as_ref(), reuse_hashes).unwrap();
1111 let plain = tok.encode(input).unwrap();
1112 assert_eq!(
1113 got,
1114 plain.token_ids(),
1115 "tail-segment assembly must be exact"
1116 );
1117 }
1118 }
1119
1120 #[test]
1121 fn miss_encode_failure_retains_only_completed_prefixes() {
1122 let input = "<s>世界</s><s>tail";
1123 let boundaries = find_special_token_boundaries(input, SPECIALS);
1124 for fail_at in 0..=boundaries.len() {
1125 for reuse_hashes in [false, true] {
1126 let cache = test_cache(8 * 1024 * 1024);
1127 let encoder = FailAt {
1128 call: 0.into(),
1129 fail_at,
1130 };
1131 let error = populate_miss(&cache, input, &encoder, reuse_hashes).unwrap_err();
1132 assert_eq!(error.to_string(), "suffix failed");
1133 assert_eq!(cache.len(), fail_at);
1134 for (index, &boundary) in boundaries.iter().enumerate() {
1135 let hash = *blake3::hash(&input.as_bytes()[..boundary]).as_bytes();
1136 let saved = cache.cache.get(&hash);
1137 if index < fail_at {
1138 assert_eq!(&*saved.unwrap(), vec![1; index + 1]);
1139 } else {
1140 assert!(saved.is_none());
1141 }
1142 }
1143 }
1144 }
1145 }
1146}