fix: fix summarize agent, use session_id and add provider fn (#1552)

This commit is contained in:
Kalvin C
2025-03-06 09:35:17 -08:00
committed by GitHub
parent 5f750a5229
commit 5237277082
+21 -1
View File
@@ -4,6 +4,7 @@
use async_trait::async_trait;
use futures::stream::BoxStream;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio::sync::Mutex;
use tracing::{debug, error, instrument, warn};
@@ -20,6 +21,7 @@ use crate::providers::base::Provider;
use crate::providers::base::ProviderUsage;
use crate::providers::errors::ProviderError;
use crate::register_agent;
use crate::session;
use crate::token_counter::TokenCounter;
use crate::truncate::{truncate_messages, OldestFirstTruncation};
use anyhow::{anyhow, Result};
@@ -164,6 +166,7 @@ impl Agent for SummarizeAgent {
async fn reply(
&self,
messages: &[Message],
session_id: Option<session::Identifier>,
) -> anyhow::Result<BoxStream<'_, anyhow::Result<Message>>> {
let mut messages = messages.to_vec();
let reply_span = tracing::Span::current();
@@ -240,7 +243,19 @@ impl Agent for SummarizeAgent {
&tools,
).await {
Ok((response, usage)) => {
capabilities.record_usage(usage).await;
capabilities.record_usage(usage.clone()).await;
// record usage for the session in the session file
if let Some(session_id) = session_id.clone() {
// TODO: track session_id in langfuse tracing
let session_file = session::get_path(session_id);
let mut metadata = session::read_metadata(&session_file)?;
metadata.total_tokens = usage.usage.total_tokens;
// The message count is the number of messages in the session + 1 for the response
// The message count does not include the tool response till next iteration
metadata.message_count = messages.len() + 1;
session::update_metadata(&session_file, &metadata).await?;
}
// Reset truncation attempt
truncation_attempt = 0;
@@ -452,6 +467,11 @@ impl Agent for SummarizeAgent {
Err(anyhow!("Prompt '{}' not found", name))
}
async fn provider(&self) -> Arc<Box<dyn Provider>> {
let capabilities = self.capabilities.lock().await;
capabilities.provider()
}
}
register_agent!("summarize", SummarizeAgent);