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" | "kimi_linear" | "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, kimi_linear, 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_detect_bpe_pattern_kimi_linear() {
480        let dir = tempfile::tempdir().unwrap();
481        create_test_config(dir.path(), "kimi_linear");
482        assert_eq!(detect_bpe_pattern(dir.path()).unwrap(), KIMI_PATTERN);
483    }
484
485    #[test]
486    fn test_load_special_tokens_no_config() {
487        let dir = tempfile::tempdir().unwrap();
488        let tokens = load_special_tokens(dir.path(), 100).unwrap();
489        assert_eq!(tokens.len(), 256);
490        assert_eq!(tokens["<|reserved_token_100|>"], 100);
491        assert_eq!(tokens["<|reserved_token_355|>"], 355);
492    }
493
494    #[test]
495    fn test_load_special_tokens_with_config() {
496        let dir = tempfile::tempdir().unwrap();
497        create_test_tokenizer_config(dir.path(), 100);
498        let tokens = load_special_tokens(dir.path(), 100).unwrap();
499        assert_eq!(tokens["[BOS]"], 100);
500        assert_eq!(tokens["[EOS]"], 101);
501        // Should also have reserved tokens filling gaps
502        assert!(tokens.len() > 2);
503    }
504
505    /// Helper: create a tiktoken file that includes raw byte tokens (byte fallback tokens).
506    fn create_test_tiktoken_file_with_byte_tokens(dir: &Path) -> String {
507        let engine = base64::engine::general_purpose::STANDARD;
508        let mut content = String::new();
509
510        let tokens: Vec<(&[u8], u32)> = vec![
511            (b"h", 0),
512            (b"e", 1),
513            (b"l", 2),
514            (b"o", 3),
515            (b" ", 4),
516            (b"hello", 5),
517        ];
518
519        for (token, rank) in &tokens {
520            let encoded = engine.encode(token);
521            content.push_str(&format!("{encoded} {rank}\n"));
522        }
523
524        // Byte-fallback tokens: individual bytes that form CJK character "δ½ " (U+4F60)
525        // UTF-8 encoding: 0xE4 0xBD 0xA0
526        let byte_tokens: Vec<(Vec<u8>, u32)> =
527            vec![(vec![0xE4], 100), (vec![0xBD], 101), (vec![0xA0], 102)];
528
529        for (token, rank) in &byte_tokens {
530            let encoded = engine.encode(token);
531            content.push_str(&format!("{encoded} {rank}\n"));
532        }
533
534        // Bytes for emoji "πŸ˜€" (U+1F600) β€” 4-byte UTF-8: 0xF0 0x9F 0x98 0x80
535        let emoji_tokens: Vec<(Vec<u8>, u32)> = vec![
536            (vec![0xF0], 200),
537            (vec![0x9F], 201),
538            (vec![0x98], 202),
539            (vec![0x80], 203),
540        ];
541
542        for (token, rank) in &emoji_tokens {
543            let encoded = engine.encode(token);
544            content.push_str(&format!("{encoded} {rank}\n"));
545        }
546
547        // Legitimate U+FFFD token: valid UTF-8 bytes EF BF BD (replacement character
548        // as an actual vocabulary entry, not an artifact of lossy conversion)
549        let fffd_token: Vec<(Vec<u8>, u32)> = vec![(vec![0xEF, 0xBF, 0xBD], 300)];
550
551        for (token, rank) in &fffd_token {
552            let encoded = engine.encode(token);
553            content.push_str(&format!("{encoded} {rank}\n"));
554        }
555
556        let file_path = dir.join("tiktoken.model");
557        let mut file = std::fs::File::create(&file_path).unwrap();
558        file.write_all(content.as_bytes()).unwrap();
559        file_path.to_str().unwrap().to_string()
560    }
561
562    fn create_byte_token_tokenizer(dir: &Path) -> TikTokenTokenizer {
563        let file_path = create_test_tiktoken_file_with_byte_tokens(dir);
564        let special_tokens = FxHashMap::default();
565        let pattern = r"[\w]+|[^\w\s]+|\s+";
566        TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap()
567    }
568
569    /// Reproduces the original panic: decoding a single byte-fallback token that is
570    /// part of a multi-byte UTF-8 character. Before the fix, CoreBPE::decode() would
571    /// call String::from_utf8() on [0xE4] and error with "incomplete utf-8 byte sequence".
572    #[test]
573    fn test_decode_single_incomplete_utf8_byte_does_not_error() {
574        let dir = tempfile::tempdir().unwrap();
575        let tokenizer = create_byte_token_tokenizer(dir.path());
576
577        let result = tokenizer.decode(&[100], false);
578        assert!(
579            result.is_ok(),
580            "decode() should not error on incomplete UTF-8 bytes"
581        );
582        let decode_result = result.unwrap();
583        assert!(
584            decode_result.is_partial(),
585            "incomplete UTF-8 byte should produce DecodeResult::Partial, got: {:?}",
586            decode_result
587        );
588    }
589
590    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
591    #[test]
592    fn test_decode_two_of_three_utf8_bytes_does_not_error() {
593        let dir = tempfile::tempdir().unwrap();
594        let tokenizer = create_byte_token_tokenizer(dir.path());
595
596        let result = tokenizer.decode(&[100, 101], false);
597        assert!(result.is_ok());
598        let decode_result = result.unwrap();
599        assert!(
600            decode_result.is_partial(),
601            "incomplete 2-of-3 UTF-8 bytes should produce DecodeResult::Partial, got: {:?}",
602            decode_result
603        );
604    }
605
606    /// When all bytes of a multi-byte character are present, the concatenated bytes form
607    /// valid UTF-8, so this test passes both before and after the fix. It serves as a
608    /// correctness check that the lossy conversion doesn't corrupt complete characters.
609    #[test]
610    fn test_decode_complete_multibyte_utf8_produces_correct_char() {
611        let dir = tempfile::tempdir().unwrap();
612        let tokenizer = create_byte_token_tokenizer(dir.path());
613
614        let result = tokenizer.decode(&[100, 101, 102], false);
615        assert!(result.is_ok());
616        assert_eq!(String::from(result.unwrap()), "δ½ ");
617    }
618
619    /// All 4 emoji bytes together form valid UTF-8, so this passes both before and after
620    /// the fix. Validates that lossy conversion doesn't alter complete multi-byte sequences.
621    #[test]
622    fn test_decode_complete_4byte_emoji_from_byte_tokens() {
623        let dir = tempfile::tempdir().unwrap();
624        let tokenizer = create_byte_token_tokenizer(dir.path());
625
626        let result = tokenizer.decode(&[200, 201, 202, 203], false);
627        assert!(result.is_ok());
628        assert_eq!(String::from(result.unwrap()), "πŸ˜€");
629    }
630
631    /// Regression test: a vocabulary token whose raw bytes are EF BF BD (the valid
632    /// UTF-8 encoding of U+FFFD) must decode as `Complete`, not `Partial`. Before the
633    /// from_utf8 fast-path fix, from_utf8_lossy + the trailing-FFFD heuristic would
634    /// misclassify this as Partial, causing the incremental decoder to suppress it.
635    #[test]
636    fn test_decode_legitimate_replacement_char_token_is_complete() {
637        let dir = tempfile::tempdir().unwrap();
638        let tokenizer = create_byte_token_tokenizer(dir.path());
639
640        let result = tokenizer.decode(&[300], false);
641        assert!(result.is_ok());
642        let decode_result = result.unwrap();
643        assert!(
644            decode_result.is_complete(),
645            "legitimate U+FFFD vocab token must be Complete, got: {:?}",
646            decode_result
647        );
648        assert_eq!(decode_result.as_str(), "\u{FFFD}");
649    }
650
651    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
652    #[test]
653    fn test_decode_partial_emoji_does_not_error() {
654        let dir = tempfile::tempdir().unwrap();
655        let tokenizer = create_byte_token_tokenizer(dir.path());
656
657        let result = tokenizer.decode(&[200], false);
658        assert!(result.is_ok());
659        assert!(result.unwrap().is_partial());
660    }
661
662    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
663    #[test]
664    fn test_decode_mixed_ascii_and_incomplete_bytes() {
665        let dir = tempfile::tempdir().unwrap();
666        let tokenizer = create_byte_token_tokenizer(dir.path());
667
668        let result = tokenizer.decode(&[5, 100], false);
669        assert!(result.is_ok());
670        let decode_result = result.unwrap();
671        assert!(
672            decode_result.is_partial(),
673            "trailing incomplete byte should produce DecodeResult::Partial"
674        );
675        let text: String = decode_result.into();
676        assert!(
677            text.starts_with("hello"),
678            "should start with 'hello', got: {:?}",
679            text
680        );
681    }
682
683    /// End-to-end incremental detokenization: DecodeStream buffers partial bytes,
684    /// emits the complete character once all bytes arrive.
685    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
686    #[test]
687    fn test_decode_stream_incremental_multibyte_reassembly() {
688        let dir = tempfile::tempdir().unwrap();
689        let tokenizer = create_byte_token_tokenizer(dir.path());
690        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
691
692        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
693
694        let r1 = stream.step(100).unwrap();
695        assert_eq!(r1, None, "first byte of 3-byte char should be buffered");
696
697        let r2 = stream.step(101).unwrap();
698        assert_eq!(r2, None, "second byte of 3-byte char should be buffered");
699
700        let r3 = stream.step(102).unwrap();
701        assert!(r3.is_some(), "third byte should complete the character");
702        assert_eq!(r3.unwrap(), "δ½ ");
703    }
704
705    /// Without the fix, fails with "incomplete utf-8 byte sequence" from CoreBPE::decode().
706    #[test]
707    fn test_decode_stream_incremental_emoji_reassembly() {
708        let dir = tempfile::tempdir().unwrap();
709        let tokenizer = create_byte_token_tokenizer(dir.path());
710        let tokenizer_arc: Arc<dyn crate::traits::Tokenizer> = Arc::new(tokenizer);
711
712        let mut stream = DecodeStream::new(tokenizer_arc, &[5], false);
713
714        let r1 = stream.step(200).unwrap();
715        assert_eq!(r1, None, "byte 1/4 of emoji should be buffered");
716
717        let r2 = stream.step(201).unwrap();
718        assert_eq!(r2, None, "byte 2/4 of emoji should be buffered");
719
720        let r3 = stream.step(202).unwrap();
721        assert_eq!(r3, None, "byte 3/4 of emoji should be buffered");
722
723        let r4 = stream.step(203).unwrap();
724        assert!(r4.is_some(), "byte 4/4 should complete the emoji");
725        assert_eq!(r4.unwrap(), "πŸ˜€");
726    }
727
728    #[test]
729    fn test_tiktoken_encode_batch() {
730        let dir = tempfile::tempdir().unwrap();
731        let file_path = create_test_tiktoken_file(dir.path());
732
733        let special_tokens = FxHashMap::default();
734        let pattern = r"[\w]+|[^\w\s]+|\s+";
735
736        let tokenizer = TikTokenTokenizer::from_file(&file_path, pattern, special_tokens).unwrap();
737
738        let inputs = &["hello", "world"];
739        let encodings = tokenizer.encode_batch(inputs).unwrap();
740        assert_eq!(encodings.len(), 2);
741
742        for (encoding, input) in encodings.iter().zip(inputs.iter()) {
743            let decoded: String = tokenizer
744                .decode(encoding.token_ids(), false)
745                .unwrap()
746                .into();
747            assert_eq!(decoded, *input);
748        }
749    }
750
751    /// Helper: create a tiktoken file containing all 256 single-byte tokens (ranks 0..255).
752    /// This gives a complete byte-level base vocabulary so any ASCII string can be encoded.
753    fn create_byte_level_tiktoken_file(dir: &Path) -> String {
754        let engine = base64::engine::general_purpose::STANDARD;
755        let mut content = String::new();
756        for byte_val in 0u16..256 {
757            let encoded = engine.encode([byte_val as u8]);
758            content.push_str(&format!("{encoded} {byte_val}\n"));
759        }
760        let file_path = dir.join("tiktoken.model");
761        std::fs::write(&file_path, &content).unwrap();
762        file_path.to_str().unwrap().to_string()
763    }
764
765    fn has_nontrivial_self_overlap(token: &str) -> bool {
766        let bytes = token.as_bytes();
767        (1..bytes.len()).any(|overlap| bytes[bytes.len() - overlap..] == bytes[..overlap])
768    }
769
770    fn have_ambiguous_overlap(a: &str, b: &str) -> bool {
771        let a = a.as_bytes();
772        let b = b.as_bytes();
773
774        if a.windows(b.len()).any(|window| window == b)
775            || b.windows(a.len()).any(|window| window == a)
776        {
777            return true;
778        }
779
780        let max_overlap = a.len().min(b.len());
781        (1..max_overlap).any(|overlap| {
782            a[a.len() - overlap..] == b[..overlap] || b[b.len() - overlap..] == a[..overlap]
783        })
784    }
785
786    fn assert_special_split_invariant(
787        tokenizer: &TikTokenTokenizer,
788        left: &str,
789        special: &str,
790        right: &str,
791    ) {
792        let whole = tokenizer
793            .encode(&format!("{left}{special}{right}"))
794            .unwrap()
795            .token_ids()
796            .to_vec();
797        let mut split = tokenizer
798            .encode(&format!("{left}{special}"))
799            .unwrap()
800            .token_ids()
801            .to_vec();
802        split.extend_from_slice(tokenizer.encode(right).unwrap().token_ids());
803        assert_eq!(
804            whole, split,
805            "splitting after registered special token {special:?} must preserve tokenization"
806        );
807    }
808
809    #[test]
810    fn test_registered_special_tokens_are_cache_safe_boundaries() {
811        let dir = tempfile::tempdir().unwrap();
812        let file_path = create_byte_level_tiktoken_file(dir.path());
813        create_test_config(dir.path(), "kimi");
814        create_test_tokenizer_config(dir.path(), 256);
815
816        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
817        let specials = tokenizer.special_tokens();
818        assert!(!specials.is_empty());
819
820        for (index, special) in specials.iter().enumerate() {
821            assert!(
822                !has_nontrivial_self_overlap(special),
823                "registered special token has an ambiguous self-overlap: {special:?}"
824            );
825            for other in &specials[index + 1..] {
826                assert!(
827                    !have_ambiguous_overlap(special, other),
828                    "registered special tokens overlap ambiguously: {special:?} and {other:?}"
829                );
830            }
831
832            for (left, right) in [
833                ("ordinary prefix ", " ordinary suffix"),
834                ("line before\n", "\tUnicode after: εŒ—δΊ¬ πŸ˜€"),
835                ("", " trailing text"),
836            ] {
837                assert_special_split_invariant(&tokenizer, left, special, right);
838            }
839        }
840
841        let bos = specials
842            .iter()
843            .find(|token| token.as_str() == "[BOS]")
844            .unwrap();
845        let eos = specials
846            .iter()
847            .find(|token| token.as_str() == "[EOS]")
848            .unwrap();
849        assert_special_split_invariant(
850            &tokenizer,
851            "prefix ",
852            bos,
853            &format!("{eos}<|reserved_token_258|>Unicode 尾部"),
854        );
855    }
856
857    /// Regression test for Kimi K2.5 tokenizer inflation.
858    ///
859    /// Python's tokenization_kimi.py names unnamed reserved tokens by absolute ID:
860    ///   `<|reserved_token_{absolute_id}|>`
861    ///
862    /// The Rust code previously used relative offsets (0..255) as the naming index,
863    /// so when a prompt contained `<|reserved_token_258|>` the Rust tokenizer did NOT
864    /// recognize it as a special token. Each occurrence was encoded as multiple BPE
865    /// tokens instead of 1, inflating an 8192-token prompt to 9038 tokens and causing
866    /// TRT-LLM to reject the request.
867    #[test]
868    fn test_reserved_token_absolute_id_naming_kimi_k25_regression() {
869        let dir = tempfile::tempdir().unwrap();
870        let file_path = create_byte_level_tiktoken_file(dir.path());
871
872        // config.json with kimi model type -> triggers KIMI_PATTERN
873        create_test_config(dir.path(), "kimi");
874
875        // tokenizer_config.json: BOS at 256, EOS at 257. Base vocab is IDs 0..255.
876        create_test_tokenizer_config(dir.path(), 256);
877
878        let tokenizer = TikTokenTokenizer::from_file_auto(&file_path).unwrap();
879
880        // ID 256 = [BOS], ID 257 = [EOS].
881        // ID 258 = first UNNAMED reserved token.
882        //   With fix:   named <|reserved_token_258|>  (absolute ID)
883        //   Before fix: named <|reserved_token_2|>    (relative offset)
884
885        // Single unnamed reserved token should be recognized as 1 special token.
886        let single = "<|reserved_token_258|>";
887        let enc = tokenizer.encode(single).unwrap();
888        assert_eq!(
889            enc.token_ids().len(),
890            1,
891            "'{single}' should be 1 special token, got {} tokens: {:?}. \
892             This means fallback naming still uses relative offsets instead of absolute IDs.",
893            enc.token_ids().len(),
894            enc.token_ids()
895        );
896        assert_eq!(enc.token_ids()[0], 258);
897
898        // Multiple unnamed reserved tokens in sequence (mini version of the benchmark).
899        // IDs 258..268 are all unnamed; with the fix they're <|reserved_token_258|>..267.
900        let multi: String = (258u32..268)
901            .map(|id| format!("<|reserved_token_{id}|>"))
902            .collect();
903        let enc_multi = tokenizer.encode(&multi).unwrap();
904        assert_eq!(
905            enc_multi.token_ids().len(),
906            10,
907            "10 reserved token strings should produce exactly 10 tokens, got {}: {:?}",
908            enc_multi.token_ids().len(),
909            enc_multi.token_ids()
910        );
911        let expected_ids: Vec<u32> = (258..268).collect();
912        assert_eq!(enc_multi.token_ids(), &expected_ids);
913    }
914
915    /// Confirm that the old relative-offset naming would cause token inflation.
916    /// Manually builds a tokenizer whose special-token map uses the WRONG names
917    /// (relative offsets), then shows the same string encodes as many tokens.
918    #[test]
919    fn test_relative_offset_naming_causes_inflation() {
920        let dir = tempfile::tempdir().unwrap();
921        let file_path = create_byte_level_tiktoken_file(dir.path());
922
923        let _encoder = parse_tiktoken_file(&file_path).unwrap();
924        let num_base_tokens = 256usize;
925
926        // Build special tokens with the OLD (buggy) relative-offset naming
927        let mut bad_special_tokens: FxHashMap<String, u32> = FxHashMap::default();
928        bad_special_tokens.insert("[BOS]".to_string(), 256);
929        bad_special_tokens.insert("[EOS]".to_string(), 257);
930        for i in 0..DEFAULT_NUM_RESERVED_SPECIAL_TOKENS {
931            let id = num_base_tokens as u32 + i;
932            if id != 256 && id != 257 {
933                // OLD naming: relative offset i, not absolute id
934                bad_special_tokens.insert(format!("<|reserved_token_{i}|>"), id);
935            }
936        }
937
938        let bad_tokenizer =
939            TikTokenTokenizer::from_file(&file_path, KIMI_PATTERN, bad_special_tokens).unwrap();
940
941        // With the wrong naming, <|reserved_token_258|> is NOT recognized as special.
942        // It gets split into byte-level BPE tokens -> many more than 1.
943        let input = "<|reserved_token_258|>";
944        let enc = bad_tokenizer.encode(input).unwrap();
945        assert!(
946            enc.token_ids().len() > 1,
947            "With buggy relative-offset naming, '{}' should NOT be recognized as a \
948             single special token. Got {} token(s): {:?}",
949            input,
950            enc.token_ids().len(),
951            enc.token_ids()
952        );
953
954        // Show the inflation: 10 reserved tokens produce far more than 10 BPE tokens.
955        let multi: String = (258u32..268)
956            .map(|id| format!("<|reserved_token_{id}|>"))
957            .collect();
958        let enc_multi = bad_tokenizer.encode(&multi).unwrap();
959        assert!(
960            enc_multi.token_ids().len() > 10,
961            "With buggy naming, 10 reserved token strings should inflate beyond 10 tokens. \
962             Got {}",
963            enc_multi.token_ids().len(),
964        );
965    }
966}