Skip to main content

gproxy_tokenize/tokenize/
registry.rs

1//! Global HF-tokenizer registry (ยง6.3): bundled deepseek vocab, persisted
2//! vocabs through the [`PersistenceBackend`] (the native database backend uses
3//! BLOB rows), and a fire-and-forget
4//! background hydrate/HF-download path through the shared [`UpstreamClient`].
5//! Native-only (`count-local` feature); tiktoken builtins are handled
6//! directly by [`super::count`] and never live here.
7
8use 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    /// Isolate a persisted tokenizer that cannot be safely parsed. Backends
22    /// may override this to move/delete the bad row; the default is a no-op.
23    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/// Bundled DeepSeek vocab, vendored from `deepseek-ai/DeepSeek-V4-Pro`
34/// (`tokenizer.json`).
35#[cfg(feature = "bundled-fallback")]
36static DEEPSEEK: &[u8] = include_bytes!("../../assets/tokenizers/deepseek-v4-pro.tokenizer.json");
37/// Names the bundled vocab answers to.
38#[cfg(feature = "bundled-fallback")]
39const BUNDLED_NAMES: &[&str] = &["deepseek", "deepseek-v4-pro"];
40
41/// Bundled vocab, parsed AT MOST ONCE per process. Parsing the 6.3MB JSON
42/// costs ~100ms; the `OnceLock` both caches the result and dedupes concurrent
43/// first accesses (losers wait on the same init instead of re-parsing).
44/// `None` is sticky on a parse failure โ€” the asset is compile-time fixed, so
45/// retrying cannot succeed.
46#[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/// Where a vocab comes from.
63#[derive(Debug, Clone, Copy, PartialEq, Eq)]
64pub enum VocabSource {
65    BuiltinTiktoken,
66    Bundled,
67    Downloaded,
68}
69
70/// Listing entry for the admin surface.
71#[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
90/// Global tokenizer registry living on `AppState`.
91pub struct TokenizerRegistry {
92    /// Persisted vocab tier (BLOBs in the native database backend).
93    store: Arc<dyn TokenizerStore>,
94    /// Mirrors `instance_settings.enable_tokenizer_download`.
95    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    /// Builtins + bundled + persisted vocabs (admin surface; async because it
122    /// asks the persistence backend).
123    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    /// memory โ†’ bundled name โ†’ `None`. Persisted/downloaded vocabs only show
147    /// up after a background [`request_load`](Self::request_load) hydrates
148    /// them into memory; a miss here never blocks the request.
149    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    /// Fire-and-forget warm-up of the bundled vocab on the blocking pool, so
165    /// the first count request never pays the parse inline. Call once at boot.
166    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    /// Fire-and-forget load pipeline, deduped per name: hydrate from the
188    /// persistence backend; when absent there, downloads are enabled, and the
189    /// name is an HF repo path (`org/repo`), download
190    /// `hf.co/{name}/resolve/main/tokenizer.json` through the shared upstream
191    /// client and persist it. Never blocks the calling request.
192    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    /// Resolve immediately or wait for persistence/download hydration.
229    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
262/// Hydrate `name` from the store, falling back to an HF download.
263async 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}