Skip to main content

dynamo_tokenizers/
tiktoken.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use 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
17/// Number of reserved special-token slots to generate when filling gaps in the vocabulary.
18/// Most tiktoken-based models reserve 256 IDs above the base vocabulary for special tokens.
19const DEFAULT_NUM_RESERVED_SPECIAL_TOKENS: u32 = 256;
20
21/// Kimi BPE pattern from moonshotai/Kimi-K2-Instruct/tokenization_kimi.py
22const 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    /// Create a TikTokenTokenizer from a tiktoken model file.
39    ///
40    /// # Arguments
41    /// * `path` - Path to the `.model` or `.tiktoken` file (base64 rank-per-line format)
42    /// * `pattern` - BPE regex pattern string
43    /// * `special_tokens` - Map of special token strings to their IDs
44    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    /// Create a TikTokenTokenizer from a tiktoken model file, auto-detecting
64    /// the BPE pattern from `config.json` and special tokens from `tokenizer_config.json`.
65    ///
66    /// The tiktoken file and config files must be in the same directory.
67    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        // Use max rank + 1 (not len) to avoid ID collisions with sparse/non-contiguous ranks
76        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    /// Atomic special-token strings registered with the underlying TikToken BPE.
92    ///
93    /// [`CoreBPE::encode_with_special_tokens`] recognizes these strings outside ordinary
94    /// BPE. For non-overlapping token spellings, splitting immediately after a match
95    /// preserves tokenization.
96    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        // Try strict UTF-8 first: valid bytes get `Complete` with zero extra allocation
125        // (takes ownership of the Vec). This correctly handles vocabulary tokens whose
126        // raw bytes are EF BF BD (legitimate U+FFFD) -- they are valid UTF-8 and must
127        // not be confused with incomplete multi-byte sequences.
128        //
129        // On failure, fall back to lossy conversion so partial multi-byte sequences
130        // become U+FFFD, then classify via the trailing-FFFD heuristic. This path is
131        // only hit during incremental detokenization of byte-fallback tokens.
132        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
149/// Parse a tiktoken model file (base64-encoded token + rank per line).
150fn 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
183/// Detect the BPE pattern for a model by reading `model_type` from `config.json`.
184fn 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        // baseten-admin/Kimi-2.5-text-nvfp4-v3 model has model_type: "deepseek_v3" in its config.json
192        // because Kimi K2.5 is built on the DeepSeek V3 architecture.
193        // it still ships the Kimi tiktoken tokenizer file, so the KIMI_PATTERN BPE regex is the
194        // correct pattern to use.  No pure DeepSeek V3 model uses tiktoken.model files
195        // (they use tokenizer.json instead) so this match is safe.
196        "kimi" | "kimi_k2" | "kimi_k25" | "deepseek_v3" => Ok(KIMI_PATTERN),
197        _ => Err(Error::msg(format!(
198            "Unsupported tiktoken model_type '{model_type}'. \
199             Currently supported: kimi, kimi_k2, kimi_k25, deepseek_v3. \
200             To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
201             with the appropriate BPE regex pattern. \
202             Alternatively, provide a tokenizer.json (HuggingFace format) instead."
203        ))),
204    }
205}
206
207/// Load special tokens from `tokenizer_config.json` in the model directory.
208///
209/// Reads the `added_tokens_decoder` field which maps string token IDs to token definitions.
210/// Falls back to generating `<|reserved_token_{id}|>` names for unmapped IDs.
211fn 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        // No tokenizer_config.json β€” generate default reserved tokens
217        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                    // This shouldn't happen in well-formed configs, but handle gracefully
246                    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        // Fill in any gaps with reserved tokens for the expected range
256        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        // No added_tokens_decoder β€” generate default reserved tokens
265        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        // Create some simple token entries: single bytes with sequential ranks
286        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        // Use a simple pattern for testing
379        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        // Test encode
389        let encoding = tokenizer.encode("hello world").unwrap();
390        let ids = encoding.token_ids();
391        assert!(!ids.is_empty());
392
393        // Test decode roundtrip
394        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        // Verify it produces the Sp variant
411        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        // Encode hello and prepend/append special tokens
431        let encoding = tokenizer.encode("hello").unwrap();
432        let mut ids = vec![21u32]; // [BOS]
433        ids.extend(encoding.token_ids());
434        ids.push(22); // [EOS]
435
436        // Decode with skip_special_tokens=true should strip special tokens
437        let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
438        assert_eq!(decoded_skip, "hello");
439
440        // Decode with skip_special_tokens=false should include them
441        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        // Basic encode/decode roundtrip
462        let encoding = tokenizer.encode("hello world").unwrap();
463        let ids = encoding.token_ids();
464        assert!(!ids.is_empty());
465
466        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
467        assert_eq!(decoded, "hello world");
468    }
469
470    #[test]
471    fn test_detect_bpe_pattern_unknown() {
472        let dir = tempfile::tempdir().unwrap();
473        create_test_config(dir.path(), "unknown_model");
474        let result = detect_bpe_pattern(dir.path());
475        assert!(result.is_err());
476    }
477
478    #[test]
479    fn test_load_special_tokens_no_config() {
480        let dir = tempfile::tempdir().unwrap();
481        let tokens = load_special_tokens(dir.path(), 100).unwrap();
482        assert_eq!(tokens.len(), 256);
483        assert_eq!(tokens["<|reserved_token_100|>"], 100);
484        assert_eq!(tokens["<|reserved_token_355|>"], 355);
485    }
486
487    #[test]
488    fn test_load_special_tokens_with_config() {
489        let dir = tempfile::tempdir().unwrap();
490        create_test_tokenizer_config(dir.path(), 100);
491        let tokens = load_special_tokens(dir.path(), 100).unwrap();
492        assert_eq!(tokens["[BOS]"], 100);
493        assert_eq!(tokens["[EOS]"], 101);
494        // Should also have reserved tokens filling gaps
495        assert!(tokens.len() > 2);
496    }
497
498    /// Helper: create a tiktoken file that includes raw byte tokens (byte fallback tokens).
499    fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
500        let engine = base64::engine::general_purpose::STANDARD;
501        let mut content = String::new();
502
503        let tokens: Vec<(&[u8], u32)> = vec![
504            (b"h", 0),
505            (b"e", 1),
506            (b"l", 2),
507            (b"o", 3),
508            (b" ", 4),
509            (b"hello", 5),
510        ];
511
512        for (token, rank) in &tokens {
513            let encoded = engine.encode(token);
514            content.push_str(&format!("{encoded} {rank}\n"));
515        }
516
517        // Byte-fallback tokens: individual bytes that form CJK character "δ½ " (U+4F60)
518        // UTF-8 encoding: 0xE4 0xBD 0xA0
519        let byte_tokens: Vec<(Vec<u8>, u32)> =
520            vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
521
522        for (token, rank) in &byte_tokens {
523            let encoded = engine.encode(token);
524            content.push_str(&format!("{encoded} {rank}\n"));
525        }
526
527        // Bytes for emoji "πŸ˜€" (U+1F600) β€” 4-byte UTF-8: 0xF0 0x9F 0x98 0x80
528        let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
529            (vec![0xF0], 200),
530            (vec![0x9F], 201),
531            (vec![0x98], 202),
532            (vec![0x80], 203),
533        ];
534
535        for (token, rank) in &emoji_tokens {
536            let encoded = engine.encode(token);
537            content.push_str(&format!("{encoded} {rank}\n"));
538        }
539
540        // Legitimate U+FFFD token: valid UTF-8 bytes EF BF BD (replacement character
541        // as an actual vocabulary entry, not an artifact of lossy conversion)
542        let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
543
544        for (token, rank) in &fffd_token {
545            let encoded = engine.encode(token);
546            content.push_str(&format!("{encoded} {rank}\n"));
547        }
548
549        let file_path = dir.join("tiktoken.model");
550        let mut file = std::fs::File::create(&file_path).unwrap();
551        file.write_all(content.as_bytes()).unwrap();
552        file_path.to_str().unwrap().to_string()
553    }
554
555    fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
556        let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
557        let special_tokens = FxHashMap::default();
558        let pattern = r"[\w]+|[^\w\s]+|\s+";
559        TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
560    }
561
562    /// Reproduces the original panic: decoding a single byte-fallback token that is
563    /// part of a multi-byte UTF-8 character. Before the fix, CoreBPE::decode() would
564    /// call String::from_utf8() on [0xE4] and error with "incomplete utf-8 byte sequence".
565    #[test]
566    fn test_decode_single_incomplete_utf8_byte_does_not_error() {
567        let dir = tempfile::tempdir().unwrap();
568        let tokenizer = create_byte_token_tokenizer(dir.path());
569
570        let result = tokenizer.decode(&[100], false);
571        assert!(
572            result.is_ok(),
573            "decode() should not error on incomplete UTF-8 bytes"
574        );
575        let decode_result = result.unwrap();
576        assert!(
577            decode_result.is_partial(),
578            "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
579            decode_result
580        );
581    }
582
583    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
584    #[test]
585    fn test_decode_two_of_three_utf8_bytes_does_not_error() {
586        let dir = tempfile::tempdir().unwrap();
587        let tokenizer = create_byte_token_tokenizer(dir.path());
588
589        let result = tokenizer.decode(&[100, 101], false);
590        assert!(result.is_ok());
591        let decode_result = result.unwrap();
592        assert!(
593            decode_result.is_partial(),
594            "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
595            decode_result
596        );
597    }
598
599    /// When all bytes of a multi-byte character are present, the concatenated bytes form
600    /// valid UTF-8, so this test passes both before and after the fix. It serves as a
601    /// correctness check that the lossy conversion doesn't corrupt complete characters.
602    #[test]
603    fn test_decode_complete_multibyte_utf8_produces_correct_char() {
604        let dir = tempfile::tempdir().unwrap();
605        let tokenizer = create_byte_token_tokenizer(dir.path());
606
607        let result = tokenizer.decode(&[100, 101, 102], false);
608        assert!(result.is_ok());
609        assert_eq!(String::from(result.unwrap()), "δ½ ");
610    }
611
612    /// All 4 emoji bytes together form valid UTF-8, so this passes both before and after
613    /// the fix. Validates that lossy conversion doesn't alter complete multi-byte sequences.
614    #[test]
615    fn test_decode_complete_4byte_emoji_from_byte_tokens() {
616        let dir = tempfile::tempdir().unwrap();
617        let tokenizer = create_byte_token_tokenizer(dir.path());
618
619        let result = tokenizer.decode(&[200, 201, 202, 203], false);
620        assert!(result.is_ok());
621        assert_eq!(String::from(result.unwrap()), "πŸ˜€");
622    }
623
624    /// Regression test: a vocabulary token whose raw bytes are EF BF BD (the valid
625    /// UTF-8 encoding of U+FFFD) must decode as `Complete`, not `Partial`. Before the
626    /// from_utf8 fast-path fix, from_utf8_lossy + the trailing-FFFD heuristic would
627    /// misclassify this as Partial, causing the incremental decoder to suppress it.
628    #[test]
629    fn test_decode_legitimate_replacement_char_token_is_complete() {
630        let dir = tempfile::tempdir().unwrap();
631        let tokenizer = create_byte_token_tokenizer(dir.path());
632
633        let result = tokenizer.decode(&[300], false);
634        assert!(result.is_ok());
635        let decode_result = result.unwrap();
636        assert!(
637            decode_result.is_complete(),
638            "legitimate U+FFFD vocab token must be Complete, got: {:?}",
639            decode_result
640        );
641        assert_eq!(decode_result.as_str(), "\u{FFFD}");
642    }
643
644    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
645    #[test]
646    fn test_decode_partial_emoji_does_not_error() {
647        let dir = tempfile::tempdir().unwrap();
648        let tokenizer = create_byte_token_tokenizer(dir.path());
649
650        let result = tokenizer.decode(&[200], false);
651        assert!(result.is_ok());
652        assert!(result.unwrap().is_partial());
653    }
654
655    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
656    #[test]
657    fn test_decode_mixed_ascii_and_incomplete_bytes() {
658        let dir = tempfile::tempdir().unwrap();
659        let tokenizer = create_byte_token_tokenizer(dir.path());
660
661        let result = tokenizer.decode(&[5, 100], false);
662        assert!(result.is_ok());
663        let decode_result = result.unwrap();
664        assert!(
665            decode_result.is_partial(),
666            "trailing incomplete byte should produce DecodeResult::Partial"
667        );
668        let text: String = decode_result.into();
669        assert!(
670            text.starts_with("hello"),
671            "should start with 'hello', got: {:?}",
672            text
673        );
674    }
675
676    /// End-to-end incremental detokenization: DecodeStream buffers partial bytes,
677    /// emits the complete character once all bytes arrive.
678    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
679    #[test]
680    fn test_decode_stream_incremental_multibyte_reassembly() {
681        let dir = tempfile::tempdir().unwrap();
682        let tokenizer = create_byte_token_tokenizer(dir.path());
683        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
684
685        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
686
687        let r1 = stream.step(100).unwrap();
688        assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
689
690        let r2 = stream.step(101).unwrap();
691        assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
692
693        let r3 = stream.step(102).unwrap();
694        assert!(r3.is_some(), "third byte should complete the character");
695        assert_eq!(r3.unwrap(), "δ½ ");
696    }
697
698    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
699    #[test]
700    fn test_decode_stream_incremental_emoji_reassembly() {
701        let dir = tempfile::tempdir().unwrap();
702        let tokenizer = create_byte_token_tokenizer(dir.path());
703        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
704
705        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
706
707        let r1 = stream.step(200).unwrap();
708        assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
709
710        let r2 = stream.step(201).unwrap();
711        assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
712
713        let r3 = stream.step(202).unwrap();
714        assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
715
716        let r4 = stream.step(203).unwrap();
717        assert!(r4.is_some(), "byte 4/4 should complete the emoji");
718        assert_eq!(r4.unwrap(), "πŸ˜€");
719    }
720
721    #[test]
722    fn test_tiktoken_encode_batch() {
723        let dir = tempfile::tempdir().unwrap();
724        let file_path = create_test_tiktoken_file(dir.path());
725
726        let special_tokens = FxHashMap::default();
727        let pattern = r"[\w]+|[^\w\s]+|\s+";
728
729        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
730
731        let inputs = &["hello", "world"];
732        let encodings = tokenizer.encode_batch(inputs).unwrap();
733        assert_eq!(encodings.len(), 2);
734
735        for (encoding, input) in encodings.iter().zip(inputs.iter()) {
736            let decoded: String = tokenizer
737                .decode(encoding.token_ids(), false)
738                .unwrap()
739                .into();
740            assert_eq!(decoded, *input);
741        }
742    }
743
744    /// Helper: create a tiktoken file containing all 256 single-byte tokens (ranks 0..255).
745    /// This gives a complete byte-level base vocabulary so any ASCII string can be encoded.
746    fn create_byte_level_tiktoken_file(dir: &Path) -> String {
747        let engine = base64::engine::general_purpose::STANDARD;
748        let mut content = String::new();
749        for byte_val in 0u16..256 {
750            let encoded = engine.encode([byte_val as u8]);
751            content.push_str(&format!("{encoded} {byte_val}\n"));
752        }
753        let file_path = dir.join("tiktoken.model");
754        std::fs::write(&file_path, &content).unwrap();
755        file_path.to_str().unwrap().to_string()
756    }
757
758    fn has_nontrivial_self_overlap(token: &str) -> bool {
759        let bytes = token.as_bytes();
760        (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
761    }
762
763    fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
764        let a = a.as_bytes();
765        let b = b.as_bytes();
766
767        if a.windows(b.len()).any(|window| window == b)
768            || b.windows(a.len()).any(|window| window == a)
769        {
770            return true;
771        }
772
773        let max_overlap = a.len().min(b.len());
774        (1..max_overlap).any(|overlap| {
775            a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
776        })
777    }
778
779    fn assert_special_split_invariant(
780        tokenizer: &TikTokenTokenizer,
781        left: &str,
782        special: &str,
783        right: &str,
784    ) {
785        let whole = tokenizer
786            .encode(&format!("{left}{special}{right}"))
787            .unwrap()
788            .token_ids()
789            .to_vec();
790        let mut split = tokenizer
791            .encode(&format!("{left}{special}"))
792            .unwrap()
793            .token_ids()
794            .to_vec();
795        split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
796        assert_eq!(
797            whole, split,
798            "splitting after registered special token {special:?} must preserve tokenization"
799        );
800    }
801
802    #[test]
803    fn test_registered_special_tokens_are_cache_safe_boundaries() {
804        let dir = tempfile::tempdir().unwrap();
805        let file_path = create_byte_level_tiktoken_file(dir.path());
806        create_test_config(dir.path(), "kimi");
807        create_test_tokenizer_config(dir.path(), 256);
808
809        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
810        let specials = tokenizer.special_tokens();
811        assert!(!specials.is_empty());
812
813        for (index, special) in specials.iter().enumerate() {
814            assert!(
815                !has_nontrivial_self_overlap(special),
816                "registered special token has an ambiguous self-overlap: {special:?}"
817            );
818            for other in &specials[index + 1..] {
819                assert!(
820                    !have_ambiguous_overlap(special, other),
821                    "registered special tokens overlap ambiguously: {special:?} and {other:?}"
822                );
823            }
824
825            for (left, right) in [
826                ("ordinary prefix ", " ordinary suffix"),
827                ("line before\n", "\tUnicode after: εŒ—δΊ¬ πŸ˜€"),
828                ("", " trailing text"),
829            ] {
830                assert_special_split_invariant(&tokenizer, left, special, right);
831            }
832        }
833
834        let bos = specials
835            .iter()
836            .find(|token| token.as_str() == "[BOS]")
837            .unwrap();
838        let eos = specials
839            .iter()
840            .find(|token| token.as_str() == "[EOS]")
841            .unwrap();
842        assert_special_split_invariant(
843            &tokenizer,
844            "prefix ",
845            bos,
846            &format!("{eos}<|reserved_token_258|>Unicode 尾部"),
847        );
848    }
849
850    /// Regression test for Kimi K2.5 tokenizer inflation.
851    ///
852    /// Python's tokenization_kimi.py names unnamed reserved tokens by absolute ID:
853    ///   `<|reserved_token_{absolute_id}|>`
854    ///
855    /// The Rust code previously used relative offsets (0..255) as the naming index,
856    /// so when a prompt contained `<|reserved_token_258|>` the Rust tokenizer did NOT
857    /// recognize it as a special token. Each occurrence was encoded as multiple BPE
858    /// tokens instead of 1, inflating an 8192-token prompt to 9038 tokens and causing
859    /// TRT-LLM to reject the request.
860    #[test]
861    fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
862        let dir = tempfile::tempdir().unwrap();
863        let file_path = create_byte_level_tiktoken_file(dir.path());
864
865        // config.json with kimi model type -> triggers KIMI_PATTERN
866        create_test_config(dir.path(), "kimi");
867
868        // tokenizer_config.json: BOS at 256, EOS at 257. Base vocab is IDs 0..255.
869        create_test_tokenizer_config(dir.path(), 256);
870
871        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
872
873        // ID 256 = [BOS], ID 257 = [EOS].
874        // ID 258 = first UNNAMED reserved token.
875        //   With fix:   named <|reserved_token_258|>  (absolute ID)
876        //   Before fix: named <|reserved_token_2|>    (relative offset)
877
878        // Single unnamed reserved token should be recognized as 1 special token.
879        let single = "<|reserved_token_258|>";
880        let enc = tokenizer.encode(single).unwrap();
881        assert_eq!(
882            enc.token_ids().len(),
883            1,
884            "'{single}' should be 1 special token, got {} tokens: {:?}. \
885             This means fallback naming still uses relative offsets instead of absolute IDs.",
886            enc.token_ids().len(),
887            enc.token_ids()
888        );
889        assert_eq!(enc.token_ids()[0], 258);
890
891        // Multiple unnamed reserved tokens in sequence (mini version of the benchmark).
892        // IDs 258..268 are all unnamed; with the fix they're <|reserved_token_258|>..267.
893        let multi: String = (258u32..268)
894            .map(|id| format!("<|reserved_token_{id}|>"))
895            .collect();
896        let enc_multi = tokenizer.encode(&multi).unwrap();
897        assert_eq!(
898            enc_multi.token_ids().len(),
899            10,
900            "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
901            enc_multi.token_ids().len(),
902            enc_multi.token_ids()
903        );
904        let expected_ids: Vec<u32> = (258..268).collect();
905        assert_eq!(enc_multi.token_ids(), &expected_ids);
906    }
907
908    /// Confirm that the old relative-offset naming would cause token inflation.
909    /// Manually builds a tokenizer whose special-token map uses the WRONG names
910    /// (relative offsets), then shows the same string encodes as many tokens.
911    #[test]
912    fn test_relative_offset_naming_causes_inflation() {
913        let dir = tempfile::tempdir().unwrap();
914        let file_path = create_byte_level_tiktoken_file(dir.path());
915
916        let _encoder = parse_tiktoken_file(&file_path).unwrap();
917        let num_base_tokens = 256usize;
918
919        // Build special tokens with the OLD (buggy) relative-offset naming
920        let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
921        bad_special_tokens.insert("[BOS]".to_string(), 256);
922        bad_special_tokens.insert("[EOS]".to_string(), 257);
923        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
924            let id = num_base_tokens as u32 + i;
925            if id != 256 && id != 257 {
926                // OLD naming: relative offset i, not absolute id
927                bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
928            }
929        }
930
931        let bad_tokenizer =
932            TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
933
934        // With the wrong naming, <|reserved_token_258|> is NOT recognized as special.
935        // It gets split into byte-level BPE tokens -> many more than 1.
936        let input = "<|reserved_token_258|>";
937        let enc = bad_tokenizer.encode(input).unwrap();
938        assert!(
939            enc.token_ids().len() > 1,
940            "With buggy relative-offset naming, '{}' should NOT be recognized as a \
941             single special token. Got {} token(s): {:?}",
942            input,
943            enc.token_ids().len(),
944            enc.token_ids()
945        );
946
947        // Show the inflation: 10 reserved tokens produce far more than 10 BPE tokens.
948        let multi: String = (258u32..268)
949            .map(|id| format!("<|reserved_token_{id}|>"))
950            .collect();
951        let enc_multi = bad_tokenizer.encode(&multi).unwrap();
952        assert!(
953            enc_multi.token_ids().len() > 10,
954            "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
955             Got {}",
956            enc_multi.token_ids().len(),
957        );
958    }
959}