1use std::path::Path;
11
12use super::{
13 EncodeSegment, Encoding, Error, Result, TokenIdType, TokenizerOptions,
14 traits::{DecodeResult, Decoder, Encoder, Tokenizer},
15};
16
17pub struct BasetenTokenizer {
19 tokenizer: basetenkenizer::Tokenizer,
20 options: TokenizerOptions,
21}
22
23impl BasetenTokenizer {
24 pub fn from_file(path: &str) -> Result<Self> {
26 let path = Path::new(path);
27 let raw = std::fs::read_to_string(path)
28 .map_err(|e| Error::msg(format!("Error reading Baseten tokenizer: {e}")))?;
29 let mut json: serde_json::Value = serde_json::from_str(&raw)
30 .map_err(|e| Error::msg(format!("Error parsing Baseten tokenizer: {e}")))?;
31 if let Some(parent) = path.parent() {
32 merge_special_tokens_from_config(&mut json, parent);
33 }
34 let tokenizer = basetenkenizer::Tokenizer::from_json(json)
35 .map_err(|e| Error::msg(format!("Error loading Baseten tokenizer: {e}")))?;
36 Ok(Self {
37 tokenizer,
38 options: TokenizerOptions::default(),
39 })
40 }
41}
42
43fn merge_special_tokens_from_config(json: &mut serde_json::Value, model_dir: &Path) {
44 let config_path = model_dir.join("tokenizer_config.json");
45 let Ok(raw) = std::fs::read_to_string(&config_path) else {
46 return;
47 };
48 let config: serde_json::Value = match serde_json::from_str(&raw) {
49 Ok(value) => value,
50 Err(error) => {
51 tracing::debug!(
52 target: "tokenizer",
53 path = %config_path.display(),
54 error = %error,
55 "tokenizer_config.json parse failed; skipping special-token merge"
56 );
57 return;
58 }
59 };
60 let Some(decoder) = config
61 .get("added_tokens_decoder")
62 .and_then(serde_json::Value::as_object)
63 else {
64 return;
65 };
66
67 if json.get("added_tokens").is_none() {
68 json["added_tokens"] = serde_json::json!([]);
69 }
70 let Some(added_tokens) = json
71 .get_mut("added_tokens")
72 .and_then(serde_json::Value::as_array_mut)
73 else {
74 return;
75 };
76
77 for (id, spec) in decoder {
78 let Some(id) = id.parse::<u32>().ok() else {
79 continue;
80 };
81 let Some(spec) = spec.as_object() else {
82 continue;
83 };
84 if spec.get("special").and_then(serde_json::Value::as_bool) != Some(true) {
85 continue;
86 }
87 let Some(content) = spec
88 .get("content")
89 .and_then(serde_json::Value::as_str)
90 .filter(|content| !content.is_empty())
91 else {
92 continue;
93 };
94
95 if let Some(existing) = added_tokens
96 .iter_mut()
97 .find(|token| token.get("content").and_then(serde_json::Value::as_str) == Some(content))
98 {
99 existing["special"] = serde_json::Value::Bool(true);
100 continue;
101 }
102
103 let mut token = serde_json::Map::from_iter([
104 ("id".to_string(), serde_json::json!(id)),
105 ("content".to_string(), serde_json::json!(content)),
106 ("special".to_string(), serde_json::Value::Bool(true)),
107 ]);
108 for field in ["single_word", "lstrip", "rstrip", "normalized"] {
109 token.insert(
110 field.to_string(),
111 serde_json::Value::Bool(
112 spec.get(field)
113 .and_then(serde_json::Value::as_bool)
114 .unwrap_or(false),
115 ),
116 );
117 }
118 added_tokens.push(serde_json::Value::Object(token));
119 }
120}
121
122impl Encoder for BasetenTokenizer {
123 fn encode(&self, input: &str) -> Result<Encoding> {
124 let ids = self
125 .tokenizer
126 .encode_with_special_tokens(input, self.options.add_special_tokens)
127 .map_err(|e| Error::msg(format!("Baseten tokenizer encode error: {e}")))?;
128 Ok(Encoding::Sp(ids))
129 }
130
131 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>> {
132 self.tokenizer
133 .encode_batch(inputs, self.options.add_special_tokens)
134 .map(|ids| ids.into_iter().map(Encoding::Sp).collect())
135 .map_err(|e| Error::msg(format!("Baseten tokenizer batch encode error: {e}")))
136 }
137
138 fn encode_segments(&self, segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
139 let segments = segments
140 .iter()
141 .map(|segment| (segment.text, segment.allow_special));
142 let ids = self
143 .tokenizer
144 .encode_segments_tiktoken_safe(segments, self.options.add_special_tokens)
145 .map_err(|e| Error::msg(format!("Baseten tokenizer segment encode error: {e}")))?;
146 Ok(Encoding::Sp(ids))
147 }
148}
149
150impl Decoder for BasetenTokenizer {
151 fn decode(&self, token_ids: &[TokenIdType], skip_special_tokens: bool) -> Result<DecodeResult> {
152 self.tokenizer
153 .decode(token_ids, skip_special_tokens)
154 .map(DecodeResult::from)
155 .map_err(|e| Error::msg(format!("Baseten tokenizer decode error: {e}")))
156 }
157}
158
159impl Tokenizer for BasetenTokenizer {
160 fn validate_prefix_cache(&self) -> Result<()> {
161 if self.options.add_special_tokens {
162 return Err(Error::msg(
163 "Baseten tokenizers configured with add_special_tokens=true must remain uncached",
164 ));
165 }
166 Ok(())
167 }
168
169 fn with_options(mut self, options: TokenizerOptions) -> Self {
170 self.options = options;
171 self
172 }
173}
174
175#[cfg(test)]
176mod tests {
177 use std::sync::Arc;
178
179 use super::*;
180 use crate::{
181 HuggingFaceTokenizer, Tokenizer as TokenizerWrapper, traits::Tokenizer as TokenizerTrait,
182 };
183
184 const TOKENIZER_PATH: &str = concat!(
185 env!("CARGO_MANIFEST_DIR"),
186 "/tests/data/minimal-bpe/tokenizer.json"
187 );
188
189 #[test]
190 fn encode_matches_hugging_face() {
191 let baseten = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
192 let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
193
194 for text in ["Hello, world!", "Hello", " world", "He llo"] {
195 let baseten_ids = baseten.encode(text).unwrap();
196 let hf_ids = hf.encode(text).unwrap();
197 assert_eq!(
198 baseten_ids.token_ids(),
199 hf_ids.token_ids(),
200 "Baseten and Hugging Face must produce identical token IDs for '{text}'"
201 );
202 }
203 }
204
205 #[test]
206 fn batch_encode_matches_sequential_encode() {
207 let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
208 let inputs = ["Hello", " world", "Hello, world!"];
209 let batch = tokenizer.encode_batch(&inputs).unwrap();
210
211 for (encoding, input) in batch.iter().zip(inputs) {
212 let sequential = tokenizer.encode(input).unwrap();
213 assert_eq!(encoding.token_ids(), sequential.token_ids());
214 }
215 }
216
217 #[test]
218 fn encode_decode_roundtrip() {
219 let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
220 let encoding = tokenizer.encode("Hello, world!").unwrap();
221 let decoded = tokenizer.decode(encoding.token_ids(), true).unwrap();
222
223 assert_eq!(decoded.as_str(), "Hello, world!");
224 }
225
226 #[test]
227 fn works_with_decode_stream() {
228 let tokenizer = Arc::new(BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap());
229 let wrapper = TokenizerWrapper::from(tokenizer);
230 let prompt_ids = wrapper.encode("Hello").unwrap().token_ids().to_vec();
231 let continuation_ids = wrapper.encode(", world!").unwrap().token_ids().to_vec();
232 let mut stream = wrapper.decode_stream(&prompt_ids, true);
233 let mut accumulated = String::new();
234
235 for id in &continuation_ids {
236 if let Some(chunk) = stream.step(*id).unwrap() {
237 accumulated.push_str(&chunk);
238 }
239 }
240
241 let mut all_ids = prompt_ids.clone();
242 all_ids.extend_from_slice(&continuation_ids);
243 let full_text: String = wrapper.decode(&all_ids, true).unwrap().into();
244 let prompt_text: String = wrapper.decode(&prompt_ids, true).unwrap().into();
245 assert_eq!(accumulated, full_text[prompt_text.len()..]);
246 }
247
248 #[test]
249 fn prefix_cache_rejects_special_token_post_processing() {
250 let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
251 assert!(plain.validate_prefix_cache().is_ok());
252
253 let with_special_tokens = BasetenTokenizer::from_file(TOKENIZER_PATH)
254 .unwrap()
255 .with_options(TokenizerOptions {
256 add_special_tokens: true,
257 });
258 assert!(with_special_tokens.validate_prefix_cache().is_err());
259 }
260
261 #[test]
262 fn segments_preserve_special_token_trust_boundary() {
263 let temp = tempfile::tempdir().unwrap();
264 let path = temp.path().join("tokenizer.json");
265 let mut json: serde_json::Value =
266 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
267 let vocab = json["model"]["vocab"].as_object_mut().unwrap();
268 vocab.insert("<".to_string(), serde_json::json!(23));
269 vocab.insert(">".to_string(), serde_json::json!(24));
270 vocab.insert("c".to_string(), serde_json::json!(25));
271 json["added_tokens"] = serde_json::json!([{
272 "id": 26,
273 "content": "<ctl>",
274 "single_word": false,
275 "lstrip": false,
276 "rstrip": false,
277 "normalized": false,
278 "special": true
279 }]);
280 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
281
282 let tokenizer: Arc<dyn TokenizerTrait> =
283 Arc::new(BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap());
284 let segments = [
285 EncodeSegment::new("Hello", true),
286 EncodeSegment::new("<ctl>", false),
287 EncodeSegment::new(" world!", true),
288 ];
289
290 let segmented = tokenizer.encode_segments(&segments).unwrap();
291 let flattened = tokenizer.encode("Hello<ctl> world!").unwrap();
292
293 assert_ne!(
294 segmented.token_ids(),
295 flattened.token_ids(),
296 "untrusted control-token-looking text must not become an added token"
297 );
298 assert!(flattened.token_ids().contains(&26));
299 assert!(!segmented.token_ids().contains(&26));
300 }
301
302 #[test]
303 fn segments_honor_add_special_tokens_option() {
304 let temp = tempfile::tempdir().unwrap();
305 let path = temp.path().join("tokenizer.json");
306 let mut json: serde_json::Value =
307 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
308 json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
309 json["added_tokens"] = serde_json::json!([{
310 "id": 23,
311 "content": "<bos>",
312 "single_word": false,
313 "lstrip": false,
314 "rstrip": false,
315 "normalized": false,
316 "special": true
317 }]);
318 json["post_processor"] = serde_json::json!({
319 "type": "TemplateProcessing",
320 "single": [
321 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
322 {"Sequence": {"id": "A", "type_id": 0}}
323 ],
324 "pair": [
325 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
326 {"Sequence": {"id": "A", "type_id": 0}},
327 {"Sequence": {"id": "B", "type_id": 0}}
328 ],
329 "special_tokens": {
330 "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
331 }
332 });
333 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
334 let segments = [EncodeSegment::new("Hello", false)];
335
336 let plain = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
337 let plain_ids = plain
338 .encode_segments(&segments)
339 .unwrap()
340 .token_ids()
341 .to_vec();
342
343 let with_bos = BasetenTokenizer::from_file(path.to_str().unwrap())
344 .unwrap()
345 .with_options(TokenizerOptions {
346 add_special_tokens: true,
347 });
348 assert_eq!(
349 with_bos.encode_segments(&segments).unwrap().token_ids(),
350 [&[23], plain_ids.as_slice()].concat()
351 );
352 }
353
354 #[test]
355 fn merges_config_only_special_tokens() {
356 let temp = tempfile::tempdir().unwrap();
357 let tokenizer_path = temp.path().join("tokenizer.json");
358 std::fs::copy(TOKENIZER_PATH, &tokenizer_path).unwrap();
359 std::fs::write(
360 temp.path().join("tokenizer_config.json"),
361 serde_json::json!({
362 "added_tokens_decoder": {
363 "23": {
364 "content": "<ctl>",
365 "special": true,
366 "single_word": false,
367 "lstrip": false,
368 "rstrip": false,
369 "normalized": false
370 }
371 }
372 })
373 .to_string(),
374 )
375 .unwrap();
376
377 let tokenizer = BasetenTokenizer::from_file(tokenizer_path.to_str().unwrap()).unwrap();
378 let encoding = tokenizer.encode("<ctl>").unwrap();
379 assert_eq!(encoding.token_ids(), &[23]);
380 assert_eq!(tokenizer.decode(&[23], false).unwrap().as_str(), "<ctl>");
381 assert_eq!(tokenizer.decode(&[23], true).unwrap().as_str(), "");
382 }
383}