use paladin_ports::output::garrison_port::GarrisonError;
use std::collections::HashMap;
use std::sync::RwLock;
use tiktoken_rs::{CoreBPE, get_bpe_from_model};
pub trait TokenCounter: Send + Sync {
fn count_tokens(&self, text: &str) -> Result<u32, GarrisonError>;
fn model_name(&self) -> &str;
}
pub struct TiktokenCounter {
bpe: CoreBPE,
model_name: String,
cache: RwLock<HashMap<String, u32>>,
}
impl TiktokenCounter {
pub fn new(model_name: &str) -> Result<Self, GarrisonError> {
let bpe = get_bpe_from_model(model_name).map_err(|e| {
GarrisonError::TokenizationError(format!(
"Failed to initialize tokenizer for model '{}': {}",
model_name, e
))
})?;
Ok(Self {
bpe,
model_name: model_name.to_string(),
cache: RwLock::new(HashMap::new()),
})
}
pub fn clear_cache(&self) {
if let Ok(mut cache) = self.cache.write() {
cache.clear();
}
}
pub fn cache_size(&self) -> usize {
self.cache.read().map(|c| c.len()).unwrap_or(0)
}
}
impl TokenCounter for TiktokenCounter {
fn count_tokens(&self, text: &str) -> Result<u32, GarrisonError> {
if let Ok(cache) = self.cache.read()
&& let Some(&count) = cache.get(text)
{
return Ok(count);
}
let tokens = self.bpe.encode_with_special_tokens(text);
let count = tokens.len() as u32;
if let Ok(mut cache) = self.cache.write() {
if cache.len() >= 1000 {
cache.clear();
}
cache.insert(text.to_string(), count);
}
Ok(count)
}
fn model_name(&self) -> &str {
&self.model_name
}
}
pub struct TokenCounterFactory;
impl TokenCounterFactory {
pub fn for_model(model_name: &str) -> Result<Box<dyn TokenCounter>, GarrisonError> {
let counter = TiktokenCounter::new(model_name)?;
Ok(Box::new(counter))
}
pub fn supported_models() -> Vec<&'static str> {
vec![
"gpt-4",
"gpt-4-32k",
"gpt-4-turbo",
"gpt-4o",
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"text-embedding-ada-002",
"text-embedding-3-small",
"text-embedding-3-large",
]
}
pub fn is_supported(model_name: &str) -> bool {
Self::supported_models().contains(&model_name)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tiktoken_counter_creation() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
assert_eq!(counter.model_name(), "gpt-4");
}
#[test]
fn test_tiktoken_counter_unsupported_model() {
let result = TiktokenCounter::new("unsupported-model-xyz");
assert!(result.is_err());
}
#[test]
fn test_count_tokens_simple() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
let count = counter.count_tokens("Hello, world!").unwrap();
assert!(count > 0);
assert!(count < 10); }
#[test]
fn test_count_tokens_empty_string() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
let count = counter.count_tokens("").unwrap();
assert_eq!(count, 0);
}
#[test]
fn test_count_tokens_caching() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
let text = "This is a test message for caching.";
let count1 = counter.count_tokens(text).unwrap();
let count2 = counter.count_tokens(text).unwrap();
assert_eq!(count1, count2);
assert_eq!(counter.cache_size(), 1);
}
#[test]
fn test_cache_clearing() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
counter.count_tokens("Test 1").unwrap();
counter.count_tokens("Test 2").unwrap();
assert_eq!(counter.cache_size(), 2);
counter.clear_cache();
assert_eq!(counter.cache_size(), 0);
}
#[test]
fn test_factory_for_model() {
let counter = TokenCounterFactory::for_model("gpt-4").unwrap();
let count = counter.count_tokens("Test").unwrap();
assert!(count > 0);
}
#[test]
fn test_factory_supported_models() {
let models = TokenCounterFactory::supported_models();
assert!(models.contains(&"gpt-4"));
assert!(models.contains(&"gpt-3.5-turbo"));
}
#[test]
fn test_factory_is_supported() {
assert!(TokenCounterFactory::is_supported("gpt-4"));
assert!(TokenCounterFactory::is_supported("gpt-3.5-turbo"));
assert!(!TokenCounterFactory::is_supported("unknown-model"));
}
#[test]
fn test_factory_error_on_unknown_model() {
let result = TokenCounterFactory::for_model("not-a-real-model-xyz");
assert!(result.is_err(), "Expected error for unsupported model");
}
#[test]
fn test_longer_text_token_count() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
let long_text = "This is a longer piece of text that should result in more tokens. \
It contains multiple sentences and should demonstrate that the token \
counter is working correctly for larger inputs.";
let count = counter.count_tokens(long_text).unwrap();
assert!(count > 20); }
#[test]
fn test_special_characters() {
let counter = TiktokenCounter::new("gpt-4").unwrap();
let text = "Hello! 你好! مرحبا! 👋";
let count = counter.count_tokens(text).unwrap();
assert!(count > 0);
}
#[test]
fn test_multiple_models() {
let gpt4 = TiktokenCounter::new("gpt-4").unwrap();
let gpt35 = TiktokenCounter::new("gpt-3.5-turbo").unwrap();
let text = "Test message";
let count_gpt4 = gpt4.count_tokens(text).unwrap();
let count_gpt35 = gpt35.count_tokens(text).unwrap();
assert!(count_gpt4 > 0);
assert!(count_gpt35 > 0);
}
}