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