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
145/// Parse a tiktoken model file (base64-encoded token + rank per line).
146fn parse_tiktoken_file(path: &str) -> Result<FxHashMap<Vec<u8>, u32>> {
147    let contents = std::fs::read_to_string(path)
148        .map_err(|err| Error::msg(format!("Failed to read tiktoken file '{path}': {err}")))?;
149
150    let engine = base64::engine::general_purpose::STANDARD;
151    let mut encoder = FxHashMap::default();
152
153    for line in contents.lines() {
154        let line = line.trim();
155        if line.is_empty() {
156            continue;
157        }
158        let mut parts = line.split_whitespace();
159        let token_b64 = parts
160            .next()
161            .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no token): {line}")))?;
162        let rank_str = parts
163            .next()
164            .ok_or_else(|| Error::msg(format!("Invalid tiktoken line (no rank): {line}")))?;
165
166        let token_bytes = engine
167            .decode(token_b64)
168            .map_err(|err| Error::msg(format!("Invalid base64 in tiktoken file: {err}")))?;
169        let rank: u32 = rank_str
170            .parse()
171            .map_err(|err| Error::msg(format!("Invalid rank in tiktoken file: {err}")))?;
172
173        encoder.insert(token_bytes, rank);
174    }
175
176    Ok(encoder)
177}
178
179/// Detect the BPE pattern for a model by reading `model_type` from `config.json`.
180fn detect_bpe_pattern(directory: &Path) -> Result<&'static str> {
181    let model_type: String = crate::file_json_field(&directory.join("config.json"), "model_type")
182        .map_err(|err| {
183        Error::msg(format!("Failed to read model_type from config.json: {err}"))
184    })?;
185
186    match model_type.as_str() {
187        // baseten-admin/Kimi-2.5-text-nvfp4-v3 model has model_type: "deepseek_v3" in its config.json
188        // because Kimi K2.5 is built on the DeepSeek V3 architecture.
189        // it still ships the Kimi tiktoken tokenizer file, so the KIMI_PATTERN BPE regex is the
190        // correct pattern to use.  No pure DeepSeek V3 model uses tiktoken.model files
191        // (they use tokenizer.json instead) so this match is safe.
192        "kimi" | "kimi_k2" | "kimi_k25" | "deepseek_v3" => Ok(KIMI_PATTERN),
193        _ => Err(Error::msg(format!(
194            "Unsupported tiktoken model_type '{model_type}'. \
195             Currently supported: kimi, kimi_k2, kimi_k25, deepseek_v3. \
196             To add a new model type, extend detect_bpe_pattern() in lib/tokenizers/src/tiktoken.rs \
197             with the appropriate BPE regex pattern. \
198             Alternatively, provide a tokenizer.json (HuggingFace format) instead."
199        ))),
200    }
201}
202
203/// Load special tokens from `tokenizer_config.json` in the model directory.
204///
205/// Reads the `added_tokens_decoder` field which maps string token IDs to token definitions.
206/// Falls back to generating `<|reserved_token_{id}|>` names for unmapped IDs.
207fn load_special_tokens(directory: &Path, num_base_tokens: usize) -> Result<FxHashMap<String, u32>> {
208    let config_path = directory.join("tokenizer_config.json");
209    let mut special_tokens = FxHashMap::default();
210
211    if !config_path.exists() {
212        // No tokenizer_config.json β€” generate default reserved tokens
213        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
214            let id = num_base_tokens as u32 + i;
215            special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
216        }
217        return Ok(special_tokens);
218    }
219
220    let contents = std::fs::read_to_string(&config_path)
221        .map_err(|err| Error::msg(format!("Failed to read tokenizer_config.json: {err}")))?;
222
223    let config: serde_json::Value = serde_json::from_str(&contents)
224        .map_err(|err| Error::msg(format!("Failed to parse tokenizer_config.json: {err}")))?;
225
226    if let Some(added_tokens) = config
227        .get("added_tokens_decoder")
228        .and_then(|v| v.as_object())
229    {
230        for (id_str, token_def) in added_tokens {
231            let id: u32 = id_str.parse().map_err(|err| {
232                Error::msg(format!(
233                    "Invalid token ID '{id_str}' in added_tokens_decoder: {err}"
234                ))
235            })?;
236
237            let content = token_def
238                .get("content")
239                .and_then(|v| v.as_str())
240                .unwrap_or_else(|| {
241                    // This shouldn't happen in well-formed configs, but handle gracefully
242                    tracing::warn!("Missing 'content' field for token ID {id}");
243                    ""
244                });
245
246            if !content.is_empty() {
247                special_tokens.insert(content.to_string(), id);
248            }
249        }
250
251        // Fill in any gaps with reserved tokens for the expected range
252        let used_ids: HashSet<u32> = special_tokens.values().copied().collect();
253        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
254            let id = num_base_tokens as u32 + i;
255            if !used_ids.contains(&id) {
256                special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
257            }
258        }
259    } else {
260        // No added_tokens_decoder β€” generate default reserved tokens
261        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
262            let id = num_base_tokens as u32 + i;
263            special_tokens.insert(format!("<|reserved_token_{id}|>"), id);
264        }
265    }
266
267    Ok(special_tokens)
268}
269
270#[cfg(test)]
271mod tests {
272    use super::*;
273    use crate::DecodeStream;
274    use std::io::Write;
275    use std::sync::Arc;
276
277    fn create_test_tiktoken_file(dir: &Path) -> String {
278        let engine = base64::engine::general_purpose::STANDARD;
279        let mut content = String::new();
280
281        // Create some simple token entries: single bytes with sequential ranks
282        let tokens: Vec<(&[u8], u32)> = vec![
283            (b"h", 0),
284            (b"e", 1),
285            (b"l", 2),
286            (b"o", 3),
287            (b" ", 4),
288            (b"w", 5),
289            (b"r", 6),
290            (b"d", 7),
291            (b"he", 8),
292            (b"ll", 9),
293            (b"lo", 10),
294            (b"wo", 11),
295            (b"rl", 12),
296            (b"hel", 13),
297            (b"llo", 14),
298            (b"wor", 15),
299            (b"hell", 16),
300            (b"ello", 17),
301            (b"worl", 18),
302            (b"hello", 19),
303            (b"world", 20),
304        ];
305
306        for (token, rank) in tokens {
307            let encoded = engine.encode(token);
308            content.push_str(&format!("{encoded} {rank}\n"));
309        }
310
311        let file_path = dir.join("tiktoken.model");
312        let mut file = std::fs::File::create(&file_path).unwrap();
313        file.write_all(content.as_bytes()).unwrap();
314        file_path.to_str().unwrap().to_string()
315    }
316
317    fn create_test_config(dir: &Path, model_type: &str) {
318        let config = serde_json::json!({
319            "model_type": model_type,
320            "max_position_embeddings": 32768,
321            "eos_token_id": [21]
322        });
323        let file_path = dir.join("config.json");
324        std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
325    }
326
327    fn create_test_tokenizer_config(dir: &Path, num_base_tokens: usize) {
328        let mut added_tokens = serde_json::Map::new();
329        let bos_id = num_base_tokens;
330        let eos_id = num_base_tokens + 1;
331
332        added_tokens.insert(
333            bos_id.to_string(),
334            serde_json::json!({"content": "[BOS]", "special": true}),
335        );
336        added_tokens.insert(
337            eos_id.to_string(),
338            serde_json::json!({"content": "[EOS]", "special": true}),
339        );
340
341        let config = serde_json::json!({
342            "added_tokens_decoder": added_tokens
343        });
344
345        let file_path = dir.join("tokenizer_config.json");
346        std::fs::write(file_path, serde_json::to_string_pretty(&config).unwrap()).unwrap();
347    }
348
349    #[test]
350    fn test_parse_tiktoken_file() {
351        let dir = tempfile::tempdir().unwrap();
352        let file_path = create_test_tiktoken_file(dir.path());
353        let encoder = parse_tiktoken_file(&file_path).unwrap();
354        assert_eq!(encoder.len(), 21);
355        assert_eq!(encoder[b"hello".as_slice()], 19);
356        assert_eq!(encoder[b"world".as_slice()], 20);
357    }
358
359    #[test]
360    fn test_parse_tiktoken_file_missing() {
361        let result = parse_tiktoken_file("/nonexistent/path/tiktoken.model");
362        assert!(result.is_err());
363    }
364
365    #[test]
366    fn test_tiktoken_from_file() {
367        let dir = tempfile::tempdir().unwrap();
368        let file_path = create_test_tiktoken_file(dir.path());
369
370        let mut special_tokens = FxHashMap::default();
371        special_tokens.insert("[BOS]".to_string(), 21_u32);
372        special_tokens.insert("[EOS]".to_string(), 22_u32);
373
374        // Use a simple pattern for testing
375        let pattern = r"[\w]+|[^\w\s]+|\s+";
376
377        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
378
379        assert_eq!(
380            tokenizer.special_tokens(),
381            &["[BOS]".to_string(), "[EOS]".to_string()]
382        );
383
384        // Test encode
385        let encoding = tokenizer.encode("hello world").unwrap();
386        let ids = encoding.token_ids();
387        assert!(!ids.is_empty());
388
389        // Test decode roundtrip
390        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
391        assert_eq!(decoded, "hello world");
392    }
393
394    #[test]
395    fn test_tiktoken_encoding_variant() {
396        let dir = tempfile::tempdir().unwrap();
397        let file_path = create_test_tiktoken_file(dir.path());
398
399        let special_tokens = FxHashMap::default();
400        let pattern = r"[\w]+|[^\w\s]+|\s+";
401
402        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
403        assert!(tokenizer.special_tokens().is_empty());
404        let encoding = tokenizer.encode("hello").unwrap();
405
406        // Verify it produces the Sp variant
407        match &encoding {
408            Encoding::Sp(_) => {}
409            other => panic!("Expected Encoding::Sp, got {:?}", other),
410        }
411    }
412
413    #[test]
414    fn test_tiktoken_skip_special_tokens() {
415        let dir = tempfile::tempdir().unwrap();
416        let file_path = create_test_tiktoken_file(dir.path());
417
418        let mut special_tokens = FxHashMap::default();
419        special_tokens.insert("[BOS]".to_string(), 21_u32);
420        special_tokens.insert("[EOS]".to_string(), 22_u32);
421
422        let pattern = r"[\w]+|[^\w\s]+|\s+";
423
424        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
425
426        // Encode hello and prepend/append special tokens
427        let encoding = tokenizer.encode("hello").unwrap();
428        let mut ids = vec![21u32]; // [BOS]
429        ids.extend(encoding.token_ids());
430        ids.push(22); // [EOS]
431
432        // Decode with skip_special_tokens=true should strip special tokens
433        let decoded_skip: String = tokenizer.decode(&ids, true).unwrap().into();
434        assert_eq!(decoded_skip, "hello");
435
436        // Decode with skip_special_tokens=false should include them
437        let decoded_all: String = tokenizer.decode(&ids, false).unwrap().into();
438        assert!(decoded_all.contains("hello"));
439    }
440
441    #[test]
442    fn test_tiktoken_from_file_auto() {
443        let dir = tempfile::tempdir().unwrap();
444        let file_path = create_test_tiktoken_file(dir.path());
445
446        create_test_config(dir.path(), "kimi");
447        create_test_tokenizer_config(dir.path(), 21);
448
449        let mut expected_specials: Vec<String> = load_special_tokens(dir.path(), 21)
450            .unwrap()
451            .into_keys()
452            .collect();
453        expected_specials.sort();
454        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
455        assert_eq!(tokenizer.special_tokens(), expected_specials);
456
457        // Basic encode/decode roundtrip
458        let encoding = tokenizer.encode("hello world").unwrap();
459        let ids = encoding.token_ids();
460        assert!(!ids.is_empty());
461
462        let decoded: String = tokenizer.decode(ids, false).unwrap().into();
463        assert_eq!(decoded, "hello world");
464    }
465
466    #[test]
467    fn test_detect_bpe_pattern_unknown() {
468        let dir = tempfile::tempdir().unwrap();
469        create_test_config(dir.path(), "unknown_model");
470        let result = detect_bpe_pattern(dir.path());
471        assert!(result.is_err());
472    }
473
474    #[test]
475    fn test_load_special_tokens_no_config() {
476        let dir = tempfile::tempdir().unwrap();
477        let tokens = load_special_tokens(dir.path(), 100).unwrap();
478        assert_eq!(tokens.len(), 256);
479        assert_eq!(tokens["<|reserved_token_100|>"], 100);
480        assert_eq!(tokens["<|reserved_token_355|>"], 355);
481    }
482
483    #[test]
484    fn test_load_special_tokens_with_config() {
485        let dir = tempfile::tempdir().unwrap();
486        create_test_tokenizer_config(dir.path(), 100);
487        let tokens = load_special_tokens(dir.path(), 100).unwrap();
488        assert_eq!(tokens["[BOS]"], 100);
489        assert_eq!(tokens["[EOS]"], 101);
490        // Should also have reserved tokens filling gaps
491        assert!(tokens.len() > 2);
492    }
493
494    /// Helper: create a tiktoken file that includes raw byte tokens (byte fallback tokens).
495    fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
496        let engine = base64::engine::general_purpose::STANDARD;
497        let mut content = String::new();
498
499        let tokens: Vec<(&[u8], u32)> = vec![
500            (b"h", 0),
501            (b"e", 1),
502            (b"l", 2),
503            (b"o", 3),
504            (b" ", 4),
505            (b"hello", 5),
506        ];
507
508        for (token, rank) in &tokens {
509            let encoded = engine.encode(token);
510            content.push_str(&format!("{encoded} {rank}\n"));
511        }
512
513        // Byte-fallback tokens: individual bytes that form CJK character "δ½ " (U+4F60)
514        // UTF-8 encoding: 0xE4 0xBD 0xA0
515        let byte_tokens: Vec<(Vec<u8>, u32)> =
516            vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
517
518        for (token, rank) in &byte_tokens {
519            let encoded = engine.encode(token);
520            content.push_str(&format!("{encoded} {rank}\n"));
521        }
522
523        // Bytes for emoji "πŸ˜€" (U+1F600) β€” 4-byte UTF-8: 0xF0 0x9F 0x98 0x80
524        let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
525            (vec![0xF0], 200),
526            (vec![0x9F], 201),
527            (vec![0x98], 202),
528            (vec![0x80], 203),
529        ];
530
531        for (token, rank) in &emoji_tokens {
532            let encoded = engine.encode(token);
533            content.push_str(&format!("{encoded} {rank}\n"));
534        }
535
536        // Legitimate U+FFFD token: valid UTF-8 bytes EF BF BD (replacement character
537        // as an actual vocabulary entry, not an artifact of lossy conversion)
538        let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
539
540        for (token, rank) in &fffd_token {
541            let encoded = engine.encode(token);
542            content.push_str(&format!("{encoded} {rank}\n"));
543        }
544
545        let file_path = dir.join("tiktoken.model");
546        let mut file = std::fs::File::create(&file_path).unwrap();
547        file.write_all(content.as_bytes()).unwrap();
548        file_path.to_str().unwrap().to_string()
549    }
550
551    fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
552        let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
553        let special_tokens = FxHashMap::default();
554        let pattern = r"[\w]+|[^\w\s]+|\s+";
555        TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
556    }
557
558    /// Reproduces the original panic: decoding a single byte-fallback token that is
559    /// part of a multi-byte UTF-8 character. Before the fix, CoreBPE::decode() would
560    /// call String::from_utf8() on [0xE4] and error with "incomplete utf-8 byte sequence".
561    #[test]
562    fn test_decode_single_incomplete_utf8_byte_does_not_error() {
563        let dir = tempfile::tempdir().unwrap();
564        let tokenizer = create_byte_token_tokenizer(dir.path());
565
566        let result = tokenizer.decode(&[100], false);
567        assert!(
568            result.is_ok(),
569            "decode() should not error on incomplete UTF-8 bytes"
570        );
571        let decode_result = result.unwrap();
572        assert!(
573            decode_result.is_partial(),
574            "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
575            decode_result
576        );
577    }
578
579    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
580    #[test]
581    fn test_decode_two_of_three_utf8_bytes_does_not_error() {
582        let dir = tempfile::tempdir().unwrap();
583        let tokenizer = create_byte_token_tokenizer(dir.path());
584
585        let result = tokenizer.decode(&[100, 101], false);
586        assert!(result.is_ok());
587        let decode_result = result.unwrap();
588        assert!(
589            decode_result.is_partial(),
590            "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
591            decode_result
592        );
593    }
594
595    /// When all bytes of a multi-byte character are present, the concatenated bytes form
596    /// valid UTF-8, so this test passes both before and after the fix. It serves as a
597    /// correctness check that the lossy conversion doesn't corrupt complete characters.
598    #[test]
599    fn test_decode_complete_multibyte_utf8_produces_correct_char() {
600        let dir = tempfile::tempdir().unwrap();
601        let tokenizer = create_byte_token_tokenizer(dir.path());
602
603        let result = tokenizer.decode(&[100, 101, 102], false);
604        assert!(result.is_ok());
605        assert_eq!(String::from(result.unwrap()), "δ½ ");
606    }
607
608    /// All 4 emoji bytes together form valid UTF-8, so this passes both before and after
609    /// the fix. Validates that lossy conversion doesn't alter complete multi-byte sequences.
610    #[test]
611    fn test_decode_complete_4byte_emoji_from_byte_tokens() {
612        let dir = tempfile::tempdir().unwrap();
613        let tokenizer = create_byte_token_tokenizer(dir.path());
614
615        let result = tokenizer.decode(&[200, 201, 202, 203], false);
616        assert!(result.is_ok());
617        assert_eq!(String::from(result.unwrap()), "πŸ˜€");
618    }
619
620    /// Regression test: a vocabulary token whose raw bytes are EF BF BD (the valid
621    /// UTF-8 encoding of U+FFFD) must decode as `Complete`, not `Partial`. Before the
622    /// from_utf8 fast-path fix, from_utf8_lossy + the trailing-FFFD heuristic would
623    /// misclassify this as Partial, causing the incremental decoder to suppress it.
624    #[test]
625    fn test_decode_legitimate_replacement_char_token_is_complete() {
626        let dir = tempfile::tempdir().unwrap();
627        let tokenizer = create_byte_token_tokenizer(dir.path());
628
629        let result = tokenizer.decode(&[300], false);
630        assert!(result.is_ok());
631        let decode_result = result.unwrap();
632        assert!(
633            decode_result.is_complete(),
634            "legitimate U+FFFD vocab token must be Complete, got: {:?}",
635            decode_result
636        );
637        assert_eq!(decode_result.as_str(), "\u{FFFD}");
638    }
639
640    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
641    #[test]
642    fn test_decode_partial_emoji_does_not_error() {
643        let dir = tempfile::tempdir().unwrap();
644        let tokenizer = create_byte_token_tokenizer(dir.path());
645
646        let result = tokenizer.decode(&[200], false);
647        assert!(result.is_ok());
648        assert!(result.unwrap().is_partial());
649    }
650
651    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
652    #[test]
653    fn test_decode_mixed_ascii_and_incomplete_bytes() {
654        let dir = tempfile::tempdir().unwrap();
655        let tokenizer = create_byte_token_tokenizer(dir.path());
656
657        let result = tokenizer.decode(&[5, 100], false);
658        assert!(result.is_ok());
659        let decode_result = result.unwrap();
660        assert!(
661            decode_result.is_partial(),
662            "trailing incomplete byte should produce DecodeResult::Partial"
663        );
664        let text: String = decode_result.into();
665        assert!(
666            text.starts_with("hello"),
667            "should start with 'hello', got: {:?}",
668            text
669        );
670    }
671
672    /// End-to-end incremental detokenization: DecodeStream buffers partial bytes,
673    /// emits the complete character once all bytes arrive.
674    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
675    #[test]
676    fn test_decode_stream_incremental_multibyte_reassembly() {
677        let dir = tempfile::tempdir().unwrap();
678        let tokenizer = create_byte_token_tokenizer(dir.path());
679        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
680
681        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
682
683        let r1 = stream.step(100).unwrap();
684        assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
685
686        let r2 = stream.step(101).unwrap();
687        assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
688
689        let r3 = stream.step(102).unwrap();
690        assert!(r3.is_some(), "third byte should complete the character");
691        assert_eq!(r3.unwrap(), "δ½ ");
692    }
693
694    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
695    #[test]
696    fn test_decode_stream_incremental_emoji_reassembly() {
697        let dir = tempfile::tempdir().unwrap();
698        let tokenizer = create_byte_token_tokenizer(dir.path());
699        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
700
701        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
702
703        let r1 = stream.step(200).unwrap();
704        assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
705
706        let r2 = stream.step(201).unwrap();
707        assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
708
709        let r3 = stream.step(202).unwrap();
710        assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
711
712        let r4 = stream.step(203).unwrap();
713        assert!(r4.is_some(), "byte 4/4 should complete the emoji");
714        assert_eq!(r4.unwrap(), "πŸ˜€");
715    }
716
717    #[test]
718    fn test_tiktoken_encode_batch() {
719        let dir = tempfile::tempdir().unwrap();
720        let file_path = create_test_tiktoken_file(dir.path());
721
722        let special_tokens = FxHashMap::default();
723        let pattern = r"[\w]+|[^\w\s]+|\s+";
724
725        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
726
727        let inputs = &["hello", "world"];
728        let encodings = tokenizer.encode_batch(inputs).unwrap();
729        assert_eq!(encodings.len(), 2);
730
731        for (encoding, input) in encodings.iter().zip(inputs.iter()) {
732            let decoded: String = tokenizer
733                .decode(encoding.token_ids(), false)
734                .unwrap()
735                .into();
736            assert_eq!(decoded, *input);
737        }
738    }
739
740    /// Helper: create a tiktoken file containing all 256 single-byte tokens (ranks 0..255).
741    /// This gives a complete byte-level base vocabulary so any ASCII string can be encoded.
742    fn create_byte_level_tiktoken_file(dir: &Path) -> String {
743        let engine = base64::engine::general_purpose::STANDARD;
744        let mut content = String::new();
745        for byte_val in 0u16..256 {
746            let encoded = engine.encode([byte_val as u8]);
747            content.push_str(&format!("{encoded} {byte_val}\n"));
748        }
749        let file_path = dir.join("tiktoken.model");
750        std::fs::write(&file_path, &content).unwrap();
751        file_path.to_str().unwrap().to_string()
752    }
753
754    fn has_nontrivial_self_overlap(token: &str) -> bool {
755        let bytes = token.as_bytes();
756        (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
757    }
758
759    fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
760        let a = a.as_bytes();
761        let b = b.as_bytes();
762
763        if a.windows(b.len()).any(|window| window == b)
764            || b.windows(a.len()).any(|window| window == a)
765        {
766            return true;
767        }
768
769        let max_overlap = a.len().min(b.len());
770        (1..max_overlap).any(|overlap| {
771            a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
772        })
773    }
774
775    fn assert_special_split_invariant(
776        tokenizer: &TikTokenTokenizer,
777        left: &str,
778        special: &str,
779        right: &str,
780    ) {
781        let whole = tokenizer
782            .encode(&format!("{left}{special}{right}"))
783            .unwrap()
784            .token_ids()
785            .to_vec();
786        let mut split = tokenizer
787            .encode(&format!("{left}{special}"))
788            .unwrap()
789            .token_ids()
790            .to_vec();
791        split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
792        assert_eq!(
793            whole, split,
794            "splitting after registered special token {special:?} must preserve tokenization"
795        );
796    }
797
798    #[test]
799    fn test_registered_special_tokens_are_cache_safe_boundaries() {
800        let dir = tempfile::tempdir().unwrap();
801        let file_path = create_byte_level_tiktoken_file(dir.path());
802        create_test_config(dir.path(), "kimi");
803        create_test_tokenizer_config(dir.path(), 256);
804
805        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
806        let specials = tokenizer.special_tokens();
807        assert!(!specials.is_empty());
808
809        for (index, special) in specials.iter().enumerate() {
810            assert!(
811                !has_nontrivial_self_overlap(special),
812                "registered special token has an ambiguous self-overlap: {special:?}"
813            );
814            for other in &specials[index + 1..] {
815                assert!(
816                    !have_ambiguous_overlap(special, other),
817                    "registered special tokens overlap ambiguously: {special:?} and {other:?}"
818                );
819            }
820
821            for (left, right) in [
822                ("ordinary prefix ", " ordinary suffix"),
823                ("line before\n", "\tUnicode after: εŒ—δΊ¬ πŸ˜€"),
824                ("", " trailing text"),
825            ] {
826                assert_special_split_invariant(&tokenizer, left, special, right);
827            }
828        }
829
830        let bos = specials
831            .iter()
832            .find(|token| token.as_str() == "[BOS]")
833            .unwrap();
834        let eos = specials
835            .iter()
836            .find(|token| token.as_str() == "[EOS]")
837            .unwrap();
838        assert_special_split_invariant(
839            &tokenizer,
840            "prefix ",
841            bos,
842            &format!("{eos}<|reserved_token_258|>Unicode 尾部"),
843        );
844    }
845
846    /// Regression test for Kimi K2.5 tokenizer inflation.
847    ///
848    /// Python's tokenization_kimi.py names unnamed reserved tokens by absolute ID:
849    ///   `<|reserved_token_{absolute_id}|>`
850    ///
851    /// The Rust code previously used relative offsets (0..255) as the naming index,
852    /// so when a prompt contained `<|reserved_token_258|>` the Rust tokenizer did NOT
853    /// recognize it as a special token. Each occurrence was encoded as multiple BPE
854    /// tokens instead of 1, inflating an 8192-token prompt to 9038 tokens and causing
855    /// TRT-LLM to reject the request.
856    #[test]
857    fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
858        let dir = tempfile::tempdir().unwrap();
859        let file_path = create_byte_level_tiktoken_file(dir.path());
860
861        // config.json with kimi model type -> triggers KIMI_PATTERN
862        create_test_config(dir.path(), "kimi");
863
864        // tokenizer_config.json: BOS at 256, EOS at 257. Base vocab is IDs 0..255.
865        create_test_tokenizer_config(dir.path(), 256);
866
867        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
868
869        // ID 256 = [BOS], ID 257 = [EOS].
870        // ID 258 = first UNNAMED reserved token.
871        //   With fix:   named <|reserved_token_258|>  (absolute ID)
872        //   Before fix: named <|reserved_token_2|>    (relative offset)
873
874        // Single unnamed reserved token should be recognized as 1 special token.
875        let single = "<|reserved_token_258|>";
876        let enc = tokenizer.encode(single).unwrap();
877        assert_eq!(
878            enc.token_ids().len(),
879            1,
880            "'{single}' should be 1 special token, got {} tokens: {:?}. \
881             This means fallback naming still uses relative offsets instead of absolute IDs.",
882            enc.token_ids().len(),
883            enc.token_ids()
884        );
885        assert_eq!(enc.token_ids()[0], 258);
886
887        // Multiple unnamed reserved tokens in sequence (mini version of the benchmark).
888        // IDs 258..268 are all unnamed; with the fix they're <|reserved_token_258|>..267.
889        let multi: String = (258u32..268)
890            .map(|id| format!("<|reserved_token_{id}|>"))
891            .collect();
892        let enc_multi = tokenizer.encode(&multi).unwrap();
893        assert_eq!(
894            enc_multi.token_ids().len(),
895            10,
896            "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
897            enc_multi.token_ids().len(),
898            enc_multi.token_ids()
899        );
900        let expected_ids: Vec<u32> = (258..268).collect();
901        assert_eq!(enc_multi.token_ids(), &expected_ids);
902    }
903
904    /// Confirm that the old relative-offset naming would cause token inflation.
905    /// Manually builds a tokenizer whose special-token map uses the WRONG names
906    /// (relative offsets), then shows the same string encodes as many tokens.
907    #[test]
908    fn test_relative_offset_naming_causes_inflation() {
909        let dir = tempfile::tempdir().unwrap();
910        let file_path = create_byte_level_tiktoken_file(dir.path());
911
912        let _encoder = parse_tiktoken_file(&file_path).unwrap();
913        let num_base_tokens = 256usize;
914
915        // Build special tokens with the OLD (buggy) relative-offset naming
916        let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
917        bad_special_tokens.insert("[BOS]".to_string(), 256);
918        bad_special_tokens.insert("[EOS]".to_string(), 257);
919        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
920            let id = num_base_tokens as u32 + i;
921            if id != 256 && id != 257 {
922                // OLD naming: relative offset i, not absolute id
923                bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
924            }
925        }
926
927        let bad_tokenizer =
928            TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
929
930        // With the wrong naming, <|reserved_token_258|> is NOT recognized as special.
931        // It gets split into byte-level BPE tokens -> many more than 1.
932        let input = "<|reserved_token_258|>";
933        let enc = bad_tokenizer.encode(input).unwrap();
934        assert!(
935            enc.token_ids().len() > 1,
936            "With buggy relative-offset naming, '{}' should NOT be recognized as a \
937             single special token. Got {} token(s): {:?}",
938            input,
939            enc.token_ids().len(),
940            enc.token_ids()
941        );
942
943        // Show the inflation: 10 reserved tokens produce far more than 10 BPE tokens.
944        let multi: String = (258u32..268)
945            .map(|id| format!("<|reserved_token_{id}|>"))
946            .collect();
947        let enc_multi = bad_tokenizer.encode(&multi).unwrap();
948        assert!(
949            enc_multi.token_ids().len() > 10,
950            "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
951             Got {}",
952            enc_multi.token_ids().len(),
953        );
954    }
955}