1use 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
13pub 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_markers: Vec<(u32, &'static str)>,
27 decode_cache: RwLock<DecodeCache>,
29}
30
31const 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 markers.extend(
47 ModelOutputProtocol::HarmonyGptOss
48 .preserved_special_token_texts()
49 .iter()
50 .filter_map(|text| tokenizer.token_to_id(text).map(|id| (id, *text))),
51 );
52 markers
53}
54
55#[derive(Debug, Clone, Default)]
57pub struct IncrementalState {
58 tokens: Vec<TokenId>,
60 text: String,
62}
63
64#[derive(Debug, Default)]
66struct DecodeCache {
67 cache: std::collections::HashMap<Vec<TokenId>, String>,
68 max_size: usize,
69}
70
71impl DecodeCache {
72 fn new(max_size: usize) -> Self {
73 Self {
74 cache: std::collections::HashMap::new(),
75 max_size,
76 }
77 }
78
79 fn get(&self, tokens: &[TokenId]) -> Option<&String> {
80 self.cache.get(tokens)
81 }
82
83 fn insert(&mut self, tokens: Vec<TokenId>, text: String) {
84 if self.cache.len() >= self.max_size {
85 let to_remove: Vec<_> = self
86 .cache
87 .keys()
88 .take(self.cache.len() / 2)
89 .cloned()
90 .collect();
91 for key in to_remove {
92 self.cache.remove(&key);
93 }
94 }
95 self.cache.insert(tokens, text);
96 }
97}
98
99fn decoded_incremental_delta(previous_text: &str, full_text: &str) -> Result<String> {
100 full_text
101 .strip_prefix(previous_text)
102 .map(ToOwned::to_owned)
103 .ok_or_else(|| {
104 ferrum_types::FerrumError::tokenizer(
105 "Incremental decode changed the previously emitted text prefix",
106 )
107 })
108}
109
110impl HuggingFaceTokenizer {
111 pub async fn new(tokenizer: HfTokenizer) -> Result<Self> {
113 let vocab_size = tokenizer.get_vocab_size(false);
114 let id_to_token = build_id_to_token(&tokenizer);
115
116 let special_tokens = extract_special_tokens(&tokenizer)?;
118
119 let info = TokenizerInfo {
120 tokenizer_type: TokenizerType::BPE, vocab_size,
122 special_tokens: special_tokens.clone(),
123 supports_incremental: true,
124 supports_chat_template: false, max_token_length: None, model_name: None, };
128
129 debug!(
130 "Created HuggingFace tokenizer with vocab size {}",
131 vocab_size
132 );
133
134 let semantic_markers = probe_semantic_markers(&tokenizer);
135 let byte_level_decoder = tokenizer.get_decoder().is_some_and(decoder_uses_byte_level);
136
137 Ok(Self {
138 tokenizer: Arc::new(tokenizer),
139 special_tokens,
140 info,
141 id_to_token,
142 byte_level_decoder,
143 semantic_markers,
144 decode_cache: RwLock::new(DecodeCache::new(1000)),
145 })
146 }
147
148 pub async fn from_file(path: &str) -> Result<Self> {
152 let tokenizer = HfTokenizer::from_file(path).map_err(|e| {
153 ferrum_types::FerrumError::tokenizer(format!("Failed to load tokenizer: {}", e))
154 })?;
155 let overrides =
156 special_token_overrides_from_configs(std::path::Path::new(path), &tokenizer);
157 let mut this = Self::new(tokenizer).await?;
158 this.apply_special_token_overrides(overrides);
159 Ok(this)
160 }
161
162 pub async fn from_source_bytes(
166 tokenizer_json: &[u8],
167 tokenizer_config_json: Option<&[u8]>,
168 generation_config_json: Option<&[u8]>,
169 ) -> Result<Self> {
170 let tokenizer = HfTokenizer::from_bytes(tokenizer_json).map_err(|error| {
171 ferrum_types::FerrumError::tokenizer(format!(
172 "Failed to load tokenizer from resolved source bytes: {error}"
173 ))
174 })?;
175 let tokenizer_config =
176 parse_optional_config_bytes(tokenizer_config_json, "tokenizer_config.json")?;
177 let generation_config =
178 parse_optional_config_bytes(generation_config_json, "generation_config.json")?;
179 let overrides = special_token_overrides_from_values(
180 generation_config.as_ref(),
181 tokenizer_config.as_ref(),
182 &tokenizer,
183 );
184 let mut this = Self::new(tokenizer).await?;
185 this.apply_special_token_overrides(overrides);
186 Ok(this)
187 }
188
189 fn apply_special_token_overrides(&mut self, overrides: SpecialTokenOverrides) {
190 if overrides.bos.is_some() {
191 self.special_tokens.bos_token = overrides.bos;
192 }
193 if overrides.eos.is_some() {
194 self.special_tokens.eos_token = overrides.eos;
195 }
196 if !overrides.extra_eos.is_empty() {
197 self.special_tokens.extra_eos_tokens = overrides.extra_eos;
198 }
199 self.info.special_tokens = self.special_tokens.clone();
200 }
201
202 pub async fn from_pretrained(repo_id: &str, _revision: Option<&str>) -> Result<Self> {
204 let api = hf_hub::api::tokio::Api::new().map_err(|e| {
205 ferrum_types::FerrumError::tokenizer(format!("Failed to create HF API: {}", e))
206 })?;
207
208 let repo = api.repo(hf_hub::Repo::model(repo_id.to_string()));
209
210 let tokenizer_file = repo.get("tokenizer.json").await.map_err(|e| {
213 ferrum_types::FerrumError::tokenizer(format!("Failed to download tokenizer: {}", e))
214 })?;
215
216 let tokenizer = HfTokenizer::from_file(&tokenizer_file).map_err(|e| {
217 ferrum_types::FerrumError::tokenizer(format!("Failed to load tokenizer: {}", e))
218 })?;
219
220 Self::new(tokenizer).await
221 }
222}
223
224impl Tokenizer for HuggingFaceTokenizer {
225 fn encode(&self, text: &str, add_special: bool) -> Result<Vec<TokenId>> {
226 let encoding = self
227 .tokenizer
228 .encode(text, add_special)
229 .map_err(|e| ferrum_types::FerrumError::tokenizer(format!("Encoding failed: {}", e)))?;
230
231 Ok(encoding
232 .get_ids()
233 .iter()
234 .map(|&id| TokenId::new(id))
235 .collect())
236 }
237
238 fn decode(&self, tokens: &[TokenId], skip_special: bool) -> Result<String> {
239 let token_ids: Vec<u32> = tokens.iter().map(|t| t.get()).collect();
240
241 if skip_special
245 && !self.semantic_markers.is_empty()
246 && token_ids
247 .iter()
248 .any(|id| self.semantic_markers.iter().any(|(mid, _)| mid == id))
249 {
250 let mut out = String::new();
251 let mut segment: Vec<u32> = Vec::with_capacity(token_ids.len());
252 for id in &token_ids {
253 if let Some((_, canonical)) =
254 self.semantic_markers.iter().find(|(mid, _)| mid == id)
255 {
256 if !segment.is_empty() {
257 out.push_str(&self.tokenizer.decode(&segment, true).map_err(|e| {
258 ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e))
259 })?);
260 segment.clear();
261 }
262 out.push_str(canonical);
263 } else {
264 segment.push(*id);
265 }
266 }
267 if !segment.is_empty() {
268 out.push_str(&self.tokenizer.decode(&segment, true).map_err(|e| {
269 ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e))
270 })?);
271 }
272 return Ok(out);
273 }
274
275 let text = self
276 .tokenizer
277 .decode(&token_ids, skip_special)
278 .map_err(|e| ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e)))?;
279
280 Ok(text)
281 }
282
283 fn decode_incremental(&self, prev: &[TokenId], next: TokenId) -> Result<String> {
284 let cached_prev = { self.decode_cache.read().get(prev).cloned() };
287 if let Some(cached_prev) = cached_prev {
288 let mut all_tokens = prev.to_vec();
289 all_tokens.push(next);
290 let full_text = self.decode(&all_tokens, true)?;
291
292 self.decode_cache
293 .write()
294 .insert(all_tokens, full_text.clone());
295
296 return decoded_incremental_delta(&cached_prev, &full_text);
297 }
298
299 let prev_text = if prev.is_empty() {
301 String::new()
302 } else {
303 self.decode(prev, true)?
304 };
305
306 let mut all_tokens = prev.to_vec();
307 all_tokens.push(next);
308 let full_text = self.decode(&all_tokens, true)?;
309
310 {
312 let mut cache = self.decode_cache.write();
313 if !prev.is_empty() {
314 cache.insert(prev.to_vec(), prev_text.clone());
315 }
316 cache.insert(all_tokens, full_text.clone());
317 }
318
319 decoded_incremental_delta(&prev_text, &full_text)
320 }
321
322 fn vocab_size(&self) -> usize {
323 self.info.vocab_size
324 }
325
326 fn special_tokens(&self) -> &SpecialTokens {
327 &self.special_tokens
328 }
329
330 fn token_id(&self, text: &str) -> Option<TokenId> {
331 self.tokenizer.token_to_id(text).map(TokenId::new)
332 }
333
334 fn token_text(&self, token_id: TokenId) -> Option<&str> {
335 self.id_to_token
336 .get(token_id.get() as usize)
337 .and_then(|value| value.as_deref())
338 }
339
340 fn token_bytes(&self, token_id: TokenId) -> Option<Vec<u8>> {
341 if self.byte_level_decoder {
342 return self.token_text(token_id).map(byte_level_token_bytes);
343 }
344 self.decode(&[token_id], false)
345 .ok()
346 .map(String::into_bytes)
347 .or_else(|| {
348 self.token_text(token_id)
349 .map(|text| text.as_bytes().to_vec())
350 })
351 }
352
353 fn apply_chat_template(
354 &self,
355 messages: &[ferrum_interfaces::tokenizer::ChatMessage],
356 ) -> Result<String> {
357 let mut result = String::new();
359 for msg in messages {
360 result.push_str(&format!("{}: {}\n", msg.role, msg.content));
361 }
362 Ok(result.trim_end().to_string())
363 }
364
365 fn info(&self) -> TokenizerInfo {
366 self.info.clone()
367 }
368}
369
370impl IncrementalTokenizer for HuggingFaceTokenizer {
371 type State = IncrementalState;
372
373 fn create_state(&self) -> Self::State {
374 IncrementalState::default()
375 }
376
377 fn decode_incremental_with_state(
378 &self,
379 state: &mut Self::State,
380 token: TokenId,
381 ) -> Result<String> {
382 state.tokens.push(token);
383
384 let full_text = self.decode(&state.tokens, true)?;
386
387 let delta = decoded_incremental_delta(&state.text, &full_text)?;
389
390 state.text = full_text;
392
393 Ok(delta)
394 }
395
396 fn reset_state(&self, state: &mut Self::State) {
397 state.tokens.clear();
398 state.text.clear();
399 }
400
401 fn get_decoded_text(&self, state: &Self::State) -> String {
402 state.text.clone()
403 }
404}
405
406#[derive(Debug, Clone, Default)]
408pub struct HuggingFaceTokenizerFactory;
409
410impl HuggingFaceTokenizerFactory {
411 pub fn new() -> Self {
412 Self
413 }
414}
415
416#[async_trait]
417impl TokenizerFactory for HuggingFaceTokenizerFactory {
418 async fn load_from_file(&self, path: &str) -> Result<Box<dyn Tokenizer>> {
419 let tokenizer = HuggingFaceTokenizer::from_file(path).await?;
420 Ok(Box::new(tokenizer))
421 }
422
423 async fn load_from_bytes(&self, data: &[u8]) -> Result<Box<dyn Tokenizer>> {
424 let tokenizer = HfTokenizer::from_bytes(data).map_err(|e| {
425 ferrum_types::FerrumError::tokenizer(format!(
426 "Failed to load tokenizer from bytes: {}",
427 e
428 ))
429 })?;
430 let tokenizer = HuggingFaceTokenizer::new(tokenizer).await?;
431 Ok(Box::new(tokenizer))
432 }
433
434 async fn load_from_hub(
435 &self,
436 repo_id: &str,
437 revision: Option<&str>,
438 ) -> Result<Box<dyn Tokenizer>> {
439 let tokenizer = HuggingFaceTokenizer::from_pretrained(repo_id, revision).await?;
440 Ok(Box::new(tokenizer))
441 }
442
443 async fn create_from_config(
444 &self,
445 config: &ferrum_interfaces::tokenizer::TokenizerConfig,
446 ) -> Result<Box<dyn Tokenizer>> {
447 self.load_from_file(&config.path).await
449 }
450
451 fn supported_types(&self) -> Vec<TokenizerType> {
452 vec![
453 TokenizerType::BPE,
454 TokenizerType::WordPiece,
455 TokenizerType::SentencePiece,
456 ]
457 }
458}
459
460fn build_id_to_token(tokenizer: &HfTokenizer) -> Vec<Option<String>> {
465 let vocab = tokenizer.get_vocab(true);
466 let Some(max_id) = vocab.values().copied().max() else {
467 return Vec::new();
468 };
469 let mut id_to_token = vec![None; max_id as usize + 1];
470 for (token, id) in vocab {
471 let slot = &mut id_to_token[id as usize];
472 if slot.is_none() {
473 *slot = Some(token);
474 }
475 }
476 id_to_token
477}
478
479fn decoder_uses_byte_level(decoder: &DecoderWrapper) -> bool {
480 match decoder {
481 DecoderWrapper::ByteLevel(_) => true,
482 DecoderWrapper::Sequence(sequence) => {
483 sequence.get_decoders().iter().any(decoder_uses_byte_level)
484 }
485 _ => false,
486 }
487}
488
489fn byte_level_char_bytes() -> &'static HashMap<char, u8> {
490 static CHAR_BYTES: OnceLock<HashMap<char, u8>> = OnceLock::new();
491 CHAR_BYTES.get_or_init(|| {
492 let mut direct = Vec::with_capacity(256);
493 direct.extend(b'!'..=b'~');
494 direct.extend(b'\xA1'..=b'\xAC');
495 direct.extend(b'\xAE'..=b'\xFF');
496
497 let mut next_codepoint = 256u32;
498 let mut mapping = HashMap::with_capacity(256);
499 for byte in 0..=u8::MAX {
500 let codepoint = if direct.contains(&byte) {
501 byte as u32
502 } else {
503 let codepoint = next_codepoint;
504 next_codepoint += 1;
505 codepoint
506 };
507 let character = char::from_u32(codepoint)
508 .expect("GPT-2 byte alphabet uses valid Unicode scalar values");
509 mapping.insert(character, byte);
510 }
511 mapping
512 })
513}
514
515fn byte_level_token_bytes(token: &str) -> Vec<u8> {
516 let mapping = byte_level_char_bytes();
517 token
518 .chars()
519 .map(|character| mapping.get(&character).copied())
520 .collect::<Option<Vec<_>>>()
521 .unwrap_or_else(|| token.as_bytes().to_vec())
522}
523
524fn extract_special_tokens(tokenizer: &HfTokenizer) -> Result<SpecialTokens> {
526 let _vocab = tokenizer.get_vocab(false);
527
528 let bos_token = tokenizer
529 .token_to_id("<s>")
530 .or_else(|| tokenizer.token_to_id("[BOS]"))
531 .or_else(|| tokenizer.token_to_id("<bos>"))
532 .map(TokenId::new);
533
534 let eos_token = tokenizer
535 .token_to_id("</s>")
536 .or_else(|| tokenizer.token_to_id("[EOS]"))
537 .or_else(|| tokenizer.token_to_id("<eos>"))
538 .map(TokenId::new);
539
540 let unk_token = tokenizer
541 .token_to_id("<unk>")
542 .or_else(|| tokenizer.token_to_id("[UNK]"))
543 .map(TokenId::new);
544
545 let pad_token = tokenizer
546 .token_to_id("<pad>")
547 .or_else(|| tokenizer.token_to_id("[PAD]"))
548 .map(TokenId::new);
549
550 let sep_token = tokenizer
551 .token_to_id("[SEP]")
552 .or_else(|| tokenizer.token_to_id("<sep>"))
553 .map(TokenId::new);
554
555 let cls_token = tokenizer
556 .token_to_id("[CLS]")
557 .or_else(|| tokenizer.token_to_id("<cls>"))
558 .map(TokenId::new);
559
560 let mask_token = tokenizer
561 .token_to_id("[MASK]")
562 .or_else(|| tokenizer.token_to_id("<mask>"))
563 .map(TokenId::new);
564
565 Ok(SpecialTokens {
566 bos_token,
567 eos_token,
568 unk_token,
569 pad_token,
570 sep_token,
571 cls_token,
572 mask_token,
573 extra_eos_tokens: Vec::new(),
574 })
575}
576
577#[derive(Debug, Default)]
586struct SpecialTokenOverrides {
587 bos: Option<TokenId>,
588 eos: Option<TokenId>,
589 extra_eos: Vec<TokenId>,
590}
591
592fn special_token_overrides_from_configs(
593 tokenizer_json: &std::path::Path,
594 tokenizer: &HfTokenizer,
595) -> SpecialTokenOverrides {
596 let Some(dir) = tokenizer_json.parent() else {
597 return SpecialTokenOverrides::default();
598 };
599 let generation_config = read_json(&dir.join("generation_config.json"));
600 let tokenizer_config = read_json(&dir.join("tokenizer_config.json"));
601 special_token_overrides_from_values(
602 generation_config.as_ref(),
603 tokenizer_config.as_ref(),
604 tokenizer,
605 )
606}
607
608fn special_token_overrides_from_values(
609 generation_config: Option<&serde_json::Value>,
610 tokenizer_config: Option<&serde_json::Value>,
611 tokenizer: &HfTokenizer,
612) -> SpecialTokenOverrides {
613 let mut overrides = SpecialTokenOverrides::default();
614
615 if let Some(gen) = generation_config {
616 let mut eos_ids = token_id_list(gen.get("eos_token_id"));
617 if !eos_ids.is_empty() {
618 overrides.eos = Some(eos_ids.remove(0));
619 overrides.extra_eos = eos_ids;
620 }
621 if let Some(bos) = token_id_list(gen.get("bos_token_id")).into_iter().next() {
622 overrides.bos = Some(bos);
623 }
624 }
625
626 if let Some(tok_cfg) = tokenizer_config {
627 if overrides.eos.is_none() {
628 overrides.eos = token_from_config_value(tok_cfg.get("eos_token"), tokenizer);
629 }
630 if overrides.bos.is_none() {
631 overrides.bos = token_from_config_value(tok_cfg.get("bos_token"), tokenizer);
632 }
633 }
634
635 overrides
636}
637
638fn parse_optional_config_bytes(
639 bytes: Option<&[u8]>,
640 source_file: &str,
641) -> Result<Option<serde_json::Value>> {
642 bytes
643 .map(|bytes| {
644 serde_json::from_slice(bytes).map_err(|error| {
645 ferrum_types::FerrumError::tokenizer(format!(
646 "Failed to parse resolved {source_file}: {error}"
647 ))
648 })
649 })
650 .transpose()
651}
652
653fn read_json(path: &std::path::Path) -> Option<serde_json::Value> {
654 let text = std::fs::read_to_string(path).ok()?;
655 serde_json::from_str(&text).ok()
656}
657
658fn token_id_list(value: Option<&serde_json::Value>) -> Vec<TokenId> {
661 match value {
662 Some(serde_json::Value::Number(n)) => n
663 .as_u64()
664 .map(|v| vec![TokenId::new(v as u32)])
665 .unwrap_or_default(),
666 Some(serde_json::Value::Array(items)) => items
667 .iter()
668 .filter_map(|v| v.as_u64())
669 .map(|v| TokenId::new(v as u32))
670 .collect(),
671 _ => Vec::new(),
672 }
673}
674
675fn token_from_config_value(
678 value: Option<&serde_json::Value>,
679 tokenizer: &HfTokenizer,
680) -> Option<TokenId> {
681 let text = match value? {
682 serde_json::Value::String(s) => s.as_str(),
683 serde_json::Value::Object(obj) => obj.get("content")?.as_str()?,
684 _ => return None,
685 };
686 tokenizer.token_to_id(text).map(TokenId::new)
687}
688
689#[cfg(test)]
690mod tests {
691 use super::*;
692
693 #[test]
694 fn test_decode_cache_creation() {
695 let cache = DecodeCache::new(100);
696 assert_eq!(cache.max_size, 100);
697 assert_eq!(cache.cache.len(), 0);
698 }
699
700 #[test]
701 fn test_decode_cache_insert_and_get() {
702 let mut cache = DecodeCache::new(10);
703 let tokens = vec![TokenId::new(1), TokenId::new(2)];
704 let text = "hello".to_string();
705
706 cache.insert(tokens.clone(), text.clone());
707
708 let result = cache.get(&tokens);
709 assert!(result.is_some());
710 assert_eq!(result.unwrap(), &text);
711 }
712
713 #[test]
714 fn test_decode_cache_eviction() {
715 let mut cache = DecodeCache::new(2);
716
717 cache.insert(vec![TokenId::new(1)], "a".to_string());
719 cache.insert(vec![TokenId::new(2)], "b".to_string());
720
721 assert_eq!(cache.cache.len(), 2);
722
723 cache.insert(vec![TokenId::new(3)], "c".to_string());
725
726 assert!(cache.cache.len() <= 2);
728 }
729
730 #[test]
731 fn incremental_delta_rejects_a_rewritten_prefix() {
732 assert!(decoded_incremental_delta("stable", "changed").is_err());
733 }
734
735 #[tokio::test]
736 async fn incremental_decode_cache_hit_after_thinking_whitespace_does_not_deadlock() {
737 use tokenizers::models::bpe::{Vocab, BPE};
738 use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
739
740 let vocab: Vocab = [
741 ("</think>".to_string(), 0),
742 ("\n".to_string(), 1),
743 ("payload".to_string(), 2),
744 ("<unk>".to_string(), 3),
745 ]
746 .into_iter()
747 .collect();
748 let bpe = BPE::builder()
749 .vocab_and_merges(vocab, vec![])
750 .unk_token("<unk>".to_string())
751 .build()
752 .unwrap();
753 let mut hf_tokenizer = HfTokenizer::new(bpe);
754 hf_tokenizer.add_special_tokens(&[AddedToken::from("</think>", true)]);
755 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
756
757 let delimiter = tokenizer.token_id("</think>").unwrap();
758 let whitespace = tokenizer.token_id("\n").unwrap();
759 let payload = tokenizer.token_id("payload").unwrap();
760 let delimiter_prefix = vec![delimiter];
761 let whitespace_prefix = vec![delimiter, whitespace];
762
763 assert_eq!(
764 tokenizer
765 .decode_incremental(&delimiter_prefix, whitespace)
766 .unwrap(),
767 "\n"
768 );
769 assert!(tokenizer
770 .decode_cache
771 .read()
772 .get(&whitespace_prefix)
773 .is_some());
774 let payload_delta = tokenizer
775 .decode_incremental(&whitespace_prefix, payload)
776 .unwrap();
777 assert_eq!(payload_delta.trim_start(), "payload");
778 }
779
780 #[test]
781 fn test_incremental_state_default() {
782 let state = IncrementalState::default();
783 let debug_str = format!("{:?}", state);
784 assert!(debug_str.contains("IncrementalState"));
785 }
786
787 #[test]
788 fn test_incremental_state_clone() {
789 let state = IncrementalState::default();
790 let cloned = state.clone();
791
792 let state_str = format!("{:?}", state);
794 let cloned_str = format!("{:?}", cloned);
795 assert_eq!(state_str, cloned_str);
796 }
797
798 #[test]
799 fn test_huggingface_tokenizer_factory_creation() {
800 let factory = HuggingFaceTokenizerFactory::new();
801 let debug_str = format!("{:?}", factory);
802 assert!(debug_str.contains("HuggingFaceTokenizerFactory"));
803 }
804
805 #[test]
806 fn test_huggingface_tokenizer_factory_default() {
807 let factory = HuggingFaceTokenizerFactory;
808 let debug_str = format!("{:?}", factory);
809 assert!(debug_str.contains("HuggingFaceTokenizerFactory"));
810 }
811
812 #[test]
813 fn test_huggingface_tokenizer_factory_clone() {
814 let factory = HuggingFaceTokenizerFactory::new();
815 let cloned = factory.clone();
816
817 let factory_str = format!("{:?}", factory);
818 let cloned_str = format!("{:?}", cloned);
819 assert_eq!(factory_str, cloned_str);
820 }
821
822 #[test]
823 fn test_huggingface_tokenizer_factory_supported_types() {
824 let factory = HuggingFaceTokenizerFactory::new();
825 let types = factory.supported_types();
826
827 assert!(!types.is_empty());
828 assert!(types.contains(&TokenizerType::BPE));
829 }
830
831 #[test]
832 fn test_extract_special_tokens_with_mock_tokenizer() {
833 use tokenizers::models::bpe::{Vocab, BPE};
834 use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
835
836 let vocab: Vocab = [
838 ("hello".to_string(), 0),
839 ("<s>".to_string(), 1),
840 ("</s>".to_string(), 2),
841 ("<unk>".to_string(), 3),
842 ("<pad>".to_string(), 4),
843 ]
844 .into_iter()
845 .collect();
846
847 let merges = vec![];
848 let bpe = BPE::builder()
849 .vocab_and_merges(vocab, merges)
850 .unk_token("<unk>".to_string())
851 .build()
852 .unwrap();
853
854 let mut tokenizer = HfTokenizer::new(bpe);
855 tokenizer.add_special_tokens(&[
856 AddedToken::from("<s>", true),
857 AddedToken::from("</s>", true),
858 AddedToken::from("<unk>", true),
859 AddedToken::from("<pad>", true),
860 ]);
861
862 let result = extract_special_tokens(&tokenizer);
864 assert!(result.is_ok());
865
866 let special_tokens = result.unwrap();
867 assert!(special_tokens.bos_token.is_some());
868 assert!(special_tokens.eos_token.is_some());
869 assert!(special_tokens.unk_token.is_some());
870 assert!(special_tokens.pad_token.is_some());
871 }
872
873 #[tokio::test]
874 async fn test_huggingface_tokenizer_with_mock() {
875 use tokenizers::models::bpe::{Vocab, BPE};
876 use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
877
878 let vocab: Vocab = [
879 ("hello".to_string(), 0),
880 ("world".to_string(), 1),
881 ("<s>".to_string(), 2),
882 ("</s>".to_string(), 3),
883 ("<unk>".to_string(), 4),
884 ]
885 .into_iter()
886 .collect();
887
888 let merges = vec![];
889 let bpe = BPE::builder()
890 .vocab_and_merges(vocab, merges)
891 .unk_token("<unk>".to_string())
892 .build()
893 .unwrap();
894
895 let mut hf_tokenizer = HfTokenizer::new(bpe);
896 hf_tokenizer.add_special_tokens(&[
897 AddedToken::from("<s>", true),
898 AddedToken::from("</s>", true),
899 AddedToken::from("<unk>", true),
900 ]);
901
902 let result = HuggingFaceTokenizer::new(hf_tokenizer).await;
904 assert!(result.is_ok());
905
906 let tokenizer = result.unwrap();
907 assert_eq!(tokenizer.vocab_size(), 5);
908 }
909
910 #[tokio::test]
911 async fn test_tokenizer_encode_decode() {
912 use tokenizers::models::bpe::{Vocab, BPE};
913 use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
914
915 let vocab: Vocab = [
916 ("hello".to_string(), 0),
917 ("world".to_string(), 1),
918 ("<s>".to_string(), 2),
919 ("</s>".to_string(), 3),
920 ("<unk>".to_string(), 4),
921 ]
922 .into_iter()
923 .collect();
924
925 let merges = vec![];
926 let bpe = BPE::builder()
927 .vocab_and_merges(vocab, merges)
928 .unk_token("<unk>".to_string())
929 .build()
930 .unwrap();
931
932 let mut hf_tokenizer = HfTokenizer::new(bpe);
933 hf_tokenizer.add_special_tokens(&[
934 AddedToken::from("<s>", true),
935 AddedToken::from("</s>", true),
936 AddedToken::from("<unk>", true),
937 ]);
938
939 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
940
941 let result = tokenizer.encode("hello", false);
943 assert!(result.is_ok());
944
945 let _tokens = result.unwrap();
946 let decoded = tokenizer.decode(&[], false);
951 assert!(decoded.is_ok());
952 }
953
954 #[tokio::test]
955 async fn test_tokenizer_special_tokens() {
956 use tokenizers::models::bpe::{Vocab, BPE};
957 use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
958
959 let vocab: Vocab = [
960 ("hello".to_string(), 0),
961 ("<s>".to_string(), 1),
962 ("</s>".to_string(), 2),
963 ]
964 .into_iter()
965 .collect();
966
967 let merges = vec![];
968 let bpe = BPE::builder()
969 .vocab_and_merges(vocab, merges)
970 .build()
971 .unwrap();
972
973 let mut hf_tokenizer = HfTokenizer::new(bpe);
974 hf_tokenizer.add_special_tokens(&[
975 AddedToken::from("<s>", true),
976 AddedToken::from("</s>", true),
977 ]);
978
979 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
980 let special_tokens = tokenizer.special_tokens();
981
982 assert!(special_tokens.bos_token.is_some() || special_tokens.eos_token.is_some());
984 }
985
986 #[tokio::test]
987 async fn test_tokenizer_token_id_lookup() {
988 use tokenizers::models::bpe::{Vocab, BPE};
989 use tokenizers::Tokenizer as HfTokenizer;
990
991 let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
992 .into_iter()
993 .collect();
994
995 let merges = vec![];
996 let bpe = BPE::builder()
997 .vocab_and_merges(vocab, merges)
998 .build()
999 .unwrap();
1000
1001 let hf_tokenizer = HfTokenizer::new(bpe);
1002 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1003
1004 let token_id = tokenizer.token_id("hello");
1006 assert!(token_id.is_some());
1007 assert_eq!(token_id.unwrap().get(), 0);
1008 }
1009
1010 #[tokio::test]
1011 async fn test_tokenizer_token_text_reverse_lookup() {
1012 use tokenizers::models::bpe::{Vocab, BPE};
1013 use tokenizers::Tokenizer as HfTokenizer;
1014
1015 let vocab: Vocab = [
1016 ("hello".to_string(), 0),
1017 ("[PAD151935]".to_string(), 1),
1018 ("</think>".to_string(), 2),
1019 ]
1020 .into_iter()
1021 .collect();
1022
1023 let merges = vec![];
1024 let bpe = BPE::builder()
1025 .vocab_and_merges(vocab, merges)
1026 .build()
1027 .unwrap();
1028
1029 let hf_tokenizer = HfTokenizer::new(bpe);
1030 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1031
1032 assert_eq!(tokenizer.token_text(TokenId::new(1)), Some("[PAD151935]"));
1033 assert_eq!(tokenizer.token_text(TokenId::new(2)), Some("</think>"));
1034 assert_eq!(tokenizer.token_text(TokenId::new(99)), None);
1035 }
1036
1037 #[tokio::test]
1038 async fn byte_level_token_bytes_preserve_split_utf8_fragments() {
1039 use tokenizers::decoders::byte_level::ByteLevel;
1040 use tokenizers::models::bpe::{Vocab, BPE};
1041 use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
1042
1043 let vocab: Vocab = [
1046 ("\u{00f0}\u{0141}".to_string(), 0),
1047 ("\u{0136}\u{00a5}".to_string(), 1),
1048 ("<eos>".to_string(), 2),
1049 ]
1050 .into_iter()
1051 .collect();
1052 let bpe = BPE::builder()
1053 .vocab_and_merges(vocab, vec![])
1054 .build()
1055 .unwrap();
1056 let mut hf_tokenizer = HfTokenizer::new(bpe);
1057 hf_tokenizer.with_decoder(Some(ByteLevel::default()));
1058 hf_tokenizer.add_special_tokens(&[AddedToken::from("<eos>", true)]);
1059
1060 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1061
1062 assert!(tokenizer
1063 .decode(&[TokenId::new(0)], false)
1064 .unwrap()
1065 .contains('\u{fffd}'));
1066 assert_eq!(
1067 tokenizer
1068 .decode(&[TokenId::new(0), TokenId::new(1)], false)
1069 .unwrap(),
1070 "\u{1f525}"
1071 );
1072 assert_eq!(
1073 tokenizer.token_bytes(TokenId::new(0)),
1074 Some(vec![0xf0, 0x9f])
1075 );
1076 assert_eq!(
1077 tokenizer.token_bytes(TokenId::new(1)),
1078 Some(vec![0x94, 0xa5])
1079 );
1080 assert_eq!(tokenizer.token_bytes(TokenId::new(99)), None);
1081 }
1082
1083 #[tokio::test]
1084 async fn test_tokenizer_info() {
1085 use tokenizers::models::bpe::{Vocab, BPE};
1086 use tokenizers::Tokenizer as HfTokenizer;
1087
1088 let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
1089 .into_iter()
1090 .collect();
1091
1092 let merges = vec![];
1093 let bpe = BPE::builder()
1094 .vocab_and_merges(vocab, merges)
1095 .build()
1096 .unwrap();
1097
1098 let hf_tokenizer = HfTokenizer::new(bpe);
1099 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1100
1101 let info = tokenizer.info();
1102 assert_eq!(info.vocab_size, 2);
1103 assert!(info.supports_incremental);
1104 assert_eq!(info.tokenizer_type, TokenizerType::BPE);
1105 }
1106
1107 #[tokio::test]
1108 async fn test_incremental_tokenizer_interface() {
1109 use tokenizers::models::bpe::{Vocab, BPE};
1110 use tokenizers::Tokenizer as HfTokenizer;
1111
1112 let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
1113 .into_iter()
1114 .collect();
1115
1116 let merges = vec![];
1117 let bpe = BPE::builder()
1118 .vocab_and_merges(vocab, merges)
1119 .build()
1120 .unwrap();
1121
1122 let hf_tokenizer = HfTokenizer::new(bpe);
1123 let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
1124
1125 let mut state = tokenizer.create_state();
1127
1128 let result = tokenizer.decode_incremental_with_state(&mut state, TokenId::new(0));
1130 assert!(result.is_ok());
1131
1132 tokenizer.reset_state(&mut state);
1134 let text = tokenizer.get_decoded_text(&state);
1135 assert!(text.is_empty());
1136 }
1137
1138 fn tiny_tokenizer_with_specials(specials: &[&str]) -> HfTokenizer {
1139 use tokenizers::models::bpe::{Vocab, BPE};
1140 use tokenizers::AddedToken;
1141
1142 let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
1143 .into_iter()
1144 .collect();
1145 let bpe = BPE::builder()
1146 .vocab_and_merges(vocab, vec![])
1147 .unk_token("hello".to_string())
1148 .build()
1149 .unwrap();
1150 let mut tokenizer = HfTokenizer::new(bpe);
1151 tokenizer.add_special_tokens(
1152 &specials
1153 .iter()
1154 .map(|s| AddedToken::from(*s, true))
1155 .collect::<Vec<_>>(),
1156 );
1157 tokenizer
1158 }
1159
1160 #[tokio::test]
1161 async fn eos_comes_from_generation_config_not_name_probing() {
1162 let tokenizer =
1165 tiny_tokenizer_with_specials(&["<|end▁of▁sentence|>", "<|User|>", "<|Assistant|>"]);
1166 let eos_id = tokenizer.token_to_id("<|end▁of▁sentence|>").unwrap();
1167 let dir = tempfile::tempdir().unwrap();
1168 let path = dir.path().join("tokenizer.json");
1169 tokenizer.save(&path, false).unwrap();
1170 std::fs::write(
1171 dir.path().join("generation_config.json"),
1172 format!("{{\"bos_token_id\": null, \"eos_token_id\": {eos_id}}}"),
1173 )
1174 .unwrap();
1175
1176 let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1177 .await
1178 .unwrap();
1179 assert_eq!(
1180 loaded.special_tokens().eos_token.map(|t| t.get()),
1181 Some(eos_id)
1182 );
1183 assert!(loaded.special_tokens().extra_eos_tokens.is_empty());
1184 }
1185
1186 #[tokio::test]
1187 async fn immutable_source_bytes_preserve_generation_config_eos() {
1188 let tokenizer = tiny_tokenizer_with_specials(&["<|end_of_text|>", "<|end_of_turn|>"]);
1189 let primary = tokenizer.token_to_id("<|end_of_text|>").unwrap();
1190 let extra = tokenizer.token_to_id("<|end_of_turn|>").unwrap();
1191 let tokenizer_json = tokenizer.to_string(false).unwrap();
1192 let generation_config = format!(r#"{{"eos_token_id":[{primary},{extra}]}}"#);
1193
1194 let loaded = HuggingFaceTokenizer::from_source_bytes(
1195 tokenizer_json.as_bytes(),
1196 None,
1197 Some(generation_config.as_bytes()),
1198 )
1199 .await
1200 .unwrap();
1201
1202 assert_eq!(
1203 loaded.special_tokens().eos_token.map(|token| token.get()),
1204 Some(primary)
1205 );
1206 assert_eq!(
1207 loaded
1208 .special_tokens()
1209 .extra_eos_tokens
1210 .iter()
1211 .map(|token| token.get())
1212 .collect::<Vec<_>>(),
1213 vec![extra]
1214 );
1215 }
1216
1217 #[tokio::test]
1218 async fn multi_eos_ids_land_in_extra_eos_tokens() {
1219 let tokenizer = tiny_tokenizer_with_specials(&["<|eot_id|>", "<|end_of_text|>"]);
1220 let primary = tokenizer.token_to_id("<|end_of_text|>").unwrap();
1221 let extra = tokenizer.token_to_id("<|eot_id|>").unwrap();
1222 let dir = tempfile::tempdir().unwrap();
1223 let path = dir.path().join("tokenizer.json");
1224 tokenizer.save(&path, false).unwrap();
1225 std::fs::write(
1226 dir.path().join("generation_config.json"),
1227 format!("{{\"eos_token_id\": [{primary}, {extra}]}}"),
1228 )
1229 .unwrap();
1230
1231 let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1232 .await
1233 .unwrap();
1234 assert_eq!(
1235 loaded.special_tokens().eos_token.map(|t| t.get()),
1236 Some(primary)
1237 );
1238 assert_eq!(
1239 loaded
1240 .special_tokens()
1241 .extra_eos_tokens
1242 .iter()
1243 .map(|t| t.get())
1244 .collect::<Vec<_>>(),
1245 vec![extra]
1246 );
1247 }
1248
1249 #[tokio::test]
1250 async fn tokenizer_config_eos_string_is_fallback_without_generation_config() {
1251 let tokenizer = tiny_tokenizer_with_specials(&["<|end▁of▁sentence|>"]);
1252 let eos_id = tokenizer.token_to_id("<|end▁of▁sentence|>").unwrap();
1253 let dir = tempfile::tempdir().unwrap();
1254 let path = dir.path().join("tokenizer.json");
1255 tokenizer.save(&path, false).unwrap();
1256 std::fs::write(
1257 dir.path().join("tokenizer_config.json"),
1258 "{\"eos_token\": {\"content\": \"<|end▁of▁sentence|>\"}}",
1259 )
1260 .unwrap();
1261
1262 let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1263 .await
1264 .unwrap();
1265 assert_eq!(
1266 loaded.special_tokens().eos_token.map(|t| t.get()),
1267 Some(eos_id)
1268 );
1269 }
1270
1271 #[tokio::test]
1272 async fn bare_tokenizer_json_still_uses_name_probing() {
1273 let tokenizer = tiny_tokenizer_with_specials(&["<s>", "</s>"]);
1274 let eos_id = tokenizer.token_to_id("</s>").unwrap();
1275 let dir = tempfile::tempdir().unwrap();
1276 let path = dir.path().join("tokenizer.json");
1277 tokenizer.save(&path, false).unwrap();
1278
1279 let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
1280 .await
1281 .unwrap();
1282 assert_eq!(
1283 loaded.special_tokens().eos_token.map(|t| t.get()),
1284 Some(eos_id)
1285 );
1286 }
1287
1288 #[tokio::test]
1289 async fn skip_special_decode_preserves_and_normalizes_think_markers() {
1290 let tokenizer = tiny_tokenizer_with_specials(&["[THINK]", "[/THINK]", "<eos>"]);
1293 let think = tokenizer.token_to_id("[THINK]").unwrap();
1294 let end_think = tokenizer.token_to_id("[/THINK]").unwrap();
1295 let eos = tokenizer.token_to_id("<eos>").unwrap();
1296 let hello = tokenizer.token_to_id("hello").unwrap();
1297 let world = tokenizer.token_to_id("world").unwrap();
1298
1299 let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
1300 let tokens: Vec<TokenId> = [think, hello, end_think, world, eos]
1301 .into_iter()
1302 .map(TokenId::new)
1303 .collect();
1304 let text = loaded.decode(&tokens, true).unwrap();
1305
1306 assert_eq!(text, "<think>hello</think>world");
1307 }
1308
1309 #[tokio::test]
1310 async fn skip_special_decode_preserves_typed_harmony_markers_only() {
1311 let harmony_markers = ModelOutputProtocol::HarmonyGptOss.preserved_special_token_texts();
1312 let mut specials = harmony_markers.to_vec();
1313 specials.push("<|endoftext|>");
1314 let tokenizer = tiny_tokenizer_with_specials(&specials);
1315
1316 let mut token_ids = vec![tokenizer.token_to_id("hello").unwrap()];
1317 token_ids.extend(
1318 harmony_markers
1319 .iter()
1320 .map(|marker| tokenizer.token_to_id(marker).unwrap()),
1321 );
1322 token_ids.push(tokenizer.token_to_id("world").unwrap());
1323 token_ids.push(tokenizer.token_to_id("<|endoftext|>").unwrap());
1324
1325 let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
1326 let tokens: Vec<TokenId> = token_ids.into_iter().map(TokenId::new).collect();
1327 let text = loaded.decode(&tokens, true).unwrap();
1328
1329 assert_eq!(
1330 text,
1331 format!("hello{}world", harmony_markers.concat()),
1332 "the typed Harmony markers must survive while unrelated special tokens stay skipped"
1333 );
1334 }
1335
1336 #[tokio::test]
1337 async fn skip_special_decode_without_markers_is_unchanged() {
1338 let tokenizer = tiny_tokenizer_with_specials(&["<eos>"]);
1339 let eos = tokenizer.token_to_id("<eos>").unwrap();
1340 let hello = tokenizer.token_to_id("hello").unwrap();
1341
1342 let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
1343 let tokens: Vec<TokenId> = [hello, eos].into_iter().map(TokenId::new).collect();
1344
1345 assert_eq!(loaded.decode(&tokens, true).unwrap(), "hello");
1346 }
1347}