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 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 = r#"[\p{Han}]+|[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?|[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"#;
23
24pub struct TikTokenTokenizer {
25 bpe: CoreBPE,
26 special_token_ids: HashSet<u32>,
27 special_tokens: Vec<String>,
28}
29
30fn 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
112impl Decoder for TikTokenTokenizer {
113 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
114 let ids: Vec<u32> = if skip_special_tokens {
115 token_ids
116 .iter()
117 .filter(|&&id| !self.special_token_ids.contains(&id))
118 .copied()
119 .collect()
120 } else {
121 token_ids.to_vec()
122 };
123
124 let bytes: Vec<u8> = self.bpe._decode_native_and_split(ids).flatten().collect();
133 match String::from_utf8(bytes) {
134 Ok(text) => Ok(DecodeResult::Complete(text)),
135 Err(e) => {
136 let text = String::from_utf8_lossy(e.as_bytes()).into_owned();
137 Ok(DecodeResult::from_decoded(text))
138 }
139 }
140 }
141}
142
143impl Tokenizer for TikTokenTokenizer {
144 fn validate_prefix_cache(&self) -> Result<()> {
145 Ok(())
146 }
147}
148
149fn parse_tiktoken_file(path: &str) -> Result<FxHashMap<Vec<u8>, u32>> {
151 let contents = std::fs::read_to_string(path)
152 .map_err(|err| Error::msg(format!("Failed to read tiktoken file '{path}': {err}")))?;
153
154 let engine = base64::engine::general_purpose::STANDARD;
155 let mut encoder = FxHashMap::default();
156
157 for line in contents.lines() {
158 let line = line.trim();
159 if line.is_empty() {
160 continue;
161 }
162 let mut parts = line.split_whitespace();
163 let token_b64 = parts
164 .next()
165 .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no token): {line}")))?;
166 let rank_str = parts
167 .next()
168 .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no rank): {line}")))?;
169
170 let token_bytes = engine
171 .decode(token_b64)
172 .map_err(|err| Error::msg(format!("Invalid base64 in tiktoken file: {err}")))?;
173 let rank: u32 = rank_str
174 .parse()
175 .map_err(|err| Error::msg(format!("Invalid rank in tiktoken file: {err}")))?;
176
177 encoder.insert(token_bytes, rank);
178 }
179
180 Ok(encoder)
181}
182
183fn detect_bpe_pattern(directory: &Path) -> Result<&'static str> {
185 let model_type: String = crate::file_json_field(&directory.join("config.json"), "model_type")
186 .map_err(|err| {
187 Error::msg(format!("Failed to read model_type from config.json: {err}"))
188 })?;
189
190 match model_type.as_str() {
191 "kimi" | "kimi_k2" | "kimi_k25" | "deepseek_v3" => Ok(KIMI_PATTERN),
197 _ => Err(Error::msg(format!(
198 "Unsupported tiktoken model_type '{model_type}'. \
199 Currently supported: kimi, kimi_k2, kimi_k25, deepseek_v3. \
200 To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
201 with the appropriate BPE regex pattern. \
202 Alternatively, provide a tokenizer.json (HuggingFace format) instead."
203 ))),
204 }
205}
206
207fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHashMap<String, u32>> {
212 let config_path = directory.join("tokenizer_config.json");
213 let mut special_tokens = FxHashMap::default();
214
215 if !config_path.exists() {
216 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
218 let id = num_base_tokens as u32 + i;
219 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
220 }
221 return Ok(special_tokens);
222 }
223
224 let contents = std::fs::read_to_string(&config_path)
225 .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
226
227 let config: serde_json::Value = serde_json::from_str(&contents)
228 .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
229
230 if let Some(added_tokens) = config
231 .get("added_tokens_decoder")
232 .and_then(|v| v.as_object())
233 {
234 for (id_str, token_def) in added_tokens {
235 let id: u32 = id_str.parse().map_err(|err| {
236 Error::msg(format!(
237 "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
238 ))
239 })?;
240
241 let content = token_def
242 .get("content")
243 .and_then(|v| v.as_str())
244 .unwrap_or_else(|| {
245 tracing::warn!("Missing 'content' field for token ID {id}");
247 ""
248 });
249
250 if !content.is_empty() {
251 special_tokens.insert(content.to_string(), id);
252 }
253 }
254
255 let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
257 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
258 let id = num_base_tokens as u32 + i;
259 if !used_ids.contains(&id) {
260 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
261 }
262 }
263 } else {
264 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
266 let id = num_base_tokens as u32 + i;
267 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
268 }
269 }
270
271 Ok(special_tokens)
272}
273
274#[cfg(test)]
275mod tests {
276 use super::*;
277 use crate::DecodeStream;
278 use std::io::Write;
279 use std::sync::Arc;
280
281 fn create_test_tiktoken_file(dir: &Path) -> String {
282 let engine = base64::engine::general_purpose::STANDARD;
283 let mut content = String::new();
284
285 let tokens: Vec<(&[u8], u32)> = vec![
287 (b"h", 0),
288 (b"e", 1),
289 (b"l", 2),
290 (b"o", 3),
291 (b" ", 4),
292 (b"w", 5),
293 (b"r", 6),
294 (b"d", 7),
295 (b"he", 8),
296 (b"ll", 9),
297 (b"lo", 10),
298 (b"wo", 11),
299 (b"rl", 12),
300 (b"hel", 13),
301 (b"llo", 14),
302 (b"wor", 15),
303 (b"hell", 16),
304 (b"ello", 17),
305 (b"worl", 18),
306 (b"hello", 19),
307 (b"world", 20),
308 ];
309
310 for (token, rank) in tokens {
311 let encoded = engine.encode(token);
312 content.push_str(&format!("{encoded} {rank}\n"));
313 }
314
315 let file_path = dir.join("tiktoken.model");
316 let mut file = std::fs::File::create(&file_path).unwrap();
317 file.write_all(content.as_bytes()).unwrap();
318 file_path.to_str().unwrap().to_string()
319 }
320
321 fn create_test_config(dir: &Path, model_type: &str) {
322 let config = serde_json::json!({
323 "model_type": model_type,
324 "max_position_embeddings": 32768,
325 "eos_token_id": [21]
326 });
327 let file_path = dir.join("config.json");
328 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
329 }
330
331 fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
332 let mut added_tokens = serde_json::Map::new();
333 let bos_id = num_base_tokens;
334 let eos_id = num_base_tokens + 1;
335
336 added_tokens.insert(
337 bos_id.to_string(),
338 serde_json::json!({"content": "[BOS]", "special": true}),
339 );
340 added_tokens.insert(
341 eos_id.to_string(),
342 serde_json::json!({"content": "[EOS]", "special": true}),
343 );
344
345 let config = serde_json::json!({
346 "added_tokens_decoder": added_tokens
347 });
348
349 let file_path = dir.join("tokenizer_config.json");
350 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
351 }
352
353 #[test]
354 fn test_parse_tiktoken_file() {
355 let dir = tempfile::tempdir().unwrap();
356 let file_path = create_test_tiktoken_file(dir.path());
357 let encoder = parse_tiktoken_file(&file_path).unwrap();
358 assert_eq!(encoder.len(), 21);
359 assert_eq!(encoder[b"hello".as_slice()], 19);
360 assert_eq!(encoder[b"world".as_slice()], 20);
361 }
362
363 #[test]
364 fn test_parse_tiktoken_file_missing() {
365 let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
366 assert!(result.is_err());
367 }
368
369 #[test]
370 fn test_tiktoken_from_file() {
371 let dir = tempfile::tempdir().unwrap();
372 let file_path = create_test_tiktoken_file(dir.path());
373
374 let mut special_tokens = FxHashMap::default();
375 special_tokens.insert("[BOS]".to_string(), 21_u32);
376 special_tokens.insert("[EOS]".to_string(), 22_u32);
377
378 let pattern = r"[\w]+|[^\w\s]+|\s+";
380
381 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
382
383 assert_eq!(
384 tokenizer.special_tokens(),
385 &["[BOS]".to_string(), "[EOS]".to_string()]
386 );
387
388 let encoding = tokenizer.encode("hello world").unwrap();
390 let ids = encoding.token_ids();
391 assert!(!ids.is_empty());
392
393 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
395 assert_eq!(decoded, "hello world");
396 }
397
398 #[test]
399 fn test_tiktoken_encoding_variant() {
400 let dir = tempfile::tempdir().unwrap();
401 let file_path = create_test_tiktoken_file(dir.path());
402
403 let special_tokens = FxHashMap::default();
404 let pattern = r"[\w]+|[^\w\s]+|\s+";
405
406 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
407 assert!(tokenizer.special_tokens().is_empty());
408 let encoding = tokenizer.encode("hello").unwrap();
409
410 match &encoding {
412 Encoding::Sp(_) => {}
413 other => panic!("Expected Encoding::Sp, got {:?}", other),
414 }
415 }
416
417 #[test]
418 fn test_tiktoken_skip_special_tokens() {
419 let dir = tempfile::tempdir().unwrap();
420 let file_path = create_test_tiktoken_file(dir.path());
421
422 let mut special_tokens = FxHashMap::default();
423 special_tokens.insert("[BOS]".to_string(), 21_u32);
424 special_tokens.insert("[EOS]".to_string(), 22_u32);
425
426 let pattern = r"[\w]+|[^\w\s]+|\s+";
427
428 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
429
430 let encoding = tokenizer.encode("hello").unwrap();
432 let mut ids = vec![21u32]; ids.extend(encoding.token_ids());
434 ids.push(22); let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
438 assert_eq!(decoded_skip, "hello");
439
440 let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
442 assert!(decoded_all.contains("hello"));
443 }
444
445 #[test]
446 fn test_tiktoken_from_file_auto() {
447 let dir = tempfile::tempdir().unwrap();
448 let file_path = create_test_tiktoken_file(dir.path());
449
450 create_test_config(dir.path(), "kimi");
451 create_test_tokenizer_config(dir.path(), 21);
452
453 let mut expected_specials: Vec<String> = load_special_tokens(dir.path(), 21)
454 .unwrap()
455 .into_keys()
456 .collect();
457 expected_specials.sort();
458 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
459 assert_eq!(tokenizer.special_tokens(), expected_specials);
460
461 let encoding = tokenizer.encode("hello world").unwrap();
463 let ids = encoding.token_ids();
464 assert!(!ids.is_empty());
465
466 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
467 assert_eq!(decoded, "hello world");
468 }
469
470 #[test]
471 fn test_detect_bpe_pattern_unknown() {
472 let dir = tempfile::tempdir().unwrap();
473 create_test_config(dir.path(), "unknown_model");
474 let result = detect_bpe_pattern(dir.path());
475 assert!(result.is_err());
476 }
477
478 #[test]
479 fn test_load_special_tokens_no_config() {
480 let dir = tempfile::tempdir().unwrap();
481 let tokens = load_special_tokens(dir.path(), 100).unwrap();
482 assert_eq!(tokens.len(), 256);
483 assert_eq!(tokens["<|reserved_token_100|>"], 100);
484 assert_eq!(tokens["<|reserved_token_355|>"], 355);
485 }
486
487 #[test]
488 fn test_load_special_tokens_with_config() {
489 let dir = tempfile::tempdir().unwrap();
490 create_test_tokenizer_config(dir.path(), 100);
491 let tokens = load_special_tokens(dir.path(), 100).unwrap();
492 assert_eq!(tokens["[BOS]"], 100);
493 assert_eq!(tokens["[EOS]"], 101);
494 assert!(tokens.len() > 2);
496 }
497
498 fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
500 let engine = base64::engine::general_purpose::STANDARD;
501 let mut content = String::new();
502
503 let tokens: Vec<(&[u8], u32)> = vec![
504 (b"h", 0),
505 (b"e", 1),
506 (b"l", 2),
507 (b"o", 3),
508 (b" ", 4),
509 (b"hello", 5),
510 ];
511
512 for (token, rank) in &tokens {
513 let encoded = engine.encode(token);
514 content.push_str(&format!("{encoded} {rank}\n"));
515 }
516
517 let byte_tokens: Vec<(Vec<u8>, u32)> =
520 vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
521
522 for (token, rank) in &byte_tokens {
523 let encoded = engine.encode(token);
524 content.push_str(&format!("{encoded} {rank}\n"));
525 }
526
527 let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
529 (vec![0xF0], 200),
530 (vec![0x9F], 201),
531 (vec![0x98], 202),
532 (vec![0x80], 203),
533 ];
534
535 for (token, rank) in &emoji_tokens {
536 let encoded = engine.encode(token);
537 content.push_str(&format!("{encoded} {rank}\n"));
538 }
539
540 let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
543
544 for (token, rank) in &fffd_token {
545 let encoded = engine.encode(token);
546 content.push_str(&format!("{encoded} {rank}\n"));
547 }
548
549 let file_path = dir.join("tiktoken.model");
550 let mut file = std::fs::File::create(&file_path).unwrap();
551 file.write_all(content.as_bytes()).unwrap();
552 file_path.to_str().unwrap().to_string()
553 }
554
555 fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
556 let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
557 let special_tokens = FxHashMap::default();
558 let pattern = r"[\w]+|[^\w\s]+|\s+";
559 TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
560 }
561
562 #[test]
566 fn test_decode_single_incomplete_utf8_byte_does_not_error() {
567 let dir = tempfile::tempdir().unwrap();
568 let tokenizer = create_byte_token_tokenizer(dir.path());
569
570 let result = tokenizer.decode(&[100], false);
571 assert!(
572 result.is_ok(),
573 "decode() should not error on incomplete UTF-8 bytes"
574 );
575 let decode_result = result.unwrap();
576 assert!(
577 decode_result.is_partial(),
578 "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
579 decode_result
580 );
581 }
582
583 #[test]
585 fn test_decode_two_of_three_utf8_bytes_does_not_error() {
586 let dir = tempfile::tempdir().unwrap();
587 let tokenizer = create_byte_token_tokenizer(dir.path());
588
589 let result = tokenizer.decode(&[100, 101], false);
590 assert!(result.is_ok());
591 let decode_result = result.unwrap();
592 assert!(
593 decode_result.is_partial(),
594 "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
595 decode_result
596 );
597 }
598
599 #[test]
603 fn test_decode_complete_multibyte_utf8_produces_correct_char() {
604 let dir = tempfile::tempdir().unwrap();
605 let tokenizer = create_byte_token_tokenizer(dir.path());
606
607 let result = tokenizer.decode(&[100, 101, 102], false);
608 assert!(result.is_ok());
609 assert_eq!(String::from(result.unwrap()), "δ½ ");
610 }
611
612 #[test]
615 fn test_decode_complete_4byte_emoji_from_byte_tokens() {
616 let dir = tempfile::tempdir().unwrap();
617 let tokenizer = create_byte_token_tokenizer(dir.path());
618
619 let result = tokenizer.decode(&[200, 201, 202, 203], false);
620 assert!(result.is_ok());
621 assert_eq!(String::from(result.unwrap()), "π");
622 }
623
624 #[test]
629 fn test_decode_legitimate_replacement_char_token_is_complete() {
630 let dir = tempfile::tempdir().unwrap();
631 let tokenizer = create_byte_token_tokenizer(dir.path());
632
633 let result = tokenizer.decode(&[300], false);
634 assert!(result.is_ok());
635 let decode_result = result.unwrap();
636 assert!(
637 decode_result.is_complete(),
638 "legitimate U+FFFD vocab token must be Complete, got: {:?}",
639 decode_result
640 );
641 assert_eq!(decode_result.as_str(), "\u{FFFD}");
642 }
643
644 #[test]
646 fn test_decode_partial_emoji_does_not_error() {
647 let dir = tempfile::tempdir().unwrap();
648 let tokenizer = create_byte_token_tokenizer(dir.path());
649
650 let result = tokenizer.decode(&[200], false);
651 assert!(result.is_ok());
652 assert!(result.unwrap().is_partial());
653 }
654
655 #[test]
657 fn test_decode_mixed_ascii_and_incomplete_bytes() {
658 let dir = tempfile::tempdir().unwrap();
659 let tokenizer = create_byte_token_tokenizer(dir.path());
660
661 let result = tokenizer.decode(&[5, 100], false);
662 assert!(result.is_ok());
663 let decode_result = result.unwrap();
664 assert!(
665 decode_result.is_partial(),
666 "trailing incomplete byte should produce DecodeResult::Partial"
667 );
668 let text: String = decode_result.into();
669 assert!(
670 text.starts_with("hello"),
671 "should start with 'hello', got: {:?}",
672 text
673 );
674 }
675
676 #[test]
680 fn test_decode_stream_incremental_multibyte_reassembly() {
681 let dir = tempfile::tempdir().unwrap();
682 let tokenizer = create_byte_token_tokenizer(dir.path());
683 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
684
685 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
686
687 let r1 = stream.step(100).unwrap();
688 assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
689
690 let r2 = stream.step(101).unwrap();
691 assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
692
693 let r3 = stream.step(102).unwrap();
694 assert!(r3.is_some(), "third byte should complete the character");
695 assert_eq!(r3.unwrap(), "δ½ ");
696 }
697
698 #[test]
700 fn test_decode_stream_incremental_emoji_reassembly() {
701 let dir = tempfile::tempdir().unwrap();
702 let tokenizer = create_byte_token_tokenizer(dir.path());
703 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
704
705 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
706
707 let r1 = stream.step(200).unwrap();
708 assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
709
710 let r2 = stream.step(201).unwrap();
711 assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
712
713 let r3 = stream.step(202).unwrap();
714 assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
715
716 let r4 = stream.step(203).unwrap();
717 assert!(r4.is_some(), "byte 4/4 should complete the emoji");
718 assert_eq!(r4.unwrap(), "π");
719 }
720
721 #[test]
722 fn test_tiktoken_encode_batch() {
723 let dir = tempfile::tempdir().unwrap();
724 let file_path = create_test_tiktoken_file(dir.path());
725
726 let special_tokens = FxHashMap::default();
727 let pattern = r"[\w]+|[^\w\s]+|\s+";
728
729 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
730
731 let inputs = &["hello", "world"];
732 let encodings = tokenizer.encode_batch(inputs).unwrap();
733 assert_eq!(encodings.len(), 2);
734
735 for (encoding, input) in encodings.iter().zip(inputs.iter()) {
736 let decoded: String = tokenizer
737 .decode(encoding.token_ids(), false)
738 .unwrap()
739 .into();
740 assert_eq!(decoded, *input);
741 }
742 }
743
744 fn create_byte_level_tiktoken_file(dir: &Path) -> String {
747 let engine = base64::engine::general_purpose::STANDARD;
748 let mut content = String::new();
749 for byte_val in 0u16..256 {
750 let encoded = engine.encode([byte_val as u8]);
751 content.push_str(&format!("{encoded} {byte_val}\n"));
752 }
753 let file_path = dir.join("tiktoken.model");
754 std::fs::write(&file_path, &content).unwrap();
755 file_path.to_str().unwrap().to_string()
756 }
757
758 fn has_nontrivial_self_overlap(token: &str) -> bool {
759 let bytes = token.as_bytes();
760 (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
761 }
762
763 fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
764 let a = a.as_bytes();
765 let b = b.as_bytes();
766
767 if a.windows(b.len()).any(|window| window == b)
768 || b.windows(a.len()).any(|window| window == a)
769 {
770 return true;
771 }
772
773 let max_overlap = a.len().min(b.len());
774 (1..max_overlap).any(|overlap| {
775 a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
776 })
777 }
778
779 fn assert_special_split_invariant(
780 tokenizer: &TikTokenTokenizer,
781 left: &str,
782 special: &str,
783 right: &str,
784 ) {
785 let whole = tokenizer
786 .encode(&format!("{left}{special}{right}"))
787 .unwrap()
788 .token_ids()
789 .to_vec();
790 let mut split = tokenizer
791 .encode(&format!("{left}{special}"))
792 .unwrap()
793 .token_ids()
794 .to_vec();
795 split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
796 assert_eq!(
797 whole, split,
798 "splitting after registered special token {special:?} must preserve tokenization"
799 );
800 }
801
802 #[test]
803 fn test_registered_special_tokens_are_cache_safe_boundaries() {
804 let dir = tempfile::tempdir().unwrap();
805 let file_path = create_byte_level_tiktoken_file(dir.path());
806 create_test_config(dir.path(), "kimi");
807 create_test_tokenizer_config(dir.path(), 256);
808
809 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
810 let specials = tokenizer.special_tokens();
811 assert!(!specials.is_empty());
812
813 for (index, special) in specials.iter().enumerate() {
814 assert!(
815 !has_nontrivial_self_overlap(special),
816 "registered special token has an ambiguous self-overlap: {special:?}"
817 );
818 for other in &specials[index + 1..] {
819 assert!(
820 !have_ambiguous_overlap(special, other),
821 "registered special tokens overlap ambiguously: {special:?} and {other:?}"
822 );
823 }
824
825 for (left, right) in [
826 ("ordinary prefix ", " ordinary suffix"),
827 ("line before\n", "\tUnicode after: εδΊ¬ π"),
828 ("", " trailing text"),
829 ] {
830 assert_special_split_invariant(&tokenizer, left, special, right);
831 }
832 }
833
834 let bos = specials
835 .iter()
836 .find(|token| token.as_str() == "[BOS]")
837 .unwrap();
838 let eos = specials
839 .iter()
840 .find(|token| token.as_str() == "[EOS]")
841 .unwrap();
842 assert_special_split_invariant(
843 &tokenizer,
844 "prefix ",
845 bos,
846 &format!("{eos}<|reserved_token_258|>Unicode ε°Ύι¨"),
847 );
848 }
849
850 #[test]
861 fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
862 let dir = tempfile::tempdir().unwrap();
863 let file_path = create_byte_level_tiktoken_file(dir.path());
864
865 create_test_config(dir.path(), "kimi");
867
868 create_test_tokenizer_config(dir.path(), 256);
870
871 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
872
873 let single = "<|reserved_token_258|>";
880 let enc = tokenizer.encode(single).unwrap();
881 assert_eq!(
882 enc.token_ids().len(),
883 1,
884 "'{single}' should be 1 special token, got {} tokens: {:?}. \
885 This means fallback naming still uses relative offsets instead of absolute IDs.",
886 enc.token_ids().len(),
887 enc.token_ids()
888 );
889 assert_eq!(enc.token_ids()[0], 258);
890
891 let multi: String = (258u32..268)
894 .map(|id| format!("<|reserved_token_{id}|>"))
895 .collect();
896 let enc_multi = tokenizer.encode(&multi).unwrap();
897 assert_eq!(
898 enc_multi.token_ids().len(),
899 10,
900 "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
901 enc_multi.token_ids().len(),
902 enc_multi.token_ids()
903 );
904 let expected_ids: Vec<u32> = (258..268).collect();
905 assert_eq!(enc_multi.token_ids(), &expected_ids);
906 }
907
908 #[test]
912 fn test_relative_offset_naming_causes_inflation() {
913 let dir = tempfile::tempdir().unwrap();
914 let file_path = create_byte_level_tiktoken_file(dir.path());
915
916 let _encoder = parse_tiktoken_file(&file_path).unwrap();
917 let num_base_tokens = 256usize;
918
919 let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
921 bad_special_tokens.insert("[BOS]".to_string(), 256);
922 bad_special_tokens.insert("[EOS]".to_string(), 257);
923 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
924 let id = num_base_tokens as u32 + i;
925 if id != 256 && id != 257 {
926 bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
928 }
929 }
930
931 let bad_tokenizer =
932 TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
933
934 let input = "<|reserved_token_258|>";
937 let enc = bad_tokenizer.encode(input).unwrap();
938 assert!(
939 enc.token_ids().len() > 1,
940 "With buggy relative-offset naming, '{}' should NOT be recognized as a \
941 single special token. Got {} token(s): {:?}",
942 input,
943 enc.token_ids().len(),
944 enc.token_ids()
945 );
946
947 let multi: String = (258u32..268)
949 .map(|id| format!("<|reserved_token_{id}|>"))
950 .collect();
951 let enc_multi = bad_tokenizer.encode(&multi).unwrap();
952 assert!(
953 enc_multi.token_ids().len() > 10,
954 "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
955 Got {}",
956 enc_multi.token_ids().len(),
957 );
958 }
959}