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
+26 -18
View File
@@ -843,7 +843,13 @@ impl Agent {
&self,
messages: &[Message],
session: &Option<SessionConfig>,
) -> Result<Option<(Conversation, String)>> {
) -> Result<
Option<(
Conversation,
String,
Option<crate::providers::base::ProviderUsage>,
)>,
> {
// Try to get session metadata for more accurate token counts
let session_metadata = if let Some(session_config) = session {
match session::storage::get_path(session_config.id.clone()) {
@@ -865,21 +871,23 @@ impl Agent {
if compact_result.compacted {
let compacted_messages = compact_result.messages;
// Create compaction notification message
let compaction_msg = if let (Some(before), Some(after)) =
(compact_result.tokens_before, compact_result.tokens_after)
{
format!(
"Auto-compacted context: {}{} tokens ({:.0}% reduction)\n\n",
before,
after,
(1.0 - (after as f64 / before as f64)) * 100.0
)
} else {
"Auto-compacted context to reduce token usage\n\n".to_string()
};
// Get threshold from config to include in message
let config = crate::config::Config::global();
let threshold = config
.get_param::<f64>("GOOSE_AUTO_COMPACT_THRESHOLD")
.unwrap_or(0.8); // Default to 80%
let threshold_percentage = (threshold * 100.0) as u32;
return Ok(Some((compacted_messages, compaction_msg)));
let compaction_msg = format!(
"Exceeded auto-compact threshold of {}%. Context has been summarized and reduced.\n\n",
threshold_percentage
);
return Ok(Some((
compacted_messages,
compaction_msg,
compact_result.summarization_usage,
)));
}
Ok(None)
@@ -893,16 +901,16 @@ impl Agent {
cancel_token: Option<CancellationToken>,
) -> Result<BoxStream<'_, Result<AgentEvent>>> {
// Handle auto-compaction before processing
let (messages, compaction_msg) = match self
let (messages, compaction_msg, _summarization_usage) = match self
.handle_auto_compaction(unfixed_conversation.messages(), &session)
.await?
{
Some((compacted_messages, msg)) => (compacted_messages, Some(msg)),
Some((compacted_messages, msg, usage)) => (compacted_messages, Some(msg), usage),
None => {
let context = self
.prepare_reply_context(unfixed_conversation, &session)
.await?;
(context.messages, None)
(context.messages, None, None)
}
};
+37 -24
View File
@@ -4,7 +4,7 @@ use crate::conversation::message::Message;
use crate::conversation::Conversation;
use crate::token_counter::create_async_token_counter;
use crate::context_mgmt::summarize::summarize_messages_async;
use crate::context_mgmt::summarize::summarize_messages;
use crate::context_mgmt::truncate::{truncate_messages, OldestFirstTruncation};
use crate::context_mgmt::{estimate_target_context_limit, get_messages_token_counts_async};
@@ -49,40 +49,53 @@ impl Agent {
}
/// Public API to summarize the conversation so that its token count is within the allowed context limit.
/// Returns the summarized messages, token counts, and the ProviderUsage from summarization
pub async fn summarize_context(
&self,
messages: &[Message], // last message is a user msg that led to assistant message with_context_length_exceeded
) -> Result<(Conversation, Vec<usize>), anyhow::Error> {
) -> Result<
(
Conversation,
Vec<usize>,
Option<crate::providers::base::ProviderUsage>,
),
anyhow::Error,
> {
let provider = self.provider().await?;
let token_counter = create_async_token_counter()
.await
.map_err(|e| anyhow::anyhow!("Failed to create token counter: {}", e))?;
let target_context_limit = estimate_target_context_limit(provider.clone());
let summary_result = summarize_messages(provider.clone(), messages).await?;
let (mut new_messages, mut new_token_counts) =
summarize_messages_async(provider, messages, &token_counter, target_context_limit)
.await?;
let (mut new_messages, mut new_token_counts, summarization_usage) = match summary_result {
Some((summary_message, provider_usage)) => {
// For token counting purposes, we use the output tokens (the actual summary content)
// since that's what will be in the context going forward
let total_tokens = provider_usage.usage.output_tokens.unwrap_or(0) as usize;
(
vec![summary_message],
vec![total_tokens],
Some(provider_usage),
)
}
None => {
// No summary was generated (empty input)
tracing::warn!("Summarization failed. Returning empty messages.");
return Ok((Conversation::empty(), vec![], None));
}
};
// If the summarized messages only contains one message, it means no tool request and response message in the summarized messages,
// Add an assistant message to the summarized messages to ensure the assistant's response is included in the context.
if new_messages.len() == 1 {
let assistant_message = Message::assistant().with_text(
"I had run into a context length exceeded error so I summarized our conversation.",
"I ran into a context length exceeded error so I summarized our conversation.",
);
let assistant_tokens =
token_counter.count_chat_tokens("", &[assistant_message.clone()], &[]);
let current_total: usize = new_token_counts.iter().sum();
if current_total + assistant_tokens <= target_context_limit {
new_messages.push(assistant_message);
new_token_counts.push(assistant_tokens);
} else {
// If we can't fit the assistant message, at least log what happened
tracing::warn!("Cannot add summarization notice message due to context limits. Current: {}, Assistant: {}, Limit: {}",
current_total, assistant_tokens, target_context_limit);
}
let assistant_message_tokens: usize = 14;
new_messages.push(assistant_message);
new_token_counts.push(assistant_message_tokens);
}
Ok((new_messages, new_token_counts))
Ok((
Conversation::new_unvalidated(new_messages),
new_token_counts,
summarization_usage,
))
}
}
+24 -2
View File
@@ -15,6 +15,7 @@ use crate::providers::toolshim::{
augment_message_with_tool_calls, convert_tool_messages_to_text,
modify_system_prompt_for_tool_json, OllamaInterpreter,
};
use crate::session;
use rmcp::model::Tool;
@@ -131,10 +132,20 @@ impl Agent {
};
// Call the provider to get a response
let (mut response, usage) = provider
let (mut response, mut usage) = provider
.complete(system_prompt, messages_for_provider.messages(), tools)
.await?;
// Ensure we have token counts, estimating if necessary
usage
.ensure_tokens(
system_prompt,
messages_for_provider.messages(),
&response,
tools,
)
.await?;
crate::providers::base::set_current_model(&usage.model);
if config.toolshim {
@@ -177,13 +188,24 @@ impl Agent {
)
.await?
} else {
let (message, usage) = provider
let (message, mut usage) = provider
.complete(
system_prompt.as_str(),
messages_for_provider.messages(),
&tools,
)
.await?;
// Ensure we have token counts for non-streaming case
usage
.ensure_tokens(
system_prompt.as_str(),
messages_for_provider.messages(),
&message,
&tools,
)
.await?;
stream_from_single_message(message, usage)
};