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 async fn quarantine_tokenizer_vocab(&self, _name: &str, _reason: &str) -> anyhow::Result<()> {
24 Ok(())
25 }
26}
27
28#[async_trait::async_trait]
29pub trait TokenizerClient: Send + Sync {
30 async fn send(&self, req: http::Request<Bytes>) -> anyhow::Result<http::Response<Bytes>>;
31}
32
33#[cfg(feature = "bundled-fallback")]
36static DEEPSEEK: &[u8] = include_bytes!("../../assets/tokenizers/deepseek-v4-pro.tokenizer.json");
37#[cfg(feature = "bundled-fallback")]
39const BUNDLED_NAMES: &[&str] = &["deepseek", "deepseek-v4-pro"];
40
41#[cfg(feature = "bundled-fallback")]
47static BUNDLED: std::sync::OnceLock<Option<Arc<Tokenizer>>> = std::sync::OnceLock::new();
48
49#[cfg(feature = "bundled-fallback")]
50fn bundled_tokenizer() -> Option<Arc<Tokenizer>> {
51 BUNDLED
52 .get_or_init(|| match Tokenizer::from_bytes(DEEPSEEK) {
53 Ok(t) => Some(Arc::new(t)),
54 Err(e) => {
55 tracing::error!(error = %e, "bundled tokenizer failed to parse");
56 None
57 }
58 })
59 .clone()
60}
61
62#[derive(Debug, Clone, Copy, PartialEq, Eq)]
64pub enum VocabSource {
65 BuiltinTiktoken,
66 Bundled,
67 Downloaded,
68}
69
70#[derive(Debug, Clone)]
72pub struct VocabInfo {
73 pub name: String,
74 pub source: VocabSource,
75 pub loaded: bool,
76}
77
78type LoadedMap = Arc<DashMap<String, Arc<Tokenizer>>>;
79
80pub const MAX_TOKENIZER_BYTES: usize = 16 * 1024 * 1024;
81
82#[derive(Debug, Clone, Copy, PartialEq, Eq)]
83pub enum LoadRequestStatus {
84 Scheduled,
85 AlreadyInFlight,
86 NegativeCached,
87 NoRuntime,
88}
89
90pub struct TokenizerRegistry {
92 store: Arc<dyn TokenizerStore>,
94 download_enabled: AtomicBool,
96 upstream: Arc<dyn TokenizerClient>,
97 loaded: LoadedMap,
98 inflight: Arc<DashMap<String, ()>>,
99 negative: Arc<DashMap<String, ()>>,
100}
101
102impl TokenizerRegistry {
103 pub fn new(store: Arc<dyn TokenizerStore>, upstream: Arc<dyn TokenizerClient>) -> Self {
104 Self {
105 store,
106 download_enabled: AtomicBool::new(false),
107 upstream,
108 loaded: Arc::new(DashMap::new()),
109 inflight: Arc::new(DashMap::new()),
110 negative: Arc::new(DashMap::new()),
111 }
112 }
113
114 pub fn set_download_enabled(&self, on: bool) {
115 self.download_enabled.store(on, Ordering::Relaxed);
116 if on {
117 self.negative.clear();
118 }
119 }
120
121 pub async fn list(&self) -> Vec<VocabInfo> {
124 let mut out = vec![
125 info("o200k_base", VocabSource::BuiltinTiktoken, true),
126 info("cl100k_base", VocabSource::BuiltinTiktoken, true),
127 ];
128 #[cfg(feature = "bundled-fallback")]
129 out.push(info(
130 BUNDLED_NAMES[0],
131 VocabSource::Bundled,
132 self.loaded.contains_key(BUNDLED_NAMES[0]),
133 ));
134 match self.store.list_tokenizer_vocabs().await {
135 Ok(names) => {
136 for name in names {
137 let loaded = self.loaded.contains_key(&name);
138 out.push(info(&name, VocabSource::Downloaded, loaded));
139 }
140 }
141 Err(e) => tracing::warn!(error = %e, "listing persisted tokenizer vocabs failed"),
142 }
143 out
144 }
145
146 pub fn resolve(&self, name: &str) -> Option<Arc<Tokenizer>> {
150 if let Some(t) = self.loaded.get(name) {
151 return Some(Arc::clone(&t));
152 }
153 #[cfg(feature = "bundled-fallback")]
154 if BUNDLED_NAMES.contains(&name) {
155 let tok = bundled_tokenizer()?;
156 for n in BUNDLED_NAMES {
157 self.loaded.insert((*n).to_owned(), Arc::clone(&tok));
158 }
159 return Some(tok);
160 }
161 None
162 }
163
164 pub fn preheat(&self) -> LoadRequestStatus {
167 #[cfg(not(feature = "bundled-fallback"))]
168 return LoadRequestStatus::NegativeCached;
169 #[cfg(feature = "bundled-fallback")]
170 let Ok(runtime) = tokio::runtime::Handle::try_current() else {
171 return LoadRequestStatus::NoRuntime;
172 };
173 #[cfg(feature = "bundled-fallback")]
174 let loaded = Arc::clone(&self.loaded);
175 #[cfg(feature = "bundled-fallback")]
176 runtime.spawn_blocking(move || {
177 if let Some(tok) = bundled_tokenizer() {
178 for n in BUNDLED_NAMES {
179 loaded.insert((*n).to_owned(), Arc::clone(&tok));
180 }
181 }
182 });
183 #[cfg(feature = "bundled-fallback")]
184 return LoadRequestStatus::Scheduled;
185 }
186
187 pub fn request_load(&self, name: &str) -> LoadRequestStatus {
193 if self.negative.contains_key(name) {
194 return LoadRequestStatus::NegativeCached;
195 }
196 if self.inflight.insert(name.to_owned(), ()).is_some() {
197 return LoadRequestStatus::AlreadyInFlight;
198 }
199 let Ok(runtime) = tokio::runtime::Handle::try_current() else {
200 self.inflight.remove(name);
201 return LoadRequestStatus::NoRuntime;
202 };
203 let store = Arc::clone(&self.store);
204 let upstream = Arc::clone(&self.upstream);
205 let loaded = Arc::clone(&self.loaded);
206 let inflight = Arc::clone(&self.inflight);
207 let negative = Arc::clone(&self.negative);
208 let download_enabled = self.download_enabled.load(Ordering::Relaxed);
209 let name = name.to_owned();
210 runtime.spawn(async move {
211 match load(store, upstream, &name, &loaded, download_enabled).await {
212 Ok(LoadOutcome::Loaded) => {
213 negative.remove(&name);
214 }
215 Ok(LoadOutcome::Missing) => {
216 negative.insert(name.clone(), ());
217 }
218 Err(e) => {
219 negative.insert(name.clone(), ());
220 tracing::warn!(name, error = %e, "tokenizer load failed");
221 }
222 }
223 inflight.remove(&name);
224 });
225 LoadRequestStatus::Scheduled
226 }
227
228 pub async fn resolve_or_load(&self, name: &str) -> anyhow::Result<Option<Arc<Tokenizer>>> {
230 if let Some(tokenizer) = self.resolve(name) {
231 return Ok(Some(tokenizer));
232 }
233 if self.negative.contains_key(name) {
234 return Ok(None);
235 }
236 match load(
237 Arc::clone(&self.store),
238 Arc::clone(&self.upstream),
239 name,
240 &self.loaded,
241 self.download_enabled.load(Ordering::Relaxed),
242 )
243 .await?
244 {
245 LoadOutcome::Loaded => {
246 self.negative.remove(name);
247 Ok(self.resolve(name))
248 }
249 LoadOutcome::Missing => {
250 self.negative.insert(name.to_owned(), ());
251 Ok(None)
252 }
253 }
254 }
255}
256
257enum LoadOutcome {
258 Loaded,
259 Missing,
260}
261
262async fn load(
264 store: Arc<dyn TokenizerStore>,
265 upstream: Arc<dyn TokenizerClient>,
266 name: &str,
267 loaded: &LoadedMap,
268 download_enabled: bool,
269) -> anyhow::Result<LoadOutcome> {
270 if let Some(bytes) = store.get_tokenizer_vocab(name).await? {
271 let parsed = if bytes.len() > MAX_TOKENIZER_BYTES {
272 Err(anyhow::anyhow!(
273 "persisted vocab exceeds {} bytes",
274 MAX_TOKENIZER_BYTES
275 ))
276 } else {
277 Tokenizer::from_bytes(&bytes).map_err(|e| anyhow::anyhow!("bad persisted vocab: {e}"))
278 };
279 match parsed {
280 Ok(tokenizer) => {
281 loaded.insert(name.to_owned(), Arc::new(tokenizer));
282 return Ok(LoadOutcome::Loaded);
283 }
284 Err(error) => {
285 store
286 .quarantine_tokenizer_vocab(name, &error.to_string())
287 .await?;
288 tracing::warn!(name, error = %error, "persisted tokenizer quarantined");
289 if !download_enabled {
290 return Err(error);
291 }
292 }
293 }
294 }
295 if !download_enabled {
296 return Ok(LoadOutcome::Missing);
297 }
298
299 validate_hf_repo_id(name)?;
300
301 let url = format!("https://huggingface.co/{name}/resolve/main/tokenizer.json");
302 let req = http::Request::builder()
303 .method(http::Method::GET)
304 .uri(&url)
305 .body(Bytes::new())?;
306 let resp = upstream.send(req).await?;
307 anyhow::ensure!(resp.status().is_success(), "HTTP {}", resp.status());
308 let body = resp.into_body();
309 anyhow::ensure!(
310 body.len() <= MAX_TOKENIZER_BYTES,
311 "downloaded vocab exceeds {} bytes",
312 MAX_TOKENIZER_BYTES
313 );
314 let tok = Tokenizer::from_bytes(&body).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
315
316 store.put_tokenizer_vocab(name, &body).await?;
317 loaded.insert(name.to_owned(), Arc::new(tok));
318 tracing::info!(name, "tokenizer downloaded");
319 Ok(LoadOutcome::Loaded)
320}
321
322fn validate_hf_repo_id(name: &str) -> anyhow::Result<()> {
323 anyhow::ensure!(name.len() <= 200, "HF repo id is too long");
324 let parts: Vec<_> = name.split('/').collect();
325 anyhow::ensure!(parts.len() == 2, "HF repo id must be `owner/repository`");
326 for part in parts {
327 anyhow::ensure!(
328 !part.is_empty() && part.len() <= 96,
329 "invalid HF repo segment"
330 );
331 anyhow::ensure!(
332 part.bytes()
333 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')),
334 "invalid character in HF repo id"
335 );
336 anyhow::ensure!(
337 !part.starts_with('.')
338 && !part.starts_with('-')
339 && !part.ends_with('.')
340 && !part.ends_with('-'),
341 "invalid HF repo segment boundary"
342 );
343 anyhow::ensure!(!part.contains(".."), "invalid HF repo traversal sequence");
344 }
345 Ok(())
346}
347
348fn info(name: &str, source: VocabSource, loaded: bool) -> VocabInfo {
349 VocabInfo {
350 name: name.to_owned(),
351 source,
352 loaded,
353 }
354}
355
356#[cfg(test)]
357mod tests {
358 use std::sync::Arc;
359 use std::sync::atomic::{AtomicUsize, Ordering};
360
361 use super::*;
362
363 struct CountingStore(AtomicUsize);
364
365 #[async_trait::async_trait]
366 impl TokenizerStore for CountingStore {
367 async fn list_tokenizer_vocabs(&self) -> anyhow::Result<Vec<String>> {
368 Ok(Vec::new())
369 }
370
371 async fn get_tokenizer_vocab(&self, _: &str) -> anyhow::Result<Option<Vec<u8>>> {
372 self.0.fetch_add(1, Ordering::Relaxed);
373 Ok(None)
374 }
375
376 async fn put_tokenizer_vocab(&self, _: &str, _: &[u8]) -> anyhow::Result<()> {
377 Ok(())
378 }
379 }
380
381 struct NoClient;
382
383 #[async_trait::async_trait]
384 impl TokenizerClient for NoClient {
385 async fn send(&self, _: http::Request<Bytes>) -> anyhow::Result<http::Response<Bytes>> {
386 anyhow::bail!("network should not be used")
387 }
388 }
389
390 #[tokio::test]
391 async fn negative_cache_avoids_repeated_store_misses() {
392 let store = Arc::new(CountingStore(AtomicUsize::new(0)));
393 let registry = TokenizerRegistry::new(store.clone(), Arc::new(NoClient));
394 assert!(registry.resolve_or_load("unknown").await.unwrap().is_none());
395 assert!(registry.resolve_or_load("unknown").await.unwrap().is_none());
396 assert_eq!(store.0.load(Ordering::Relaxed), 1);
397 }
398
399 #[test]
400 fn validates_hugging_face_repo_ids() {
401 assert!(validate_hf_repo_id("owner/model-name").is_ok());
402 assert!(validate_hf_repo_id("owner/model/extra").is_err());
403 assert!(validate_hf_repo_id("../model").is_err());
404 assert!(validate_hf_repo_id("owner/model?revision=main").is_err());
405 }
406}