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