Compaction overhaul (#5186)
Co-authored-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: David Katz <dkatz@squareup.com> Co-authored-by: Alex Hancock <alexhancock@block.xyz>
This commit is contained in:
@@ -20,16 +20,12 @@ const ENUM_INIT: isize = -3;
|
||||
const ENUM_ITEM: usize = 3;
|
||||
const FUNC_END: usize = 12;
|
||||
|
||||
pub struct AsyncTokenCounter {
|
||||
pub struct TokenCounter {
|
||||
tokenizer: Arc<CoreBPE>,
|
||||
token_cache: Arc<DashMap<u64, usize>>,
|
||||
}
|
||||
|
||||
pub struct TokenCounter {
|
||||
tokenizer: Arc<CoreBPE>,
|
||||
}
|
||||
|
||||
impl AsyncTokenCounter {
|
||||
impl TokenCounter {
|
||||
pub async fn new() -> Result<Self, String> {
|
||||
let tokenizer = get_tokenizer().await?;
|
||||
Ok(Self {
|
||||
@@ -130,6 +126,9 @@ impl AsyncTokenCounter {
|
||||
}
|
||||
|
||||
for message in messages {
|
||||
if !message.metadata.agent_visible {
|
||||
continue;
|
||||
}
|
||||
num_tokens += tokens_per_message;
|
||||
for content in &message.content {
|
||||
if let Some(content_text) = content.as_text() {
|
||||
@@ -183,136 +182,6 @@ impl AsyncTokenCounter {
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for TokenCounter {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl TokenCounter {
|
||||
pub fn new() -> Self {
|
||||
let tokenizer = get_tokenizer_blocking().expect("Failed to initialize tokenizer");
|
||||
Self { tokenizer }
|
||||
}
|
||||
|
||||
pub fn count_tokens(&self, text: &str) -> usize {
|
||||
let tokens = self.tokenizer.encode_with_special_tokens(text);
|
||||
tokens.len()
|
||||
}
|
||||
|
||||
pub fn count_tokens_for_tools(&self, tools: &[Tool]) -> usize {
|
||||
let mut func_token_count = 0;
|
||||
if !tools.is_empty() {
|
||||
for tool in tools {
|
||||
func_token_count += FUNC_INIT;
|
||||
let name = &tool.name;
|
||||
let description = &tool
|
||||
.description
|
||||
.as_ref()
|
||||
.map(|d| d.as_ref())
|
||||
.unwrap_or_default()
|
||||
.trim_end_matches('.');
|
||||
let line = format!("{}:{}", name, description);
|
||||
func_token_count += self.count_tokens(&line);
|
||||
|
||||
if let Some(serde_json::Value::Object(properties)) =
|
||||
tool.input_schema.get("properties")
|
||||
{
|
||||
if !properties.is_empty() {
|
||||
func_token_count += PROP_INIT;
|
||||
for (key, value) in properties {
|
||||
func_token_count += PROP_KEY;
|
||||
let p_name = key;
|
||||
let p_type = value.get("type").and_then(|v| v.as_str()).unwrap_or("");
|
||||
let p_desc = value
|
||||
.get("description")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.trim_end_matches('.');
|
||||
let line = format!("{}:{}:{}", p_name, p_type, p_desc);
|
||||
func_token_count += self.count_tokens(&line);
|
||||
if let Some(enum_values) = value.get("enum").and_then(|v| v.as_array())
|
||||
{
|
||||
func_token_count =
|
||||
func_token_count.saturating_add_signed(ENUM_INIT);
|
||||
for item in enum_values {
|
||||
if let Some(item_str) = item.as_str() {
|
||||
func_token_count += ENUM_ITEM;
|
||||
func_token_count += self.count_tokens(item_str);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
func_token_count += FUNC_END;
|
||||
}
|
||||
|
||||
func_token_count
|
||||
}
|
||||
|
||||
pub fn count_chat_tokens(
|
||||
&self,
|
||||
system_prompt: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
) -> usize {
|
||||
let tokens_per_message = 4;
|
||||
|
||||
let mut num_tokens = 0;
|
||||
if !system_prompt.is_empty() {
|
||||
num_tokens += self.count_tokens(system_prompt) + tokens_per_message;
|
||||
}
|
||||
|
||||
for message in messages {
|
||||
num_tokens += tokens_per_message;
|
||||
for content in &message.content {
|
||||
if let Some(content_text) = content.as_text() {
|
||||
num_tokens += self.count_tokens(content_text);
|
||||
} else if let Some(tool_request) = content.as_tool_request() {
|
||||
if let Ok(tool_call) = tool_request.tool_call.as_ref() {
|
||||
let text = format!(
|
||||
"{}:{}:{:?}",
|
||||
tool_request.id, tool_call.name, tool_call.arguments
|
||||
);
|
||||
num_tokens += self.count_tokens(&text);
|
||||
}
|
||||
} else if let Some(tool_response_text) = content.as_tool_response_text() {
|
||||
num_tokens += self.count_tokens(&tool_response_text);
|
||||
} else {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !tools.is_empty() {
|
||||
num_tokens += self.count_tokens_for_tools(tools);
|
||||
}
|
||||
|
||||
num_tokens += 3;
|
||||
|
||||
num_tokens
|
||||
}
|
||||
|
||||
pub fn count_everything(
|
||||
&self,
|
||||
system_prompt: &str,
|
||||
messages: &[Message],
|
||||
tools: &[Tool],
|
||||
resources: &[String],
|
||||
) -> usize {
|
||||
let mut num_tokens = self.count_chat_tokens(system_prompt, messages, tools);
|
||||
|
||||
if !resources.is_empty() {
|
||||
for resource in resources {
|
||||
num_tokens += self.count_tokens(resource);
|
||||
}
|
||||
}
|
||||
num_tokens
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_tokenizer() -> Result<Arc<CoreBPE>, String> {
|
||||
let tokenizer = TOKENIZER
|
||||
.get_or_init(|| async {
|
||||
@@ -325,99 +194,17 @@ async fn get_tokenizer() -> Result<Arc<CoreBPE>, String> {
|
||||
Ok(tokenizer.clone())
|
||||
}
|
||||
|
||||
fn get_tokenizer_blocking() -> Result<Arc<CoreBPE>, String> {
|
||||
if let Some(tokenizer) = TOKENIZER.get() {
|
||||
return Ok(tokenizer.clone());
|
||||
}
|
||||
|
||||
match tiktoken_rs::o200k_base() {
|
||||
Ok(bpe) => {
|
||||
let tokenizer = Arc::new(bpe);
|
||||
let _ = TOKENIZER.set(tokenizer.clone());
|
||||
Ok(tokenizer)
|
||||
}
|
||||
Err(e) => Err(format!("Failed to initialize o200k_base tokenizer: {}", e)),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn create_async_token_counter() -> Result<AsyncTokenCounter, String> {
|
||||
AsyncTokenCounter::new().await
|
||||
pub async fn create_token_counter() -> Result<TokenCounter, String> {
|
||||
TokenCounter::new().await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use rmcp::model::{Role, Tool};
|
||||
use rmcp::object;
|
||||
|
||||
#[test]
|
||||
fn test_count_chat_tokens() {
|
||||
let counter = TokenCounter::new();
|
||||
|
||||
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!("Total tokens without tools: {}", token_count_without_tools);
|
||||
|
||||
let token_count_with_tools = counter.count_chat_tokens(system_prompt, &messages, &tools);
|
||||
println!("Total tokens with tools: {}", token_count_with_tools);
|
||||
|
||||
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]
|
||||
async fn test_async_token_caching() {
|
||||
let counter = create_async_token_counter().await.unwrap();
|
||||
async fn test_token_caching() {
|
||||
let counter = create_token_counter().await.unwrap();
|
||||
|
||||
let text = "This is a test for caching functionality";
|
||||
|
||||
@@ -434,8 +221,8 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_async_cache_management() {
|
||||
let counter = create_async_token_counter().await.unwrap();
|
||||
async fn test_cache_management() {
|
||||
let counter = create_token_counter().await.unwrap();
|
||||
|
||||
counter.count_tokens("First text");
|
||||
counter.count_tokens("Second text");
|
||||
@@ -454,7 +241,7 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn test_concurrent_token_counter_creation() {
|
||||
let handles: Vec<_> = (0..10)
|
||||
.map(|_| tokio::spawn(async { create_async_token_counter().await.unwrap() }))
|
||||
.map(|_| tokio::spawn(async { create_token_counter().await.unwrap() }))
|
||||
.collect();
|
||||
|
||||
let counters: Vec<_> = futures::future::join_all(handles)
|
||||
@@ -473,7 +260,7 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_eviction_behavior() {
|
||||
let counter = create_async_token_counter().await.unwrap();
|
||||
let counter = create_token_counter().await.unwrap();
|
||||
|
||||
let mut cached_texts = Vec::new();
|
||||
for i in 0..50 {
|
||||
@@ -493,7 +280,7 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
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_token_counter().await.unwrap());
|
||||
|
||||
let handles: Vec<_> = (0..20)
|
||||
.map(|i| {
|
||||
@@ -518,24 +305,4 @@ mod tests {
|
||||
assert!(counter.cache_size() > 0);
|
||||
assert!(counter.cache_size() <= MAX_TOKEN_CACHE_SIZE);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tokenizer_consistency() {
|
||||
let sync_counter = TokenCounter::new();
|
||||
let text = "This is a test for tokenizer consistency";
|
||||
let sync_count = sync_counter.count_tokens(text);
|
||||
|
||||
assert!(sync_count > 0, "Sync tokenizer should produce tokens");
|
||||
|
||||
let short_text = "Hi";
|
||||
let long_text = "This is a much longer text that should produce significantly more tokens than the short text";
|
||||
|
||||
let short_count = sync_counter.count_tokens(short_text);
|
||||
let long_count = sync_counter.count_tokens(long_text);
|
||||
|
||||
assert!(
|
||||
short_count < long_count,
|
||||
"Longer text should have more tokens"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user