Token counting reliability + summarization integration. (#3721)

This commit is contained in:
David Katz
2025-08-14 11:23:37 -04:00
committed by GitHub
parent f2e335cb81
commit 80826c2b23
15 changed files with 359 additions and 891 deletions
+28
View File
@@ -217,6 +217,34 @@ impl ProviderUsage {
pub fn new(model: String, usage: Usage) -> Self {
Self { model, usage }
}
/// Ensures this ProviderUsage has token counts, estimating them if necessary
pub async fn ensure_tokens(
&mut self,
system_prompt: &str,
request_messages: &[Message],
response: &Message,
tools: &[Tool],
) -> Result<(), ProviderError> {
crate::providers::usage_estimator::ensure_usage_tokens(
self,
system_prompt,
request_messages,
response,
tools,
)
.await
.map_err(|e| ProviderError::ExecutionError(format!("Failed to ensure usage tokens: {}", e)))
}
/// Combine this ProviderUsage with another, adding their token counts
/// Uses the model from this ProviderUsage
pub fn combine_with(&self, other: &ProviderUsage) -> ProviderUsage {
ProviderUsage {
model: self.model.clone(),
usage: self.usage + other.usage,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Default, Copy)]
+1
View File
@@ -28,6 +28,7 @@ pub mod sagemaker_tgi;
pub mod snowflake;
pub mod testprovider;
pub mod toolshim;
pub mod usage_estimator;
pub mod utils;
pub mod utils_universal_openai_stream;
pub mod venice;
@@ -0,0 +1,128 @@
use crate::conversation::message::Message;
use crate::providers::base::ProviderUsage;
use crate::token_counter::create_async_token_counter;
use anyhow::Result;
use rmcp::model::Tool;
/// Ensures that ProviderUsage has token counts, estimating them if necessary.
/// This provides a single place to handle the fallback logic for providers that don't return usage data.
pub async fn ensure_usage_tokens(
provider_usage: &mut ProviderUsage,
system_prompt: &str,
request_messages: &[Message],
response: &Message,
tools: &[Tool],
) -> Result<()> {
if provider_usage.usage.input_tokens.is_some() && provider_usage.usage.output_tokens.is_some() {
return Ok(());
}
let token_counter = create_async_token_counter()
.await
.map_err(|e| anyhow::anyhow!("Failed to create token counter: {}", e))?;
if provider_usage.usage.input_tokens.is_none() {
let input_count = token_counter.count_chat_tokens(system_prompt, request_messages, tools);
provider_usage.usage.input_tokens = Some(input_count as i32);
}
if provider_usage.usage.output_tokens.is_none() {
let response_text = response
.content
.iter()
.map(|c| format!("{}", c))
.collect::<Vec<_>>()
.join(" ");
let output_count = token_counter.count_tokens(&response_text);
provider_usage.usage.output_tokens = Some(output_count as i32);
}
if let (Some(input), Some(output)) = (
provider_usage.usage.input_tokens,
provider_usage.usage.output_tokens,
) {
provider_usage.usage.total_tokens = Some(input + output);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::conversation::message::Message;
use crate::providers::base::Usage;
#[tokio::test]
async fn test_ensure_usage_tokens_already_complete() {
let mut usage = ProviderUsage::new(
"test-model".to_string(),
Usage::new(Some(100), Some(50), Some(150)),
);
let response = Message::assistant().with_text("Test response");
ensure_usage_tokens(&mut usage, "system", &[], &response, &[])
.await
.unwrap();
// Should remain unchanged
assert_eq!(usage.usage.input_tokens, Some(100));
assert_eq!(usage.usage.output_tokens, Some(50));
assert_eq!(usage.usage.total_tokens, Some(150));
}
#[tokio::test]
async fn test_ensure_usage_tokens_missing_all() {
let mut usage = ProviderUsage::new("test-model".to_string(), Usage::default());
let response = Message::assistant().with_text("Test response");
let messages = vec![Message::user().with_text("Hello")];
ensure_usage_tokens(
&mut usage,
"You are a helpful assistant",
&messages,
&response,
&[],
)
.await
.unwrap();
// Should have estimated values
assert!(usage.usage.input_tokens.is_some());
assert!(usage.usage.output_tokens.is_some());
assert!(usage.usage.total_tokens.is_some());
// Basic sanity checks
assert!(usage.usage.input_tokens.unwrap() > 0);
assert!(usage.usage.output_tokens.unwrap() > 0);
assert_eq!(
usage.usage.total_tokens.unwrap(),
usage.usage.input_tokens.unwrap() + usage.usage.output_tokens.unwrap()
);
}
#[tokio::test]
async fn test_ensure_usage_tokens_partial() {
let mut usage =
ProviderUsage::new("test-model".to_string(), Usage::new(Some(100), None, None));
let response = Message::assistant().with_text("Test response");
ensure_usage_tokens(&mut usage, "system", &[], &response, &[])
.await
.unwrap();
// Input should remain unchanged
assert_eq!(usage.usage.input_tokens, Some(100));
// Output should be estimated
assert!(usage.usage.output_tokens.is_some());
assert!(usage.usage.output_tokens.unwrap() > 0);
// Total should be calculated
assert_eq!(
usage.usage.total_tokens.unwrap(),
usage.usage.input_tokens.unwrap() + usage.usage.output_tokens.unwrap()
);
}
}