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
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
27/// Bundled DeepSeek vocab, vendored from `deepseek-ai/DeepSeek-V4-Pro`
28/// (`tokenizer.json`).
29static DEEPSEEK: &[u8] = include_bytes!("../../assets/tokenizers/deepseek-v4-pro.tokenizer.json");
30/// Names the bundled vocab answers to.
31const BUNDLED_NAMES: &[&str] = &["deepseek", "deepseek-v4-pro"];
32
33/// Where a vocab comes from.
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub enum VocabSource {
36    BuiltinTiktoken,
37    Bundled,
38    Downloaded,
39}
40
41/// Listing entry for the admin surface.
42#[derive(Debug, Clone)]
43pub struct VocabInfo {
44    pub name: String,
45    pub source: VocabSource,
46    pub loaded: bool,
47}
48
49type LoadedMap = Arc<DashMap<String, Arc<Tokenizer>>>;
50
51/// Global tokenizer registry living on `AppState`.
52pub struct TokenizerRegistry {
53    /// Persisted vocab tier (BLOBs in the native database backend).
54    store: Arc<dyn TokenizerStore>,
55    /// Mirrors `instance_settings.enable_tokenizer_download`.
56    download_enabled: AtomicBool,
57    upstream: Arc<dyn TokenizerClient>,
58    loaded: LoadedMap,
59    inflight: Arc<DashMap<String, ()>>,
60}
61
62impl TokenizerRegistry {
63    pub fn new(store: Arc<dyn TokenizerStore>, upstream: Arc<dyn TokenizerClient>) -> Self {
64        Self {
65            store,
66            download_enabled: AtomicBool::new(false),
67            upstream,
68            loaded: Arc::new(DashMap::new()),
69            inflight: Arc::new(DashMap::new()),
70        }
71    }
72
73    pub fn set_download_enabled(&self, on: bool) {
74        self.download_enabled.store(on, Ordering::Relaxed);
75    }
76
77    /// Builtins + bundled + persisted vocabs (admin surface; async because it
78    /// asks the persistence backend).
79    pub async fn list(&self) -> Vec<VocabInfo> {
80        let mut out = vec![
81            info("o200k_base", VocabSource::BuiltinTiktoken, true),
82            info("cl100k_base", VocabSource::BuiltinTiktoken, true),
83            info(
84                BUNDLED_NAMES[0],
85                VocabSource::Bundled,
86                self.loaded.contains_key(BUNDLED_NAMES[0]),
87            ),
88        ];
89        match self.store.list_tokenizer_vocabs().await {
90            Ok(names) => {
91                for name in names {
92                    let loaded = self.loaded.contains_key(&name);
93                    out.push(info(&name, VocabSource::Downloaded, loaded));
94                }
95            }
96            Err(e) => tracing::warn!(error = %e, "listing persisted tokenizer vocabs failed"),
97        }
98        out
99    }
100
101    /// memory → bundled name → `None`. Persisted/downloaded vocabs only show
102    /// up after a background [`request_load`](Self::request_load) hydrates
103    /// them into memory; a miss here never blocks the request.
104    pub fn resolve(&self, name: &str) -> Option<Arc<Tokenizer>> {
105        if let Some(t) = self.loaded.get(name) {
106            return Some(Arc::clone(&t));
107        }
108        if BUNDLED_NAMES.contains(&name) {
109            let tok = Arc::new(Tokenizer::from_bytes(DEEPSEEK).ok()?);
110            for n in BUNDLED_NAMES {
111                self.loaded.insert((*n).to_owned(), Arc::clone(&tok));
112            }
113            return Some(tok);
114        }
115        None
116    }
117
118    /// Fire-and-forget load pipeline, deduped per name: hydrate from the
119    /// persistence backend; when absent there, downloads are enabled, and the
120    /// name is an HF repo path (`org/repo`), download
121    /// `hf.co/{name}/resolve/main/tokenizer.json` through the shared upstream
122    /// client and persist it. Never blocks the calling request.
123    pub fn request_load(&self, name: &str) {
124        if self.inflight.insert(name.to_owned(), ()).is_some() {
125            return;
126        }
127        let store = Arc::clone(&self.store);
128        let upstream = Arc::clone(&self.upstream);
129        let loaded = Arc::clone(&self.loaded);
130        let inflight = Arc::clone(&self.inflight);
131        let download_enabled = self.download_enabled.load(Ordering::Relaxed);
132        let name = name.to_owned();
133        tokio::spawn(async move {
134            if let Err(e) = load(store, upstream, &name, &loaded, download_enabled).await {
135                tracing::warn!(name, error = %e, "tokenizer load failed");
136            }
137            inflight.remove(&name);
138        });
139    }
140}
141
142/// Hydrate `name` from the store, falling back to an HF download.
143async fn load(
144    store: Arc<dyn TokenizerStore>,
145    upstream: Arc<dyn TokenizerClient>,
146    name: &str,
147    loaded: &LoadedMap,
148    download_enabled: bool,
149) -> anyhow::Result<()> {
150    if let Some(bytes) = store.get_tokenizer_vocab(name).await? {
151        let tok = Tokenizer::from_bytes(&bytes).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
152        loaded.insert(name.to_owned(), Arc::new(tok));
153        return Ok(());
154    }
155    if !download_enabled || !name.contains('/') {
156        return Ok(());
157    }
158
159    let url = format!("https://huggingface.co/{name}/resolve/main/tokenizer.json");
160    let req = http::Request::builder()
161        .method(http::Method::GET)
162        .uri(&url)
163        .body(Bytes::new())?;
164    let resp = upstream.send(req).await?;
165    anyhow::ensure!(resp.status().is_success(), "HTTP {}", resp.status());
166    let body = resp.into_body();
167    let tok = Tokenizer::from_bytes(&body).map_err(|e| anyhow::anyhow!("bad vocab: {e}"))?;
168
169    store.put_tokenizer_vocab(name, &body).await?;
170    loaded.insert(name.to_owned(), Arc::new(tok));
171    tracing::info!(name, "tokenizer downloaded");
172    Ok(())
173}
174
175fn info(name: &str, source: VocabSource, loaded: bool) -> VocabInfo {
176    VocabInfo {
177        name: name.to_owned(),
178        source,
179        loaded,
180    }
181}