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