1use 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
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 think_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_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#[derive(Debug, Clone, Default)]
49pub struct IncrementalState {
50 tokens: Vec<TokenId>,
52 text: String,
54}
55
56#[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 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 let special_tokens = extract_special_tokens(&tokenizer)?;
110
111 let info = TokenizerInfo {
112 tokenizer_type: TokenizerType::BPE, vocab_size,
114 special_tokens: special_tokens.clone(),
115 supports_incremental: true,
116 supports_chat_template: false, max_token_length: None, model_name: None, };
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 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 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 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 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 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 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 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 {
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 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 let full_text = self.decode(&state.tokens, true)?;
376
377 let delta = decoded_incremental_delta(&state.text, &full_text)?;
379
380 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#[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 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
450fn 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
514fn 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#[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
648fn 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
665fn 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 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 cache.insert(vec![TokenId::new(3)], "c".to_string());
715
716 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 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 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 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 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 let result = tokenizer.encode("hello", false);
933 assert!(result.is_ok());
934
935 let _tokens = result.unwrap();
936 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 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 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 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 let mut state = tokenizer.create_state();
1117
1118 let result = tokenizer.decode_incremental_with_state(&mut state, TokenId::new(0));
1120 assert!(result.is_ok());
1121
1122 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 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 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}