Skip to main content

ferrum_tokenizer/implementations/
huggingface.rs

1//! HuggingFace tokenizer implementation
2
3use crate::{IncrementalTokenizer, Tokenizer, TokenizerFactory, TokenizerInfo, TokenizerType};
4use async_trait::async_trait;
5use ferrum_types::{Result, SpecialTokens, TokenId};
6use parking_lot::RwLock;
7use std::collections::HashMap;
8use std::sync::{Arc, OnceLock};
9use tokenizers::decoders::DecoderWrapper;
10use tokenizers::Tokenizer as HfTokenizer;
11use tracing::debug;
12
13/// HuggingFace tokenizer wrapper
14pub struct HuggingFaceTokenizer {
15    tokenizer: Arc<HfTokenizer>,
16    special_tokens: SpecialTokens,
17    info: TokenizerInfo,
18    id_to_token: Vec<Option<String>>,
19    byte_level_decoder: bool,
20    /// Reasoning-marker token ids mapped to the canonical tag emitted in
21    /// their place. Some vocabs mark think tags `special: true`
22    /// (Magistral's `[THINK]`/`[/THINK]`), so a skip-special decode would
23    /// silently drop them and the thinking text would leak into content.
24    /// Decode preserves these ids and normalizes every dialect to
25    /// `<think>`/`</think>`, which is what the serving layer splits on.
26    think_markers: Vec<(u32, &'static str)>,
27    /// Incremental decode cache for efficiency
28    decode_cache: RwLock<DecodeCache>,
29}
30
31/// Marker-string dialects mapped to the canonical tags. Probed against the
32/// vocab at construction; absent dialects cost nothing.
33const THINK_MARKER_DIALECTS: [(&str, &'static str); 4] = [
34    ("<think>", "<think>"),
35    ("</think>", "</think>"),
36    ("[THINK]", "<think>"),
37    ("[/THINK]", "</think>"),
38];
39
40fn probe_think_markers(tokenizer: &HfTokenizer) -> Vec<(u32, &'static str)> {
41    THINK_MARKER_DIALECTS
42        .iter()
43        .filter_map(|(text, canonical)| tokenizer.token_to_id(text).map(|id| (id, *canonical)))
44        .collect()
45}
46
47/// Incremental decoding state
48#[derive(Debug, Clone, Default)]
49pub struct IncrementalState {
50    /// Accumulated tokens
51    tokens: Vec<TokenId>,
52    /// Decoded text so far
53    text: String,
54}
55
56/// Cache for decoded token sequences
57#[derive(Debug, Default)]
58struct DecodeCache {
59    cache: std::collections::HashMap<Vec<TokenId>, String>,
60    max_size: usize,
61}
62
63impl DecodeCache {
64    fn new(max_size: usize) -> Self {
65        Self {
66            cache: std::collections::HashMap::new(),
67            max_size,
68        }
69    }
70
71    fn get(&self, tokens: &[TokenId]) -> Option<&String> {
72        self.cache.get(tokens)
73    }
74
75    fn insert(&mut self, tokens: Vec<TokenId>, text: String) {
76        if self.cache.len() >= self.max_size {
77            let to_remove: Vec<_> = self
78                .cache
79                .keys()
80                .take(self.cache.len() / 2)
81                .cloned()
82                .collect();
83            for key in to_remove {
84                self.cache.remove(&key);
85            }
86        }
87        self.cache.insert(tokens, text);
88    }
89}
90
91fn decoded_incremental_delta(previous_text: &str, full_text: &str) -> Result<String> {
92    full_text
93        .strip_prefix(previous_text)
94        .map(ToOwned::to_owned)
95        .ok_or_else(|| {
96            ferrum_types::FerrumError::tokenizer(
97                "Incremental decode changed the previously emitted text prefix",
98            )
99        })
100}
101
102impl HuggingFaceTokenizer {
103    /// Create new HuggingFace tokenizer
104    pub async fn new(tokenizer: HfTokenizer) -> Result<Self> {
105        let vocab_size = tokenizer.get_vocab_size(false);
106        let id_to_token = build_id_to_token(&tokenizer);
107
108        // Extract special tokens
109        let special_tokens = extract_special_tokens(&tokenizer)?;
110
111        let info = TokenizerInfo {
112            tokenizer_type: TokenizerType::BPE, // Most HF tokenizers use BPE
113            vocab_size,
114            special_tokens: special_tokens.clone(),
115            supports_incremental: true,
116            supports_chat_template: false, // MVP: chat template support disabled
117            max_token_length: None,        // HF tokenizers don't expose this directly
118            model_name: None,              // Can be set externally
119        };
120
121        debug!(
122            "Created HuggingFace tokenizer with vocab size {}",
123            vocab_size
124        );
125
126        let think_markers = probe_think_markers(&tokenizer);
127        let byte_level_decoder = tokenizer.get_decoder().is_some_and(decoder_uses_byte_level);
128
129        Ok(Self {
130            tokenizer: Arc::new(tokenizer),
131            special_tokens,
132            info,
133            id_to_token,
134            byte_level_decoder,
135            think_markers,
136            decode_cache: RwLock::new(DecodeCache::new(1000)),
137        })
138    }
139
140    /// Create from file path. Special tokens are resolved from the sibling
141    /// `generation_config.json` / `tokenizer_config.json` when present;
142    /// vocab name probing is only the fallback for bare tokenizer.json files.
143    pub async fn from_file(path: &str) -> Result<Self> {
144        let tokenizer = HfTokenizer::from_file(path).map_err(|e| {
145            ferrum_types::FerrumError::tokenizer(format!("Failed to load tokenizer: {}", e))
146        })?;
147        let overrides =
148            special_token_overrides_from_configs(std::path::Path::new(path), &tokenizer);
149        let mut this = Self::new(tokenizer).await?;
150        this.apply_special_token_overrides(overrides);
151        Ok(this)
152    }
153
154    /// Construct from the immutable bytes retained by product source
155    /// resolution. Config bytes are parsed from the same source bundle so
156    /// special-token semantics cannot drift through a second filesystem read.
157    pub async fn from_source_bytes(
158        tokenizer_json: &[u8],
159        tokenizer_config_json: Option<&[u8]>,
160        generation_config_json: Option<&[u8]>,
161    ) -> Result<Self> {
162        let tokenizer = HfTokenizer::from_bytes(tokenizer_json).map_err(|error| {
163            ferrum_types::FerrumError::tokenizer(format!(
164                "Failed to load tokenizer from resolved source bytes: {error}"
165            ))
166        })?;
167        let tokenizer_config =
168            parse_optional_config_bytes(tokenizer_config_json, "tokenizer_config.json")?;
169        let generation_config =
170            parse_optional_config_bytes(generation_config_json, "generation_config.json")?;
171        let overrides = special_token_overrides_from_values(
172            generation_config.as_ref(),
173            tokenizer_config.as_ref(),
174            &tokenizer,
175        );
176        let mut this = Self::new(tokenizer).await?;
177        this.apply_special_token_overrides(overrides);
178        Ok(this)
179    }
180
181    fn apply_special_token_overrides(&mut self, overrides: SpecialTokenOverrides) {
182        if overrides.bos.is_some() {
183            self.special_tokens.bos_token = overrides.bos;
184        }
185        if overrides.eos.is_some() {
186            self.special_tokens.eos_token = overrides.eos;
187        }
188        if !overrides.extra_eos.is_empty() {
189            self.special_tokens.extra_eos_tokens = overrides.extra_eos;
190        }
191        self.info.special_tokens = self.special_tokens.clone();
192    }
193
194    /// Create from HuggingFace Hub
195    pub async fn from_pretrained(repo_id: &str, _revision: Option<&str>) -> Result<Self> {
196        let api = hf_hub::api::tokio::Api::new().map_err(|e| {
197            ferrum_types::FerrumError::tokenizer(format!("Failed to create HF API: {}", e))
198        })?;
199
200        let repo = api.repo(hf_hub::Repo::model(repo_id.to_string()));
201
202        // Note: hf_hub::api::tokio::ApiRepo doesn't have set_revision in newer versions
203        // Revision is handled via the Repo struct or api.model_with_revision
204        let tokenizer_file = repo.get("tokenizer.json").await.map_err(|e| {
205            ferrum_types::FerrumError::tokenizer(format!("Failed to download tokenizer: {}", e))
206        })?;
207
208        let tokenizer = HfTokenizer::from_file(&tokenizer_file).map_err(|e| {
209            ferrum_types::FerrumError::tokenizer(format!("Failed to load tokenizer: {}", e))
210        })?;
211
212        Self::new(tokenizer).await
213    }
214}
215
216impl Tokenizer for HuggingFaceTokenizer {
217    fn encode(&self, text: &str, add_special: bool) -> Result<Vec<TokenId>> {
218        let encoding = self
219            .tokenizer
220            .encode(text, add_special)
221            .map_err(|e| ferrum_types::FerrumError::tokenizer(format!("Encoding failed: {}", e)))?;
222
223        Ok(encoding
224            .get_ids()
225            .iter()
226            .map(|&id| TokenId::new(id))
227            .collect())
228    }
229
230    fn decode(&self, tokens: &[TokenId], skip_special: bool) -> Result<String> {
231        let token_ids: Vec<u32> = tokens.iter().map(|t| t.get()).collect();
232
233        // Skip-special decode must not swallow reasoning markers: split at
234        // marker ids, decode the segments, and splice the canonical tags
235        // back in. The common no-marker case stays a single decode call.
236        if skip_special
237            && !self.think_markers.is_empty()
238            && token_ids
239                .iter()
240                .any(|id| self.think_markers.iter().any(|(mid, _)| mid == id))
241        {
242            let mut out = String::new();
243            let mut segment: Vec<u32> = Vec::with_capacity(token_ids.len());
244            for id in &token_ids {
245                if let Some((_, canonical)) = self.think_markers.iter().find(|(mid, _)| mid == id) {
246                    if !segment.is_empty() {
247                        out.push_str(&self.tokenizer.decode(&segment, true).map_err(|e| {
248                            ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e))
249                        })?);
250                        segment.clear();
251                    }
252                    out.push_str(canonical);
253                } else {
254                    segment.push(*id);
255                }
256            }
257            if !segment.is_empty() {
258                out.push_str(&self.tokenizer.decode(&segment, true).map_err(|e| {
259                    ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e))
260                })?);
261            }
262            return Ok(out);
263        }
264
265        let text = self
266            .tokenizer
267            .decode(&token_ids, skip_special)
268            .map_err(|e| ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e)))?;
269
270        Ok(text)
271    }
272
273    fn decode_incremental(&self, prev: &[TokenId], next: TokenId) -> Result<String> {
274        // Clone under the read lock so the guard is released before this
275        // request inserts the extended sequence under the write lock.
276        let cached_prev = { self.decode_cache.read().get(prev).cloned() };
277        if let Some(cached_prev) = cached_prev {
278            let mut all_tokens = prev.to_vec();
279            all_tokens.push(next);
280            let full_text = self.decode(&all_tokens, true)?;
281
282            self.decode_cache
283                .write()
284                .insert(all_tokens, full_text.clone());
285
286            return decoded_incremental_delta(&cached_prev, &full_text);
287        }
288
289        // No cache hit, decode both
290        let prev_text = if prev.is_empty() {
291            String::new()
292        } else {
293            self.decode(prev, true)?
294        };
295
296        let mut all_tokens = prev.to_vec();
297        all_tokens.push(next);
298        let full_text = self.decode(&all_tokens, true)?;
299
300        // Update cache
301        {
302            let mut cache = self.decode_cache.write();
303            if !prev.is_empty() {
304                cache.insert(prev.to_vec(), prev_text.clone());
305            }
306            cache.insert(all_tokens, full_text.clone());
307        }
308
309        decoded_incremental_delta(&prev_text, &full_text)
310    }
311
312    fn vocab_size(&self) -> usize {
313        self.info.vocab_size
314    }
315
316    fn special_tokens(&self) -> &SpecialTokens {
317        &self.special_tokens
318    }
319
320    fn token_id(&self, text: &str) -> Option<TokenId> {
321        self.tokenizer.token_to_id(text).map(TokenId::new)
322    }
323
324    fn token_text(&self, token_id: TokenId) -> Option<&str> {
325        self.id_to_token
326            .get(token_id.get() as usize)
327            .and_then(|value| value.as_deref())
328    }
329
330    fn token_bytes(&self, token_id: TokenId) -> Option<Vec<u8>> {
331        if self.byte_level_decoder {
332            return self.token_text(token_id).map(byte_level_token_bytes);
333        }
334        self.decode(&[token_id], false)
335            .ok()
336            .map(String::into_bytes)
337            .or_else(|| {
338                self.token_text(token_id)
339                    .map(|text| text.as_bytes().to_vec())
340            })
341    }
342
343    fn apply_chat_template(
344        &self,
345        messages: &[ferrum_interfaces::tokenizer::ChatMessage],
346    ) -> Result<String> {
347        // MVP: simple concatenation
348        let mut result = String::new();
349        for msg in messages {
350            result.push_str(&format!("{}: {}\n", msg.role, msg.content));
351        }
352        Ok(result.trim_end().to_string())
353    }
354
355    fn info(&self) -> TokenizerInfo {
356        self.info.clone()
357    }
358}
359
360impl IncrementalTokenizer for HuggingFaceTokenizer {
361    type State = IncrementalState;
362
363    fn create_state(&self) -> Self::State {
364        IncrementalState::default()
365    }
366
367    fn decode_incremental_with_state(
368        &self,
369        state: &mut Self::State,
370        token: TokenId,
371    ) -> Result<String> {
372        state.tokens.push(token);
373
374        // Decode all tokens
375        let full_text = self.decode(&state.tokens, true)?;
376
377        // Calculate delta before updating the append-only decoded prefix.
378        let delta = decoded_incremental_delta(&state.text, &full_text)?;
379
380        // Update state
381        state.text = full_text;
382
383        Ok(delta)
384    }
385
386    fn reset_state(&self, state: &mut Self::State) {
387        state.tokens.clear();
388        state.text.clear();
389    }
390
391    fn get_decoded_text(&self, state: &Self::State) -> String {
392        state.text.clone()
393    }
394}
395
396/// HuggingFace tokenizer factory
397#[derive(Debug, Clone, Default)]
398pub struct HuggingFaceTokenizerFactory;
399
400impl HuggingFaceTokenizerFactory {
401    pub fn new() -> Self {
402        Self
403    }
404}
405
406#[async_trait]
407impl TokenizerFactory for HuggingFaceTokenizerFactory {
408    async fn load_from_file(&self, path: &str) -> Result<Box<dyn Tokenizer>> {
409        let tokenizer = HuggingFaceTokenizer::from_file(path).await?;
410        Ok(Box::new(tokenizer))
411    }
412
413    async fn load_from_bytes(&self, data: &[u8]) -> Result<Box<dyn Tokenizer>> {
414        let tokenizer = HfTokenizer::from_bytes(data).map_err(|e| {
415            ferrum_types::FerrumError::tokenizer(format!(
416                "Failed to load tokenizer from bytes: {}",
417                e
418            ))
419        })?;
420        let tokenizer = HuggingFaceTokenizer::new(tokenizer).await?;
421        Ok(Box::new(tokenizer))
422    }
423
424    async fn load_from_hub(
425        &self,
426        repo_id: &str,
427        revision: Option<&str>,
428    ) -> Result<Box<dyn Tokenizer>> {
429        let tokenizer = HuggingFaceTokenizer::from_pretrained(repo_id, revision).await?;
430        Ok(Box::new(tokenizer))
431    }
432
433    async fn create_from_config(
434        &self,
435        config: &ferrum_interfaces::tokenizer::TokenizerConfig,
436    ) -> Result<Box<dyn Tokenizer>> {
437        // Load from path specified in config
438        self.load_from_file(&config.path).await
439    }
440
441    fn supported_types(&self) -> Vec<TokenizerType> {
442        vec![
443            TokenizerType::BPE,
444            TokenizerType::WordPiece,
445            TokenizerType::SentencePiece,
446        ]
447    }
448}
449
450// ============================================================================
451// Helper Functions
452// ============================================================================
453
454fn build_id_to_token(tokenizer: &HfTokenizer) -> Vec<Option<String>> {
455    let vocab = tokenizer.get_vocab(true);
456    let Some(max_id) = vocab.values().copied().max() else {
457        return Vec::new();
458    };
459    let mut id_to_token = vec![None; max_id as usize + 1];
460    for (token, id) in vocab {
461        let slot = &mut id_to_token[id as usize];
462        if slot.is_none() {
463            *slot = Some(token);
464        }
465    }
466    id_to_token
467}
468
469fn decoder_uses_byte_level(decoder: &DecoderWrapper) -> bool {
470    match decoder {
471        DecoderWrapper::ByteLevel(_) => true,
472        DecoderWrapper::Sequence(sequence) => {
473            sequence.get_decoders().iter().any(decoder_uses_byte_level)
474        }
475        _ => false,
476    }
477}
478
479fn byte_level_char_bytes() -> &'static HashMap<char, u8> {
480    static CHAR_BYTES: OnceLock<HashMap<char, u8>> = OnceLock::new();
481    CHAR_BYTES.get_or_init(|| {
482        let mut direct = Vec::with_capacity(256);
483        direct.extend(b'!'..=b'~');
484        direct.extend(b'\xA1'..=b'\xAC');
485        direct.extend(b'\xAE'..=b'\xFF');
486
487        let mut next_codepoint = 256u32;
488        let mut mapping = HashMap::with_capacity(256);
489        for byte in 0..=u8::MAX {
490            let codepoint = if direct.contains(&byte) {
491                byte as u32
492            } else {
493                let codepoint = next_codepoint;
494                next_codepoint += 1;
495                codepoint
496            };
497            let character = char::from_u32(codepoint)
498                .expect("GPT-2 byte alphabet uses valid Unicode scalar values");
499            mapping.insert(character, byte);
500        }
501        mapping
502    })
503}
504
505fn byte_level_token_bytes(token: &str) -> Vec<u8> {
506    let mapping = byte_level_char_bytes();
507    token
508        .chars()
509        .map(|character| mapping.get(&character).copied())
510        .collect::<Option<Vec<_>>>()
511        .unwrap_or_else(|| token.as_bytes().to_vec())
512}
513
514/// Extract special tokens from HF tokenizer
515fn extract_special_tokens(tokenizer: &HfTokenizer) -> Result<SpecialTokens> {
516    let _vocab = tokenizer.get_vocab(false);
517
518    let bos_token = tokenizer
519        .token_to_id("<s>")
520        .or_else(|| tokenizer.token_to_id("[BOS]"))
521        .or_else(|| tokenizer.token_to_id("<bos>"))
522        .map(TokenId::new);
523
524    let eos_token = tokenizer
525        .token_to_id("</s>")
526        .or_else(|| tokenizer.token_to_id("[EOS]"))
527        .or_else(|| tokenizer.token_to_id("<eos>"))
528        .map(TokenId::new);
529
530    let unk_token = tokenizer
531        .token_to_id("<unk>")
532        .or_else(|| tokenizer.token_to_id("[UNK]"))
533        .map(TokenId::new);
534
535    let pad_token = tokenizer
536        .token_to_id("<pad>")
537        .or_else(|| tokenizer.token_to_id("[PAD]"))
538        .map(TokenId::new);
539
540    let sep_token = tokenizer
541        .token_to_id("[SEP]")
542        .or_else(|| tokenizer.token_to_id("<sep>"))
543        .map(TokenId::new);
544
545    let cls_token = tokenizer
546        .token_to_id("[CLS]")
547        .or_else(|| tokenizer.token_to_id("<cls>"))
548        .map(TokenId::new);
549
550    let mask_token = tokenizer
551        .token_to_id("[MASK]")
552        .or_else(|| tokenizer.token_to_id("<mask>"))
553        .map(TokenId::new);
554
555    Ok(SpecialTokens {
556        bos_token,
557        eos_token,
558        unk_token,
559        pad_token,
560        sep_token,
561        cls_token,
562        mask_token,
563        extra_eos_tokens: Vec::new(),
564    })
565}
566
567/// EOS/BOS overrides read from the model's config files next to
568/// `tokenizer.json`. The vocabulary alone cannot tell which token a model is
569/// trained to emit as EOS — name probing (`</s>`-style) breaks on models that
570/// rename their special-token slots (e.g. DeepSeek-R1 distills rename
571/// `<|endoftext|>` to `<|end▁of▁sentence|>`), so the config files are
572/// authoritative: `generation_config.json` `eos_token_id` (int or list)
573/// first, then `tokenizer_config.json` `eos_token` / `bos_token` (string or
574/// `{"content": ...}` AddedToken object).
575#[derive(Debug, Default)]
576struct SpecialTokenOverrides {
577    bos: Option<TokenId>,
578    eos: Option<TokenId>,
579    extra_eos: Vec<TokenId>,
580}
581
582fn special_token_overrides_from_configs(
583    tokenizer_json: &std::path::Path,
584    tokenizer: &HfTokenizer,
585) -> SpecialTokenOverrides {
586    let Some(dir) = tokenizer_json.parent() else {
587        return SpecialTokenOverrides::default();
588    };
589    let generation_config = read_json(&dir.join("generation_config.json"));
590    let tokenizer_config = read_json(&dir.join("tokenizer_config.json"));
591    special_token_overrides_from_values(
592        generation_config.as_ref(),
593        tokenizer_config.as_ref(),
594        tokenizer,
595    )
596}
597
598fn special_token_overrides_from_values(
599    generation_config: Option<&serde_json::Value>,
600    tokenizer_config: Option<&serde_json::Value>,
601    tokenizer: &HfTokenizer,
602) -> SpecialTokenOverrides {
603    let mut overrides = SpecialTokenOverrides::default();
604
605    if let Some(gen) = generation_config {
606        let mut eos_ids = token_id_list(gen.get("eos_token_id"));
607        if !eos_ids.is_empty() {
608            overrides.eos = Some(eos_ids.remove(0));
609            overrides.extra_eos = eos_ids;
610        }
611        if let Some(bos) = token_id_list(gen.get("bos_token_id")).into_iter().next() {
612            overrides.bos = Some(bos);
613        }
614    }
615
616    if let Some(tok_cfg) = tokenizer_config {
617        if overrides.eos.is_none() {
618            overrides.eos = token_from_config_value(tok_cfg.get("eos_token"), tokenizer);
619        }
620        if overrides.bos.is_none() {
621            overrides.bos = token_from_config_value(tok_cfg.get("bos_token"), tokenizer);
622        }
623    }
624
625    overrides
626}
627
628fn parse_optional_config_bytes(
629    bytes: Option<&[u8]>,
630    source_file: &str,
631) -> Result<Option<serde_json::Value>> {
632    bytes
633        .map(|bytes| {
634            serde_json::from_slice(bytes).map_err(|error| {
635                ferrum_types::FerrumError::tokenizer(format!(
636                    "Failed to parse resolved {source_file}: {error}"
637                ))
638            })
639        })
640        .transpose()
641}
642
643fn read_json(path: &std::path::Path) -> Option<serde_json::Value> {
644    let text = std::fs::read_to_string(path).ok()?;
645    serde_json::from_str(&text).ok()
646}
647
648/// `eos_token_id` / `bos_token_id` in generation_config.json is either a
649/// number or a list of numbers.
650fn token_id_list(value: Option<&serde_json::Value>) -> Vec<TokenId> {
651    match value {
652        Some(serde_json::Value::Number(n)) => n
653            .as_u64()
654            .map(|v| vec![TokenId::new(v as u32)])
655            .unwrap_or_default(),
656        Some(serde_json::Value::Array(items)) => items
657            .iter()
658            .filter_map(|v| v.as_u64())
659            .map(|v| TokenId::new(v as u32))
660            .collect(),
661        _ => Vec::new(),
662    }
663}
664
665/// `eos_token` / `bos_token` in tokenizer_config.json is either a plain
666/// string or an AddedToken object `{"content": "...", ...}`.
667fn token_from_config_value(
668    value: Option<&serde_json::Value>,
669    tokenizer: &HfTokenizer,
670) -> Option<TokenId> {
671    let text = match value? {
672        serde_json::Value::String(s) => s.as_str(),
673        serde_json::Value::Object(obj) => obj.get("content")?.as_str()?,
674        _ => return None,
675    };
676    tokenizer.token_to_id(text).map(TokenId::new)
677}
678
679#[cfg(test)]
680mod tests {
681    use super::*;
682
683    #[test]
684    fn test_decode_cache_creation() {
685        let cache = DecodeCache::new(100);
686        assert_eq!(cache.max_size, 100);
687        assert_eq!(cache.cache.len(), 0);
688    }
689
690    #[test]
691    fn test_decode_cache_insert_and_get() {
692        let mut cache = DecodeCache::new(10);
693        let tokens = vec![TokenId::new(1), TokenId::new(2)];
694        let text = "hello".to_string();
695
696        cache.insert(tokens.clone(), text.clone());
697
698        let result = cache.get(&tokens);
699        assert!(result.is_some());
700        assert_eq!(result.unwrap(), &text);
701    }
702
703    #[test]
704    fn test_decode_cache_eviction() {
705        let mut cache = DecodeCache::new(2);
706
707        // 填满缓存
708        cache.insert(vec![TokenId::new(1)], "a".to_string());
709        cache.insert(vec![TokenId::new(2)], "b".to_string());
710
711        assert_eq!(cache.cache.len(), 2);
712
713        // 触发驱逐
714        cache.insert(vec![TokenId::new(3)], "c".to_string());
715
716        // 应该已经清理了一些旧条目
717        assert!(cache.cache.len() <= 2);
718    }
719
720    #[test]
721    fn incremental_delta_rejects_a_rewritten_prefix() {
722        assert!(decoded_incremental_delta("stable", "changed").is_err());
723    }
724
725    #[tokio::test]
726    async fn incremental_decode_cache_hit_after_thinking_whitespace_does_not_deadlock() {
727        use tokenizers::models::bpe::{Vocab, BPE};
728        use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
729
730        let vocab: Vocab = [
731            ("</think>".to_string(), 0),
732            ("\n".to_string(), 1),
733            ("payload".to_string(), 2),
734            ("<unk>".to_string(), 3),
735        ]
736        .into_iter()
737        .collect();
738        let bpe = BPE::builder()
739            .vocab_and_merges(vocab, vec![])
740            .unk_token("<unk>".to_string())
741            .build()
742            .unwrap();
743        let mut hf_tokenizer = HfTokenizer::new(bpe);
744        hf_tokenizer.add_special_tokens(&[AddedToken::from("</think>", true)]);
745        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
746
747        let delimiter = tokenizer.token_id("</think>").unwrap();
748        let whitespace = tokenizer.token_id("\n").unwrap();
749        let payload = tokenizer.token_id("payload").unwrap();
750        let delimiter_prefix = vec![delimiter];
751        let whitespace_prefix = vec![delimiter, whitespace];
752
753        assert_eq!(
754            tokenizer
755                .decode_incremental(&delimiter_prefix, whitespace)
756                .unwrap(),
757            "\n"
758        );
759        assert!(tokenizer
760            .decode_cache
761            .read()
762            .get(&whitespace_prefix)
763            .is_some());
764        let payload_delta = tokenizer
765            .decode_incremental(&whitespace_prefix, payload)
766            .unwrap();
767        assert_eq!(payload_delta.trim_start(), "payload");
768    }
769
770    #[test]
771    fn test_incremental_state_default() {
772        let state = IncrementalState::default();
773        let debug_str = format!("{:?}", state);
774        assert!(debug_str.contains("IncrementalState"));
775    }
776
777    #[test]
778    fn test_incremental_state_clone() {
779        let state = IncrementalState::default();
780        let cloned = state.clone();
781
782        // 验证克隆成功
783        let state_str = format!("{:?}", state);
784        let cloned_str = format!("{:?}", cloned);
785        assert_eq!(state_str, cloned_str);
786    }
787
788    #[test]
789    fn test_huggingface_tokenizer_factory_creation() {
790        let factory = HuggingFaceTokenizerFactory::new();
791        let debug_str = format!("{:?}", factory);
792        assert!(debug_str.contains("HuggingFaceTokenizerFactory"));
793    }
794
795    #[test]
796    fn test_huggingface_tokenizer_factory_default() {
797        let factory = HuggingFaceTokenizerFactory;
798        let debug_str = format!("{:?}", factory);
799        assert!(debug_str.contains("HuggingFaceTokenizerFactory"));
800    }
801
802    #[test]
803    fn test_huggingface_tokenizer_factory_clone() {
804        let factory = HuggingFaceTokenizerFactory::new();
805        let cloned = factory.clone();
806
807        let factory_str = format!("{:?}", factory);
808        let cloned_str = format!("{:?}", cloned);
809        assert_eq!(factory_str, cloned_str);
810    }
811
812    #[test]
813    fn test_huggingface_tokenizer_factory_supported_types() {
814        let factory = HuggingFaceTokenizerFactory::new();
815        let types = factory.supported_types();
816
817        assert!(!types.is_empty());
818        assert!(types.contains(&TokenizerType::BPE));
819    }
820
821    #[test]
822    fn test_extract_special_tokens_with_mock_tokenizer() {
823        use tokenizers::models::bpe::{Vocab, BPE};
824        use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
825
826        // 创建一个简单的 mock tokenizer
827        let vocab: Vocab = [
828            ("hello".to_string(), 0),
829            ("<s>".to_string(), 1),
830            ("</s>".to_string(), 2),
831            ("<unk>".to_string(), 3),
832            ("<pad>".to_string(), 4),
833        ]
834        .into_iter()
835        .collect();
836
837        let merges = vec![];
838        let bpe = BPE::builder()
839            .vocab_and_merges(vocab, merges)
840            .unk_token("<unk>".to_string())
841            .build()
842            .unwrap();
843
844        let mut tokenizer = HfTokenizer::new(bpe);
845        tokenizer.add_special_tokens(&[
846            AddedToken::from("<s>", true),
847            AddedToken::from("</s>", true),
848            AddedToken::from("<unk>", true),
849            AddedToken::from("<pad>", true),
850        ]);
851
852        // 测试提取特殊 tokens
853        let result = extract_special_tokens(&tokenizer);
854        assert!(result.is_ok());
855
856        let special_tokens = result.unwrap();
857        assert!(special_tokens.bos_token.is_some());
858        assert!(special_tokens.eos_token.is_some());
859        assert!(special_tokens.unk_token.is_some());
860        assert!(special_tokens.pad_token.is_some());
861    }
862
863    #[tokio::test]
864    async fn test_huggingface_tokenizer_with_mock() {
865        use tokenizers::models::bpe::{Vocab, BPE};
866        use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
867
868        let vocab: Vocab = [
869            ("hello".to_string(), 0),
870            ("world".to_string(), 1),
871            ("<s>".to_string(), 2),
872            ("</s>".to_string(), 3),
873            ("<unk>".to_string(), 4),
874        ]
875        .into_iter()
876        .collect();
877
878        let merges = vec![];
879        let bpe = BPE::builder()
880            .vocab_and_merges(vocab, merges)
881            .unk_token("<unk>".to_string())
882            .build()
883            .unwrap();
884
885        let mut hf_tokenizer = HfTokenizer::new(bpe);
886        hf_tokenizer.add_special_tokens(&[
887            AddedToken::from("<s>", true),
888            AddedToken::from("</s>", true),
889            AddedToken::from("<unk>", true),
890        ]);
891
892        // 测试创建 HuggingFaceTokenizer
893        let result = HuggingFaceTokenizer::new(hf_tokenizer).await;
894        assert!(result.is_ok());
895
896        let tokenizer = result.unwrap();
897        assert_eq!(tokenizer.vocab_size(), 5);
898    }
899
900    #[tokio::test]
901    async fn test_tokenizer_encode_decode() {
902        use tokenizers::models::bpe::{Vocab, BPE};
903        use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
904
905        let vocab: Vocab = [
906            ("hello".to_string(), 0),
907            ("world".to_string(), 1),
908            ("<s>".to_string(), 2),
909            ("</s>".to_string(), 3),
910            ("<unk>".to_string(), 4),
911        ]
912        .into_iter()
913        .collect();
914
915        let merges = vec![];
916        let bpe = BPE::builder()
917            .vocab_and_merges(vocab, merges)
918            .unk_token("<unk>".to_string())
919            .build()
920            .unwrap();
921
922        let mut hf_tokenizer = HfTokenizer::new(bpe);
923        hf_tokenizer.add_special_tokens(&[
924            AddedToken::from("<s>", true),
925            AddedToken::from("</s>", true),
926            AddedToken::from("<unk>", true),
927        ]);
928
929        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
930
931        // 测试 encode - 即使无法编码,也会返回 UNK token
932        let result = tokenizer.encode("hello", false);
933        assert!(result.is_ok());
934
935        let _tokens = result.unwrap();
936        // Tokenizer 可能返回空数组或 UNK tokens
937        // 我们只验证结果是 Ok
938
939        // 测试 decode with empty tokens
940        let decoded = tokenizer.decode(&[], false);
941        assert!(decoded.is_ok());
942    }
943
944    #[tokio::test]
945    async fn test_tokenizer_special_tokens() {
946        use tokenizers::models::bpe::{Vocab, BPE};
947        use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
948
949        let vocab: Vocab = [
950            ("hello".to_string(), 0),
951            ("<s>".to_string(), 1),
952            ("</s>".to_string(), 2),
953        ]
954        .into_iter()
955        .collect();
956
957        let merges = vec![];
958        let bpe = BPE::builder()
959            .vocab_and_merges(vocab, merges)
960            .build()
961            .unwrap();
962
963        let mut hf_tokenizer = HfTokenizer::new(bpe);
964        hf_tokenizer.add_special_tokens(&[
965            AddedToken::from("<s>", true),
966            AddedToken::from("</s>", true),
967        ]);
968
969        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
970        let special_tokens = tokenizer.special_tokens();
971
972        // 应该能找到一些特殊 tokens
973        assert!(special_tokens.bos_token.is_some() || special_tokens.eos_token.is_some());
974    }
975
976    #[tokio::test]
977    async fn test_tokenizer_token_id_lookup() {
978        use tokenizers::models::bpe::{Vocab, BPE};
979        use tokenizers::Tokenizer as HfTokenizer;
980
981        let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
982            .into_iter()
983            .collect();
984
985        let merges = vec![];
986        let bpe = BPE::builder()
987            .vocab_and_merges(vocab, merges)
988            .build()
989            .unwrap();
990
991        let hf_tokenizer = HfTokenizer::new(bpe);
992        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
993
994        // 测试 token_id 查找
995        let token_id = tokenizer.token_id("hello");
996        assert!(token_id.is_some());
997        assert_eq!(token_id.unwrap().get(), 0);
998    }
999
1000    #[tokio::test]
1001    async fn test_tokenizer_token_text_reverse_lookup() {
1002        use tokenizers::models::bpe::{Vocab, BPE};
1003        use tokenizers::Tokenizer as HfTokenizer;
1004
1005        let vocab: Vocab = [
1006            ("hello".to_string(), 0),
1007            ("[PAD151935]".to_string(), 1),
1008            ("</think>".to_string(), 2),
1009        ]
1010        .into_iter()
1011        .collect();
1012
1013        let merges = vec![];
1014        let bpe = BPE::builder()
1015            .vocab_and_merges(vocab, merges)
1016            .build()
1017            .unwrap();
1018
1019        let hf_tokenizer = HfTokenizer::new(bpe);
1020        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1021
1022        assert_eq!(tokenizer.token_text(TokenId::new(1)), Some("[PAD151935]"));
1023        assert_eq!(tokenizer.token_text(TokenId::new(2)), Some("</think>"));
1024        assert_eq!(tokenizer.token_text(TokenId::new(99)), None);
1025    }
1026
1027    #[tokio::test]
1028    async fn byte_level_token_bytes_preserve_split_utf8_fragments() {
1029        use tokenizers::decoders::byte_level::ByteLevel;
1030        use tokenizers::models::bpe::{Vocab, BPE};
1031        use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
1032
1033        // GPT-2 byte-alphabet spellings for [f0, 9f] and [94, a5]. Each
1034        // fragment is invalid UTF-8 alone, but together they encode U+1F525.
1035        let vocab: Vocab = [
1036            ("\u{00f0}\u{0141}".to_string(), 0),
1037            ("\u{0136}\u{00a5}".to_string(), 1),
1038            ("<eos>".to_string(), 2),
1039        ]
1040        .into_iter()
1041        .collect();
1042        let bpe = BPE::builder()
1043            .vocab_and_merges(vocab, vec![])
1044            .build()
1045            .unwrap();
1046        let mut hf_tokenizer = HfTokenizer::new(bpe);
1047        hf_tokenizer.with_decoder(Some(ByteLevel::default()));
1048        hf_tokenizer.add_special_tokens(&[AddedToken::from("<eos>", true)]);
1049
1050        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1051
1052        assert!(tokenizer
1053            .decode(&[TokenId::new(0)], false)
1054            .unwrap()
1055            .contains('\u{fffd}'));
1056        assert_eq!(
1057            tokenizer
1058                .decode(&[TokenId::new(0), TokenId::new(1)], false)
1059                .unwrap(),
1060            "\u{1f525}"
1061        );
1062        assert_eq!(
1063            tokenizer.token_bytes(TokenId::new(0)),
1064            Some(vec![0xf0, 0x9f])
1065        );
1066        assert_eq!(
1067            tokenizer.token_bytes(TokenId::new(1)),
1068            Some(vec![0x94, 0xa5])
1069        );
1070        assert_eq!(tokenizer.token_bytes(TokenId::new(99)), None);
1071    }
1072
1073    #[tokio::test]
1074    async fn test_tokenizer_info() {
1075        use tokenizers::models::bpe::{Vocab, BPE};
1076        use tokenizers::Tokenizer as HfTokenizer;
1077
1078        let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
1079            .into_iter()
1080            .collect();
1081
1082        let merges = vec![];
1083        let bpe = BPE::builder()
1084            .vocab_and_merges(vocab, merges)
1085            .build()
1086            .unwrap();
1087
1088        let hf_tokenizer = HfTokenizer::new(bpe);
1089        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1090
1091        let info = tokenizer.info();
1092        assert_eq!(info.vocab_size, 2);
1093        assert!(info.supports_incremental);
1094        assert_eq!(info.tokenizer_type, TokenizerType::BPE);
1095    }
1096
1097    #[tokio::test]
1098    async fn test_incremental_tokenizer_interface() {
1099        use tokenizers::models::bpe::{Vocab, BPE};
1100        use tokenizers::Tokenizer as HfTokenizer;
1101
1102        let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
1103            .into_iter()
1104            .collect();
1105
1106        let merges = vec![];
1107        let bpe = BPE::builder()
1108            .vocab_and_merges(vocab, merges)
1109            .build()
1110            .unwrap();
1111
1112        let hf_tokenizer = HfTokenizer::new(bpe);
1113        let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1114
1115        // 测试增量解码接口
1116        let mut state = tokenizer.create_state();
1117
1118        // 添加一个 token
1119        let result = tokenizer.decode_incremental_with_state(&mut state, TokenId::new(0));
1120        assert!(result.is_ok());
1121
1122        // 重置状态
1123        tokenizer.reset_state(&mut state);
1124        let text = tokenizer.get_decoded_text(&state);
1125        assert!(text.is_empty());
1126    }
1127
1128    fn tiny_tokenizer_with_specials(specials: &[&str]) -> HfTokenizer {
1129        use tokenizers::models::bpe::{Vocab, BPE};
1130        use tokenizers::AddedToken;
1131
1132        let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
1133            .into_iter()
1134            .collect();
1135        let bpe = BPE::builder()
1136            .vocab_and_merges(vocab, vec![])
1137            .unk_token("hello".to_string())
1138            .build()
1139            .unwrap();
1140        let mut tokenizer = HfTokenizer::new(bpe);
1141        tokenizer.add_special_tokens(
1142            &specials
1143                .iter()
1144                .map(|s| AddedToken::from(*s, true))
1145                .collect::<Vec<_>>(),
1146        );
1147        tokenizer
1148    }
1149
1150    #[tokio::test]
1151    async fn eos_comes_from_generation_config_not_name_probing() {
1152        // DeepSeek-R1-distill style: special-token slots renamed, no
1153        // `</s>` / `<|endoftext|>`-style names anywhere in the vocab.
1154        let tokenizer =
1155            tiny_tokenizer_with_specials(&["<|end▁of▁sentence|>", "<|User|>", "<|Assistant|>"]);
1156        let eos_id = tokenizer.token_to_id("<|end▁of▁sentence|>").unwrap();
1157        let dir = tempfile::tempdir().unwrap();
1158        let path = dir.path().join("tokenizer.json");
1159        tokenizer.save(&path, false).unwrap();
1160        std::fs::write(
1161            dir.path().join("generation_config.json"),
1162            format!("{{\"bos_token_id\": null, \"eos_token_id\": {eos_id}}}"),
1163        )
1164        .unwrap();
1165
1166        let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1167            .await
1168            .unwrap();
1169        assert_eq!(
1170            loaded.special_tokens().eos_token.map(|t| t.get()),
1171            Some(eos_id)
1172        );
1173        assert!(loaded.special_tokens().extra_eos_tokens.is_empty());
1174    }
1175
1176    #[tokio::test]
1177    async fn immutable_source_bytes_preserve_generation_config_eos() {
1178        let tokenizer = tiny_tokenizer_with_specials(&["<|end_of_text|>", "<|end_of_turn|>"]);
1179        let primary = tokenizer.token_to_id("<|end_of_text|>").unwrap();
1180        let extra = tokenizer.token_to_id("<|end_of_turn|>").unwrap();
1181        let tokenizer_json = tokenizer.to_string(false).unwrap();
1182        let generation_config = format!(r#"{{"eos_token_id":[{primary},{extra}]}}"#);
1183
1184        let loaded = HuggingFaceTokenizer::from_source_bytes(
1185            tokenizer_json.as_bytes(),
1186            None,
1187            Some(generation_config.as_bytes()),
1188        )
1189        .await
1190        .unwrap();
1191
1192        assert_eq!(
1193            loaded.special_tokens().eos_token.map(|token| token.get()),
1194            Some(primary)
1195        );
1196        assert_eq!(
1197            loaded
1198                .special_tokens()
1199                .extra_eos_tokens
1200                .iter()
1201                .map(|token| token.get())
1202                .collect::<Vec<_>>(),
1203            vec![extra]
1204        );
1205    }
1206
1207    #[tokio::test]
1208    async fn multi_eos_ids_land_in_extra_eos_tokens() {
1209        let tokenizer = tiny_tokenizer_with_specials(&["<|eot_id|>", "<|end_of_text|>"]);
1210        let primary = tokenizer.token_to_id("<|end_of_text|>").unwrap();
1211        let extra = tokenizer.token_to_id("<|eot_id|>").unwrap();
1212        let dir = tempfile::tempdir().unwrap();
1213        let path = dir.path().join("tokenizer.json");
1214        tokenizer.save(&path, false).unwrap();
1215        std::fs::write(
1216            dir.path().join("generation_config.json"),
1217            format!("{{\"eos_token_id\": [{primary}, {extra}]}}"),
1218        )
1219        .unwrap();
1220
1221        let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1222            .await
1223            .unwrap();
1224        assert_eq!(
1225            loaded.special_tokens().eos_token.map(|t| t.get()),
1226            Some(primary)
1227        );
1228        assert_eq!(
1229            loaded
1230                .special_tokens()
1231                .extra_eos_tokens
1232                .iter()
1233                .map(|t| t.get())
1234                .collect::<Vec<_>>(),
1235            vec![extra]
1236        );
1237    }
1238
1239    #[tokio::test]
1240    async fn tokenizer_config_eos_string_is_fallback_without_generation_config() {
1241        let tokenizer = tiny_tokenizer_with_specials(&["<|end▁of▁sentence|>"]);
1242        let eos_id = tokenizer.token_to_id("<|end▁of▁sentence|>").unwrap();
1243        let dir = tempfile::tempdir().unwrap();
1244        let path = dir.path().join("tokenizer.json");
1245        tokenizer.save(&path, false).unwrap();
1246        std::fs::write(
1247            dir.path().join("tokenizer_config.json"),
1248            "{\"eos_token\": {\"content\": \"<|end▁of▁sentence|>\"}}",
1249        )
1250        .unwrap();
1251
1252        let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1253            .await
1254            .unwrap();
1255        assert_eq!(
1256            loaded.special_tokens().eos_token.map(|t| t.get()),
1257            Some(eos_id)
1258        );
1259    }
1260
1261    #[tokio::test]
1262    async fn bare_tokenizer_json_still_uses_name_probing() {
1263        let tokenizer = tiny_tokenizer_with_specials(&["<s>", "</s>"]);
1264        let eos_id = tokenizer.token_to_id("</s>").unwrap();
1265        let dir = tempfile::tempdir().unwrap();
1266        let path = dir.path().join("tokenizer.json");
1267        tokenizer.save(&path, false).unwrap();
1268
1269        let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1270            .await
1271            .unwrap();
1272        assert_eq!(
1273            loaded.special_tokens().eos_token.map(|t| t.get()),
1274            Some(eos_id)
1275        );
1276    }
1277
1278    #[tokio::test]
1279    async fn skip_special_decode_preserves_and_normalizes_think_markers() {
1280        // Magistral-style: [THINK]/[/THINK] are `special: true`, so a plain
1281        // skip-special decode would drop them and leak thinking into content.
1282        let tokenizer = tiny_tokenizer_with_specials(&["[THINK]", "[/THINK]", "<eos>"]);
1283        let think = tokenizer.token_to_id("[THINK]").unwrap();
1284        let end_think = tokenizer.token_to_id("[/THINK]").unwrap();
1285        let eos = tokenizer.token_to_id("<eos>").unwrap();
1286        let hello = tokenizer.token_to_id("hello").unwrap();
1287        let world = tokenizer.token_to_id("world").unwrap();
1288
1289        let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
1290        let tokens: Vec<TokenId> = [think, hello, end_think, world, eos]
1291            .into_iter()
1292            .map(TokenId::new)
1293            .collect();
1294        let text = loaded.decode(&tokens, true).unwrap();
1295
1296        assert_eq!(text, "<think>hello</think>world");
1297    }
1298
1299    #[tokio::test]
1300    async fn skip_special_decode_without_markers_is_unchanged() {
1301        let tokenizer = tiny_tokenizer_with_specials(&["<eos>"]);
1302        let eos = tokenizer.token_to_id("<eos>").unwrap();
1303        let hello = tokenizer.token_to_id("hello").unwrap();
1304
1305        let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
1306        let tokens: Vec<TokenId> = [hello, eos].into_iter().map(TokenId::new).collect();
1307
1308        assert_eq!(loaded.decode(&tokens, true).unwrap(), "hello");
1309    }
1310}