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}
28
29impl TikTokenTokenizer {
30    /// Create a TikTokenTokenizer from a tiktoken model file.
31    ///
32    /// # Arguments
33    /// * `path` - Path to the `.model` or `.tiktoken` file (base64 rank-per-line format)
34    /// * `pattern` - BPE regex pattern string
35    /// * `special_tokens` - Map of special token strings to their IDs
36    pub fn from_file(
37        path: &str,
38        pattern: &str,
39        special_tokens: FxHashMap<String, u32>,
40    ) -> Result<Self> {
41        let encoder = parse_tiktoken_file(path)?;
42        let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
43
44        let bpe = CoreBPE::new(encoder, special_tokens, pattern)
45            .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
46
47        Ok(Self {
48            bpe,
49            special_token_ids,
50        })
51    }
52
53    /// Create a TikTokenTokenizer from a tiktoken model file, auto-detecting
54    /// the BPE pattern from `config.json` and special tokens from `tokenizer_config.json`.
55    ///
56    /// The tiktoken file and config files must be in the same directory.
57    pub fn from_file_auto(path: &str) -> Result<Self> {
58        let file_path = Path::new(path);
59        let directory = file_path
60            .parent()
61            .ok_or_else(|| Error::msg("Cannot determine parent directory of tiktoken file"))?;
62
63        let pattern = detect_bpe_pattern(directory)?;
64        let encoder = parse_tiktoken_file(path)?;
65        // Use max rank + 1 (not len) to avoid ID collisions with sparse/non-contiguous ranks
66        let num_base_tokens = encoder.values().max().map_or(0, |&m| m + 1) as usize;
67        let special_tokens = load_special_tokens(directory, num_base_tokens)?;
68        let special_token_ids: HashSet<u32> = special_tokens.values().copied().collect();
69
70        let bpe = CoreBPE::new(encoder, special_tokens, pattern)
71            .map_err(|err| Error::msg(format!("Error creating tiktoken BPE: {err}")))?;
72
73        Ok(Self {
74            bpe,
75            special_token_ids,
76        })
77    }
78}
79
80impl Encoder for TikTokenTokenizer {
81    fn encode(&self, input: &str) -> Result<Encoding> {
82        let token_ids: Vec<u32> = self.bpe.encode_with_special_tokens(input);
83        Ok(Encoding::Sp(token_ids))
84    }
85
86    fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
87        inputs.par_iter().map(|input| self.encode(input)).collect()
88    }
89}
90
91impl Decoder for TikTokenTokenizer {
92    fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
93        let ids: Vec<u32> = if skip_special_tokens {
94            token_ids
95                .iter()
96                .filter(|&&id| !self.special_token_ids.contains(&id))
97                .copied()
98                .collect()
99        } else {
100            token_ids.to_vec()
101        };
102
103        // Try strict UTF-8 first: valid bytes get `Complete` with zero extra allocation
104        // (takes ownership of the Vec). This correctly handles vocabulary tokens whose
105        // raw bytes are EF BF BD (legitimate U+FFFD) -- they are valid UTF-8 and must
106        // not be confused with incomplete multi-byte sequences.
107        //
108        // On failure, fall back to lossy conversion so partial multi-byte sequences
109        // become U+FFFD, then classify via the trailing-FFFD heuristic. This path is
110        // only hit during incremental detokenization of byte-fallback tokens.
111        let bytes: Vec<u8> = self.bpe._decode_native_and_split(ids).flatten().collect();
112        match String::from_utf8(bytes) {
113            Ok(text) => Ok(DecodeResult::Complete(text)),
114            Err(e) => {
115                let text = String::from_utf8_lossy(e.as_bytes()).into_owned();
116                Ok(DecodeResult::from_decoded(text))
117            }
118        }
119    }
120}
121
122impl Tokenizer for TikTokenTokenizer {}
123
124/// Parse a tiktoken model file (base64-encoded token + rank per line).
125fn parse_tiktoken_file(path: &str) -> Result<FxHashMap<Vec<u8>, u32>> {
126    let contents = std::fs::read_to_string(path)
127        .map_err(|err| Error::msg(format!("Failed to read tiktoken file '{path}': {err}")))?;
128
129    let engine = base64::engine::general_purpose::STANDARD;
130    let mut encoder = FxHashMap::default();
131
132    for line in contents.lines() {
133        let line = line.trim();
134        if line.is_empty() {
135            continue;
136        }
137        let mut parts = line.split_whitespace();
138        let token_b64 = parts
139            .next()
140            .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no token): {line}")))?;
141        let rank_str = parts
142            .next()
143            .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no rank): {line}")))?;
144
145        let token_bytes = engine
146            .decode(token_b64)
147            .map_err(|err| Error::msg(format!("Invalid base64 in tiktoken file: {err}")))?;
148        let rank: u32 = rank_str
149            .parse()
150            .map_err(|err| Error::msg(format!("Invalid rank in tiktoken file: {err}")))?;
151
152        encoder.insert(token_bytes, rank);
153    }
154
155    Ok(encoder)
156}
157
158/// Detect the BPE pattern for a model by reading `model_type` from `config.json`.
159fn detect_bpe_pattern(directory: &Path) -> Result<&'static str> {
160    let model_type: String = crate::file_json_field(&directory.join("config.json"), "model_type")
161        .map_err(|err| {
162        Error::msg(format!("Failed to read model_type from config.json: {err}"))
163    })?;
164
165    match model_type.as_str() {
166        // baseten-admin/Kimi-2.5-text-nvfp4-v3 model has model_type: "deepseek_v3" in its config.json
167        // because Kimi K2.5 is built on the DeepSeek V3 architecture.
168        // it still ships the Kimi tiktoken tokenizer file, so the KIMI_PATTERN BPE regex is the
169        // correct pattern to use.  No pure DeepSeek V3 model uses tiktoken.model files
170        // (they use tokenizer.json instead) so this match is safe.
171        "kimi" | "kimi_k2" | "kimi_k25" | "deepseek_v3" => Ok(KIMI_PATTERN),
172        _ => Err(Error::msg(format!(
173            "Unsupported tiktoken model_type '{model_type}'. \
174             Currently supported: kimi, kimi_k2, kimi_k25, deepseek_v3. \
175             To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
176             with the appropriate BPE regex pattern. \
177             Alternatively, provide a tokenizer.json (HuggingFace format) instead."
178        ))),
179    }
180}
181
182/// Load special tokens from `tokenizer_config.json` in the model directory.
183///
184/// Reads the `added_tokens_decoder` field which maps string token IDs to token definitions.
185/// Falls back to generating `<|reserved_token_{id}|>` names for unmapped IDs.
186fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHashMap<String, u32>> {
187    let config_path = directory.join("tokenizer_config.json");
188    let mut special_tokens = FxHashMap::default();
189
190    if !config_path.exists() {
191        // No tokenizer_config.json — generate default reserved tokens
192        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
193            let id = num_base_tokens as u32 + i;
194            special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
195        }
196        return Ok(special_tokens);
197    }
198
199    let contents = std::fs::read_to_string(&config_path)
200        .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
201
202    let config: serde_json::Value = serde_json::from_str(&contents)
203        .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
204
205    if let Some(added_tokens) = config
206        .get("added_tokens_decoder")
207        .and_then(|v| v.as_object())
208    {
209        for (id_str, token_def) in added_tokens {
210            let id: u32 = id_str.parse().map_err(|err| {
211                Error::msg(format!(
212                    "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
213                ))
214            })?;
215
216            let content = token_def
217                .get("content")
218                .and_then(|v| v.as_str())
219                .unwrap_or_else(|| {
220                    // This shouldn't happen in well-formed configs, but handle gracefully
221                    tracing::warn!("Missing 'content' field for token ID {id}");
222                    ""
223                });
224
225            if !content.is_empty() {
226                special_tokens.insert(content.to_string(), id);
227            }
228        }
229
230        // Fill in any gaps with reserved tokens for the expected range
231        let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
232        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
233            let id = num_base_tokens as u32 + i;
234            if !used_ids.contains(&id) {
235                special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
236            }
237        }
238    } else {
239        // No added_tokens_decoder — generate default reserved tokens
240        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
241            let id = num_base_tokens as u32 + i;
242            special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
243        }
244    }
245
246    Ok(special_tokens)
247}
248
249#[cfg(test)]
250mod tests {
251    use super::*;
252    use crate::DecodeStream;
253    use std::io::Write;
254    use std::sync::Arc;
255
256    fn create_test_tiktoken_file(dir: &Path) -> String {
257        let engine = base64::engine::general_purpose::STANDARD;
258        let mut content = String::new();
259
260        // Create some simple token entries: single bytes with sequential ranks
261        let tokens: Vec<(&[u8], u32)> = vec![
262            (b"h", 0),
263            (b"e", 1),
264            (b"l", 2),
265            (b"o", 3),
266            (b" ", 4),
267            (b"w", 5),
268            (b"r", 6),
269            (b"d", 7),
270            (b"he", 8),
271            (b"ll", 9),
272            (b"lo", 10),
273            (b"wo", 11),
274            (b"rl", 12),
275            (b"hel", 13),
276            (b"llo", 14),
277            (b"wor", 15),
278            (b"hell", 16),
279            (b"ello", 17),
280            (b"worl", 18),
281            (b"hello", 19),
282            (b"world", 20),
283        ];
284
285        for (token, rank) in tokens {
286            let encoded = engine.encode(token);
287            content.push_str(&format!("{encoded} {rank}\n"));
288        }
289
290        let file_path = dir.join("tiktoken.model");
291        let mut file = std::fs::File::create(&file_path).unwrap();
292        file.write_all(content.as_bytes()).unwrap();
293        file_path.to_str().unwrap().to_string()
294    }
295
296    fn create_test_config(dir: &Path, model_type: &str) {
297        let config = serde_json::json!({
298            "model_type": model_type,
299            "max_position_embeddings": 32768,
300            "eos_token_id": [21]
301        });
302        let file_path = dir.join("config.json");
303        std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
304    }
305
306    fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
307        let mut added_tokens = serde_json::Map::new();
308        let bos_id = num_base_tokens;
309        let eos_id = num_base_tokens + 1;
310
311        added_tokens.insert(
312            bos_id.to_string(),
313            serde_json::json!({"content": "[BOS]", "special": true}),
314        );
315        added_tokens.insert(
316            eos_id.to_string(),
317            serde_json::json!({"content": "[EOS]", "special": true}),
318        );
319
320        let config = serde_json::json!({
321            "added_tokens_decoder": added_tokens
322        });
323
324        let file_path = dir.join("tokenizer_config.json");
325        std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
326    }
327
328    #[test]
329    fn test_parse_tiktoken_file() {
330        let dir = tempfile::tempdir().unwrap();
331        let file_path = create_test_tiktoken_file(dir.path());
332        let encoder = parse_tiktoken_file(&file_path).unwrap();
333        assert_eq!(encoder.len(), 21);
334        assert_eq!(encoder[b"hello".as_slice()], 19);
335        assert_eq!(encoder[b"world".as_slice()], 20);
336    }
337
338    #[test]
339    fn test_parse_tiktoken_file_missing() {
340        let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
341        assert!(result.is_err());
342    }
343
344    #[test]
345    fn test_tiktoken_from_file() {
346        let dir = tempfile::tempdir().unwrap();
347        let file_path = create_test_tiktoken_file(dir.path());
348
349        let mut special_tokens = FxHashMap::default();
350        special_tokens.insert("[BOS]".to_string(), 21_u32);
351        special_tokens.insert("[EOS]".to_string(), 22_u32);
352
353        // Use a simple pattern for testing
354        let pattern = r"[\w]+|[^\w\s]+|\s+";
355
356        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
357
358        // Test encode
359        let encoding = tokenizer.encode("hello world").unwrap();
360        let ids = encoding.token_ids();
361        assert!(!ids.is_empty());
362
363        // Test decode roundtrip
364        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
365        assert_eq!(decoded, "hello world");
366    }
367
368    #[test]
369    fn test_tiktoken_encoding_variant() {
370        let dir = tempfile::tempdir().unwrap();
371        let file_path = create_test_tiktoken_file(dir.path());
372
373        let special_tokens = FxHashMap::default();
374        let pattern = r"[\w]+|[^\w\s]+|\s+";
375
376        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
377        let encoding = tokenizer.encode("hello").unwrap();
378
379        // Verify it produces the Sp variant
380        match &encoding {
381            Encoding::Sp(_) => {}
382            other => panic!("Expected Encoding::Sp, got {:?}", other),
383        }
384    }
385
386    #[test]
387    fn test_tiktoken_skip_special_tokens() {
388        let dir = tempfile::tempdir().unwrap();
389        let file_path = create_test_tiktoken_file(dir.path());
390
391        let mut special_tokens = FxHashMap::default();
392        special_tokens.insert("[BOS]".to_string(), 21_u32);
393        special_tokens.insert("[EOS]".to_string(), 22_u32);
394
395        let pattern = r"[\w]+|[^\w\s]+|\s+";
396
397        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
398
399        // Encode hello and prepend/append special tokens
400        let encoding = tokenizer.encode("hello").unwrap();
401        let mut ids = vec![21u32]; // [BOS]
402        ids.extend(encoding.token_ids());
403        ids.push(22); // [EOS]
404
405        // Decode with skip_special_tokens=true should strip special tokens
406        let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
407        assert_eq!(decoded_skip, "hello");
408
409        // Decode with skip_special_tokens=false should include them
410        let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
411        assert!(decoded_all.contains("hello"));
412    }
413
414    #[test]
415    fn test_tiktoken_from_file_auto() {
416        let dir = tempfile::tempdir().unwrap();
417        let file_path = create_test_tiktoken_file(dir.path());
418
419        create_test_config(dir.path(), "kimi");
420        create_test_tokenizer_config(dir.path(), 21);
421
422        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
423
424        // Basic encode/decode roundtrip
425        let encoding = tokenizer.encode("hello world").unwrap();
426        let ids = encoding.token_ids();
427        assert!(!ids.is_empty());
428
429        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
430        assert_eq!(decoded, "hello world");
431    }
432
433    #[test]
434    fn test_detect_bpe_pattern_unknown() {
435        let dir = tempfile::tempdir().unwrap();
436        create_test_config(dir.path(), "unknown_model");
437        let result = detect_bpe_pattern(dir.path());
438        assert!(result.is_err());
439    }
440
441    #[test]
442    fn test_load_special_tokens_no_config() {
443        let dir = tempfile::tempdir().unwrap();
444        let tokens = load_special_tokens(dir.path(), 100).unwrap();
445        assert_eq!(tokens.len(), 256);
446        assert_eq!(tokens["<|reserved_token_100|>"], 100);
447        assert_eq!(tokens["<|reserved_token_355|>"], 355);
448    }
449
450    #[test]
451    fn test_load_special_tokens_with_config() {
452        let dir = tempfile::tempdir().unwrap();
453        create_test_tokenizer_config(dir.path(), 100);
454        let tokens = load_special_tokens(dir.path(), 100).unwrap();
455        assert_eq!(tokens["[BOS]"], 100);
456        assert_eq!(tokens["[EOS]"], 101);
457        // Should also have reserved tokens filling gaps
458        assert!(tokens.len() > 2);
459    }
460
461    /// Helper: create a tiktoken file that includes raw byte tokens (byte fallback tokens).
462    fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
463        let engine = base64::engine::general_purpose::STANDARD;
464        let mut content = String::new();
465
466        let tokens: Vec<(&[u8], u32)> = vec![
467            (b"h", 0),
468            (b"e", 1),
469            (b"l", 2),
470            (b"o", 3),
471            (b" ", 4),
472            (b"hello", 5),
473        ];
474
475        for (token, rank) in &tokens {
476            let encoded = engine.encode(token);
477            content.push_str(&format!("{encoded} {rank}\n"));
478        }
479
480        // Byte-fallback tokens: individual bytes that form CJK character "你" (U+4F60)
481        // UTF-8 encoding: 0xE4 0xBD 0xA0
482        let byte_tokens: Vec<(Vec<u8>, u32)> =
483            vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
484
485        for (token, rank) in &byte_tokens {
486            let encoded = engine.encode(token);
487            content.push_str(&format!("{encoded} {rank}\n"));
488        }
489
490        // Bytes for emoji "😀" (U+1F600) — 4-byte UTF-8: 0xF0 0x9F 0x98 0x80
491        let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
492            (vec![0xF0], 200),
493            (vec![0x9F], 201),
494            (vec![0x98], 202),
495            (vec![0x80], 203),
496        ];
497
498        for (token, rank) in &emoji_tokens {
499            let encoded = engine.encode(token);
500            content.push_str(&format!("{encoded} {rank}\n"));
501        }
502
503        // Legitimate U+FFFD token: valid UTF-8 bytes EF BF BD (replacement character
504        // as an actual vocabulary entry, not an artifact of lossy conversion)
505        let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
506
507        for (token, rank) in &fffd_token {
508            let encoded = engine.encode(token);
509            content.push_str(&format!("{encoded} {rank}\n"));
510        }
511
512        let file_path = dir.join("tiktoken.model");
513        let mut file = std::fs::File::create(&file_path).unwrap();
514        file.write_all(content.as_bytes()).unwrap();
515        file_path.to_str().unwrap().to_string()
516    }
517
518    fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
519        let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
520        let special_tokens = FxHashMap::default();
521        let pattern = r"[\w]+|[^\w\s]+|\s+";
522        TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
523    }
524
525    /// Reproduces the original panic: decoding a single byte-fallback token that is
526    /// part of a multi-byte UTF-8 character. Before the fix, CoreBPE::decode() would
527    /// call String::from_utf8() on [0xE4] and error with "incomplete utf-8 byte sequence".
528    #[test]
529    fn test_decode_single_incomplete_utf8_byte_does_not_error() {
530        let dir = tempfile::tempdir().unwrap();
531        let tokenizer = create_byte_token_tokenizer(dir.path());
532
533        let result = tokenizer.decode(&[100], false);
534        assert!(
535            result.is_ok(),
536            "decode() should not error on incomplete UTF-8 bytes"
537        );
538        let decode_result = result.unwrap();
539        assert!(
540            decode_result.is_partial(),
541            "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
542            decode_result
543        );
544    }
545
546    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
547    #[test]
548    fn test_decode_two_of_three_utf8_bytes_does_not_error() {
549        let dir = tempfile::tempdir().unwrap();
550        let tokenizer = create_byte_token_tokenizer(dir.path());
551
552        let result = tokenizer.decode(&[100, 101], false);
553        assert!(result.is_ok());
554        let decode_result = result.unwrap();
555        assert!(
556            decode_result.is_partial(),
557            "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
558            decode_result
559        );
560    }
561
562    /// When all bytes of a multi-byte character are present, the concatenated bytes form
563    /// valid UTF-8, so this test passes both before and after the fix. It serves as a
564    /// correctness check that the lossy conversion doesn't corrupt complete characters.
565    #[test]
566    fn test_decode_complete_multibyte_utf8_produces_correct_char() {
567        let dir = tempfile::tempdir().unwrap();
568        let tokenizer = create_byte_token_tokenizer(dir.path());
569
570        let result = tokenizer.decode(&[100, 101, 102], false);
571        assert!(result.is_ok());
572        assert_eq!(String::from(result.unwrap()), "你");
573    }
574
575    /// All 4 emoji bytes together form valid UTF-8, so this passes both before and after
576    /// the fix. Validates that lossy conversion doesn't alter complete multi-byte sequences.
577    #[test]
578    fn test_decode_complete_4byte_emoji_from_byte_tokens() {
579        let dir = tempfile::tempdir().unwrap();
580        let tokenizer = create_byte_token_tokenizer(dir.path());
581
582        let result = tokenizer.decode(&[200, 201, 202, 203], false);
583        assert!(result.is_ok());
584        assert_eq!(String::from(result.unwrap()), "😀");
585    }
586
587    /// Regression test: a vocabulary token whose raw bytes are EF BF BD (the valid
588    /// UTF-8 encoding of U+FFFD) must decode as `Complete`, not `Partial`. Before the
589    /// from_utf8 fast-path fix, from_utf8_lossy + the trailing-FFFD heuristic would
590    /// misclassify this as Partial, causing the incremental decoder to suppress it.
591    #[test]
592    fn test_decode_legitimate_replacement_char_token_is_complete() {
593        let dir = tempfile::tempdir().unwrap();
594        let tokenizer = create_byte_token_tokenizer(dir.path());
595
596        let result = tokenizer.decode(&[300], false);
597        assert!(result.is_ok());
598        let decode_result = result.unwrap();
599        assert!(
600            decode_result.is_complete(),
601            "legitimate U+FFFD vocab token must be Complete, got: {:?}",
602            decode_result
603        );
604        assert_eq!(decode_result.as_str(), "\u{FFFD}");
605    }
606
607    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
608    #[test]
609    fn test_decode_partial_emoji_does_not_error() {
610        let dir = tempfile::tempdir().unwrap();
611        let tokenizer = create_byte_token_tokenizer(dir.path());
612
613        let result = tokenizer.decode(&[200], false);
614        assert!(result.is_ok());
615        assert!(result.unwrap().is_partial());
616    }
617
618    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
619    #[test]
620    fn test_decode_mixed_ascii_and_incomplete_bytes() {
621        let dir = tempfile::tempdir().unwrap();
622        let tokenizer = create_byte_token_tokenizer(dir.path());
623
624        let result = tokenizer.decode(&[5, 100], false);
625        assert!(result.is_ok());
626        let decode_result = result.unwrap();
627        assert!(
628            decode_result.is_partial(),
629            "trailing incomplete byte should produce DecodeResult::Partial"
630        );
631        let text: String = decode_result.into();
632        assert!(
633            text.starts_with("hello"),
634            "should start with 'hello', got: {:?}",
635            text
636        );
637    }
638
639    /// End-to-end incremental detokenization: DecodeStream buffers partial bytes,
640    /// emits the complete character once all bytes arrive.
641    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
642    #[test]
643    fn test_decode_stream_incremental_multibyte_reassembly() {
644        let dir = tempfile::tempdir().unwrap();
645        let tokenizer = create_byte_token_tokenizer(dir.path());
646        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
647
648        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
649
650        let r1 = stream.step(100).unwrap();
651        assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
652
653        let r2 = stream.step(101).unwrap();
654        assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
655
656        let r3 = stream.step(102).unwrap();
657        assert!(r3.is_some(), "third byte should complete the character");
658        assert_eq!(r3.unwrap(), "你");
659    }
660
661    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
662    #[test]
663    fn test_decode_stream_incremental_emoji_reassembly() {
664        let dir = tempfile::tempdir().unwrap();
665        let tokenizer = create_byte_token_tokenizer(dir.path());
666        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
667
668        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
669
670        let r1 = stream.step(200).unwrap();
671        assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
672
673        let r2 = stream.step(201).unwrap();
674        assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
675
676        let r3 = stream.step(202).unwrap();
677        assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
678
679        let r4 = stream.step(203).unwrap();
680        assert!(r4.is_some(), "byte 4/4 should complete the emoji");
681        assert_eq!(r4.unwrap(), "😀");
682    }
683
684    #[test]
685    fn test_tiktoken_encode_batch() {
686        let dir = tempfile::tempdir().unwrap();
687        let file_path = create_test_tiktoken_file(dir.path());
688
689        let special_tokens = FxHashMap::default();
690        let pattern = r"[\w]+|[^\w\s]+|\s+";
691
692        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
693
694        let inputs = &["hello", "world"];
695        let encodings = tokenizer.encode_batch(inputs).unwrap();
696        assert_eq!(encodings.len(), 2);
697
698        for (encoding, input) in encodings.iter().zip(inputs.iter()) {
699            let decoded: String = tokenizer
700                .decode(encoding.token_ids(), false)
701                .unwrap()
702                .into();
703            assert_eq!(decoded, *input);
704        }
705    }
706
707    /// Helper: create a tiktoken file containing all 256 single-byte tokens (ranks 0..255).
708    /// This gives a complete byte-level base vocabulary so any ASCII string can be encoded.
709    fn create_byte_level_tiktoken_file(dir: &Path) -> String {
710        let engine = base64::engine::general_purpose::STANDARD;
711        let mut content = String::new();
712        for byte_val in 0u16..256 {
713            let encoded = engine.encode([byte_val as u8]);
714            content.push_str(&format!("{encoded} {byte_val}\n"));
715        }
716        let file_path = dir.join("tiktoken.model");
717        std::fs::write(&file_path, &content).unwrap();
718        file_path.to_str().unwrap().to_string()
719    }
720
721    /// Regression test for Kimi K2.5 tokenizer inflation.
722    ///
723    /// Python's tokenization_kimi.py names unnamed reserved tokens by absolute ID:
724    ///   `<|reserved_token_{absolute_id}|>`
725    ///
726    /// The Rust code previously used relative offsets (0..255) as the naming index,
727    /// so when a prompt contained `<|reserved_token_258|>` the Rust tokenizer did NOT
728    /// recognize it as a special token. Each occurrence was encoded as multiple BPE
729    /// tokens instead of 1, inflating an 8192-token prompt to 9038 tokens and causing
730    /// TRT-LLM to reject the request.
731    #[test]
732    fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
733        let dir = tempfile::tempdir().unwrap();
734        let file_path = create_byte_level_tiktoken_file(dir.path());
735
736        // config.json with kimi model type -> triggers KIMI_PATTERN
737        create_test_config(dir.path(), "kimi");
738
739        // tokenizer_config.json: BOS at 256, EOS at 257. Base vocab is IDs 0..255.
740        create_test_tokenizer_config(dir.path(), 256);
741
742        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
743
744        // ID 256 = [BOS], ID 257 = [EOS].
745        // ID 258 = first UNNAMED reserved token.
746        //   With fix:   named <|reserved_token_258|>  (absolute ID)
747        //   Before fix: named <|reserved_token_2|>    (relative offset)
748
749        // Single unnamed reserved token should be recognized as 1 special token.
750        let single = "<|reserved_token_258|>";
751        let enc = tokenizer.encode(single).unwrap();
752        assert_eq!(
753            enc.token_ids().len(),
754            1,
755            "'{single}' should be 1 special token, got {} tokens: {:?}. \
756             This means fallback naming still uses relative offsets instead of absolute IDs.",
757            enc.token_ids().len(),
758            enc.token_ids()
759        );
760        assert_eq!(enc.token_ids()[0], 258);
761
762        // Multiple unnamed reserved tokens in sequence (mini version of the benchmark).
763        // IDs 258..268 are all unnamed; with the fix they're <|reserved_token_258|>..267.
764        let multi: String = (258u32..268)
765            .map(|id| format!("<|reserved_token_{id}|>"))
766            .collect();
767        let enc_multi = tokenizer.encode(&multi).unwrap();
768        assert_eq!(
769            enc_multi.token_ids().len(),
770            10,
771            "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
772            enc_multi.token_ids().len(),
773            enc_multi.token_ids()
774        );
775        let expected_ids: Vec<u32> = (258..268).collect();
776        assert_eq!(enc_multi.token_ids(), &expected_ids);
777    }
778
779    /// Confirm that the old relative-offset naming would cause token inflation.
780    /// Manually builds a tokenizer whose special-token map uses the WRONG names
781    /// (relative offsets), then shows the same string encodes as many tokens.
782    #[test]
783    fn test_relative_offset_naming_causes_inflation() {
784        let dir = tempfile::tempdir().unwrap();
785        let file_path = create_byte_level_tiktoken_file(dir.path());
786
787        let _encoder = parse_tiktoken_file(&file_path).unwrap();
788        let num_base_tokens = 256usize;
789
790        // Build special tokens with the OLD (buggy) relative-offset naming
791        let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
792        bad_special_tokens.insert("[BOS]".to_string(), 256);
793        bad_special_tokens.insert("[EOS]".to_string(), 257);
794        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
795            let id = num_base_tokens as u32 + i;
796            if id != 256 && id != 257 {
797                // OLD naming: relative offset i, not absolute id
798                bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
799            }
800        }
801
802        let bad_tokenizer =
803            TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
804
805        // With the wrong naming, <|reserved_token_258|> is NOT recognized as special.
806        // It gets split into byte-level BPE tokens -> many more than 1.
807        let input = "<|reserved_token_258|>";
808        let enc = bad_tokenizer.encode(input).unwrap();
809        assert!(
810            enc.token_ids().len() > 1,
811            "With buggy relative-offset naming, '{}' should NOT be recognized as a \
812             single special token. Got {} token(s): {:?}",
813            input,
814            enc.token_ids().len(),
815            enc.token_ids()
816        );
817
818        // Show the inflation: 10 reserved tokens produce far more than 10 BPE tokens.
819        let multi: String = (258u32..268)
820            .map(|id| format!("<|reserved_token_{id}|>"))
821            .collect();
822        let enc_multi = bad_tokenizer.encode(&multi).unwrap();
823        assert!(
824            enc_multi.token_ids().len() > 10,
825            "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
826             Got {}",
827            enc_multi.token_ids().len(),
828        );
829    }
830}