1use std::collections::HashSet;
5use std::path::Path;
6
7use base64::Engine as _;
8use rayon::prelude::*;
9use rustc_hash::FxHashMap;
10use tiktoken_rs::CoreBPE;
11
12use super::{
13 EncodeSegment, Encoding, Error, Result, TokenIdType,
14 traits::{DecodeResult, Decoder, Encoder, Tokenizer},
15};
16
17const DEFAULT_NUM_RESERVED_SPECIAL_TOKENS: u32 = 256;
20
21const KIMI_PATTERN: &str = fastokens::tiktoken::KIMI_PATTERN;
23
24pub struct TikTokenTokenizer {
25 bpe: CoreBPE,
26 special_token_ids: HashSet<u32>,
27 special_tokens: Vec<String>,
28}
29
30pub(crate) fn sorted_special_token_strings(special_tokens: &FxHashMap<String, u32>) -> Vec<String> {
31 let mut strings: Vec<String> = special_tokens.keys().cloned().collect();
32 strings.sort();
33 strings.dedup();
34 strings
35}
36
37impl TikTokenTokenizer {
38 pub fn from_file(
45 path: &str,
46 pattern: &str,
47 special_tokens: FxHashMap<String, u32>,
48 ) -> Result<Self> {
49 let encoder = parse_tiktoken_file(path)?;
50 let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
51 let special_token_strings = sorted_special_token_strings(&special_tokens);
52
53 let bpe = CoreBPE::new(encoder, special_tokens, pattern)
54 .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
55
56 Ok(Self {
57 bpe,
58 special_token_ids,
59 special_tokens: special_token_strings,
60 })
61 }
62
63 pub fn from_file_auto(path: &str) -> Result<Self> {
68 let file_path = Path::new(path);
69 let directory = file_path
70 .parent()
71 .ok_or_else(|| Error::msg("Cannot determine parent directory of tiktoken file"))?;
72
73 let pattern = detect_bpe_pattern(directory)?;
74 let encoder = parse_tiktoken_file(path)?;
75 let num_base_tokens = encoder.values().max().map_or(0, |&m| m + 1) as usize;
77 let special_tokens = load_special_tokens(directory, num_base_tokens)?;
78 let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
79 let special_token_strings = sorted_special_token_strings(&special_tokens);
80
81 let bpe = CoreBPE::new(encoder, special_tokens, pattern)
82 .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
83
84 Ok(Self {
85 bpe,
86 special_token_ids,
87 special_tokens: special_token_strings,
88 })
89 }
90
91 pub fn special_tokens(&self) -> &[String] {
97 &self.special_tokens
98 }
99}
100
101impl Encoder for TikTokenTokenizer {
102 fn encode(&self, input: &str) -> Result<Encoding> {
103 let token_ids: Vec<u32> = self.bpe.encode_with_special_tokens(input);
104 Ok(Encoding::Sp(token_ids))
105 }
106
107 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
108 inputs.par_iter().map(|input| self.encode(input)).collect()
109 }
110
111 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
112 let mut token_ids = Vec::new();
113 for segment in segments {
114 if segment.allow_special {
115 token_ids.extend(self.bpe.encode_with_special_tokens(segment.text));
116 } else {
117 token_ids.extend(self.bpe.encode_ordinary(segment.text));
118 }
119 }
120 Ok(Encoding::Sp(token_ids))
121 }
122}
123
124impl Decoder for TikTokenTokenizer {
125 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
126 let ids: Vec<u32> = if skip_special_tokens {
127 token_ids
128 .iter()
129 .filter(|&&id| !self.special_token_ids.contains(&id))
130 .copied()
131 .collect()
132 } else {
133 token_ids.to_vec()
134 };
135
136 let bytes: Vec<u8> = self.bpe._decode_native_and_split(ids).flatten().collect();
145 match String::from_utf8(bytes) {
146 Ok(text) => Ok(DecodeResult::Complete(text)),
147 Err(e) => {
148 let text = String::from_utf8_lossy(e.as_bytes()).into_owned();
149 Ok(DecodeResult::from_decoded(text))
150 }
151 }
152 }
153}
154
155impl Tokenizer for TikTokenTokenizer {
156 fn validate_prefix_cache(&self) -> Result<()> {
157 Ok(())
158 }
159
160 fn vocab_size(&self) -> Option<usize> {
165 None
166 }
167
168 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
169 let mut ids: Vec<TokenIdType> = self.special_token_ids.iter().copied().collect();
170 ids.sort_unstable();
171 Ok(ids)
172 }
173
174 fn num_special_tokens_added(&self) -> Result<usize> {
178 Ok(0)
179 }
180}
181
182pub(crate) fn parse_tiktoken_file(path: &str) -> Result<FxHashMap<Vec<u8>, u32>> {
184 let contents = std::fs::read_to_string(path)
185 .map_err(|err| Error::msg(format!("Failed to read tiktoken file '{path}': {err}")))?;
186
187 let engine = base64::engine::general_purpose::STANDARD;
188 let mut encoder = FxHashMap::default();
189
190 for line in contents.lines() {
191 let line = line.trim();
192 if line.is_empty() {
193 continue;
194 }
195 let mut parts = line.split_whitespace();
196 let token_b64 = parts
197 .next()
198 .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no token): {line}")))?;
199 let rank_str = parts
200 .next()
201 .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no rank): {line}")))?;
202
203 let token_bytes = engine
204 .decode(token_b64)
205 .map_err(|err| Error::msg(format!("Invalid base64 in tiktoken file: {err}")))?;
206 let rank: u32 = rank_str
207 .parse()
208 .map_err(|err| Error::msg(format!("Invalid rank in tiktoken file: {err}")))?;
209
210 encoder.insert(token_bytes, rank);
211 }
212
213 Ok(encoder)
214}
215
216pub(crate) fn detect_bpe_pattern(directory: &Path) -> Result<&'static str> {
218 let model_type: String = crate::file_json_field(&directory.join("config.json"), "model_type")
219 .map_err(|err| {
220 Error::msg(format!("Failed to read model_type from config.json: {err}"))
221 })?;
222
223 match model_type.as_str() {
224 "kimi" | "kimi_k2" | "kimi_k25" | "kimi_k3" | "kimi_linear" | "deepseek_v3" => {
230 Ok(KIMI_PATTERN)
231 }
232 _ => Err(Error::msg(format!(
233 "Unsupported tiktoken model_type '{model_type}'. \
234 Currently supported: kimi, kimi_k2, kimi_k25, kimi_k3, kimi_linear, deepseek_v3. \
235 To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
236 with the appropriate BPE regex pattern. \
237 Alternatively, provide a tokenizer.json (HuggingFace format) instead."
238 ))),
239 }
240}
241
242pub(crate) fn load_special_tokens(
247 directory: &Path,
248 num_base_tokens: usize,
249) -> Result<FxHashMap<String, u32>> {
250 let config_path = directory.join("tokenizer_config.json");
251 let mut special_tokens = FxHashMap::default();
252
253 if !config_path.exists() {
254 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
256 let id = num_base_tokens as u32 + i;
257 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
258 }
259 return Ok(special_tokens);
260 }
261
262 let contents = std::fs::read_to_string(&config_path)
263 .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
264
265 let config: serde_json::Value = serde_json::from_str(&contents)
266 .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
267
268 if let Some(added_tokens) = config
269 .get("added_tokens_decoder")
270 .and_then(|v| v.as_object())
271 {
272 for (id_str, token_def) in added_tokens {
273 let id: u32 = id_str.parse().map_err(|err| {
274 Error::msg(format!(
275 "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
276 ))
277 })?;
278
279 let content = token_def
280 .get("content")
281 .and_then(|v| v.as_str())
282 .unwrap_or_else(|| {
283 tracing::warn!("Missing 'content' field for token ID {id}");
285 ""
286 });
287
288 if !content.is_empty() {
289 special_tokens.insert(content.to_string(), id);
290 }
291 }
292
293 let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
295 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
296 let id = num_base_tokens as u32 + i;
297 if !used_ids.contains(&id) {
298 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
299 }
300 }
301 } else {
302 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
304 let id = num_base_tokens as u32 + i;
305 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
306 }
307 }
308
309 Ok(special_tokens)
310}
311
312#[cfg(test)]
313mod tests {
314 use super::*;
315 use crate::DecodeStream;
316 use std::io::Write;
317 use std::sync::Arc;
318
319 fn create_test_tiktoken_file(dir: &Path) -> String {
320 let engine = base64::engine::general_purpose::STANDARD;
321 let mut content = String::new();
322
323 let tokens: Vec<(&[u8], u32)> = vec![
325 (b"h", 0),
326 (b"e", 1),
327 (b"l", 2),
328 (b"o", 3),
329 (b" ", 4),
330 (b"w", 5),
331 (b"r", 6),
332 (b"d", 7),
333 (b"he", 8),
334 (b"ll", 9),
335 (b"lo", 10),
336 (b"wo", 11),
337 (b"rl", 12),
338 (b"hel", 13),
339 (b"llo", 14),
340 (b"wor", 15),
341 (b"hell", 16),
342 (b"ello", 17),
343 (b"worl", 18),
344 (b"hello", 19),
345 (b"world", 20),
346 ];
347
348 for (token, rank) in tokens {
349 let encoded = engine.encode(token);
350 content.push_str(&format!("{encoded} {rank}\n"));
351 }
352
353 let file_path = dir.join("tiktoken.model");
354 let mut file = std::fs::File::create(&file_path).unwrap();
355 file.write_all(content.as_bytes()).unwrap();
356 file_path.to_str().unwrap().to_string()
357 }
358
359 fn create_test_config(dir: &Path, model_type: &str) {
360 let config = serde_json::json!({
361 "model_type": model_type,
362 "max_position_embeddings": 32768,
363 "eos_token_id": [21]
364 });
365 let file_path = dir.join("config.json");
366 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
367 }
368
369 fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
370 let mut added_tokens = serde_json::Map::new();
371 let bos_id = num_base_tokens;
372 let eos_id = num_base_tokens + 1;
373
374 added_tokens.insert(
375 bos_id.to_string(),
376 serde_json::json!({"content": "[BOS]", "special": true}),
377 );
378 added_tokens.insert(
379 eos_id.to_string(),
380 serde_json::json!({"content": "[EOS]", "special": true}),
381 );
382
383 let config = serde_json::json!({
384 "added_tokens_decoder": added_tokens
385 });
386
387 let file_path = dir.join("tokenizer_config.json");
388 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
389 }
390
391 #[test]
392 fn test_parse_tiktoken_file() {
393 let dir = tempfile::tempdir().unwrap();
394 let file_path = create_test_tiktoken_file(dir.path());
395 let encoder = parse_tiktoken_file(&file_path).unwrap();
396 assert_eq!(encoder.len(), 21);
397 assert_eq!(encoder[b"hello".as_slice()], 19);
398 assert_eq!(encoder[b"world".as_slice()], 20);
399 }
400
401 #[test]
402 fn test_parse_tiktoken_file_missing() {
403 let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
404 assert!(result.is_err());
405 }
406
407 #[test]
408 fn test_tiktoken_from_file() {
409 let dir = tempfile::tempdir().unwrap();
410 let file_path = create_test_tiktoken_file(dir.path());
411
412 let mut special_tokens = FxHashMap::default();
413 special_tokens.insert("[BOS]".to_string(), 21_u32);
414 special_tokens.insert("[EOS]".to_string(), 22_u32);
415
416 let pattern = r"[\w]+|[^\w\s]+|\s+";
418
419 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
420
421 assert_eq!(
422 tokenizer.special_tokens(),
423 &["[BOS]".to_string(), "[EOS]".to_string()]
424 );
425
426 let encoding = tokenizer.encode("hello world").unwrap();
428 let ids = encoding.token_ids();
429 assert!(!ids.is_empty());
430
431 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
433 assert_eq!(decoded, "hello world");
434 }
435
436 #[test]
437 fn test_tiktoken_encoding_variant() {
438 let dir = tempfile::tempdir().unwrap();
439 let file_path = create_test_tiktoken_file(dir.path());
440
441 let special_tokens = FxHashMap::default();
442 let pattern = r"[\w]+|[^\w\s]+|\s+";
443
444 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
445 assert!(tokenizer.special_tokens().is_empty());
446 let encoding = tokenizer.encode("hello").unwrap();
447
448 match &encoding {
450 Encoding::Sp(_) => {}
451 other => panic!("Expected Encoding::Sp, got {:?}", other),
452 }
453 }
454
455 #[test]
456 fn test_tiktoken_skip_special_tokens() {
457 let dir = tempfile::tempdir().unwrap();
458 let file_path = create_test_tiktoken_file(dir.path());
459
460 let mut special_tokens = FxHashMap::default();
461 special_tokens.insert("[BOS]".to_string(), 21_u32);
462 special_tokens.insert("[EOS]".to_string(), 22_u32);
463
464 let pattern = r"[\w]+|[^\w\s]+|\s+";
465
466 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
467
468 let encoding = tokenizer.encode("hello").unwrap();
470 let mut ids = vec![21u32]; ids.extend(encoding.token_ids());
472 ids.push(22); let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
476 assert_eq!(decoded_skip, "hello");
477
478 let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
480 assert!(decoded_all.contains("hello"));
481 }
482
483 #[test]
484 fn test_special_token_ids_reports_configured_ids_sorted() {
485 let dir = tempfile::tempdir().unwrap();
486 let file_path = create_test_tiktoken_file(dir.path());
487
488 let mut special_tokens = FxHashMap::default();
489 special_tokens.insert("[EOS]".to_string(), 22_u32);
490 special_tokens.insert("[BOS]".to_string(), 21_u32);
491
492 let pattern = r"[\w]+|[^\w\s]+|\s+";
493 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
494
495 assert_eq!(tokenizer.special_token_ids().unwrap(), vec![21, 22]);
496 }
497
498 #[test]
499 fn test_vocab_size_and_token_to_id_are_unsupported() {
500 let dir = tempfile::tempdir().unwrap();
501 let file_path = create_test_tiktoken_file(dir.path());
502 let pattern = r"[\w]+|[^\w\s]+|\s+";
503 let tokenizer =
504 TikTokenTokenizer::from_file(&file_path, pattern, FxHashMap::default()).unwrap();
505
506 assert_eq!(tokenizer.vocab_size(), None);
507 assert!(tokenizer.token_to_id("hello").is_err());
508 }
509
510 #[test]
511 fn test_tiktoken_from_file_auto() {
512 let dir = tempfile::tempdir().unwrap();
513 let file_path = create_test_tiktoken_file(dir.path());
514
515 create_test_config(dir.path(), "kimi");
516 create_test_tokenizer_config(dir.path(), 21);
517
518 let mut expected_specials: Vec<String> = load_special_tokens(dir.path(), 21)
519 .unwrap()
520 .into_keys()
521 .collect();
522 expected_specials.sort();
523 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
524 assert_eq!(tokenizer.special_tokens(), expected_specials);
525
526 let encoding = tokenizer.encode("hello world").unwrap();
528 let ids = encoding.token_ids();
529 assert!(!ids.is_empty());
530
531 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
532 assert_eq!(decoded, "hello world");
533 }
534
535 #[test]
536 fn test_detect_bpe_pattern_unknown() {
537 let dir = tempfile::tempdir().unwrap();
538 create_test_config(dir.path(), "unknown_model");
539 let result = detect_bpe_pattern(dir.path());
540 assert!(result.is_err());
541 }
542
543 #[test]
544 fn test_kimi_k3_uses_kimi_bpe_pattern() {
545 let dir = tempfile::tempdir().unwrap();
546 create_test_config(dir.path(), "kimi_k3");
547 assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
548 }
549
550 #[test]
551 fn test_detect_bpe_pattern_kimi_linear() {
552 let dir = tempfile::tempdir().unwrap();
553 create_test_config(dir.path(), "kimi_linear");
554 assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
555 }
556
557 #[test]
558 fn test_load_special_tokens_no_config() {
559 let dir = tempfile::tempdir().unwrap();
560 let tokens = load_special_tokens(dir.path(), 100).unwrap();
561 assert_eq!(tokens.len(), 256);
562 assert_eq!(tokens["<|reserved_token_100|>"], 100);
563 assert_eq!(tokens["<|reserved_token_355|>"], 355);
564 }
565
566 #[test]
567 fn test_load_special_tokens_with_config() {
568 let dir = tempfile::tempdir().unwrap();
569 create_test_tokenizer_config(dir.path(), 100);
570 let tokens = load_special_tokens(dir.path(), 100).unwrap();
571 assert_eq!(tokens["[BOS]"], 100);
572 assert_eq!(tokens["[EOS]"], 101);
573 assert!(tokens.len() > 2);
575 }
576
577 fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
579 let engine = base64::engine::general_purpose::STANDARD;
580 let mut content = String::new();
581
582 let tokens: Vec<(&[u8], u32)> = vec![
583 (b"h", 0),
584 (b"e", 1),
585 (b"l", 2),
586 (b"o", 3),
587 (b" ", 4),
588 (b"hello", 5),
589 ];
590
591 for (token, rank) in &tokens {
592 let encoded = engine.encode(token);
593 content.push_str(&format!("{encoded} {rank}\n"));
594 }
595
596 let byte_tokens: Vec<(Vec<u8>, u32)> =
599 vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
600
601 for (token, rank) in &byte_tokens {
602 let encoded = engine.encode(token);
603 content.push_str(&format!("{encoded} {rank}\n"));
604 }
605
606 let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
608 (vec![0xF0], 200),
609 (vec![0x9F], 201),
610 (vec![0x98], 202),
611 (vec![0x80], 203),
612 ];
613
614 for (token, rank) in &emoji_tokens {
615 let encoded = engine.encode(token);
616 content.push_str(&format!("{encoded} {rank}\n"));
617 }
618
619 let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
622
623 for (token, rank) in &fffd_token {
624 let encoded = engine.encode(token);
625 content.push_str(&format!("{encoded} {rank}\n"));
626 }
627
628 let file_path = dir.join("tiktoken.model");
629 let mut file = std::fs::File::create(&file_path).unwrap();
630 file.write_all(content.as_bytes()).unwrap();
631 file_path.to_str().unwrap().to_string()
632 }
633
634 fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
635 let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
636 let special_tokens = FxHashMap::default();
637 let pattern = r"[\w]+|[^\w\s]+|\s+";
638 TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
639 }
640
641 #[test]
645 fn test_decode_single_incomplete_utf8_byte_does_not_error() {
646 let dir = tempfile::tempdir().unwrap();
647 let tokenizer = create_byte_token_tokenizer(dir.path());
648
649 let result = tokenizer.decode(&[100], false);
650 assert!(
651 result.is_ok(),
652 "decode() should not error on incomplete UTF-8 bytes"
653 );
654 let decode_result = result.unwrap();
655 assert!(
656 decode_result.is_partial(),
657 "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
658 decode_result
659 );
660 }
661
662 #[test]
664 fn test_decode_two_of_three_utf8_bytes_does_not_error() {
665 let dir = tempfile::tempdir().unwrap();
666 let tokenizer = create_byte_token_tokenizer(dir.path());
667
668 let result = tokenizer.decode(&[100, 101], false);
669 assert!(result.is_ok());
670 let decode_result = result.unwrap();
671 assert!(
672 decode_result.is_partial(),
673 "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
674 decode_result
675 );
676 }
677
678 #[test]
682 fn test_decode_complete_multibyte_utf8_produces_correct_char() {
683 let dir = tempfile::tempdir().unwrap();
684 let tokenizer = create_byte_token_tokenizer(dir.path());
685
686 let result = tokenizer.decode(&[100, 101, 102], false);
687 assert!(result.is_ok());
688 assert_eq!(String::from(result.unwrap()), "δ½ ");
689 }
690
691 #[test]
694 fn test_decode_complete_4byte_emoji_from_byte_tokens() {
695 let dir = tempfile::tempdir().unwrap();
696 let tokenizer = create_byte_token_tokenizer(dir.path());
697
698 let result = tokenizer.decode(&[200, 201, 202, 203], false);
699 assert!(result.is_ok());
700 assert_eq!(String::from(result.unwrap()), "π");
701 }
702
703 #[test]
708 fn test_decode_legitimate_replacement_char_token_is_complete() {
709 let dir = tempfile::tempdir().unwrap();
710 let tokenizer = create_byte_token_tokenizer(dir.path());
711
712 let result = tokenizer.decode(&[300], false);
713 assert!(result.is_ok());
714 let decode_result = result.unwrap();
715 assert!(
716 decode_result.is_complete(),
717 "legitimate U+FFFD vocab token must be Complete, got: {:?}",
718 decode_result
719 );
720 assert_eq!(decode_result.as_str(), "\u{FFFD}");
721 }
722
723 #[test]
725 fn test_decode_partial_emoji_does_not_error() {
726 let dir = tempfile::tempdir().unwrap();
727 let tokenizer = create_byte_token_tokenizer(dir.path());
728
729 let result = tokenizer.decode(&[200], false);
730 assert!(result.is_ok());
731 assert!(result.unwrap().is_partial());
732 }
733
734 #[test]
736 fn test_decode_mixed_ascii_and_incomplete_bytes() {
737 let dir = tempfile::tempdir().unwrap();
738 let tokenizer = create_byte_token_tokenizer(dir.path());
739
740 let result = tokenizer.decode(&[5, 100], false);
741 assert!(result.is_ok());
742 let decode_result = result.unwrap();
743 assert!(
744 decode_result.is_partial(),
745 "trailing incomplete byte should produce DecodeResult::Partial"
746 );
747 let text: String = decode_result.into();
748 assert!(
749 text.starts_with("hello"),
750 "should start with 'hello', got: {:?}",
751 text
752 );
753 }
754
755 #[test]
759 fn test_decode_stream_incremental_multibyte_reassembly() {
760 let dir = tempfile::tempdir().unwrap();
761 let tokenizer = create_byte_token_tokenizer(dir.path());
762 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
763
764 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
765
766 let r1 = stream.step(100).unwrap();
767 assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
768
769 let r2 = stream.step(101).unwrap();
770 assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
771
772 let r3 = stream.step(102).unwrap();
773 assert!(r3.is_some(), "third byte should complete the character");
774 assert_eq!(r3.unwrap(), "δ½ ");
775 }
776
777 #[test]
779 fn test_decode_stream_incremental_emoji_reassembly() {
780 let dir = tempfile::tempdir().unwrap();
781 let tokenizer = create_byte_token_tokenizer(dir.path());
782 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
783
784 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
785
786 let r1 = stream.step(200).unwrap();
787 assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
788
789 let r2 = stream.step(201).unwrap();
790 assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
791
792 let r3 = stream.step(202).unwrap();
793 assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
794
795 let r4 = stream.step(203).unwrap();
796 assert!(r4.is_some(), "byte 4/4 should complete the emoji");
797 assert_eq!(r4.unwrap(), "π");
798 }
799
800 #[test]
801 fn test_tiktoken_encode_batch() {
802 let dir = tempfile::tempdir().unwrap();
803 let file_path = create_test_tiktoken_file(dir.path());
804
805 let special_tokens = FxHashMap::default();
806 let pattern = r"[\w]+|[^\w\s]+|\s+";
807
808 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
809
810 let inputs = &["hello", "world"];
811 let encodings = tokenizer.encode_batch(inputs).unwrap();
812 assert_eq!(encodings.len(), 2);
813
814 for (encoding, input) in encodings.iter().zip(inputs.iter()) {
815 let decoded: String = tokenizer
816 .decode(encoding.token_ids(), false)
817 .unwrap()
818 .into();
819 assert_eq!(decoded, *input);
820 }
821 }
822
823 fn create_byte_level_tiktoken_file(dir: &Path) -> String {
826 let engine = base64::engine::general_purpose::STANDARD;
827 let mut content = String::new();
828 for byte_val in 0u16..256 {
829 let encoded = engine.encode([byte_val as u8]);
830 content.push_str(&format!("{encoded} {byte_val}\n"));
831 }
832 let file_path = dir.join("tiktoken.model");
833 std::fs::write(&file_path, &content).unwrap();
834 file_path.to_str().unwrap().to_string()
835 }
836
837 #[test]
838 fn test_segment_encoding_preserves_special_vs_ordinary_boundary() {
839 let dir = tempfile::tempdir().unwrap();
840 let file_path = create_byte_level_tiktoken_file(dir.path());
841 let marker = "<|open|>";
842 let marker_id = 256;
843 let mut special_tokens = FxHashMap::default();
844 special_tokens.insert(marker.to_string(), marker_id);
845 let tokenizer =
846 TikTokenTokenizer::from_file(&file_path, r"(?s).+", special_tokens).unwrap();
847
848 let structural = tokenizer
849 .encode_segments(&[EncodeSegment::control(marker)])
850 .unwrap();
851 let user_text = tokenizer
852 .encode_segments(&[EncodeSegment::ordinary(marker)])
853 .unwrap();
854
855 assert_eq!(structural.token_ids(), &[marker_id]);
856 assert_ne!(user_text.token_ids(), &[marker_id]);
857 let decoded: String = tokenizer
858 .decode(user_text.token_ids(), false)
859 .unwrap()
860 .into();
861 assert_eq!(decoded, marker);
862 }
863
864 fn has_nontrivial_self_overlap(token: &str) -> bool {
865 let bytes = token.as_bytes();
866 (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
867 }
868
869 fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
870 let a = a.as_bytes();
871 let b = b.as_bytes();
872
873 if a.windows(b.len()).any(|window| window == b)
874 || b.windows(a.len()).any(|window| window == a)
875 {
876 return true;
877 }
878
879 let max_overlap = a.len().min(b.len());
880 (1..max_overlap).any(|overlap| {
881 a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
882 })
883 }
884
885 fn assert_special_split_invariant(
886 tokenizer: &TikTokenTokenizer,
887 left: &str,
888 special: &str,
889 right: &str,
890 ) {
891 let whole = tokenizer
892 .encode(&format!("{left}{special}{right}"))
893 .unwrap()
894 .token_ids()
895 .to_vec();
896 let mut split = tokenizer
897 .encode(&format!("{left}{special}"))
898 .unwrap()
899 .token_ids()
900 .to_vec();
901 split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
902 assert_eq!(
903 whole, split,
904 "splitting after registered special token {special:?} must preserve tokenization"
905 );
906 }
907
908 #[test]
909 fn test_registered_special_tokens_are_cache_safe_boundaries() {
910 let dir = tempfile::tempdir().unwrap();
911 let file_path = create_byte_level_tiktoken_file(dir.path());
912 create_test_config(dir.path(), "kimi");
913 create_test_tokenizer_config(dir.path(), 256);
914
915 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
916 let specials = tokenizer.special_tokens();
917 assert!(!specials.is_empty());
918
919 for (index, special) in specials.iter().enumerate() {
920 assert!(
921 !has_nontrivial_self_overlap(special),
922 "registered special token has an ambiguous self-overlap: {special:?}"
923 );
924 for other in &specials[index + 1..] {
925 assert!(
926 !have_ambiguous_overlap(special, other),
927 "registered special tokens overlap ambiguously: {special:?} and {other:?}"
928 );
929 }
930
931 for (left, right) in [
932 ("ordinary prefix ", " ordinary suffix"),
933 ("line before\n", "\tUnicode after: εδΊ¬ π"),
934 ("", " trailing text"),
935 ] {
936 assert_special_split_invariant(&tokenizer, left, special, right);
937 }
938 }
939
940 let bos = specials
941 .iter()
942 .find(|token| token.as_str() == "[BOS]")
943 .unwrap();
944 let eos = specials
945 .iter()
946 .find(|token| token.as_str() == "[EOS]")
947 .unwrap();
948 assert_special_split_invariant(
949 &tokenizer,
950 "prefix ",
951 bos,
952 &format!("{eos}<|reserved_token_258|>Unicode ε°Ύι¨"),
953 );
954 }
955
956 #[test]
967 fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
968 let dir = tempfile::tempdir().unwrap();
969 let file_path = create_byte_level_tiktoken_file(dir.path());
970
971 create_test_config(dir.path(), "kimi");
973
974 create_test_tokenizer_config(dir.path(), 256);
976
977 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
978
979 let single = "<|reserved_token_258|>";
986 let enc = tokenizer.encode(single).unwrap();
987 assert_eq!(
988 enc.token_ids().len(),
989 1,
990 "'{single}' should be 1 special token, got {} tokens: {:?}. \
991 This means fallback naming still uses relative offsets instead of absolute IDs.",
992 enc.token_ids().len(),
993 enc.token_ids()
994 );
995 assert_eq!(enc.token_ids()[0], 258);
996
997 let multi: String = (258u32..268)
1000 .map(|id| format!("<|reserved_token_{id}|>"))
1001 .collect();
1002 let enc_multi = tokenizer.encode(&multi).unwrap();
1003 assert_eq!(
1004 enc_multi.token_ids().len(),
1005 10,
1006 "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
1007 enc_multi.token_ids().len(),
1008 enc_multi.token_ids()
1009 );
1010 let expected_ids: Vec<u32> = (258..268).collect();
1011 assert_eq!(enc_multi.token_ids(), &expected_ids);
1012 }
1013
1014 #[test]
1018 fn test_relative_offset_naming_causes_inflation() {
1019 let dir = tempfile::tempdir().unwrap();
1020 let file_path = create_byte_level_tiktoken_file(dir.path());
1021
1022 let _encoder = parse_tiktoken_file(&file_path).unwrap();
1023 let num_base_tokens = 256usize;
1024
1025 let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
1027 bad_special_tokens.insert("[BOS]".to_string(), 256);
1028 bad_special_tokens.insert("[EOS]".to_string(), 257);
1029 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
1030 let id = num_base_tokens as u32 + i;
1031 if id != 256 && id != 257 {
1032 bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
1034 }
1035 }
1036
1037 let bad_tokenizer =
1038 TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
1039
1040 let input = "<|reserved_token_258|>";
1043 let enc = bad_tokenizer.encode(input).unwrap();
1044 assert!(
1045 enc.token_ids().len() > 1,
1046 "With buggy relative-offset naming, '{}' should NOT be recognized as a \
1047 single special token. Got {} token(s): {:?}",
1048 input,
1049 enc.token_ids().len(),
1050 enc.token_ids()
1051 );
1052
1053 let multi: String = (258u32..268)
1055 .map(|id| format!("<|reserved_token_{id}|>"))
1056 .collect();
1057 let enc_multi = bad_tokenizer.encode(&multi).unwrap();
1058 assert!(
1059 enc_multi.token_ids().len() > 10,
1060 "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
1061 Got {}",
1062 enc_multi.token_ids().len(),
1063 );
1064 }
1065}