use tiktoken_rs::{CoreBPE, cl100k_base, o200k_base};
pub enum TokenizerBackend {
Cl100k,
O200k,
}
pub struct Tokenizer {
bpe: CoreBPE,
}
impl Tokenizer {
pub fn new(backend: TokenizerBackend) -> Self {
let bpe = match backend {
TokenizerBackend::Cl100k => cl100k_base().expect("failed to load cl100k_base"),
TokenizerBackend::O200k => o200k_base().expect("failed to load o200k_base"),
};
Self { bpe }
}
pub fn count(&self, text: &str) -> u32 {
self.bpe.encode_ordinary(text).len() as u32
}
pub fn count_batch(&self, texts: &[&str]) -> Vec<u32> {
texts.iter().map(|t| self.count(t)).collect()
}
pub fn truncate<'a>(&self, text: &'a str, max_tokens: u32) -> &'a str {
let tokens = self.bpe.encode_ordinary(text);
if tokens.len() as u32 <= max_tokens {
return text;
}
let target = &tokens[..max_tokens as usize];
let decoded = self.bpe.decode(target.to_vec()).unwrap_or_default();
let byte_len = decoded.len().min(text.len());
let mut end = byte_len;
while end > 0 && !text.is_char_boundary(end) {
end -= 1;
}
&text[..end]
}
}
pub fn count_tokens(text: &str) -> u32 {
let t = Tokenizer::new(TokenizerBackend::Cl100k);
t.count(text)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn counts_tokens() {
let t = Tokenizer::new(TokenizerBackend::Cl100k);
let count = t.count("Hello, world!");
assert!(count > 0);
assert!(count < 10);
}
#[test]
fn truncate_respects_budget() {
let t = Tokenizer::new(TokenizerBackend::Cl100k);
let text = "The quick brown fox jumps over the lazy dog. ".repeat(100);
let truncated = t.truncate(&text, 10);
assert!(t.count(truncated) <= 10);
assert!(!truncated.is_empty());
}
#[test]
fn batch_counting() {
let t = Tokenizer::new(TokenizerBackend::Cl100k);
let counts = t.count_batch(&["hello", "world", "foo bar baz"]);
assert_eq!(counts.len(), 3);
assert!(counts.iter().all(|&c| c > 0));
}
}