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 = 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 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
182fn 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
216fn 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
242fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHashMap<String, u32>> {
247 let config_path = directory.join("tokenizer_config.json");
248 let mut special_tokens = FxHashMap::default();
249
250 if !config_path.exists() {
251 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
253 let id = num_base_tokens as u32 + i;
254 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
255 }
256 return Ok(special_tokens);
257 }
258
259 let contents = std::fs::read_to_string(&config_path)
260 .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
261
262 let config: serde_json::Value = serde_json::from_str(&contents)
263 .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
264
265 if let Some(added_tokens) = config
266 .get("added_tokens_decoder")
267 .and_then(|v| v.as_object())
268 {
269 for (id_str, token_def) in added_tokens {
270 let id: u32 = id_str.parse().map_err(|err| {
271 Error::msg(format!(
272 "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
273 ))
274 })?;
275
276 let content = token_def
277 .get("content")
278 .and_then(|v| v.as_str())
279 .unwrap_or_else(|| {
280 tracing::warn!("Missing 'content' field for token ID {id}");
282 ""
283 });
284
285 if !content.is_empty() {
286 special_tokens.insert(content.to_string(), id);
287 }
288 }
289
290 let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
292 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
293 let id = num_base_tokens as u32 + i;
294 if !used_ids.contains(&id) {
295 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
296 }
297 }
298 } else {
299 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
301 let id = num_base_tokens as u32 + i;
302 special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
303 }
304 }
305
306 Ok(special_tokens)
307}
308
309#[cfg(test)]
310mod tests {
311 use super::*;
312 use crate::DecodeStream;
313 use std::io::Write;
314 use std::sync::Arc;
315
316 fn create_test_tiktoken_file(dir: &Path) -> String {
317 let engine = base64::engine::general_purpose::STANDARD;
318 let mut content = String::new();
319
320 let tokens: Vec<(&[u8], u32)> = vec![
322 (b"h", 0),
323 (b"e", 1),
324 (b"l", 2),
325 (b"o", 3),
326 (b" ", 4),
327 (b"w", 5),
328 (b"r", 6),
329 (b"d", 7),
330 (b"he", 8),
331 (b"ll", 9),
332 (b"lo", 10),
333 (b"wo", 11),
334 (b"rl", 12),
335 (b"hel", 13),
336 (b"llo", 14),
337 (b"wor", 15),
338 (b"hell", 16),
339 (b"ello", 17),
340 (b"worl", 18),
341 (b"hello", 19),
342 (b"world", 20),
343 ];
344
345 for (token, rank) in tokens {
346 let encoded = engine.encode(token);
347 content.push_str(&format!("{encoded} {rank}\n"));
348 }
349
350 let file_path = dir.join("tiktoken.model");
351 let mut file = std::fs::File::create(&file_path).unwrap();
352 file.write_all(content.as_bytes()).unwrap();
353 file_path.to_str().unwrap().to_string()
354 }
355
356 fn create_test_config(dir: &Path, model_type: &str) {
357 let config = serde_json::json!({
358 "model_type": model_type,
359 "max_position_embeddings": 32768,
360 "eos_token_id": [21]
361 });
362 let file_path = dir.join("config.json");
363 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
364 }
365
366 fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
367 let mut added_tokens = serde_json::Map::new();
368 let bos_id = num_base_tokens;
369 let eos_id = num_base_tokens + 1;
370
371 added_tokens.insert(
372 bos_id.to_string(),
373 serde_json::json!({"content": "[BOS]", "special": true}),
374 );
375 added_tokens.insert(
376 eos_id.to_string(),
377 serde_json::json!({"content": "[EOS]", "special": true}),
378 );
379
380 let config = serde_json::json!({
381 "added_tokens_decoder": added_tokens
382 });
383
384 let file_path = dir.join("tokenizer_config.json");
385 std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
386 }
387
388 #[test]
389 fn test_parse_tiktoken_file() {
390 let dir = tempfile::tempdir().unwrap();
391 let file_path = create_test_tiktoken_file(dir.path());
392 let encoder = parse_tiktoken_file(&file_path).unwrap();
393 assert_eq!(encoder.len(), 21);
394 assert_eq!(encoder[b"hello".as_slice()], 19);
395 assert_eq!(encoder[b"world".as_slice()], 20);
396 }
397
398 #[test]
399 fn test_parse_tiktoken_file_missing() {
400 let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
401 assert!(result.is_err());
402 }
403
404 #[test]
405 fn test_tiktoken_from_file() {
406 let dir = tempfile::tempdir().unwrap();
407 let file_path = create_test_tiktoken_file(dir.path());
408
409 let mut special_tokens = FxHashMap::default();
410 special_tokens.insert("[BOS]".to_string(), 21_u32);
411 special_tokens.insert("[EOS]".to_string(), 22_u32);
412
413 let pattern = r"[\w]+|[^\w\s]+|\s+";
415
416 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
417
418 assert_eq!(
419 tokenizer.special_tokens(),
420 &["[BOS]".to_string(), "[EOS]".to_string()]
421 );
422
423 let encoding = tokenizer.encode("hello world").unwrap();
425 let ids = encoding.token_ids();
426 assert!(!ids.is_empty());
427
428 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
430 assert_eq!(decoded, "hello world");
431 }
432
433 #[test]
434 fn test_tiktoken_encoding_variant() {
435 let dir = tempfile::tempdir().unwrap();
436 let file_path = create_test_tiktoken_file(dir.path());
437
438 let special_tokens = FxHashMap::default();
439 let pattern = r"[\w]+|[^\w\s]+|\s+";
440
441 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
442 assert!(tokenizer.special_tokens().is_empty());
443 let encoding = tokenizer.encode("hello").unwrap();
444
445 match &encoding {
447 Encoding::Sp(_) => {}
448 other => panic!("Expected Encoding::Sp, got {:?}", other),
449 }
450 }
451
452 #[test]
453 fn test_tiktoken_skip_special_tokens() {
454 let dir = tempfile::tempdir().unwrap();
455 let file_path = create_test_tiktoken_file(dir.path());
456
457 let mut special_tokens = FxHashMap::default();
458 special_tokens.insert("[BOS]".to_string(), 21_u32);
459 special_tokens.insert("[EOS]".to_string(), 22_u32);
460
461 let pattern = r"[\w]+|[^\w\s]+|\s+";
462
463 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
464
465 let encoding = tokenizer.encode("hello").unwrap();
467 let mut ids = vec![21u32]; ids.extend(encoding.token_ids());
469 ids.push(22); let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
473 assert_eq!(decoded_skip, "hello");
474
475 let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
477 assert!(decoded_all.contains("hello"));
478 }
479
480 #[test]
481 fn test_special_token_ids_reports_configured_ids_sorted() {
482 let dir = tempfile::tempdir().unwrap();
483 let file_path = create_test_tiktoken_file(dir.path());
484
485 let mut special_tokens = FxHashMap::default();
486 special_tokens.insert("[EOS]".to_string(), 22_u32);
487 special_tokens.insert("[BOS]".to_string(), 21_u32);
488
489 let pattern = r"[\w]+|[^\w\s]+|\s+";
490 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
491
492 assert_eq!(tokenizer.special_token_ids().unwrap(), vec![21, 22]);
493 }
494
495 #[test]
496 fn test_vocab_size_and_token_to_id_are_unsupported() {
497 let dir = tempfile::tempdir().unwrap();
498 let file_path = create_test_tiktoken_file(dir.path());
499 let pattern = r"[\w]+|[^\w\s]+|\s+";
500 let tokenizer =
501 TikTokenTokenizer::from_file(&file_path, pattern, FxHashMap::default()).unwrap();
502
503 assert_eq!(tokenizer.vocab_size(), None);
504 assert!(tokenizer.token_to_id("hello").is_err());
505 }
506
507 #[test]
508 fn test_tiktoken_from_file_auto() {
509 let dir = tempfile::tempdir().unwrap();
510 let file_path = create_test_tiktoken_file(dir.path());
511
512 create_test_config(dir.path(), "kimi");
513 create_test_tokenizer_config(dir.path(), 21);
514
515 let mut expected_specials: Vec<String> = load_special_tokens(dir.path(), 21)
516 .unwrap()
517 .into_keys()
518 .collect();
519 expected_specials.sort();
520 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
521 assert_eq!(tokenizer.special_tokens(), expected_specials);
522
523 let encoding = tokenizer.encode("hello world").unwrap();
525 let ids = encoding.token_ids();
526 assert!(!ids.is_empty());
527
528 let decoded: String = tokenizer.decode(ids, false).unwrap().into();
529 assert_eq!(decoded, "hello world");
530 }
531
532 #[test]
533 fn test_detect_bpe_pattern_unknown() {
534 let dir = tempfile::tempdir().unwrap();
535 create_test_config(dir.path(), "unknown_model");
536 let result = detect_bpe_pattern(dir.path());
537 assert!(result.is_err());
538 }
539
540 #[test]
541 fn test_kimi_k3_uses_kimi_bpe_pattern() {
542 let dir = tempfile::tempdir().unwrap();
543 create_test_config(dir.path(), "kimi_k3");
544 assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
545 }
546
547 #[test]
548 fn test_detect_bpe_pattern_kimi_linear() {
549 let dir = tempfile::tempdir().unwrap();
550 create_test_config(dir.path(), "kimi_linear");
551 assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
552 }
553
554 #[test]
555 fn test_load_special_tokens_no_config() {
556 let dir = tempfile::tempdir().unwrap();
557 let tokens = load_special_tokens(dir.path(), 100).unwrap();
558 assert_eq!(tokens.len(), 256);
559 assert_eq!(tokens["<|reserved_token_100|>"], 100);
560 assert_eq!(tokens["<|reserved_token_355|>"], 355);
561 }
562
563 #[test]
564 fn test_load_special_tokens_with_config() {
565 let dir = tempfile::tempdir().unwrap();
566 create_test_tokenizer_config(dir.path(), 100);
567 let tokens = load_special_tokens(dir.path(), 100).unwrap();
568 assert_eq!(tokens["[BOS]"], 100);
569 assert_eq!(tokens["[EOS]"], 101);
570 assert!(tokens.len() > 2);
572 }
573
574 fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
576 let engine = base64::engine::general_purpose::STANDARD;
577 let mut content = String::new();
578
579 let tokens: Vec<(&[u8], u32)> = vec![
580 (b"h", 0),
581 (b"e", 1),
582 (b"l", 2),
583 (b"o", 3),
584 (b" ", 4),
585 (b"hello", 5),
586 ];
587
588 for (token, rank) in &tokens {
589 let encoded = engine.encode(token);
590 content.push_str(&format!("{encoded} {rank}\n"));
591 }
592
593 let byte_tokens: Vec<(Vec<u8>, u32)> =
596 vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
597
598 for (token, rank) in &byte_tokens {
599 let encoded = engine.encode(token);
600 content.push_str(&format!("{encoded} {rank}\n"));
601 }
602
603 let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
605 (vec![0xF0], 200),
606 (vec![0x9F], 201),
607 (vec![0x98], 202),
608 (vec![0x80], 203),
609 ];
610
611 for (token, rank) in &emoji_tokens {
612 let encoded = engine.encode(token);
613 content.push_str(&format!("{encoded} {rank}\n"));
614 }
615
616 let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
619
620 for (token, rank) in &fffd_token {
621 let encoded = engine.encode(token);
622 content.push_str(&format!("{encoded} {rank}\n"));
623 }
624
625 let file_path = dir.join("tiktoken.model");
626 let mut file = std::fs::File::create(&file_path).unwrap();
627 file.write_all(content.as_bytes()).unwrap();
628 file_path.to_str().unwrap().to_string()
629 }
630
631 fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
632 let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
633 let special_tokens = FxHashMap::default();
634 let pattern = r"[\w]+|[^\w\s]+|\s+";
635 TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
636 }
637
638 #[test]
642 fn test_decode_single_incomplete_utf8_byte_does_not_error() {
643 let dir = tempfile::tempdir().unwrap();
644 let tokenizer = create_byte_token_tokenizer(dir.path());
645
646 let result = tokenizer.decode(&[100], false);
647 assert!(
648 result.is_ok(),
649 "decode() should not error on incomplete UTF-8 bytes"
650 );
651 let decode_result = result.unwrap();
652 assert!(
653 decode_result.is_partial(),
654 "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
655 decode_result
656 );
657 }
658
659 #[test]
661 fn test_decode_two_of_three_utf8_bytes_does_not_error() {
662 let dir = tempfile::tempdir().unwrap();
663 let tokenizer = create_byte_token_tokenizer(dir.path());
664
665 let result = tokenizer.decode(&[100, 101], false);
666 assert!(result.is_ok());
667 let decode_result = result.unwrap();
668 assert!(
669 decode_result.is_partial(),
670 "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
671 decode_result
672 );
673 }
674
675 #[test]
679 fn test_decode_complete_multibyte_utf8_produces_correct_char() {
680 let dir = tempfile::tempdir().unwrap();
681 let tokenizer = create_byte_token_tokenizer(dir.path());
682
683 let result = tokenizer.decode(&[100, 101, 102], false);
684 assert!(result.is_ok());
685 assert_eq!(String::from(result.unwrap()), "δ½ ");
686 }
687
688 #[test]
691 fn test_decode_complete_4byte_emoji_from_byte_tokens() {
692 let dir = tempfile::tempdir().unwrap();
693 let tokenizer = create_byte_token_tokenizer(dir.path());
694
695 let result = tokenizer.decode(&[200, 201, 202, 203], false);
696 assert!(result.is_ok());
697 assert_eq!(String::from(result.unwrap()), "π");
698 }
699
700 #[test]
705 fn test_decode_legitimate_replacement_char_token_is_complete() {
706 let dir = tempfile::tempdir().unwrap();
707 let tokenizer = create_byte_token_tokenizer(dir.path());
708
709 let result = tokenizer.decode(&[300], false);
710 assert!(result.is_ok());
711 let decode_result = result.unwrap();
712 assert!(
713 decode_result.is_complete(),
714 "legitimate U+FFFD vocab token must be Complete, got: {:?}",
715 decode_result
716 );
717 assert_eq!(decode_result.as_str(), "\u{FFFD}");
718 }
719
720 #[test]
722 fn test_decode_partial_emoji_does_not_error() {
723 let dir = tempfile::tempdir().unwrap();
724 let tokenizer = create_byte_token_tokenizer(dir.path());
725
726 let result = tokenizer.decode(&[200], false);
727 assert!(result.is_ok());
728 assert!(result.unwrap().is_partial());
729 }
730
731 #[test]
733 fn test_decode_mixed_ascii_and_incomplete_bytes() {
734 let dir = tempfile::tempdir().unwrap();
735 let tokenizer = create_byte_token_tokenizer(dir.path());
736
737 let result = tokenizer.decode(&[5, 100], false);
738 assert!(result.is_ok());
739 let decode_result = result.unwrap();
740 assert!(
741 decode_result.is_partial(),
742 "trailing incomplete byte should produce DecodeResult::Partial"
743 );
744 let text: String = decode_result.into();
745 assert!(
746 text.starts_with("hello"),
747 "should start with 'hello', got: {:?}",
748 text
749 );
750 }
751
752 #[test]
756 fn test_decode_stream_incremental_multibyte_reassembly() {
757 let dir = tempfile::tempdir().unwrap();
758 let tokenizer = create_byte_token_tokenizer(dir.path());
759 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
760
761 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
762
763 let r1 = stream.step(100).unwrap();
764 assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
765
766 let r2 = stream.step(101).unwrap();
767 assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
768
769 let r3 = stream.step(102).unwrap();
770 assert!(r3.is_some(), "third byte should complete the character");
771 assert_eq!(r3.unwrap(), "δ½ ");
772 }
773
774 #[test]
776 fn test_decode_stream_incremental_emoji_reassembly() {
777 let dir = tempfile::tempdir().unwrap();
778 let tokenizer = create_byte_token_tokenizer(dir.path());
779 let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
780
781 let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
782
783 let r1 = stream.step(200).unwrap();
784 assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
785
786 let r2 = stream.step(201).unwrap();
787 assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
788
789 let r3 = stream.step(202).unwrap();
790 assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
791
792 let r4 = stream.step(203).unwrap();
793 assert!(r4.is_some(), "byte 4/4 should complete the emoji");
794 assert_eq!(r4.unwrap(), "π");
795 }
796
797 #[test]
798 fn test_tiktoken_encode_batch() {
799 let dir = tempfile::tempdir().unwrap();
800 let file_path = create_test_tiktoken_file(dir.path());
801
802 let special_tokens = FxHashMap::default();
803 let pattern = r"[\w]+|[^\w\s]+|\s+";
804
805 let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
806
807 let inputs = &["hello", "world"];
808 let encodings = tokenizer.encode_batch(inputs).unwrap();
809 assert_eq!(encodings.len(), 2);
810
811 for (encoding, input) in encodings.iter().zip(inputs.iter()) {
812 let decoded: String = tokenizer
813 .decode(encoding.token_ids(), false)
814 .unwrap()
815 .into();
816 assert_eq!(decoded, *input);
817 }
818 }
819
820 fn create_byte_level_tiktoken_file(dir: &Path) -> String {
823 let engine = base64::engine::general_purpose::STANDARD;
824 let mut content = String::new();
825 for byte_val in 0u16..256 {
826 let encoded = engine.encode([byte_val as u8]);
827 content.push_str(&format!("{encoded} {byte_val}\n"));
828 }
829 let file_path = dir.join("tiktoken.model");
830 std::fs::write(&file_path, &content).unwrap();
831 file_path.to_str().unwrap().to_string()
832 }
833
834 #[test]
835 fn test_segment_encoding_preserves_special_vs_ordinary_boundary() {
836 let dir = tempfile::tempdir().unwrap();
837 let file_path = create_byte_level_tiktoken_file(dir.path());
838 let marker = "<|open|>";
839 let marker_id = 256;
840 let mut special_tokens = FxHashMap::default();
841 special_tokens.insert(marker.to_string(), marker_id);
842 let tokenizer =
843 TikTokenTokenizer::from_file(&file_path, r"(?s).+", special_tokens).unwrap();
844
845 let structural = tokenizer
846 .encode_segments(&[EncodeSegment::control(marker)])
847 .unwrap();
848 let user_text = tokenizer
849 .encode_segments(&[EncodeSegment::ordinary(marker)])
850 .unwrap();
851
852 assert_eq!(structural.token_ids(), &[marker_id]);
853 assert_ne!(user_text.token_ids(), &[marker_id]);
854 let decoded: String = tokenizer
855 .decode(user_text.token_ids(), false)
856 .unwrap()
857 .into();
858 assert_eq!(decoded, marker);
859 }
860
861 fn has_nontrivial_self_overlap(token: &str) -> bool {
862 let bytes = token.as_bytes();
863 (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
864 }
865
866 fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
867 let a = a.as_bytes();
868 let b = b.as_bytes();
869
870 if a.windows(b.len()).any(|window| window == b)
871 || b.windows(a.len()).any(|window| window == a)
872 {
873 return true;
874 }
875
876 let max_overlap = a.len().min(b.len());
877 (1..max_overlap).any(|overlap| {
878 a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
879 })
880 }
881
882 fn assert_special_split_invariant(
883 tokenizer: &TikTokenTokenizer,
884 left: &str,
885 special: &str,
886 right: &str,
887 ) {
888 let whole = tokenizer
889 .encode(&format!("{left}{special}{right}"))
890 .unwrap()
891 .token_ids()
892 .to_vec();
893 let mut split = tokenizer
894 .encode(&format!("{left}{special}"))
895 .unwrap()
896 .token_ids()
897 .to_vec();
898 split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
899 assert_eq!(
900 whole, split,
901 "splitting after registered special token {special:?} must preserve tokenization"
902 );
903 }
904
905 #[test]
906 fn test_registered_special_tokens_are_cache_safe_boundaries() {
907 let dir = tempfile::tempdir().unwrap();
908 let file_path = create_byte_level_tiktoken_file(dir.path());
909 create_test_config(dir.path(), "kimi");
910 create_test_tokenizer_config(dir.path(), 256);
911
912 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
913 let specials = tokenizer.special_tokens();
914 assert!(!specials.is_empty());
915
916 for (index, special) in specials.iter().enumerate() {
917 assert!(
918 !has_nontrivial_self_overlap(special),
919 "registered special token has an ambiguous self-overlap: {special:?}"
920 );
921 for other in &specials[index + 1..] {
922 assert!(
923 !have_ambiguous_overlap(special, other),
924 "registered special tokens overlap ambiguously: {special:?} and {other:?}"
925 );
926 }
927
928 for (left, right) in [
929 ("ordinary prefix ", " ordinary suffix"),
930 ("line before\n", "\tUnicode after: εδΊ¬ π"),
931 ("", " trailing text"),
932 ] {
933 assert_special_split_invariant(&tokenizer, left, special, right);
934 }
935 }
936
937 let bos = specials
938 .iter()
939 .find(|token| token.as_str() == "[BOS]")
940 .unwrap();
941 let eos = specials
942 .iter()
943 .find(|token| token.as_str() == "[EOS]")
944 .unwrap();
945 assert_special_split_invariant(
946 &tokenizer,
947 "prefix ",
948 bos,
949 &format!("{eos}<|reserved_token_258|>Unicode ε°Ύι¨"),
950 );
951 }
952
953 #[test]
964 fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
965 let dir = tempfile::tempdir().unwrap();
966 let file_path = create_byte_level_tiktoken_file(dir.path());
967
968 create_test_config(dir.path(), "kimi");
970
971 create_test_tokenizer_config(dir.path(), 256);
973
974 let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
975
976 let single = "<|reserved_token_258|>";
983 let enc = tokenizer.encode(single).unwrap();
984 assert_eq!(
985 enc.token_ids().len(),
986 1,
987 "'{single}' should be 1 special token, got {} tokens: {:?}. \
988 This means fallback naming still uses relative offsets instead of absolute IDs.",
989 enc.token_ids().len(),
990 enc.token_ids()
991 );
992 assert_eq!(enc.token_ids()[0], 258);
993
994 let multi: String = (258u32..268)
997 .map(|id| format!("<|reserved_token_{id}|>"))
998 .collect();
999 let enc_multi = tokenizer.encode(&multi).unwrap();
1000 assert_eq!(
1001 enc_multi.token_ids().len(),
1002 10,
1003 "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
1004 enc_multi.token_ids().len(),
1005 enc_multi.token_ids()
1006 );
1007 let expected_ids: Vec<u32> = (258..268).collect();
1008 assert_eq!(enc_multi.token_ids(), &expected_ids);
1009 }
1010
1011 #[test]
1015 fn test_relative_offset_naming_causes_inflation() {
1016 let dir = tempfile::tempdir().unwrap();
1017 let file_path = create_byte_level_tiktoken_file(dir.path());
1018
1019 let _encoder = parse_tiktoken_file(&file_path).unwrap();
1020 let num_base_tokens = 256usize;
1021
1022 let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
1024 bad_special_tokens.insert("[BOS]".to_string(), 256);
1025 bad_special_tokens.insert("[EOS]".to_string(), 257);
1026 for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
1027 let id = num_base_tokens as u32 + i;
1028 if id != 256 && id != 257 {
1029 bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
1031 }
1032 }
1033
1034 let bad_tokenizer =
1035 TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
1036
1037 let input = "<|reserved_token_258|>";
1040 let enc = bad_tokenizer.encode(input).unwrap();
1041 assert!(
1042 enc.token_ids().len() > 1,
1043 "With buggy relative-offset naming, '{}' should NOT be recognized as a \
1044 single special token. Got {} token(s): {:?}",
1045 input,
1046 enc.token_ids().len(),
1047 enc.token_ids()
1048 );
1049
1050 let multi: String = (258u32..268)
1052 .map(|id| format!("<|reserved_token_{id}|>"))
1053 .collect();
1054 let enc_multi = bad_tokenizer.encode(&multi).unwrap();
1055 assert!(
1056 enc_multi.token_ids().len() > 10,
1057 "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
1058 Got {}",
1059 enc_multi.token_ids().len(),
1060 );
1061 }
1062}