gproxy_tokenize/tokenize/
registry.rs1use std::sync::Arc;
9use std::sync::atomic::{AtomicBool, Ordering};
10
11use bytes::Bytes;
12use dashmap::DashMap;
13use tokenizers::Tokenizer;
14
15#[async_trait::async_trait]
16pub trait TokenizerStore: Send + Sync {
17 async fn list_tokenizer_vocabs(&self) -> anyhow::Result<Vec<String>>;
18 async fn get_tokenizer_vocab(&self, name: &str) -> anyhow::Result<Option<Vec<u8>>>;
19 async fn put_tokenizer_vocab(&self, name: &str, bytes: &[u8]) -> anyhow::Result<()>;
20}
21
22#[async_trait::async_trait]
23pub trait TokenizerClient: Send + Sync {
24 async fn send(&self, req: http::Request<Bytes>) -> anyhow::Result<http::Response<Bytes>>;
25}
26
27static DEEPSEEK: &[u8] = include_bytes!("../../assets/tokenizers/deepseek-v4-pro.tokenizer.json");
30const BUNDLED_NAMES: &[&str] = &["deepseek", "deepseek-v4-pro"];
32
33static BUNDLED: std::sync::OnceLock<Option<Arc<Tokenizer>>> = std::sync::OnceLock::new();
39
40fn bundled_tokenizer() -> Option<Arc<Tokenizer>> {
41 BUNDLED
42 .get_or_init(|| match Tokenizer::from_bytes(DEEPSEEK) {
43 Ok(t) => Some(Arc::new(t)),
44 Err(e) => {
45 tracing::error!(error = %e, "bundled tokenizer failed to parse");
46 None
47 }
48 })
49 .clone()
50}
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54pub enum VocabSource {
55 BuiltinTiktoken,
56 Bundled,
57 Downloaded,
58}
59
60#[derive(Debug, Clone)]
62pub struct VocabInfo {
63 pub name: String,
64 pub source: VocabSource,
65 pub loaded: bool,
66}
67
68type LoadedMap = Arc<DashMap<String, Arc<Tokenizer>>>;
69
70pub struct TokenizerRegistry {
72 store: Arc<dyn TokenizerStore>,
74 download_enabled: AtomicBool,
76 upstream: Arc<dyn TokenizerClient>,
77 loaded: LoadedMap,
78 inflight: Arc<DashMap<String, ()>>,
79}
80
81impl TokenizerRegistry {
82 pub fn new(store: Arc<dyn TokenizerStore>, upstream: Arc<dyn TokenizerClient>) -> Self {
83 Self {
84 store,
85 download_enabled: AtomicBool::new(false),
86 upstream,
87 loaded: Arc::new(DashMap::new()),
88 inflight: Arc::new(DashMap::new()),
89 }
90 }
91
92 pub fn set_download_enabled(&self, on: bool) {
93 self.download_enabled.store(on, Ordering::Relaxed);
94 }
95
96 pub async fn list(&self) -> Vec<VocabInfo> {
99 let mut out = vec![
100 info("o200k_base", VocabSource::BuiltinTiktoken, true),
101 info("cl100k_base", VocabSource::BuiltinTiktoken, true),
102 info(
103 BUNDLED_NAMES[0],
104 VocabSource::Bundled,
105 self.loaded.contains_key(BUNDLED_NAMES[0]),
106 ),
107 ];
108 match self.store.list_tokenizer_vocabs().await {
109 Ok(names) => {
110 for name in names {
111 let loaded = self.loaded.contains_key(&name);
112 out.push(info(&name, VocabSource::Downloaded, loaded));
113 }
114 }
115 Err(e) => tracing::warn!(error = %e, "listing persisted tokenizer vocabs failed"),
116 }
117 out
118 }
119
120 pub fn resolve(&self, name: &str) -> Option<Arc<Tokenizer>> {
124 if let Some(t) = self.loaded.get(name) {
125 return Some(Arc::clone(&t));
126 }
127 if BUNDLED_NAMES.contains(&name) {
128 let tok = bundled_tokenizer()?;
129 for n in BUNDLED_NAMES {
130 self.loaded.insert((*n).to_owned(), Arc::clone(&tok));
131 }
132 return Some(tok);
133 }
134 None
135 }
136
137 pub fn preheat(&self) {
140 let loaded = Arc::clone(&self.loaded);
141 tokio::task::spawn_blocking(move || {
142 if let Some(tok) = bundled_tokenizer() {
143 for n in BUNDLED_NAMES {
144 loaded.insert((*n).to_owned(), Arc::clone(&tok));
145 }
146 }
147 });
148 }
149
150 pub fn request_load(&self, name: &str) {
156 if self.inflight.insert(name.to_owned(), ()).is_some() {
157 return;
158 }
159 let store = Arc::clone(&self.store);
160 let upstream = Arc::clone(&self.upstream);
161 let loaded = Arc::clone(&self.loaded);
162 let inflight = Arc::clone(&self.inflight);
163 let download_enabled = self.download_enabled.load(Ordering::Relaxed);
164 let name = name.to_owned();
165 tokio::spawn(async move {
166 if let Err(e) = load(store, upstream, &name, &loaded, download_enabled).await {
167 tracing::warn!(name, error = %e, "tokenizer load failed");
168 }
169 inflight.remove(&name);
170 });
171 }
172}
173
174async fn load(
176 store: Arc<dyn TokenizerStore>,
177 upstream: Arc<dyn TokenizerClient>,
178 name: &str,
179 loaded: &LoadedMap,
180 download_enabled: bool,
181) -> anyhow::Result<()> {
182 if let Some(bytes) = store.get_tokenizer_vocab(name).await? {
183 let tok = Tokenizer::from_bytes(&bytes).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
184 loaded.insert(name.to_owned(), Arc::new(tok));
185 return Ok(());
186 }
187 if !download_enabled || !name.contains('/') {
188 return Ok(());
189 }
190
191 let url = format!("https://huggingface.co/{name}/resolve/main/tokenizer.json");
192 let req = http::Request::builder()
193 .method(http::Method::GET)
194 .uri(&url)
195 .body(Bytes::new())?;
196 let resp = upstream.send(req).await?;
197 anyhow::ensure!(resp.status().is_success(), "HTTP {}", resp.status());
198 let body = resp.into_body();
199 let tok = Tokenizer::from_bytes(&body).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
200
201 store.put_tokenizer_vocab(name, &body).await?;
202 loaded.insert(name.to_owned(), Arc::new(tok));
203 tracing::info!(name, "tokenizer downloaded");
204 Ok(())
205}
206
207fn info(name: &str, source: VocabSource, loaded: bool) -> VocabInfo {
208 VocabInfo {
209 name: name.to_owned(),
210 source,
211 loaded,
212 }
213}