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" | "kimi_linear" | "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, kimi_linear, 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_detect_bpe_pattern_kimi_linear() {
480 let dir = tempfile::tempdir().unwrap();
481 create_test_config(dir.path(), "kimi_linear");
482 assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
483 }
484
485 #[test]
486 fn test_load_special_tokens_no_config() {
487 let dir = tempfile::tempdir().unwrap();
488 let tokens = load_special_tokens(dir.path(), 100).unwrap();
489 assert_eq!(tokens.len(), 256);
490 assert_eq!(tokens["<|reserved_token_100|>"], 100);
491 assert_eq!(tokens["<|reserved_token_355|>"], 355);
492 }
493
494 #[test]
495 fn test_load_special_tokens_with_config() {
496 let dir = tempfile::tempdir().unwrap();
497 create_test_tokenizer_config(dir.path(), 100);
498 let tokens = load_special_tokens(dir.path(), 100).unwrap();
499 assert_eq!(tokens["[BOS]"], 100);
500 assert_eq!(tokens["[EOS]"], 101);
501 assert!(tokens.len() > 2);
503 }
504
505 fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
507 let engine = base64::engine::general_purpose::STANDARD;
508 let mut content = String::new();
509
510 let tokens: Vec<(&[u8], u32)> = vec![
511 (b"h", 0),
512 (b"e", 1),
513 (b"l", 2),
514 (b"o", 3),
515 (b" ", 4),
516 (b"hello", 5),
517 ];
518
519 for (token, rank) in &tokens {
520 let encoded = engine.encode(token);
521 content.push_str(&format!("{encoded} {rank}\n"));
522 }
523
524 let byte_tokens: Vec<(Vec<u8>, u32)> =
527 vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
528
529 for (token, rank) in &byte_tokens {
530 let encoded = engine.encode(token);
531 content.push_str(&format!("{encoded} {rank}\n"));
532 }
533
534 let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
536 (vec![0xF0], 200),
537 (vec![0x9F], 201),
538 (vec![0x98], 202),
539 (vec![0x80], 203),
540 ];
541
542 for (token, rank) in &emoji_tokens {
543 let encoded = engine.encode(token);
544 content.push_str(&format!("{encoded} {rank}\n"));
545 }
546
547 let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
550
551 for (token, rank) in &fffd_token {
552 let encoded = engine.encode(token);
553 content.push_str(&format!("{encoded} {rank}\n"));
554 }
555
556 let file_path = dir.join("tiktoken.model");
557 let mut file = std::fs::File::create(&file_path).unwrap();
558 file.write_all(content.as_bytes()).unwrap();
559 file_path.to_str().unwrap().to_string()
560 }
561
562 fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
563 let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
564 let special_tokens = FxHashMap::default();
565 let pattern = r"[\w]+|[^\w\s]+|\s+";
566 TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
567 }
568
569 #[test]
573 fn test_decode_single_incomplete_utf8_byte_does_not_error() {
574 let dir = tempfile::tempdir().unwrap();
575 let tokenizer = create_byte_token_tokenizer(dir.path());
576
577 let result = tokenizer.decode(&[100], false);
578 assert!(
579 result.is_ok(),
580 "decode() should not error on incomplete UTF-8 bytes"
581 );
582 let decode_result = result.unwrap();
583 assert!(
584 decode_result.is_partial(),
585 "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
586 decode_result
587 );
588 }
589
590 #[test]
592 fn test_decode_two_of_three_utf8_bytes_does_not_error() {
593 let dir = tempfile::tempdir().unwrap();
594 let tokenizer = create_byte_token_tokenizer(dir.path());
595
596 let result = tokenizer.decode(&[100, 101], false);
597 assert!(result.is_ok());
598 let decode_result = result.unwrap();
599 assert!(
600 decode_result.is_partial(),
601 "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
602 decode_result
603 );
604 }
605
606 #[test]
610 fn test_decode_complete_multibyte_utf8_produces_correct_char() {
611 let dir = tempfile::tempdir().unwrap();
612 let tokenizer = create_byte_token_tokenizer(dir.path());
613
614 let result = tokenizer.decode(&[100, 101, 102], false);
615 assert!(result.is_ok());
616 assert_eq!(String::from(result.unwrap()), "δ½ ");
617 }
618
619 #[test]
622 fn test_decode_complete_4byte_emoji_from_byte_tokens() {
623 let dir = tempfile::tempdir().unwrap();
624 let tokenizer = create_byte_token_tokenizer(dir.path());
625
626 let result = tokenizer.decode(&[200, 201, 202, 203], false);
627 assert!(result.is_ok());
628 assert_eq!(String::from(result.unwrap()), "π");
629 }
630
631 #[test]
636 fn test_decode_legitimate_replacement_char_token_is_complete() {
637 let dir = tempfile::tempdir().unwrap();
638 let tokenizer = create_byte_token_tokenizer(dir.path());
639
640 let result = tokenizer.decode(&[300], false);
641 assert!(result.is_ok());
642 let decode_result = result.unwrap();
643 assert!(
644 decode_result.is_complete(),
645 "legitimate U+FFFD vocab token must be Complete, got: {:?}",
646 decode_result
647 );
648 assert_eq!(decode_result.as_str(), "\u{FFFD}");
649 }
650
651 #[test]
653 fn test_decode_partial_emoji_does_not_error() {
654 let dir = tempfile::tempdir().unwrap();
655 let tokenizer = create_byte_token_tokenizer(dir.path());
656
657 let result = tokenizer.decode(&[200], false);
658 assert!(result.is_ok());
659 assert!(result.unwrap().is_partial());
660 }
661
662 #[test]
664 fn test_decode_mixed_ascii_and_incomplete_bytes() {
665 let dir = tempfile::tempdir().unwrap();
666 let tokenizer = create_byte_token_tokenizer(dir.path());
667
668 let result = tokenizer.decode(&[5, 100], false);
669 assert!(result.is_ok());
670 let decode_result = result.unwrap();
671 assert!(
672 decode_result.is_partial(),
673 "trailing incomplete byte should produce DecodeResult::Partial"
674 );
675 let text: String = decode_result.into();
676 assert!(
677 text.starts_with("hello"),
678 "should start with 'hello', got: {:?}",
679 text
680 );
681 }
682
683 #[test]
687 fn test_decode_stream_incremental_multibyte_reassembly() {
688 let dir = tempfile::tempdir().unwrap();
689 let tokenizer = create_byte_token_tokenizer(dir.path());
690 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
691
692 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
693
694 let r1 = stream.step(100).unwrap();
695 assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
696
697 let r2 = stream.step(101).unwrap();
698 assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
699
700 let r3 = stream.step(102).unwrap();
701 assert!(r3.is_some(), "third byte should complete the character");
702 assert_eq!(r3.unwrap(), "δ½ ");
703 }
704
705 #[test]
707 fn test_decode_stream_incremental_emoji_reassembly() {
708 let dir = tempfile::tempdir().unwrap();
709 let tokenizer = create_byte_token_tokenizer(dir.path());
710 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
711
712 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
713
714 let r1 = stream.step(200).unwrap();
715 assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
716
717 let r2 = stream.step(201).unwrap();
718 assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
719
720 let r3 = stream.step(202).unwrap();
721 assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
722
723 let r4 = stream.step(203).unwrap();
724 assert!(r4.is_some(), "byte 4/4 should complete the emoji");
725 assert_eq!(r4.unwrap(), "π");
726 }
727
728 #[test]
729 fn test_tiktoken_encode_batch() {
730 let dir = tempfile::tempdir().unwrap();
731 let file_path = create_test_tiktoken_file(dir.path());
732
733 let special_tokens = FxHashMap::default();
734 let pattern = r"[\w]+|[^\w\s]+|\s+";
735
736 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
737
738 let inputs = &["hello", "world"];
739 let encodings = tokenizer.encode_batch(inputs).unwrap();
740 assert_eq!(encodings.len(), 2);
741
742 for (encoding, input) in encodings.iter().zip(inputs.iter()) {
743 let decoded: String = tokenizer
744 .decode(encoding.token_ids(), false)
745 .unwrap()
746 .into();
747 assert_eq!(decoded, *input);
748 }
749 }
750
751 fn create_byte_level_tiktoken_file(dir: &Path) -> String {
754 let engine = base64::engine::general_purpose::STANDARD;
755 let mut content = String::new();
756 for byte_val in 0u16..256 {
757 let encoded = engine.encode([byte_val as u8]);
758 content.push_str(&format!("{encoded} {byte_val}\n"));
759 }
760 let file_path = dir.join("tiktoken.model");
761 std::fs::write(&file_path, &content).unwrap();
762 file_path.to_str().unwrap().to_string()
763 }
764
765 fn has_nontrivial_self_overlap(token: &str) -> bool {
766 let bytes = token.as_bytes();
767 (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
768 }
769
770 fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
771 let a = a.as_bytes();
772 let b = b.as_bytes();
773
774 if a.windows(b.len()).any(|window| window == b)
775 || b.windows(a.len()).any(|window| window == a)
776 {
777 return true;
778 }
779
780 let max_overlap = a.len().min(b.len());
781 (1..max_overlap).any(|overlap| {
782 a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
783 })
784 }
785
786 fn assert_special_split_invariant(
787 tokenizer: &TikTokenTokenizer,
788 left: &str,
789 special: &str,
790 right: &str,
791 ) {
792 let whole = tokenizer
793 .encode(&format!("{left}{special}{right}"))
794 .unwrap()
795 .token_ids()
796 .to_vec();
797 let mut split = tokenizer
798 .encode(&format!("{left}{special}"))
799 .unwrap()
800 .token_ids()
801 .to_vec();
802 split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
803 assert_eq!(
804 whole, split,
805 "splitting after registered special token {special:?} must preserve tokenization"
806 );
807 }
808
809 #[test]
810 fn test_registered_special_tokens_are_cache_safe_boundaries() {
811 let dir = tempfile::tempdir().unwrap();
812 let file_path = create_byte_level_tiktoken_file(dir.path());
813 create_test_config(dir.path(), "kimi");
814 create_test_tokenizer_config(dir.path(), 256);
815
816 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
817 let specials = tokenizer.special_tokens();
818 assert!(!specials.is_empty());
819
820 for (index, special) in specials.iter().enumerate() {
821 assert!(
822 !has_nontrivial_self_overlap(special),
823 "registered special token has an ambiguous self-overlap: {special:?}"
824 );
825 for other in &specials[index + 1..] {
826 assert!(
827 !have_ambiguous_overlap(special, other),
828 "registered special tokens overlap ambiguously: {special:?} and {other:?}"
829 );
830 }
831
832 for (left, right) in [
833 ("ordinary prefix ", " ordinary suffix"),
834 ("line before\n", "\tUnicode after: εδΊ¬ π"),
835 ("", " trailing text"),
836 ] {
837 assert_special_split_invariant(&tokenizer, left, special, right);
838 }
839 }
840
841 let bos = specials
842 .iter()
843 .find(|token| token.as_str() == "[BOS]")
844 .unwrap();
845 let eos = specials
846 .iter()
847 .find(|token| token.as_str() == "[EOS]")
848 .unwrap();
849 assert_special_split_invariant(
850 &tokenizer,
851 "prefix ",
852 bos,
853 &format!("{eos}<|reserved_token_258|>Unicode ε°Ύι¨"),
854 );
855 }
856
857 #[test]
868 fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
869 let dir = tempfile::tempdir().unwrap();
870 let file_path = create_byte_level_tiktoken_file(dir.path());
871
872 create_test_config(dir.path(), "kimi");
874
875 create_test_tokenizer_config(dir.path(), 256);
877
878 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
879
880 let single = "<|reserved_token_258|>";
887 let enc = tokenizer.encode(single).unwrap();
888 assert_eq!(
889 enc.token_ids().len(),
890 1,
891 "'{single}' should be 1 special token, got {} tokens: {:?}. \
892 This means fallback naming still uses relative offsets instead of absolute IDs.",
893 enc.token_ids().len(),
894 enc.token_ids()
895 );
896 assert_eq!(enc.token_ids()[0], 258);
897
898 let multi: String = (258u32..268)
901 .map(|id| format!("<|reserved_token_{id}|>"))
902 .collect();
903 let enc_multi = tokenizer.encode(&multi).unwrap();
904 assert_eq!(
905 enc_multi.token_ids().len(),
906 10,
907 "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
908 enc_multi.token_ids().len(),
909 enc_multi.token_ids()
910 );
911 let expected_ids: Vec<u32> = (258..268).collect();
912 assert_eq!(enc_multi.token_ids(), &expected_ids);
913 }
914
915 #[test]
919 fn test_relative_offset_naming_causes_inflation() {
920 let dir = tempfile::tempdir().unwrap();
921 let file_path = create_byte_level_tiktoken_file(dir.path());
922
923 let _encoder = parse_tiktoken_file(&file_path).unwrap();
924 let num_base_tokens = 256usize;
925
926 let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
928 bad_special_tokens.insert("[BOS]".to_string(), 256);
929 bad_special_tokens.insert("[EOS]".to_string(), 257);
930 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
931 let id = num_base_tokens as u32 + i;
932 if id != 256 && id != 257 {
933 bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
935 }
936 }
937
938 let bad_tokenizer =
939 TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
940
941 let input = "<|reserved_token_258|>";
944 let enc = bad_tokenizer.encode(input).unwrap();
945 assert!(
946 enc.token_ids().len() > 1,
947 "With buggy relative-offset naming, '{}' should NOT be recognized as a \
948 single special token. Got {} token(s): {:?}",
949 input,
950 enc.token_ids().len(),
951 enc.token_ids()
952 );
953
954 let multi: String = (258u32..268)
956 .map(|id| format!("<|reserved_token_{id}|>"))
957 .collect();
958 let enc_multi = bad_tokenizer.encode(&multi).unwrap();
959 assert!(
960 enc_multi.token_ids().len() > 10,
961 "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
962 Got {}",
963 enc_multi.token_ids().len(),
964 );
965 }
966}