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