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