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
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub enum VocabSource {
36 BuiltinTiktoken,
37 Bundled,
38 Downloaded,
39}
40
41#[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
51pub struct TokenizerRegistry {
53 store: Arc<dyn TokenizerStore>,
55 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 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 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 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
142async 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}