use crate::{IncrementalTokenizer, Tokenizer, TokenizerFactory, TokenizerInfo, TokenizerType};
use async_trait::async_trait;
use ferrum_types::{ModelOutputProtocol, Result, SpecialTokens, TokenId};
use parking_lot::RwLock;
use std::collections::HashMap;
use std::sync::{Arc, OnceLock};
use tokenizers::decoders::DecoderWrapper;
use tokenizers::Tokenizer as HfTokenizer;
use tracing::debug;
pub struct HuggingFaceTokenizer {
tokenizer: Arc<HfTokenizer>,
special_tokens: SpecialTokens,
info: TokenizerInfo,
id_to_token: Vec<Option<String>>,
byte_level_decoder: bool,
semantic_markers: Vec<(u32, &'static str)>,
decode_cache: RwLock<DecodeCache>,
}
const THINK_MARKER_DIALECTS: [(&str, &'static str); 4] = [
("<think>", "<think>"),
("</think>", "</think>"),
("[THINK]", "<think>"),
("[/THINK]", "</think>"),
];
fn probe_semantic_markers(tokenizer: &HfTokenizer) -> Vec<(u32, &'static str)> {
let mut markers: Vec<_> = THINK_MARKER_DIALECTS
.iter()
.filter_map(|(text, canonical)| tokenizer.token_to_id(text).map(|id| (id, *canonical)))
.collect();
markers.extend(
ModelOutputProtocol::HarmonyGptOss
.preserved_special_token_texts()
.iter()
.filter_map(|text| tokenizer.token_to_id(text).map(|id| (id, *text))),
);
markers
}
#[derive(Debug, Clone, Default)]
pub struct IncrementalState {
tokens: Vec<TokenId>,
text: String,
}
#[derive(Debug, Default)]
struct DecodeCache {
cache: std::collections::HashMap<Vec<TokenId>, String>,
max_size: usize,
}
impl DecodeCache {
fn new(max_size: usize) -> Self {
Self {
cache: std::collections::HashMap::new(),
max_size,
}
}
fn get(&self, tokens: &[TokenId]) -> Option<&String> {
self.cache.get(tokens)
}
fn insert(&mut self, tokens: Vec<TokenId>, text: String) {
if self.cache.len() >= self.max_size {
let to_remove: Vec<_> = self
.cache
.keys()
.take(self.cache.len() / 2)
.cloned()
.collect();
for key in to_remove {
self.cache.remove(&key);
}
}
self.cache.insert(tokens, text);
}
}
fn decoded_incremental_delta(previous_text: &str, full_text: &str) -> Result<String> {
full_text
.strip_prefix(previous_text)
.map(ToOwned::to_owned)
.ok_or_else(|| {
ferrum_types::FerrumError::tokenizer(
"Incremental decode changed the previously emitted text prefix",
)
})
}
impl HuggingFaceTokenizer {
pub async fn new(tokenizer: HfTokenizer) -> Result<Self> {
let vocab_size = tokenizer.get_vocab_size(false);
let id_to_token = build_id_to_token(&tokenizer);
let special_tokens = extract_special_tokens(&tokenizer)?;
let info = TokenizerInfo {
tokenizer_type: TokenizerType::BPE, vocab_size,
special_tokens: special_tokens.clone(),
supports_incremental: true,
supports_chat_template: false, max_token_length: None, model_name: None, };
debug!(
"Created HuggingFace tokenizer with vocab size {}",
vocab_size
);
let semantic_markers = probe_semantic_markers(&tokenizer);
let byte_level_decoder = tokenizer.get_decoder().is_some_and(decoder_uses_byte_level);
Ok(Self {
tokenizer: Arc::new(tokenizer),
special_tokens,
info,
id_to_token,
byte_level_decoder,
semantic_markers,
decode_cache: RwLock::new(DecodeCache::new(1000)),
})
}
pub async fn from_file(path: &str) -> Result<Self> {
let tokenizer = HfTokenizer::from_file(path).map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!("Failed to load tokenizer: {}", e))
})?;
let overrides =
special_token_overrides_from_configs(std::path::Path::new(path), &tokenizer);
let mut this = Self::new(tokenizer).await?;
this.apply_special_token_overrides(overrides);
Ok(this)
}
pub async fn from_source_bytes(
tokenizer_json: &[u8],
tokenizer_config_json: Option<&[u8]>,
generation_config_json: Option<&[u8]>,
) -> Result<Self> {
let tokenizer = HfTokenizer::from_bytes(tokenizer_json).map_err(|error| {
ferrum_types::FerrumError::tokenizer(format!(
"Failed to load tokenizer from resolved source bytes: {error}"
))
})?;
let tokenizer_config =
parse_optional_config_bytes(tokenizer_config_json, "tokenizer_config.json")?;
let generation_config =
parse_optional_config_bytes(generation_config_json, "generation_config.json")?;
let overrides = special_token_overrides_from_values(
generation_config.as_ref(),
tokenizer_config.as_ref(),
&tokenizer,
);
let mut this = Self::new(tokenizer).await?;
this.apply_special_token_overrides(overrides);
Ok(this)
}
fn apply_special_token_overrides(&mut self, overrides: SpecialTokenOverrides) {
if overrides.bos.is_some() {
self.special_tokens.bos_token = overrides.bos;
}
if overrides.eos.is_some() {
self.special_tokens.eos_token = overrides.eos;
}
if !overrides.extra_eos.is_empty() {
self.special_tokens.extra_eos_tokens = overrides.extra_eos;
}
self.info.special_tokens = self.special_tokens.clone();
}
pub async fn from_pretrained(repo_id: &str, _revision: Option<&str>) -> Result<Self> {
let api = hf_hub::api::tokio::Api::new().map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!("Failed to create HF API: {}", e))
})?;
let repo = api.repo(hf_hub::Repo::model(repo_id.to_string()));
let tokenizer_file = repo.get("tokenizer.json").await.map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!("Failed to download tokenizer: {}", e))
})?;
let tokenizer = HfTokenizer::from_file(&tokenizer_file).map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!("Failed to load tokenizer: {}", e))
})?;
Self::new(tokenizer).await
}
}
impl Tokenizer for HuggingFaceTokenizer {
fn encode(&self, text: &str, add_special: bool) -> Result<Vec<TokenId>> {
let encoding = self
.tokenizer
.encode(text, add_special)
.map_err(|e| ferrum_types::FerrumError::tokenizer(format!("Encoding failed: {}", e)))?;
Ok(encoding
.get_ids()
.iter()
.map(|&id| TokenId::new(id))
.collect())
}
fn decode(&self, tokens: &[TokenId], skip_special: bool) -> Result<String> {
let token_ids: Vec<u32> = tokens.iter().map(|t| t.get()).collect();
if skip_special
&& !self.semantic_markers.is_empty()
&& token_ids
.iter()
.any(|id| self.semantic_markers.iter().any(|(mid, _)| mid == id))
{
let mut out = String::new();
let mut segment: Vec<u32> = Vec::with_capacity(token_ids.len());
for id in &token_ids {
if let Some((_, canonical)) =
self.semantic_markers.iter().find(|(mid, _)| mid == id)
{
if !segment.is_empty() {
out.push_str(&self.tokenizer.decode(&segment, true).map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e))
})?);
segment.clear();
}
out.push_str(canonical);
} else {
segment.push(*id);
}
}
if !segment.is_empty() {
out.push_str(&self.tokenizer.decode(&segment, true).map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e))
})?);
}
return Ok(out);
}
let text = self
.tokenizer
.decode(&token_ids, skip_special)
.map_err(|e| ferrum_types::FerrumError::tokenizer(format!("Decoding failed: {}", e)))?;
Ok(text)
}
fn decode_incremental(&self, prev: &[TokenId], next: TokenId) -> Result<String> {
let cached_prev = { self.decode_cache.read().get(prev).cloned() };
if let Some(cached_prev) = cached_prev {
let mut all_tokens = prev.to_vec();
all_tokens.push(next);
let full_text = self.decode(&all_tokens, true)?;
self.decode_cache
.write()
.insert(all_tokens, full_text.clone());
return decoded_incremental_delta(&cached_prev, &full_text);
}
let prev_text = if prev.is_empty() {
String::new()
} else {
self.decode(prev, true)?
};
let mut all_tokens = prev.to_vec();
all_tokens.push(next);
let full_text = self.decode(&all_tokens, true)?;
{
let mut cache = self.decode_cache.write();
if !prev.is_empty() {
cache.insert(prev.to_vec(), prev_text.clone());
}
cache.insert(all_tokens, full_text.clone());
}
decoded_incremental_delta(&prev_text, &full_text)
}
fn vocab_size(&self) -> usize {
self.info.vocab_size
}
fn special_tokens(&self) -> &SpecialTokens {
&self.special_tokens
}
fn token_id(&self, text: &str) -> Option<TokenId> {
self.tokenizer.token_to_id(text).map(TokenId::new)
}
fn token_text(&self, token_id: TokenId) -> Option<&str> {
self.id_to_token
.get(token_id.get() as usize)
.and_then(|value| value.as_deref())
}
fn token_bytes(&self, token_id: TokenId) -> Option<Vec<u8>> {
if self.byte_level_decoder {
return self.token_text(token_id).map(byte_level_token_bytes);
}
self.decode(&[token_id], false)
.ok()
.map(String::into_bytes)
.or_else(|| {
self.token_text(token_id)
.map(|text| text.as_bytes().to_vec())
})
}
fn apply_chat_template(
&self,
messages: &[ferrum_interfaces::tokenizer::ChatMessage],
) -> Result<String> {
let mut result = String::new();
for msg in messages {
result.push_str(&format!("{}: {}\n", msg.role, msg.content));
}
Ok(result.trim_end().to_string())
}
fn info(&self) -> TokenizerInfo {
self.info.clone()
}
}
impl IncrementalTokenizer for HuggingFaceTokenizer {
type State = IncrementalState;
fn create_state(&self) -> Self::State {
IncrementalState::default()
}
fn decode_incremental_with_state(
&self,
state: &mut Self::State,
token: TokenId,
) -> Result<String> {
state.tokens.push(token);
let full_text = self.decode(&state.tokens, true)?;
let delta = decoded_incremental_delta(&state.text, &full_text)?;
state.text = full_text;
Ok(delta)
}
fn reset_state(&self, state: &mut Self::State) {
state.tokens.clear();
state.text.clear();
}
fn get_decoded_text(&self, state: &Self::State) -> String {
state.text.clone()
}
}
#[derive(Debug, Clone, Default)]
pub struct HuggingFaceTokenizerFactory;
impl HuggingFaceTokenizerFactory {
pub fn new() -> Self {
Self
}
}
#[async_trait]
impl TokenizerFactory for HuggingFaceTokenizerFactory {
async fn load_from_file(&self, path: &str) -> Result<Box<dyn Tokenizer>> {
let tokenizer = HuggingFaceTokenizer::from_file(path).await?;
Ok(Box::new(tokenizer))
}
async fn load_from_bytes(&self, data: &[u8]) -> Result<Box<dyn Tokenizer>> {
let tokenizer = HfTokenizer::from_bytes(data).map_err(|e| {
ferrum_types::FerrumError::tokenizer(format!(
"Failed to load tokenizer from bytes: {}",
e
))
})?;
let tokenizer = HuggingFaceTokenizer::new(tokenizer).await?;
Ok(Box::new(tokenizer))
}
async fn load_from_hub(
&self,
repo_id: &str,
revision: Option<&str>,
) -> Result<Box<dyn Tokenizer>> {
let tokenizer = HuggingFaceTokenizer::from_pretrained(repo_id, revision).await?;
Ok(Box::new(tokenizer))
}
async fn create_from_config(
&self,
config: &ferrum_interfaces::tokenizer::TokenizerConfig,
) -> Result<Box<dyn Tokenizer>> {
self.load_from_file(&config.path).await
}
fn supported_types(&self) -> Vec<TokenizerType> {
vec![
TokenizerType::BPE,
TokenizerType::WordPiece,
TokenizerType::SentencePiece,
]
}
}
fn build_id_to_token(tokenizer: &HfTokenizer) -> Vec<Option<String>> {
let vocab = tokenizer.get_vocab(true);
let Some(max_id) = vocab.values().copied().max() else {
return Vec::new();
};
let mut id_to_token = vec![None; max_id as usize + 1];
for (token, id) in vocab {
let slot = &mut id_to_token[id as usize];
if slot.is_none() {
*slot = Some(token);
}
}
id_to_token
}
fn decoder_uses_byte_level(decoder: &DecoderWrapper) -> bool {
match decoder {
DecoderWrapper::ByteLevel(_) => true,
DecoderWrapper::Sequence(sequence) => {
sequence.get_decoders().iter().any(decoder_uses_byte_level)
}
_ => false,
}
}
fn byte_level_char_bytes() -> &'static HashMap<char, u8> {
static CHAR_BYTES: OnceLock<HashMap<char, u8>> = OnceLock::new();
CHAR_BYTES.get_or_init(|| {
let mut direct = Vec::with_capacity(256);
direct.extend(b'!'..=b'~');
direct.extend(b'\xA1'..=b'\xAC');
direct.extend(b'\xAE'..=b'\xFF');
let mut next_codepoint = 256u32;
let mut mapping = HashMap::with_capacity(256);
for byte in 0..=u8::MAX {
let codepoint = if direct.contains(&byte) {
byte as u32
} else {
let codepoint = next_codepoint;
next_codepoint += 1;
codepoint
};
let character = char::from_u32(codepoint)
.expect("GPT-2 byte alphabet uses valid Unicode scalar values");
mapping.insert(character, byte);
}
mapping
})
}
fn byte_level_token_bytes(token: &str) -> Vec<u8> {
let mapping = byte_level_char_bytes();
token
.chars()
.map(|character| mapping.get(&character).copied())
.collect::<Option<Vec<_>>>()
.unwrap_or_else(|| token.as_bytes().to_vec())
}
fn extract_special_tokens(tokenizer: &HfTokenizer) -> Result<SpecialTokens> {
let _vocab = tokenizer.get_vocab(false);
let bos_token = tokenizer
.token_to_id("<s>")
.or_else(|| tokenizer.token_to_id("[BOS]"))
.or_else(|| tokenizer.token_to_id("<bos>"))
.map(TokenId::new);
let eos_token = tokenizer
.token_to_id("</s>")
.or_else(|| tokenizer.token_to_id("[EOS]"))
.or_else(|| tokenizer.token_to_id("<eos>"))
.map(TokenId::new);
let unk_token = tokenizer
.token_to_id("<unk>")
.or_else(|| tokenizer.token_to_id("[UNK]"))
.map(TokenId::new);
let pad_token = tokenizer
.token_to_id("<pad>")
.or_else(|| tokenizer.token_to_id("[PAD]"))
.map(TokenId::new);
let sep_token = tokenizer
.token_to_id("[SEP]")
.or_else(|| tokenizer.token_to_id("<sep>"))
.map(TokenId::new);
let cls_token = tokenizer
.token_to_id("[CLS]")
.or_else(|| tokenizer.token_to_id("<cls>"))
.map(TokenId::new);
let mask_token = tokenizer
.token_to_id("[MASK]")
.or_else(|| tokenizer.token_to_id("<mask>"))
.map(TokenId::new);
Ok(SpecialTokens {
bos_token,
eos_token,
unk_token,
pad_token,
sep_token,
cls_token,
mask_token,
extra_eos_tokens: Vec::new(),
})
}
#[derive(Debug, Default)]
struct SpecialTokenOverrides {
bos: Option<TokenId>,
eos: Option<TokenId>,
extra_eos: Vec<TokenId>,
}
fn special_token_overrides_from_configs(
tokenizer_json: &std::path::Path,
tokenizer: &HfTokenizer,
) -> SpecialTokenOverrides {
let Some(dir) = tokenizer_json.parent() else {
return SpecialTokenOverrides::default();
};
let generation_config = read_json(&dir.join("generation_config.json"));
let tokenizer_config = read_json(&dir.join("tokenizer_config.json"));
special_token_overrides_from_values(
generation_config.as_ref(),
tokenizer_config.as_ref(),
tokenizer,
)
}
fn special_token_overrides_from_values(
generation_config: Option<&serde_json::Value>,
tokenizer_config: Option<&serde_json::Value>,
tokenizer: &HfTokenizer,
) -> SpecialTokenOverrides {
let mut overrides = SpecialTokenOverrides::default();
if let Some(gen) = generation_config {
let mut eos_ids = token_id_list(gen.get("eos_token_id"));
if !eos_ids.is_empty() {
overrides.eos = Some(eos_ids.remove(0));
overrides.extra_eos = eos_ids;
}
if let Some(bos) = token_id_list(gen.get("bos_token_id")).into_iter().next() {
overrides.bos = Some(bos);
}
}
if let Some(tok_cfg) = tokenizer_config {
if overrides.eos.is_none() {
overrides.eos = token_from_config_value(tok_cfg.get("eos_token"), tokenizer);
}
if overrides.bos.is_none() {
overrides.bos = token_from_config_value(tok_cfg.get("bos_token"), tokenizer);
}
}
overrides
}
fn parse_optional_config_bytes(
bytes: Option<&[u8]>,
source_file: &str,
) -> Result<Option<serde_json::Value>> {
bytes
.map(|bytes| {
serde_json::from_slice(bytes).map_err(|error| {
ferrum_types::FerrumError::tokenizer(format!(
"Failed to parse resolved {source_file}: {error}"
))
})
})
.transpose()
}
fn read_json(path: &std::path::Path) -> Option<serde_json::Value> {
let text = std::fs::read_to_string(path).ok()?;
serde_json::from_str(&text).ok()
}
fn token_id_list(value: Option<&serde_json::Value>) -> Vec<TokenId> {
match value {
Some(serde_json::Value::Number(n)) => n
.as_u64()
.map(|v| vec![TokenId::new(v as u32)])
.unwrap_or_default(),
Some(serde_json::Value::Array(items)) => items
.iter()
.filter_map(|v| v.as_u64())
.map(|v| TokenId::new(v as u32))
.collect(),
_ => Vec::new(),
}
}
fn token_from_config_value(
value: Option<&serde_json::Value>,
tokenizer: &HfTokenizer,
) -> Option<TokenId> {
let text = match value? {
serde_json::Value::String(s) => s.as_str(),
serde_json::Value::Object(obj) => obj.get("content")?.as_str()?,
_ => return None,
};
tokenizer.token_to_id(text).map(TokenId::new)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_decode_cache_creation() {
let cache = DecodeCache::new(100);
assert_eq!(cache.max_size, 100);
assert_eq!(cache.cache.len(), 0);
}
#[test]
fn test_decode_cache_insert_and_get() {
let mut cache = DecodeCache::new(10);
let tokens = vec![TokenId::new(1), TokenId::new(2)];
let text = "hello".to_string();
cache.insert(tokens.clone(), text.clone());
let result = cache.get(&tokens);
assert!(result.is_some());
assert_eq!(result.unwrap(), &text);
}
#[test]
fn test_decode_cache_eviction() {
let mut cache = DecodeCache::new(2);
cache.insert(vec![TokenId::new(1)], "a".to_string());
cache.insert(vec![TokenId::new(2)], "b".to_string());
assert_eq!(cache.cache.len(), 2);
cache.insert(vec![TokenId::new(3)], "c".to_string());
assert!(cache.cache.len() <= 2);
}
#[test]
fn incremental_delta_rejects_a_rewritten_prefix() {
assert!(decoded_incremental_delta("stable", "changed").is_err());
}
#[tokio::test]
async fn incremental_decode_cache_hit_after_thinking_whitespace_does_not_deadlock() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
let vocab: Vocab = [
("</think>".to_string(), 0),
("\n".to_string(), 1),
("payload".to_string(), 2),
("<unk>".to_string(), 3),
]
.into_iter()
.collect();
let bpe = BPE::builder()
.vocab_and_merges(vocab, vec![])
.unk_token("<unk>".to_string())
.build()
.unwrap();
let mut hf_tokenizer = HfTokenizer::new(bpe);
hf_tokenizer.add_special_tokens(&[AddedToken::from("</think>", true)]);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
let delimiter = tokenizer.token_id("</think>").unwrap();
let whitespace = tokenizer.token_id("\n").unwrap();
let payload = tokenizer.token_id("payload").unwrap();
let delimiter_prefix = vec![delimiter];
let whitespace_prefix = vec![delimiter, whitespace];
assert_eq!(
tokenizer
.decode_incremental(&delimiter_prefix, whitespace)
.unwrap(),
"\n"
);
assert!(tokenizer
.decode_cache
.read()
.get(&whitespace_prefix)
.is_some());
let payload_delta = tokenizer
.decode_incremental(&whitespace_prefix, payload)
.unwrap();
assert_eq!(payload_delta.trim_start(), "payload");
}
#[test]
fn test_incremental_state_default() {
let state = IncrementalState::default();
let debug_str = format!("{:?}", state);
assert!(debug_str.contains("IncrementalState"));
}
#[test]
fn test_incremental_state_clone() {
let state = IncrementalState::default();
let cloned = state.clone();
let state_str = format!("{:?}", state);
let cloned_str = format!("{:?}", cloned);
assert_eq!(state_str, cloned_str);
}
#[test]
fn test_huggingface_tokenizer_factory_creation() {
let factory = HuggingFaceTokenizerFactory::new();
let debug_str = format!("{:?}", factory);
assert!(debug_str.contains("HuggingFaceTokenizerFactory"));
}
#[test]
fn test_huggingface_tokenizer_factory_default() {
let factory = HuggingFaceTokenizerFactory;
let debug_str = format!("{:?}", factory);
assert!(debug_str.contains("HuggingFaceTokenizerFactory"));
}
#[test]
fn test_huggingface_tokenizer_factory_clone() {
let factory = HuggingFaceTokenizerFactory::new();
let cloned = factory.clone();
let factory_str = format!("{:?}", factory);
let cloned_str = format!("{:?}", cloned);
assert_eq!(factory_str, cloned_str);
}
#[test]
fn test_huggingface_tokenizer_factory_supported_types() {
let factory = HuggingFaceTokenizerFactory::new();
let types = factory.supported_types();
assert!(!types.is_empty());
assert!(types.contains(&TokenizerType::BPE));
}
#[test]
fn test_extract_special_tokens_with_mock_tokenizer() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
let vocab: Vocab = [
("hello".to_string(), 0),
("<s>".to_string(), 1),
("</s>".to_string(), 2),
("<unk>".to_string(), 3),
("<pad>".to_string(), 4),
]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.unk_token("<unk>".to_string())
.build()
.unwrap();
let mut tokenizer = HfTokenizer::new(bpe);
tokenizer.add_special_tokens(&[
AddedToken::from("<s>", true),
AddedToken::from("</s>", true),
AddedToken::from("<unk>", true),
AddedToken::from("<pad>", true),
]);
let result = extract_special_tokens(&tokenizer);
assert!(result.is_ok());
let special_tokens = result.unwrap();
assert!(special_tokens.bos_token.is_some());
assert!(special_tokens.eos_token.is_some());
assert!(special_tokens.unk_token.is_some());
assert!(special_tokens.pad_token.is_some());
}
#[tokio::test]
async fn test_huggingface_tokenizer_with_mock() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
let vocab: Vocab = [
("hello".to_string(), 0),
("world".to_string(), 1),
("<s>".to_string(), 2),
("</s>".to_string(), 3),
("<unk>".to_string(), 4),
]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.unk_token("<unk>".to_string())
.build()
.unwrap();
let mut hf_tokenizer = HfTokenizer::new(bpe);
hf_tokenizer.add_special_tokens(&[
AddedToken::from("<s>", true),
AddedToken::from("</s>", true),
AddedToken::from("<unk>", true),
]);
let result = HuggingFaceTokenizer::new(hf_tokenizer).await;
assert!(result.is_ok());
let tokenizer = result.unwrap();
assert_eq!(tokenizer.vocab_size(), 5);
}
#[tokio::test]
async fn test_tokenizer_encode_decode() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
let vocab: Vocab = [
("hello".to_string(), 0),
("world".to_string(), 1),
("<s>".to_string(), 2),
("</s>".to_string(), 3),
("<unk>".to_string(), 4),
]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.unk_token("<unk>".to_string())
.build()
.unwrap();
let mut hf_tokenizer = HfTokenizer::new(bpe);
hf_tokenizer.add_special_tokens(&[
AddedToken::from("<s>", true),
AddedToken::from("</s>", true),
AddedToken::from("<unk>", true),
]);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
let result = tokenizer.encode("hello", false);
assert!(result.is_ok());
let _tokens = result.unwrap();
let decoded = tokenizer.decode(&[], false);
assert!(decoded.is_ok());
}
#[tokio::test]
async fn test_tokenizer_special_tokens() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
let vocab: Vocab = [
("hello".to_string(), 0),
("<s>".to_string(), 1),
("</s>".to_string(), 2),
]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.build()
.unwrap();
let mut hf_tokenizer = HfTokenizer::new(bpe);
hf_tokenizer.add_special_tokens(&[
AddedToken::from("<s>", true),
AddedToken::from("</s>", true),
]);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
let special_tokens = tokenizer.special_tokens();
assert!(special_tokens.bos_token.is_some() || special_tokens.eos_token.is_some());
}
#[tokio::test]
async fn test_tokenizer_token_id_lookup() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::Tokenizer as HfTokenizer;
let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.build()
.unwrap();
let hf_tokenizer = HfTokenizer::new(bpe);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
let token_id = tokenizer.token_id("hello");
assert!(token_id.is_some());
assert_eq!(token_id.unwrap().get(), 0);
}
#[tokio::test]
async fn test_tokenizer_token_text_reverse_lookup() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::Tokenizer as HfTokenizer;
let vocab: Vocab = [
("hello".to_string(), 0),
("[PAD151935]".to_string(), 1),
("</think>".to_string(), 2),
]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.build()
.unwrap();
let hf_tokenizer = HfTokenizer::new(bpe);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
assert_eq!(tokenizer.token_text(TokenId::new(1)), Some("[PAD151935]"));
assert_eq!(tokenizer.token_text(TokenId::new(2)), Some("</think>"));
assert_eq!(tokenizer.token_text(TokenId::new(99)), None);
}
#[tokio::test]
async fn byte_level_token_bytes_preserve_split_utf8_fragments() {
use tokenizers::decoders::byte_level::ByteLevel;
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::{AddedToken, Tokenizer as HfTokenizer};
let vocab: Vocab = [
("\u{00f0}\u{0141}".to_string(), 0),
("\u{0136}\u{00a5}".to_string(), 1),
("<eos>".to_string(), 2),
]
.into_iter()
.collect();
let bpe = BPE::builder()
.vocab_and_merges(vocab, vec![])
.build()
.unwrap();
let mut hf_tokenizer = HfTokenizer::new(bpe);
hf_tokenizer.with_decoder(Some(ByteLevel::default()));
hf_tokenizer.add_special_tokens(&[AddedToken::from("<eos>", true)]);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
assert!(tokenizer
.decode(&[TokenId::new(0)], false)
.unwrap()
.contains('\u{fffd}'));
assert_eq!(
tokenizer
.decode(&[TokenId::new(0), TokenId::new(1)], false)
.unwrap(),
"\u{1f525}"
);
assert_eq!(
tokenizer.token_bytes(TokenId::new(0)),
Some(vec![0xf0, 0x9f])
);
assert_eq!(
tokenizer.token_bytes(TokenId::new(1)),
Some(vec![0x94, 0xa5])
);
assert_eq!(tokenizer.token_bytes(TokenId::new(99)), None);
}
#[tokio::test]
async fn test_tokenizer_info() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::Tokenizer as HfTokenizer;
let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.build()
.unwrap();
let hf_tokenizer = HfTokenizer::new(bpe);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
let info = tokenizer.info();
assert_eq!(info.vocab_size, 2);
assert!(info.supports_incremental);
assert_eq!(info.tokenizer_type, TokenizerType::BPE);
}
#[tokio::test]
async fn test_incremental_tokenizer_interface() {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::Tokenizer as HfTokenizer;
let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
.into_iter()
.collect();
let merges = vec![];
let bpe = BPE::builder()
.vocab_and_merges(vocab, merges)
.build()
.unwrap();
let hf_tokenizer = HfTokenizer::new(bpe);
let tokenizer = HuggingFaceTokenizer::new(hf_tokenizer).await.unwrap();
let mut state = tokenizer.create_state();
let result = tokenizer.decode_incremental_with_state(&mut state, TokenId::new(0));
assert!(result.is_ok());
tokenizer.reset_state(&mut state);
let text = tokenizer.get_decoded_text(&state);
assert!(text.is_empty());
}
fn tiny_tokenizer_with_specials(specials: &[&str]) -> HfTokenizer {
use tokenizers::models::bpe::{Vocab, BPE};
use tokenizers::AddedToken;
let vocab: Vocab = [("hello".to_string(), 0), ("world".to_string(), 1)]
.into_iter()
.collect();
let bpe = BPE::builder()
.vocab_and_merges(vocab, vec![])
.unk_token("hello".to_string())
.build()
.unwrap();
let mut tokenizer = HfTokenizer::new(bpe);
tokenizer.add_special_tokens(
&specials
.iter()
.map(|s| AddedToken::from(*s, true))
.collect::<Vec<_>>(),
);
tokenizer
}
#[tokio::test]
async fn eos_comes_from_generation_config_not_name_probing() {
let tokenizer =
tiny_tokenizer_with_specials(&["<|end▁of▁sentence|>", "<|User|>", "<|Assistant|>"]);
let eos_id = tokenizer.token_to_id("<|end▁of▁sentence|>").unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("tokenizer.json");
tokenizer.save(&path, false).unwrap();
std::fs::write(
dir.path().join("generation_config.json"),
format!("{{\"bos_token_id\": null, \"eos_token_id\": {eos_id}}}"),
)
.unwrap();
let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
.await
.unwrap();
assert_eq!(
loaded.special_tokens().eos_token.map(|t| t.get()),
Some(eos_id)
);
assert!(loaded.special_tokens().extra_eos_tokens.is_empty());
}
#[tokio::test]
async fn immutable_source_bytes_preserve_generation_config_eos() {
let tokenizer = tiny_tokenizer_with_specials(&["<|end_of_text|>", "<|end_of_turn|>"]);
let primary = tokenizer.token_to_id("<|end_of_text|>").unwrap();
let extra = tokenizer.token_to_id("<|end_of_turn|>").unwrap();
let tokenizer_json = tokenizer.to_string(false).unwrap();
let generation_config = format!(r#"{{"eos_token_id":[{primary},{extra}]}}"#);
let loaded = HuggingFaceTokenizer::from_source_bytes(
tokenizer_json.as_bytes(),
None,
Some(generation_config.as_bytes()),
)
.await
.unwrap();
assert_eq!(
loaded.special_tokens().eos_token.map(|token| token.get()),
Some(primary)
);
assert_eq!(
loaded
.special_tokens()
.extra_eos_tokens
.iter()
.map(|token| token.get())
.collect::<Vec<_>>(),
vec![extra]
);
}
#[tokio::test]
async fn multi_eos_ids_land_in_extra_eos_tokens() {
let tokenizer = tiny_tokenizer_with_specials(&["<|eot_id|>", "<|end_of_text|>"]);
let primary = tokenizer.token_to_id("<|end_of_text|>").unwrap();
let extra = tokenizer.token_to_id("<|eot_id|>").unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("tokenizer.json");
tokenizer.save(&path, false).unwrap();
std::fs::write(
dir.path().join("generation_config.json"),
format!("{{\"eos_token_id\": [{primary}, {extra}]}}"),
)
.unwrap();
let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
.await
.unwrap();
assert_eq!(
loaded.special_tokens().eos_token.map(|t| t.get()),
Some(primary)
);
assert_eq!(
loaded
.special_tokens()
.extra_eos_tokens
.iter()
.map(|t| t.get())
.collect::<Vec<_>>(),
vec![extra]
);
}
#[tokio::test]
async fn tokenizer_config_eos_string_is_fallback_without_generation_config() {
let tokenizer = tiny_tokenizer_with_specials(&["<|end▁of▁sentence|>"]);
let eos_id = tokenizer.token_to_id("<|end▁of▁sentence|>").unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("tokenizer.json");
tokenizer.save(&path, false).unwrap();
std::fs::write(
dir.path().join("tokenizer_config.json"),
"{\"eos_token\": {\"content\": \"<|end▁of▁sentence|>\"}}",
)
.unwrap();
let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
.await
.unwrap();
assert_eq!(
loaded.special_tokens().eos_token.map(|t| t.get()),
Some(eos_id)
);
}
#[tokio::test]
async fn bare_tokenizer_json_still_uses_name_probing() {
let tokenizer = tiny_tokenizer_with_specials(&["<s>", "</s>"]);
let eos_id = tokenizer.token_to_id("</s>").unwrap();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("tokenizer.json");
tokenizer.save(&path, false).unwrap();
let loaded = HuggingFaceTokenizer::from_file(path.to_str().unwrap())
.await
.unwrap();
assert_eq!(
loaded.special_tokens().eos_token.map(|t| t.get()),
Some(eos_id)
);
}
#[tokio::test]
async fn skip_special_decode_preserves_and_normalizes_think_markers() {
let tokenizer = tiny_tokenizer_with_specials(&["[THINK]", "[/THINK]", "<eos>"]);
let think = tokenizer.token_to_id("[THINK]").unwrap();
let end_think = tokenizer.token_to_id("[/THINK]").unwrap();
let eos = tokenizer.token_to_id("<eos>").unwrap();
let hello = tokenizer.token_to_id("hello").unwrap();
let world = tokenizer.token_to_id("world").unwrap();
let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
let tokens: Vec<TokenId> = [think, hello, end_think, world, eos]
.into_iter()
.map(TokenId::new)
.collect();
let text = loaded.decode(&tokens, true).unwrap();
assert_eq!(text, "<think>hello</think>world");
}
#[tokio::test]
async fn skip_special_decode_preserves_typed_harmony_markers_only() {
let harmony_markers = ModelOutputProtocol::HarmonyGptOss.preserved_special_token_texts();
let mut specials = harmony_markers.to_vec();
specials.push("<|endoftext|>");
let tokenizer = tiny_tokenizer_with_specials(&specials);
let mut token_ids = vec![tokenizer.token_to_id("hello").unwrap()];
token_ids.extend(
harmony_markers
.iter()
.map(|marker| tokenizer.token_to_id(marker).unwrap()),
);
token_ids.push(tokenizer.token_to_id("world").unwrap());
token_ids.push(tokenizer.token_to_id("<|endoftext|>").unwrap());
let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
let tokens: Vec<TokenId> = token_ids.into_iter().map(TokenId::new).collect();
let text = loaded.decode(&tokens, true).unwrap();
assert_eq!(
text,
format!("hello{}world", harmony_markers.concat()),
"the typed Harmony markers must survive while unrelated special tokens stay skipped"
);
}
#[tokio::test]
async fn skip_special_decode_without_markers_is_unchanged() {
let tokenizer = tiny_tokenizer_with_specials(&["<eos>"]);
let eos = tokenizer.token_to_id("<eos>").unwrap();
let hello = tokenizer.token_to_id("hello").unwrap();
let loaded = HuggingFaceTokenizer::new(tokenizer).await.unwrap();
let tokens: Vec<TokenId> = [hello, eos].into_iter().map(TokenId::new).collect();
assert_eq!(loaded.decode(&tokens, true).unwrap(), "hello");
}
}