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 fn vocab_size(&self) -> Option<usize> {
175 let model = self.tokenizer.model();
182 let model_size = model.vocab_size();
183 let extra = self.tokenizer.added_tokens().map_or(0, |added| {
184 added
185 .iter()
186 .filter(|info| model.token_to_id(info.content).is_none())
187 .count()
188 });
189 Some(model_size + extra)
190 }
191
192 fn token_to_id(&self, token: &str) -> Result<Option<TokenIdType>> {
193 Ok(self.tokenizer.token_to_id(token))
194 }
195
196 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
197 let Some(added_tokens) = self.tokenizer.added_tokens() else {
198 return Ok(Vec::new());
199 };
200 let mut ids: Vec<TokenIdType> = added_tokens
201 .iter()
202 .filter_map(|info| info.special.then_some(info.id))
203 .collect();
204 ids.sort_unstable();
205 Ok(ids)
206 }
207
208 fn num_special_tokens_added(&self) -> Result<usize> {
209 Ok(self.tokenizer.post_process(Vec::new(), true).len())
220 }
221}
222
223#[cfg(test)]
224mod tests {
225 use std::sync::Arc;
226
227 use super::*;
228 use crate::{
229 HuggingFaceTokenizer, Tokenizer as TokenizerWrapper, traits::Tokenizer as TokenizerTrait,
230 };
231
232 const TOKENIZER_PATH: &str = concat!(
233 env!("CARGO_MANIFEST_DIR"),
234 "/tests/data/minimal-bpe/tokenizer.json"
235 );
236
237 #[test]
238 fn encode_matches_hugging_face() {
239 let baseten = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
240 let hf = HuggingFaceTokenizer::from_file(TOKENIZER_PATH).unwrap();
241
242 for text in ["Hello, world!", "Hello", " world", "He llo"] {
243 let baseten_ids = baseten.encode(text).unwrap();
244 let hf_ids = hf.encode(text).unwrap();
245 assert_eq!(
246 baseten_ids.token_ids(),
247 hf_ids.token_ids(),
248 "Baseten and Hugging Face must produce identical token IDs for '{text}'"
249 );
250 }
251 }
252
253 #[test]
254 fn batch_encode_matches_sequential_encode() {
255 let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
256 let inputs = ["Hello", " world", "Hello, world!"];
257 let batch = tokenizer.encode_batch(&inputs).unwrap();
258
259 for (encoding, input) in batch.iter().zip(inputs) {
260 let sequential = tokenizer.encode(input).unwrap();
261 assert_eq!(encoding.token_ids(), sequential.token_ids());
262 }
263 }
264
265 #[test]
266 fn encode_decode_roundtrip() {
267 let tokenizer = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
268 let encoding = tokenizer.encode("Hello, world!").unwrap();
269 let decoded = tokenizer.decode(encoding.token_ids(), true).unwrap();
270
271 assert_eq!(decoded.as_str(), "Hello, world!");
272 }
273
274 #[test]
275 fn works_with_decode_stream() {
276 let tokenizer = Arc::new(BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap());
277 let wrapper = TokenizerWrapper::from(tokenizer);
278 let prompt_ids = wrapper.encode("Hello").unwrap().token_ids().to_vec();
279 let continuation_ids = wrapper.encode(", world!").unwrap().token_ids().to_vec();
280 let mut stream = wrapper.decode_stream(&prompt_ids, true);
281 let mut accumulated = String::new();
282
283 for id in &continuation_ids {
284 if let Some(chunk) = stream.step(*id).unwrap() {
285 accumulated.push_str(&chunk);
286 }
287 }
288
289 let mut all_ids = prompt_ids.clone();
290 all_ids.extend_from_slice(&continuation_ids);
291 let full_text: String = wrapper.decode(&all_ids, true).unwrap().into();
292 let prompt_text: String = wrapper.decode(&prompt_ids, true).unwrap().into();
293 assert_eq!(accumulated, full_text[prompt_text.len()..]);
294 }
295
296 #[test]
297 fn prefix_cache_rejects_special_token_post_processing() {
298 let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
299 assert!(plain.validate_prefix_cache().is_ok());
300
301 let with_special_tokens = BasetenTokenizer::from_file(TOKENIZER_PATH)
302 .unwrap()
303 .with_options(TokenizerOptions {
304 add_special_tokens: true,
305 });
306 assert!(with_special_tokens.validate_prefix_cache().is_err());
307 }
308
309 #[test]
310 fn segments_preserve_special_token_trust_boundary() {
311 let temp = tempfile::tempdir().unwrap();
312 let path = temp.path().join("tokenizer.json");
313 let mut json: serde_json::Value =
314 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
315 let vocab = json["model"]["vocab"].as_object_mut().unwrap();
316 vocab.insert("<".to_string(), serde_json::json!(23));
317 vocab.insert(">".to_string(), serde_json::json!(24));
318 vocab.insert("c".to_string(), serde_json::json!(25));
319 json["added_tokens"] = serde_json::json!([{
320 "id": 26,
321 "content": "<ctl>",
322 "single_word": false,
323 "lstrip": false,
324 "rstrip": false,
325 "normalized": false,
326 "special": true
327 }]);
328 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
329
330 let tokenizer: Arc<dyn TokenizerTrait> =
331 Arc::new(BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap());
332 let segments = [
333 EncodeSegment::new("Hello", true),
334 EncodeSegment::new("<ctl>", false),
335 EncodeSegment::new(" world!", true),
336 ];
337
338 let segmented = tokenizer.encode_segments(&segments).unwrap();
339 let flattened = tokenizer.encode("Hello<ctl> world!").unwrap();
340
341 assert_ne!(
342 segmented.token_ids(),
343 flattened.token_ids(),
344 "untrusted control-token-looking text must not become an added token"
345 );
346 assert!(flattened.token_ids().contains(&26));
347 assert!(!segmented.token_ids().contains(&26));
348 }
349
350 #[test]
351 fn segments_honor_add_special_tokens_option() {
352 let temp = tempfile::tempdir().unwrap();
353 let path = temp.path().join("tokenizer.json");
354 let mut json: serde_json::Value =
355 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
356 json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
357 json["added_tokens"] = serde_json::json!([{
358 "id": 23,
359 "content": "<bos>",
360 "single_word": false,
361 "lstrip": false,
362 "rstrip": false,
363 "normalized": false,
364 "special": true
365 }]);
366 json["post_processor"] = serde_json::json!({
367 "type": "TemplateProcessing",
368 "single": [
369 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
370 {"Sequence": {"id": "A", "type_id": 0}}
371 ],
372 "pair": [
373 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
374 {"Sequence": {"id": "A", "type_id": 0}},
375 {"Sequence": {"id": "B", "type_id": 0}}
376 ],
377 "special_tokens": {
378 "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
379 }
380 });
381 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
382 let segments = [EncodeSegment::new("Hello", false)];
383
384 let plain = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
385 let plain_ids = plain
386 .encode_segments(&segments)
387 .unwrap()
388 .token_ids()
389 .to_vec();
390
391 let with_bos = BasetenTokenizer::from_file(path.to_str().unwrap())
392 .unwrap()
393 .with_options(TokenizerOptions {
394 add_special_tokens: true,
395 });
396 assert_eq!(
397 with_bos.encode_segments(&segments).unwrap().token_ids(),
398 [&[23], plain_ids.as_slice()].concat()
399 );
400 }
401
402 #[test]
403 fn num_special_tokens_added_reflects_bos_post_processor() {
404 let temp = tempfile::tempdir().unwrap();
405 let path = temp.path().join("tokenizer.json");
406 let mut json: serde_json::Value =
407 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
408 json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
409 json["added_tokens"] = serde_json::json!([{
410 "id": 23,
411 "content": "<bos>",
412 "single_word": false,
413 "lstrip": false,
414 "rstrip": false,
415 "normalized": false,
416 "special": true
417 }]);
418 json["post_processor"] = serde_json::json!({
419 "type": "TemplateProcessing",
420 "single": [
421 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
422 {"Sequence": {"id": "A", "type_id": 0}}
423 ],
424 "pair": [
425 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
426 {"Sequence": {"id": "A", "type_id": 0}},
427 {"Sequence": {"id": "B", "type_id": 0}}
428 ],
429 "special_tokens": {
430 "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
431 }
432 });
433 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
434
435 let with_bos = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
436 assert_eq!(with_bos.num_special_tokens_added().unwrap(), 1);
437
438 let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
439 assert_eq!(plain.num_special_tokens_added().unwrap(), 0);
440 }
441
442 #[test]
443 fn num_special_tokens_added_is_length_independent_under_sequence_post_processor() {
444 let temp = tempfile::tempdir().unwrap();
445 let path = temp.path().join("tokenizer.json");
446 let mut json: serde_json::Value =
447 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
448 json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
449 json["model"]["vocab"]["<eos>"] = serde_json::json!(24);
450 json["added_tokens"] = serde_json::json!([
451 {
452 "id": 23,
453 "content": "<bos>",
454 "single_word": false,
455 "lstrip": false,
456 "rstrip": false,
457 "normalized": false,
458 "special": true
459 },
460 {
461 "id": 24,
462 "content": "<eos>",
463 "single_word": false,
464 "lstrip": false,
465 "rstrip": false,
466 "normalized": false,
467 "special": true
468 }
469 ]);
470 json["post_processor"] = serde_json::json!({
471 "type": "Sequence",
472 "processors": [
473 {
474 "type": "TemplateProcessing",
475 "single": [
476 {"SpecialToken": {"id": "<bos>", "type_id": 0}},
477 {"Sequence": {"id": "A", "type_id": 0}}
478 ],
479 "pair": [],
480 "special_tokens": {
481 "<bos>": {"id": "<bos>", "ids": [23], "tokens": ["<bos>"]}
482 }
483 },
484 {
485 "type": "TemplateProcessing",
486 "single": [
487 {"Sequence": {"id": "A", "type_id": 0}},
488 {"SpecialToken": {"id": "<eos>", "type_id": 0}}
489 ],
490 "pair": [],
491 "special_tokens": {
492 "<eos>": {"id": "<eos>", "ids": [24], "tokens": ["<eos>"]}
493 }
494 }
495 ]
496 });
497 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
498
499 let tokenizer = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
500 assert_eq!(tokenizer.num_special_tokens_added().unwrap(), 2);
501
502 let with_specials = tokenizer.with_options(TokenizerOptions {
503 add_special_tokens: true,
504 });
505 let plain = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
506
507 for text in ["h", "hello there world"] {
508 let without = plain.encode(text).unwrap().token_ids().len();
509 let with = with_specials.encode(text).unwrap().token_ids().len();
510 assert_eq!(
511 with - without,
512 2,
513 "'{text}' should grow by exactly num_special_tokens_added()"
514 );
515 }
516 }
517
518 #[test]
519 fn merges_config_only_special_tokens() {
520 let temp = tempfile::tempdir().unwrap();
521 let tokenizer_path = temp.path().join("tokenizer.json");
522 std::fs::copy(TOKENIZER_PATH, &tokenizer_path).unwrap();
523 std::fs::write(
524 temp.path().join("tokenizer_config.json"),
525 serde_json::json!({
526 "added_tokens_decoder": {
527 "23": {
528 "content": "<ctl>",
529 "special": true,
530 "single_word": false,
531 "lstrip": false,
532 "rstrip": false,
533 "normalized": false
534 }
535 }
536 })
537 .to_string(),
538 )
539 .unwrap();
540
541 let tokenizer = BasetenTokenizer::from_file(tokenizer_path.to_str().unwrap()).unwrap();
542 let encoding = tokenizer.encode("<ctl>").unwrap();
543 assert_eq!(encoding.token_ids(), &[23]);
544 assert_eq!(tokenizer.decode(&[23], false).unwrap().as_str(), "<ctl>");
545 assert_eq!(tokenizer.decode(&[23], true).unwrap().as_str(), "");
546 }
547
548 #[test]
549 fn vocab_introspection_accessors() {
550 let plain = BasetenTokenizer::from_file(TOKENIZER_PATH).unwrap();
551 assert_eq!(plain.vocab_size(), Some(23));
552 assert_eq!(plain.token_to_id("hello").unwrap(), None);
553 assert_eq!(plain.token_to_id("h").unwrap(), Some(10));
554 assert!(plain.special_token_ids().unwrap().is_empty());
555
556 let temp = tempfile::tempdir().unwrap();
557 let path = temp.path().join("tokenizer.json");
558 let mut json: serde_json::Value =
559 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
560 json["model"]["vocab"]["<bos>"] = serde_json::json!(23);
561 json["added_tokens"] = serde_json::json!([{
562 "id": 23,
563 "content": "<bos>",
564 "single_word": false,
565 "lstrip": false,
566 "rstrip": false,
567 "normalized": false,
568 "special": true
569 }]);
570 std::fs::write(&path, serde_json::to_vec(&json).unwrap()).unwrap();
571
572 let with_added = BasetenTokenizer::from_file(path.to_str().unwrap()).unwrap();
573 assert_eq!(with_added.vocab_size(), Some(24));
574 assert_eq!(with_added.token_to_id("<bos>").unwrap(), Some(23));
575 assert_eq!(with_added.special_token_ids().unwrap(), vec![23]);
576
577 let path2 = temp.path().join("tokenizer2.json");
578 let mut json2: serde_json::Value =
579 serde_json::from_str(&std::fs::read_to_string(TOKENIZER_PATH).unwrap()).unwrap();
580 json2["added_tokens"] = serde_json::json!([{
581 "id": 23,
582 "content": "<extra>",
583 "single_word": false,
584 "lstrip": false,
585 "rstrip": false,
586 "normalized": false,
587 "special": true
588 }]);
589 std::fs::write(&path2, serde_json::to_vec(&json2).unwrap()).unwrap();
590
591 let genuinely_added = BasetenTokenizer::from_file(path2.to_str().unwrap()).unwrap();
592 assert_eq!(genuinely_added.vocab_size(), Some(24));
593 assert_eq!(genuinely_added.token_to_id("<extra>").unwrap(), Some(23));
594 }
595}