Make token counter safer (#4924)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -8,25 +8,28 @@ use tokio::sync::OnceCell;
|
|||||||
|
|
||||||
use crate::conversation::message::Message;
|
use crate::conversation::message::Message;
|
||||||
|
|
||||||
// Global tokenizer instance to avoid repeated initialization
|
|
||||||
static TOKENIZER: OnceCell<Arc<CoreBPE>> = OnceCell::const_new();
|
static TOKENIZER: OnceCell<Arc<CoreBPE>> = OnceCell::const_new();
|
||||||
|
|
||||||
// Cache size limits to prevent unbounded growth
|
|
||||||
const MAX_TOKEN_CACHE_SIZE: usize = 10_000;
|
const MAX_TOKEN_CACHE_SIZE: usize = 10_000;
|
||||||
|
|
||||||
/// Async token counter with caching capabilities
|
// token use for various bits of a tool calls:
|
||||||
|
const FUNC_INIT: usize = 7;
|
||||||
|
const PROP_INIT: usize = 3;
|
||||||
|
const PROP_KEY: usize = 3;
|
||||||
|
const ENUM_INIT: isize = -3;
|
||||||
|
const ENUM_ITEM: usize = 3;
|
||||||
|
const FUNC_END: usize = 12;
|
||||||
|
|
||||||
pub struct AsyncTokenCounter {
|
pub struct AsyncTokenCounter {
|
||||||
tokenizer: Arc<CoreBPE>,
|
tokenizer: Arc<CoreBPE>,
|
||||||
token_cache: Arc<DashMap<u64, usize>>, // content hash -> token count
|
token_cache: Arc<DashMap<u64, usize>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Legacy synchronous token counter for backward compatibility
|
|
||||||
pub struct TokenCounter {
|
pub struct TokenCounter {
|
||||||
tokenizer: Arc<CoreBPE>,
|
tokenizer: Arc<CoreBPE>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl AsyncTokenCounter {
|
impl AsyncTokenCounter {
|
||||||
/// Creates a new async token counter with caching
|
|
||||||
pub async fn new() -> Result<Self, String> {
|
pub async fn new() -> Result<Self, String> {
|
||||||
let tokenizer = get_tokenizer().await?;
|
let tokenizer = get_tokenizer().await?;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
@@ -35,25 +38,19 @@ impl AsyncTokenCounter {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Count tokens with optimized caching
|
|
||||||
pub fn count_tokens(&self, text: &str) -> usize {
|
pub fn count_tokens(&self, text: &str) -> usize {
|
||||||
// Use faster AHash for better performance
|
|
||||||
let mut hasher = AHasher::default();
|
let mut hasher = AHasher::default();
|
||||||
text.hash(&mut hasher);
|
text.hash(&mut hasher);
|
||||||
let hash = hasher.finish();
|
let hash = hasher.finish();
|
||||||
|
|
||||||
// Check cache first
|
|
||||||
if let Some(count) = self.token_cache.get(&hash) {
|
if let Some(count) = self.token_cache.get(&hash) {
|
||||||
return *count;
|
return *count;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Compute and cache result with size management
|
|
||||||
let tokens = self.tokenizer.encode_with_special_tokens(text);
|
let tokens = self.tokenizer.encode_with_special_tokens(text);
|
||||||
let count = tokens.len();
|
let count = tokens.len();
|
||||||
|
|
||||||
// Manage cache size to prevent unbounded growth
|
|
||||||
if self.token_cache.len() >= MAX_TOKEN_CACHE_SIZE {
|
if self.token_cache.len() >= MAX_TOKEN_CACHE_SIZE {
|
||||||
// Simple eviction: remove a random entry
|
|
||||||
if let Some(entry) = self.token_cache.iter().next() {
|
if let Some(entry) = self.token_cache.iter().next() {
|
||||||
let old_hash = *entry.key();
|
let old_hash = *entry.key();
|
||||||
self.token_cache.remove(&old_hash);
|
self.token_cache.remove(&old_hash);
|
||||||
@@ -64,20 +61,11 @@ impl AsyncTokenCounter {
|
|||||||
count
|
count
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Count tokens for tools with optimized string handling
|
|
||||||
pub fn count_tokens_for_tools(&self, tools: &[Tool]) -> usize {
|
pub fn count_tokens_for_tools(&self, tools: &[Tool]) -> usize {
|
||||||
// Token counts for different function components
|
|
||||||
let func_init = 7; // Tokens for function initialization
|
|
||||||
let prop_init = 3; // Tokens for properties initialization
|
|
||||||
let prop_key = 3; // Tokens for each property key
|
|
||||||
let enum_init: isize = -3; // Tokens adjustment for enum list start
|
|
||||||
let enum_item = 3; // Tokens for each enum item
|
|
||||||
let func_end = 12; // Tokens for function ending
|
|
||||||
|
|
||||||
let mut func_token_count = 0;
|
let mut func_token_count = 0;
|
||||||
if !tools.is_empty() {
|
if !tools.is_empty() {
|
||||||
for tool in tools {
|
for tool in tools {
|
||||||
func_token_count += func_init;
|
func_token_count += FUNC_INIT;
|
||||||
let name = &tool.name;
|
let name = &tool.name;
|
||||||
let description = &tool
|
let description = &tool
|
||||||
.description
|
.description
|
||||||
@@ -86,7 +74,6 @@ impl AsyncTokenCounter {
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.trim_end_matches('.');
|
.trim_end_matches('.');
|
||||||
|
|
||||||
// Note: the separator (:) is likely tokenized with adjacent tokens, so we use original approach for accuracy
|
|
||||||
let line = format!("{}:{}", name, description);
|
let line = format!("{}:{}", name, description);
|
||||||
func_token_count += self.count_tokens(&line);
|
func_token_count += self.count_tokens(&line);
|
||||||
|
|
||||||
@@ -94,26 +81,27 @@ impl AsyncTokenCounter {
|
|||||||
tool.input_schema.get("properties")
|
tool.input_schema.get("properties")
|
||||||
{
|
{
|
||||||
if !properties.is_empty() {
|
if !properties.is_empty() {
|
||||||
func_token_count += prop_init;
|
func_token_count += PROP_INIT;
|
||||||
for (key, value) in properties {
|
for (key, value) in properties {
|
||||||
func_token_count += prop_key;
|
func_token_count += PROP_KEY;
|
||||||
let p_name = key;
|
let p_name = key;
|
||||||
let p_type = value["type"].as_str().unwrap_or("");
|
let p_type = value.get("type").and_then(|v| v.as_str()).unwrap_or("");
|
||||||
let p_desc = value["description"]
|
let p_desc = value
|
||||||
.as_str()
|
.get("description")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
.unwrap_or("")
|
.unwrap_or("")
|
||||||
.trim_end_matches('.');
|
.trim_end_matches('.');
|
||||||
|
|
||||||
// Note: separators are tokenized with adjacent tokens, keep original for accuracy
|
|
||||||
let line = format!("{}:{}:{}", p_name, p_type, p_desc);
|
let line = format!("{}:{}:{}", p_name, p_type, p_desc);
|
||||||
func_token_count += self.count_tokens(&line);
|
func_token_count += self.count_tokens(&line);
|
||||||
|
|
||||||
if let Some(enum_values) = value["enum"].as_array() {
|
if let Some(enum_values) = value.get("enum").and_then(|v| v.as_array())
|
||||||
|
{
|
||||||
func_token_count =
|
func_token_count =
|
||||||
func_token_count.saturating_add_signed(enum_init);
|
func_token_count.saturating_add_signed(ENUM_INIT);
|
||||||
for item in enum_values {
|
for item in enum_values {
|
||||||
if let Some(item_str) = item.as_str() {
|
if let Some(item_str) = item.as_str() {
|
||||||
func_token_count += enum_item;
|
func_token_count += ENUM_ITEM;
|
||||||
func_token_count += self.count_tokens(item_str);
|
func_token_count += self.count_tokens(item_str);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -122,13 +110,12 @@ impl AsyncTokenCounter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
func_token_count += func_end;
|
func_token_count += FUNC_END;
|
||||||
}
|
}
|
||||||
|
|
||||||
func_token_count
|
func_token_count
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Count chat tokens (using cached count_tokens)
|
|
||||||
pub fn count_chat_tokens(
|
pub fn count_chat_tokens(
|
||||||
&self,
|
&self,
|
||||||
system_prompt: &str,
|
system_prompt: &str,
|
||||||
@@ -148,13 +135,13 @@ impl AsyncTokenCounter {
|
|||||||
if let Some(content_text) = content.as_text() {
|
if let Some(content_text) = content.as_text() {
|
||||||
num_tokens += self.count_tokens(content_text);
|
num_tokens += self.count_tokens(content_text);
|
||||||
} else if let Some(tool_request) = content.as_tool_request() {
|
} else if let Some(tool_request) = content.as_tool_request() {
|
||||||
let tool_call = tool_request.tool_call.as_ref().unwrap();
|
if let Ok(tool_call) = tool_request.tool_call.as_ref() {
|
||||||
// Note: separators are tokenized with adjacent tokens, keep original for accuracy
|
let text = format!(
|
||||||
let text = format!(
|
"{}:{}:{:?}",
|
||||||
"{}:{}:{:?}",
|
tool_request.id, tool_call.name, tool_call.arguments
|
||||||
tool_request.id, tool_call.name, tool_call.arguments
|
);
|
||||||
);
|
num_tokens += self.count_tokens(&text);
|
||||||
num_tokens += self.count_tokens(&text);
|
}
|
||||||
} else if let Some(tool_response_text) = content.as_tool_response_text() {
|
} else if let Some(tool_response_text) = content.as_tool_response_text() {
|
||||||
num_tokens += self.count_tokens(&tool_response_text);
|
num_tokens += self.count_tokens(&tool_response_text);
|
||||||
}
|
}
|
||||||
@@ -170,7 +157,6 @@ impl AsyncTokenCounter {
|
|||||||
num_tokens
|
num_tokens
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Count everything including resources (using cached count_tokens)
|
|
||||||
pub fn count_everything(
|
pub fn count_everything(
|
||||||
&self,
|
&self,
|
||||||
system_prompt: &str,
|
system_prompt: &str,
|
||||||
@@ -188,7 +174,6 @@ impl AsyncTokenCounter {
|
|||||||
num_tokens
|
num_tokens
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Cache management methods
|
|
||||||
pub fn clear_cache(&self) {
|
pub fn clear_cache(&self) {
|
||||||
self.token_cache.clear();
|
self.token_cache.clear();
|
||||||
}
|
}
|
||||||
@@ -205,32 +190,21 @@ impl Default for TokenCounter {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl TokenCounter {
|
impl TokenCounter {
|
||||||
/// Creates a new `TokenCounter` using the fixed o200k_base encoding.
|
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
// Use blocking version of get_tokenizer
|
|
||||||
let tokenizer = get_tokenizer_blocking().expect("Failed to initialize tokenizer");
|
let tokenizer = get_tokenizer_blocking().expect("Failed to initialize tokenizer");
|
||||||
Self { tokenizer }
|
Self { tokenizer }
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Count tokens for a piece of text using our single tokenizer.
|
|
||||||
pub fn count_tokens(&self, text: &str) -> usize {
|
pub fn count_tokens(&self, text: &str) -> usize {
|
||||||
let tokens = self.tokenizer.encode_with_special_tokens(text);
|
let tokens = self.tokenizer.encode_with_special_tokens(text);
|
||||||
tokens.len()
|
tokens.len()
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn count_tokens_for_tools(&self, tools: &[Tool]) -> usize {
|
pub fn count_tokens_for_tools(&self, tools: &[Tool]) -> usize {
|
||||||
// Token counts for different function components
|
|
||||||
let func_init = 7; // Tokens for function initialization
|
|
||||||
let prop_init = 3; // Tokens for properties initialization
|
|
||||||
let prop_key = 3; // Tokens for each property key
|
|
||||||
let enum_init: isize = -3; // Tokens adjustment for enum list start
|
|
||||||
let enum_item = 3; // Tokens for each enum item
|
|
||||||
let func_end = 12; // Tokens for function ending
|
|
||||||
|
|
||||||
let mut func_token_count = 0;
|
let mut func_token_count = 0;
|
||||||
if !tools.is_empty() {
|
if !tools.is_empty() {
|
||||||
for tool in tools {
|
for tool in tools {
|
||||||
func_token_count += func_init; // Add tokens for start of each function
|
func_token_count += FUNC_INIT;
|
||||||
let name = &tool.name;
|
let name = &tool.name;
|
||||||
let description = &tool
|
let description = &tool
|
||||||
.description
|
.description
|
||||||
@@ -239,27 +213,31 @@ impl TokenCounter {
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.trim_end_matches('.');
|
.trim_end_matches('.');
|
||||||
let line = format!("{}:{}", name, description);
|
let line = format!("{}:{}", name, description);
|
||||||
func_token_count += self.count_tokens(&line); // Add tokens for name and description
|
func_token_count += self.count_tokens(&line);
|
||||||
|
|
||||||
if let serde_json::Value::Object(properties) = &tool.input_schema["properties"] {
|
if let Some(serde_json::Value::Object(properties)) =
|
||||||
|
tool.input_schema.get("properties")
|
||||||
|
{
|
||||||
if !properties.is_empty() {
|
if !properties.is_empty() {
|
||||||
func_token_count += prop_init; // Add tokens for start of properties
|
func_token_count += PROP_INIT;
|
||||||
for (key, value) in properties {
|
for (key, value) in properties {
|
||||||
func_token_count += prop_key; // Add tokens for each property
|
func_token_count += PROP_KEY;
|
||||||
let p_name = key;
|
let p_name = key;
|
||||||
let p_type = value["type"].as_str().unwrap_or("");
|
let p_type = value.get("type").and_then(|v| v.as_str()).unwrap_or("");
|
||||||
let p_desc = value["description"]
|
let p_desc = value
|
||||||
.as_str()
|
.get("description")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
.unwrap_or("")
|
.unwrap_or("")
|
||||||
.trim_end_matches('.');
|
.trim_end_matches('.');
|
||||||
let line = format!("{}:{}:{}", p_name, p_type, p_desc);
|
let line = format!("{}:{}:{}", p_name, p_type, p_desc);
|
||||||
func_token_count += self.count_tokens(&line);
|
func_token_count += self.count_tokens(&line);
|
||||||
if let Some(enum_values) = value["enum"].as_array() {
|
if let Some(enum_values) = value.get("enum").and_then(|v| v.as_array())
|
||||||
|
{
|
||||||
func_token_count =
|
func_token_count =
|
||||||
func_token_count.saturating_add_signed(enum_init); // Add tokens if property has enum list
|
func_token_count.saturating_add_signed(ENUM_INIT);
|
||||||
for item in enum_values {
|
for item in enum_values {
|
||||||
if let Some(item_str) = item.as_str() {
|
if let Some(item_str) = item.as_str() {
|
||||||
func_token_count += enum_item;
|
func_token_count += ENUM_ITEM;
|
||||||
func_token_count += self.count_tokens(item_str);
|
func_token_count += self.count_tokens(item_str);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -268,7 +246,7 @@ impl TokenCounter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
func_token_count += func_end;
|
func_token_count += FUNC_END;
|
||||||
}
|
}
|
||||||
|
|
||||||
func_token_count
|
func_token_count
|
||||||
@@ -280,10 +258,8 @@ impl TokenCounter {
|
|||||||
messages: &[Message],
|
messages: &[Message],
|
||||||
tools: &[Tool],
|
tools: &[Tool],
|
||||||
) -> usize {
|
) -> usize {
|
||||||
// <|im_start|>ROLE<|im_sep|>MESSAGE<|im_end|>
|
|
||||||
let tokens_per_message = 4;
|
let tokens_per_message = 4;
|
||||||
|
|
||||||
// Count tokens in the system prompt
|
|
||||||
let mut num_tokens = 0;
|
let mut num_tokens = 0;
|
||||||
if !system_prompt.is_empty() {
|
if !system_prompt.is_empty() {
|
||||||
num_tokens += self.count_tokens(system_prompt) + tokens_per_message;
|
num_tokens += self.count_tokens(system_prompt) + tokens_per_message;
|
||||||
@@ -291,33 +267,29 @@ impl TokenCounter {
|
|||||||
|
|
||||||
for message in messages {
|
for message in messages {
|
||||||
num_tokens += tokens_per_message;
|
num_tokens += tokens_per_message;
|
||||||
// Count tokens in the content
|
|
||||||
for content in &message.content {
|
for content in &message.content {
|
||||||
// content can either be text response or tool request
|
|
||||||
if let Some(content_text) = content.as_text() {
|
if let Some(content_text) = content.as_text() {
|
||||||
num_tokens += self.count_tokens(content_text);
|
num_tokens += self.count_tokens(content_text);
|
||||||
} else if let Some(tool_request) = content.as_tool_request() {
|
} else if let Some(tool_request) = content.as_tool_request() {
|
||||||
let tool_call = tool_request.tool_call.as_ref().unwrap();
|
if let Ok(tool_call) = tool_request.tool_call.as_ref() {
|
||||||
let text = format!(
|
let text = format!(
|
||||||
"{}:{}:{:?}",
|
"{}:{}:{:?}",
|
||||||
tool_request.id, tool_call.name, tool_call.arguments
|
tool_request.id, tool_call.name, tool_call.arguments
|
||||||
);
|
);
|
||||||
num_tokens += self.count_tokens(&text);
|
num_tokens += self.count_tokens(&text);
|
||||||
|
}
|
||||||
} else if let Some(tool_response_text) = content.as_tool_response_text() {
|
} else if let Some(tool_response_text) = content.as_tool_response_text() {
|
||||||
num_tokens += self.count_tokens(&tool_response_text);
|
num_tokens += self.count_tokens(&tool_response_text);
|
||||||
} else {
|
} else {
|
||||||
// unsupported content type such as image - pass
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Count tokens for tools if provided
|
|
||||||
if !tools.is_empty() {
|
if !tools.is_empty() {
|
||||||
num_tokens += self.count_tokens_for_tools(tools);
|
num_tokens += self.count_tokens_for_tools(tools);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Every reply is primed with <|start|>assistant<|message|>
|
|
||||||
num_tokens += 3;
|
num_tokens += 3;
|
||||||
|
|
||||||
num_tokens
|
num_tokens
|
||||||
@@ -341,8 +313,6 @@ impl TokenCounter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the global tokenizer instance (async version)
|
|
||||||
/// Fixed encoding for all tokenization - using o200k_base for GPT-4o and o1 models
|
|
||||||
async fn get_tokenizer() -> Result<Arc<CoreBPE>, String> {
|
async fn get_tokenizer() -> Result<Arc<CoreBPE>, String> {
|
||||||
let tokenizer = TOKENIZER
|
let tokenizer = TOKENIZER
|
||||||
.get_or_init(|| async {
|
.get_or_init(|| async {
|
||||||
@@ -355,18 +325,14 @@ async fn get_tokenizer() -> Result<Arc<CoreBPE>, String> {
|
|||||||
Ok(tokenizer.clone())
|
Ok(tokenizer.clone())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get the global tokenizer instance (blocking version for backward compatibility)
|
|
||||||
fn get_tokenizer_blocking() -> Result<Arc<CoreBPE>, String> {
|
fn get_tokenizer_blocking() -> Result<Arc<CoreBPE>, String> {
|
||||||
// For the blocking version, we need to handle the case where the tokenizer hasn't been initialized yet
|
|
||||||
if let Some(tokenizer) = TOKENIZER.get() {
|
if let Some(tokenizer) = TOKENIZER.get() {
|
||||||
return Ok(tokenizer.clone());
|
return Ok(tokenizer.clone());
|
||||||
}
|
}
|
||||||
|
|
||||||
// Initialize the tokenizer synchronously
|
|
||||||
match tiktoken_rs::o200k_base() {
|
match tiktoken_rs::o200k_base() {
|
||||||
Ok(bpe) => {
|
Ok(bpe) => {
|
||||||
let tokenizer = Arc::new(bpe);
|
let tokenizer = Arc::new(bpe);
|
||||||
// Try to set it in the OnceCell, but it's okay if another thread beat us to it
|
|
||||||
let _ = TOKENIZER.set(tokenizer.clone());
|
let _ = TOKENIZER.set(tokenizer.clone());
|
||||||
Ok(tokenizer)
|
Ok(tokenizer)
|
||||||
}
|
}
|
||||||
@@ -374,7 +340,6 @@ fn get_tokenizer_blocking() -> Result<Arc<CoreBPE>, String> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Factory function for creating async token counters with proper error handling
|
|
||||||
pub async fn create_async_token_counter() -> Result<AsyncTokenCounter, String> {
|
pub async fn create_async_token_counter() -> Result<AsyncTokenCounter, String> {
|
||||||
AsyncTokenCounter::new().await
|
AsyncTokenCounter::new().await
|
||||||
}
|
}
|
||||||
@@ -386,30 +351,6 @@ mod tests {
|
|||||||
use rmcp::model::{Role, Tool};
|
use rmcp::model::{Role, Tool};
|
||||||
use rmcp::object;
|
use rmcp::object;
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_token_counter_basic() {
|
|
||||||
let counter = TokenCounter::new();
|
|
||||||
|
|
||||||
let text = "Hello, how are you?";
|
|
||||||
let count = counter.count_tokens(text);
|
|
||||||
println!("Token count for '{}': {:?}", text, count);
|
|
||||||
|
|
||||||
// With o200k_base encoding, this should give us a reasonable count
|
|
||||||
assert!(count > 0, "Token count should be greater than 0");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_token_counter_simple_text() {
|
|
||||||
let counter = TokenCounter::new();
|
|
||||||
|
|
||||||
let text = "Hey there!";
|
|
||||||
let count = counter.count_tokens(text);
|
|
||||||
println!("Token count for '{}': {:?}", text, count);
|
|
||||||
|
|
||||||
// With o200k_base encoding, this should give us a reasonable count
|
|
||||||
assert!(count > 0, "Token count should be greater than 0");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_count_chat_tokens() {
|
fn test_count_chat_tokens() {
|
||||||
let counter = TokenCounter::new();
|
let counter = TokenCounter::new();
|
||||||
@@ -464,7 +405,6 @@ mod tests {
|
|||||||
let token_count_with_tools = counter.count_chat_tokens(system_prompt, &messages, &tools);
|
let token_count_with_tools = counter.count_chat_tokens(system_prompt, &messages, &tools);
|
||||||
println!("Total tokens with tools: {}", token_count_with_tools);
|
println!("Total tokens with tools: {}", token_count_with_tools);
|
||||||
|
|
||||||
// Basic sanity checks - with o200k_base the exact counts may differ from the old tokenizer
|
|
||||||
assert!(
|
assert!(
|
||||||
token_count_without_tools > 0,
|
token_count_without_tools > 0,
|
||||||
"Should have some tokens without tools"
|
"Should have some tokens without tools"
|
||||||
@@ -475,122 +415,37 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_async_token_counter() {
|
|
||||||
let counter = create_async_token_counter().await.unwrap();
|
|
||||||
|
|
||||||
let text = "Hello, how are you?";
|
|
||||||
let count = counter.count_tokens(text);
|
|
||||||
println!("Async token count for '{}': {:?}", text, count);
|
|
||||||
|
|
||||||
assert!(count > 0, "Async token count should be greater than 0");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_async_token_caching() {
|
async fn test_async_token_caching() {
|
||||||
let counter = create_async_token_counter().await.unwrap();
|
let counter = create_async_token_counter().await.unwrap();
|
||||||
|
|
||||||
let text = "This is a test for caching functionality";
|
let text = "This is a test for caching functionality";
|
||||||
|
|
||||||
// First call should compute and cache
|
|
||||||
let count1 = counter.count_tokens(text);
|
let count1 = counter.count_tokens(text);
|
||||||
assert_eq!(counter.cache_size(), 1);
|
assert_eq!(counter.cache_size(), 1);
|
||||||
|
|
||||||
// Second call should use cache
|
|
||||||
let count2 = counter.count_tokens(text);
|
let count2 = counter.count_tokens(text);
|
||||||
assert_eq!(count1, count2);
|
assert_eq!(count1, count2);
|
||||||
assert_eq!(counter.cache_size(), 1);
|
assert_eq!(counter.cache_size(), 1);
|
||||||
|
|
||||||
// Different text should increase cache
|
|
||||||
let count3 = counter.count_tokens("Different text");
|
let count3 = counter.count_tokens("Different text");
|
||||||
assert_eq!(counter.cache_size(), 2);
|
assert_eq!(counter.cache_size(), 2);
|
||||||
assert_ne!(count1, count3);
|
assert_ne!(count1, count3);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_async_count_chat_tokens() {
|
|
||||||
let counter = create_async_token_counter().await.unwrap();
|
|
||||||
|
|
||||||
let system_prompt =
|
|
||||||
"You are a helpful assistant that can answer questions about the weather.";
|
|
||||||
|
|
||||||
let messages = vec![
|
|
||||||
Message::new(
|
|
||||||
Role::User,
|
|
||||||
0,
|
|
||||||
vec![MessageContent::text(
|
|
||||||
"What's the weather like in San Francisco?",
|
|
||||||
)],
|
|
||||||
),
|
|
||||||
Message::new(
|
|
||||||
Role::Assistant,
|
|
||||||
1,
|
|
||||||
vec![MessageContent::text(
|
|
||||||
"Looks like it's 60 degrees Fahrenheit in San Francisco.",
|
|
||||||
)],
|
|
||||||
),
|
|
||||||
Message::new(
|
|
||||||
Role::User,
|
|
||||||
2,
|
|
||||||
vec![MessageContent::text("How about New York?")],
|
|
||||||
),
|
|
||||||
];
|
|
||||||
|
|
||||||
let tools = vec![Tool::new(
|
|
||||||
"get_current_weather",
|
|
||||||
"Get the current weather in a given location",
|
|
||||||
object!({
|
|
||||||
"properties": {
|
|
||||||
"location": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "The city and state, e.g. San Francisco, CA"
|
|
||||||
},
|
|
||||||
"unit": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "The unit of temperature to return",
|
|
||||||
"enum": ["celsius", "fahrenheit"]
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"required": ["location"]
|
|
||||||
}),
|
|
||||||
)];
|
|
||||||
|
|
||||||
let token_count_without_tools = counter.count_chat_tokens(system_prompt, &messages, &[]);
|
|
||||||
println!(
|
|
||||||
"Async total tokens without tools: {}",
|
|
||||||
token_count_without_tools
|
|
||||||
);
|
|
||||||
|
|
||||||
let token_count_with_tools = counter.count_chat_tokens(system_prompt, &messages, &tools);
|
|
||||||
println!("Async total tokens with tools: {}", token_count_with_tools);
|
|
||||||
|
|
||||||
// Basic sanity checks
|
|
||||||
assert!(
|
|
||||||
token_count_without_tools > 0,
|
|
||||||
"Should have some tokens without tools"
|
|
||||||
);
|
|
||||||
assert!(
|
|
||||||
token_count_with_tools > token_count_without_tools,
|
|
||||||
"Should have more tokens with tools"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_async_cache_management() {
|
async fn test_async_cache_management() {
|
||||||
let counter = create_async_token_counter().await.unwrap();
|
let counter = create_async_token_counter().await.unwrap();
|
||||||
|
|
||||||
// Add some items to cache
|
|
||||||
counter.count_tokens("First text");
|
counter.count_tokens("First text");
|
||||||
counter.count_tokens("Second text");
|
counter.count_tokens("Second text");
|
||||||
counter.count_tokens("Third text");
|
counter.count_tokens("Third text");
|
||||||
|
|
||||||
assert_eq!(counter.cache_size(), 3);
|
assert_eq!(counter.cache_size(), 3);
|
||||||
|
|
||||||
// Clear cache
|
|
||||||
counter.clear_cache();
|
counter.clear_cache();
|
||||||
assert_eq!(counter.cache_size(), 0);
|
assert_eq!(counter.cache_size(), 0);
|
||||||
|
|
||||||
// Re-count should work fine
|
|
||||||
let count = counter.count_tokens("First text");
|
let count = counter.count_tokens("First text");
|
||||||
assert!(count > 0);
|
assert!(count > 0);
|
||||||
assert_eq!(counter.cache_size(), 1);
|
assert_eq!(counter.cache_size(), 1);
|
||||||
@@ -598,7 +453,6 @@ mod tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_concurrent_token_counter_creation() {
|
async fn test_concurrent_token_counter_creation() {
|
||||||
// Test concurrent creation of token counters to verify no race conditions
|
|
||||||
let handles: Vec<_> = (0..10)
|
let handles: Vec<_> = (0..10)
|
||||||
.map(|_| tokio::spawn(async { create_async_token_counter().await.unwrap() }))
|
.map(|_| tokio::spawn(async { create_async_token_counter().await.unwrap() }))
|
||||||
.collect();
|
.collect();
|
||||||
@@ -609,7 +463,6 @@ mod tests {
|
|||||||
.map(|r| r.unwrap())
|
.map(|r| r.unwrap())
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// All should work and give same results
|
|
||||||
let text = "Test concurrent creation";
|
let text = "Test concurrent creation";
|
||||||
let expected_count = counters[0].count_tokens(text);
|
let expected_count = counters[0].count_tokens(text);
|
||||||
|
|
||||||
@@ -622,7 +475,6 @@ mod tests {
|
|||||||
async fn test_cache_eviction_behavior() {
|
async fn test_cache_eviction_behavior() {
|
||||||
let counter = create_async_token_counter().await.unwrap();
|
let counter = create_async_token_counter().await.unwrap();
|
||||||
|
|
||||||
// Fill cache beyond normal size to test eviction
|
|
||||||
let mut cached_texts = Vec::new();
|
let mut cached_texts = Vec::new();
|
||||||
for i in 0..50 {
|
for i in 0..50 {
|
||||||
let text = format!("Test string number {}", i);
|
let text = format!("Test string number {}", i);
|
||||||
@@ -630,14 +482,11 @@ mod tests {
|
|||||||
cached_texts.push(text);
|
cached_texts.push(text);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Cache should be bounded
|
|
||||||
assert!(counter.cache_size() <= MAX_TOKEN_CACHE_SIZE);
|
assert!(counter.cache_size() <= MAX_TOKEN_CACHE_SIZE);
|
||||||
|
|
||||||
// Earlier entries may have been evicted, but recent ones should still be cached
|
|
||||||
let recent_text = &cached_texts[cached_texts.len() - 1];
|
let recent_text = &cached_texts[cached_texts.len() - 1];
|
||||||
let start_size = counter.cache_size();
|
let start_size = counter.cache_size();
|
||||||
|
|
||||||
// This should be a cache hit (no size increase)
|
|
||||||
counter.count_tokens(recent_text);
|
counter.count_tokens(recent_text);
|
||||||
assert_eq!(counter.cache_size(), start_size);
|
assert_eq!(counter.cache_size(), start_size);
|
||||||
}
|
}
|
||||||
@@ -646,12 +495,11 @@ mod tests {
|
|||||||
async fn test_concurrent_cache_operations() {
|
async fn test_concurrent_cache_operations() {
|
||||||
let counter = std::sync::Arc::new(create_async_token_counter().await.unwrap());
|
let counter = std::sync::Arc::new(create_async_token_counter().await.unwrap());
|
||||||
|
|
||||||
// Test concurrent token counting operations
|
|
||||||
let handles: Vec<_> = (0..20)
|
let handles: Vec<_> = (0..20)
|
||||||
.map(|i| {
|
.map(|i| {
|
||||||
let counter_clone = counter.clone();
|
let counter_clone = counter.clone();
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
let text = format!("Concurrent test {}", i % 5); // Some repetition for cache hits
|
let text = format!("Concurrent test {}", i % 5);
|
||||||
counter_clone.count_tokens(&text)
|
counter_clone.count_tokens(&text)
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
@@ -663,27 +511,22 @@ mod tests {
|
|||||||
.map(|r| r.unwrap())
|
.map(|r| r.unwrap())
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// All results should be valid (> 0)
|
|
||||||
for result in results {
|
for result in results {
|
||||||
assert!(result > 0);
|
assert!(result > 0);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Cache should have some entries but be bounded
|
|
||||||
assert!(counter.cache_size() > 0);
|
assert!(counter.cache_size() > 0);
|
||||||
assert!(counter.cache_size() <= MAX_TOKEN_CACHE_SIZE);
|
assert!(counter.cache_size() <= MAX_TOKEN_CACHE_SIZE);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tokenizer_consistency() {
|
fn test_tokenizer_consistency() {
|
||||||
// Test that both sync and async versions give the same results
|
|
||||||
let sync_counter = TokenCounter::new();
|
let sync_counter = TokenCounter::new();
|
||||||
let text = "This is a test for tokenizer consistency";
|
let text = "This is a test for tokenizer consistency";
|
||||||
let sync_count = sync_counter.count_tokens(text);
|
let sync_count = sync_counter.count_tokens(text);
|
||||||
|
|
||||||
// Test that the tokenizer is working correctly
|
|
||||||
assert!(sync_count > 0, "Sync tokenizer should produce tokens");
|
assert!(sync_count > 0, "Sync tokenizer should produce tokens");
|
||||||
|
|
||||||
// Test with different text lengths
|
|
||||||
let short_text = "Hi";
|
let short_text = "Hi";
|
||||||
let long_text = "This is a much longer text that should produce significantly more tokens than the short text";
|
let long_text = "This is a much longer text that should produce significantly more tokens than the short text";
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user