use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use bytes::Bytes;
use dashmap::DashMap;
use tokenizers::Tokenizer;
#[async_trait::async_trait]
pub trait TokenizerStore: Send + Sync {
async fn list_tokenizer_vocabs(&self) -> anyhow::Result<Vec<String>>;
async fn get_tokenizer_vocab(&self, name: &str) -> anyhow::Result<Option<Vec<u8>>>;
async fn put_tokenizer_vocab(&self, name: &str, bytes: &[u8]) -> anyhow::Result<()>;
}
#[async_trait::async_trait]
pub trait TokenizerClient: Send + Sync {
async fn send(&self, req: http::Request<Bytes>) -> anyhow::Result<http::Response<Bytes>>;
}
static DEEPSEEK: &[u8] = include_bytes!("../../assets/tokenizers/deepseek-v4-pro.tokenizer.json");
const BUNDLED_NAMES: &[&str] = &["deepseek", "deepseek-v4-pro"];
static BUNDLED: std::sync::OnceLock<Option<Arc<Tokenizer>>> = std::sync::OnceLock::new();
fn bundled_tokenizer() -> Option<Arc<Tokenizer>> {
BUNDLED
.get_or_init(|| match Tokenizer::from_bytes(DEEPSEEK) {
Ok(t) => Some(Arc::new(t)),
Err(e) => {
tracing::error!(error = %e, "bundled tokenizer failed to parse");
None
}
})
.clone()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VocabSource {
BuiltinTiktoken,
Bundled,
Downloaded,
}
#[derive(Debug, Clone)]
pub struct VocabInfo {
pub name: String,
pub source: VocabSource,
pub loaded: bool,
}
type LoadedMap = Arc<DashMap<String, Arc<Tokenizer>>>;
pub struct TokenizerRegistry {
store: Arc<dyn TokenizerStore>,
download_enabled: AtomicBool,
upstream: Arc<dyn TokenizerClient>,
loaded: LoadedMap,
inflight: Arc<DashMap<String, ()>>,
}
impl TokenizerRegistry {
pub fn new(store: Arc<dyn TokenizerStore>, upstream: Arc<dyn TokenizerClient>) -> Self {
Self {
store,
download_enabled: AtomicBool::new(false),
upstream,
loaded: Arc::new(DashMap::new()),
inflight: Arc::new(DashMap::new()),
}
}
pub fn set_download_enabled(&self, on: bool) {
self.download_enabled.store(on, Ordering::Relaxed);
}
pub async fn list(&self) -> Vec<VocabInfo> {
let mut out = vec![
info("o200k_base", VocabSource::BuiltinTiktoken, true),
info("cl100k_base", VocabSource::BuiltinTiktoken, true),
info(
BUNDLED_NAMES[0],
VocabSource::Bundled,
self.loaded.contains_key(BUNDLED_NAMES[0]),
),
];
match self.store.list_tokenizer_vocabs().await {
Ok(names) => {
for name in names {
let loaded = self.loaded.contains_key(&name);
out.push(info(&name, VocabSource::Downloaded, loaded));
}
}
Err(e) => tracing::warn!(error = %e, "listing persisted tokenizer vocabs failed"),
}
out
}
pub fn resolve(&self, name: &str) -> Option<Arc<Tokenizer>> {
if let Some(t) = self.loaded.get(name) {
return Some(Arc::clone(&t));
}
if BUNDLED_NAMES.contains(&name) {
let tok = bundled_tokenizer()?;
for n in BUNDLED_NAMES {
self.loaded.insert((*n).to_owned(), Arc::clone(&tok));
}
return Some(tok);
}
None
}
pub fn preheat(&self) {
let loaded = Arc::clone(&self.loaded);
tokio::task::spawn_blocking(move || {
if let Some(tok) = bundled_tokenizer() {
for n in BUNDLED_NAMES {
loaded.insert((*n).to_owned(), Arc::clone(&tok));
}
}
});
}
pub fn request_load(&self, name: &str) {
if self.inflight.insert(name.to_owned(), ()).is_some() {
return;
}
let store = Arc::clone(&self.store);
let upstream = Arc::clone(&self.upstream);
let loaded = Arc::clone(&self.loaded);
let inflight = Arc::clone(&self.inflight);
let download_enabled = self.download_enabled.load(Ordering::Relaxed);
let name = name.to_owned();
tokio::spawn(async move {
if let Err(e) = load(store, upstream, &name, &loaded, download_enabled).await {
tracing::warn!(name, error = %e, "tokenizer load failed");
}
inflight.remove(&name);
});
}
}
async fn load(
store: Arc<dyn TokenizerStore>,
upstream: Arc<dyn TokenizerClient>,
name: &str,
loaded: &LoadedMap,
download_enabled: bool,
) -> anyhow::Result<()> {
if let Some(bytes) = store.get_tokenizer_vocab(name).await? {
let tok = Tokenizer::from_bytes(&bytes).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
loaded.insert(name.to_owned(), Arc::new(tok));
return Ok(());
}
if !download_enabled || !name.contains('/') {
return Ok(());
}
let url = format!("https://huggingface.co/{name}/resolve/main/tokenizer.json");
let req = http::Request::builder()
.method(http::Method::GET)
.uri(&url)
.body(Bytes::new())?;
let resp = upstream.send(req).await?;
anyhow::ensure!(resp.status().is_success(), "HTTP {}", resp.status());
let body = resp.into_body();
let tok = Tokenizer::from_bytes(&body).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
store.put_tokenizer_vocab(name, &body).await?;
loaded.insert(name.to_owned(), Arc::new(tok));
tracing::info!(name, "tokenizer downloaded");
Ok(())
}
fn info(name: &str, source: VocabSource, loaded: bool) -> VocabInfo {
VocabInfo {
name: name.to_owned(),
source,
loaded,
}
}