Stream token usage on every agent message (#5342)
This commit is contained in:
@@ -19,8 +19,8 @@ use goose::config::declarative_providers::{
|
||||
};
|
||||
use goose::conversation::message::{
|
||||
FrontendToolRequest, Message, MessageContent, MessageMetadata, RedactedThinkingContent,
|
||||
SystemNotificationContent, SystemNotificationType, ThinkingContent, ToolConfirmationRequest,
|
||||
ToolRequest, ToolResponse,
|
||||
SystemNotificationContent, SystemNotificationType, ThinkingContent, TokenState,
|
||||
ToolConfirmationRequest, ToolRequest, ToolResponse,
|
||||
};
|
||||
|
||||
use crate::routes::reply::MessageEvent;
|
||||
@@ -404,6 +404,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
Message,
|
||||
MessageContent,
|
||||
MessageMetadata,
|
||||
TokenState,
|
||||
ContentSchema,
|
||||
EmbeddedResourceSchema,
|
||||
ImageContentSchema,
|
||||
|
||||
@@ -8,7 +8,7 @@ use axum::{
|
||||
};
|
||||
use bytes::Bytes;
|
||||
use futures::{stream::StreamExt, Stream};
|
||||
use goose::conversation::message::{Message, MessageContent};
|
||||
use goose::conversation::message::{Message, MessageContent, TokenState};
|
||||
use goose::conversation::Conversation;
|
||||
use goose::permission::{Permission, PermissionConfirmation};
|
||||
use goose::session::SessionManager;
|
||||
@@ -126,6 +126,7 @@ impl IntoResponse for SseResponse {
|
||||
pub enum MessageEvent {
|
||||
Message {
|
||||
message: Message,
|
||||
token_state: TokenState,
|
||||
},
|
||||
Error {
|
||||
error: String,
|
||||
@@ -159,6 +160,7 @@ async fn stream_event(
|
||||
e
|
||||
)
|
||||
});
|
||||
|
||||
if tx.send(format!("data: {}\n\n", json)).await.is_err() {
|
||||
tracing::info!("client hung up");
|
||||
cancel_token.cancel();
|
||||
@@ -305,7 +307,32 @@ pub async fn reply(
|
||||
}
|
||||
|
||||
all_messages.push(message.clone());
|
||||
stream_event(MessageEvent::Message { message }, &tx, &cancel_token).await;
|
||||
|
||||
let token_state = match SessionManager::get_session(&session_id, false).await {
|
||||
Ok(session) => {
|
||||
TokenState {
|
||||
input_tokens: session.input_tokens.unwrap_or(0),
|
||||
output_tokens: session.output_tokens.unwrap_or(0),
|
||||
total_tokens: session.total_tokens.unwrap_or(0),
|
||||
accumulated_input_tokens: session.accumulated_input_tokens.unwrap_or(0),
|
||||
accumulated_output_tokens: session.accumulated_output_tokens.unwrap_or(0),
|
||||
accumulated_total_tokens: session.accumulated_total_tokens.unwrap_or(0),
|
||||
}
|
||||
},
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to fetch session for token state: {}", e);
|
||||
TokenState {
|
||||
input_tokens: 0,
|
||||
output_tokens: 0,
|
||||
total_tokens: 0,
|
||||
accumulated_input_tokens: 0,
|
||||
accumulated_output_tokens: 0,
|
||||
accumulated_total_tokens: 0,
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
stream_event(MessageEvent::Message { message, token_state }, &tx, &cancel_token).await;
|
||||
}
|
||||
Ok(Some(Ok(AgentEvent::HistoryReplaced(new_messages)))) => {
|
||||
all_messages = new_messages.clone();
|
||||
|
||||
@@ -825,9 +825,11 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
yield AgentEvent::Message(Message::assistant().with_text(
|
||||
format!("Ran into this error trying to compact: {e}.\n\nPlease try again or create a new session")
|
||||
));
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Ran into this error trying to compact: {e}.\n\nPlease try again or create a new session")
|
||||
)
|
||||
);
|
||||
}
|
||||
}
|
||||
}))
|
||||
@@ -917,7 +919,7 @@ impl Agent {
|
||||
if let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||
if final_output_tool.final_output.is_some() {
|
||||
let final_event = AgentEvent::Message(
|
||||
Message::assistant().with_text(final_output_tool.final_output.clone().unwrap()),
|
||||
Message::assistant().with_text(final_output_tool.final_output.clone().unwrap())
|
||||
);
|
||||
yield final_event;
|
||||
break;
|
||||
@@ -926,9 +928,11 @@ impl Agent {
|
||||
|
||||
turns_taken += 1;
|
||||
if turns_taken > max_turns {
|
||||
yield AgentEvent::Message(Message::assistant().with_text(
|
||||
"I've reached the maximum number of actions I can do without user input. Would you like me to continue?"
|
||||
));
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
"I've reached the maximum number of actions I can do without user input. Would you like me to continue?"
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -1178,18 +1182,22 @@ impl Agent {
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error: {}", e);
|
||||
yield AgentEvent::Message(Message::assistant().with_text(
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Ran into this error trying to compact: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||
));
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Error: {}", e);
|
||||
yield AgentEvent::Message(Message::assistant().with_text(
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Ran into this error: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||
));
|
||||
)
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
@@ -1224,9 +1232,11 @@ impl Agent {
|
||||
}
|
||||
Err(e) => {
|
||||
error!("Retry logic failed: {}", e);
|
||||
yield AgentEvent::Message(Message::assistant().with_text(
|
||||
format!("Retry logic encountered an error: {}", e)
|
||||
));
|
||||
yield AgentEvent::Message(
|
||||
Message::assistant().with_text(
|
||||
format!("Retry logic encountered an error: {}", e)
|
||||
)
|
||||
);
|
||||
exit_chat = true;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -711,6 +711,17 @@ impl Message {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct TokenState {
|
||||
pub input_tokens: i32,
|
||||
pub output_tokens: i32,
|
||||
pub total_tokens: i32,
|
||||
pub accumulated_input_tokens: i32,
|
||||
pub accumulated_output_tokens: i32,
|
||||
pub accumulated_total_tokens: i32,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::conversation::message::{Message, MessageContent, MessageMetadata};
|
||||
|
||||
@@ -278,11 +278,11 @@ impl Add for Usage {
|
||||
type Output = Self;
|
||||
|
||||
fn add(self, other: Self) -> Self {
|
||||
Self {
|
||||
input_tokens: sum_optionals(self.input_tokens, other.input_tokens),
|
||||
output_tokens: sum_optionals(self.output_tokens, other.output_tokens),
|
||||
total_tokens: sum_optionals(self.total_tokens, other.total_tokens),
|
||||
}
|
||||
Self::new(
|
||||
sum_optionals(self.input_tokens, other.input_tokens),
|
||||
sum_optionals(self.output_tokens, other.output_tokens),
|
||||
sum_optionals(self.total_tokens, other.total_tokens),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -298,10 +298,21 @@ impl Usage {
|
||||
output_tokens: Option<i32>,
|
||||
total_tokens: Option<i32>,
|
||||
) -> Self {
|
||||
let calculated_total = if total_tokens.is_none() {
|
||||
match (input_tokens, output_tokens) {
|
||||
(Some(input), Some(output)) => Some(input + output),
|
||||
(Some(input), None) => Some(input),
|
||||
(None, Some(output)) => Some(output),
|
||||
(None, None) => None,
|
||||
}
|
||||
} else {
|
||||
total_tokens
|
||||
};
|
||||
|
||||
Self {
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
total_tokens,
|
||||
total_tokens: calculated_total,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -345,11 +345,11 @@ pub fn from_bedrock_role(role: &bedrock::ConversationRole) -> Result<Role> {
|
||||
}
|
||||
|
||||
pub fn from_bedrock_usage(usage: &bedrock::TokenUsage) -> Usage {
|
||||
Usage {
|
||||
input_tokens: Some(usage.input_tokens),
|
||||
output_tokens: Some(usage.output_tokens),
|
||||
total_tokens: Some(usage.total_tokens),
|
||||
}
|
||||
Usage::new(
|
||||
Some(usage.input_tokens),
|
||||
Some(usage.output_tokens),
|
||||
Some(usage.total_tokens),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn from_bedrock_json(document: &Document) -> Result<Value> {
|
||||
|
||||
@@ -307,11 +307,11 @@ impl Provider for SageMakerTgiProvider {
|
||||
let message = self.parse_tgi_response(response)?;
|
||||
|
||||
// TGI doesn't provide usage statistics, so we estimate
|
||||
let usage = Usage {
|
||||
input_tokens: Some(0), // Would need to tokenize input to get accurate count
|
||||
output_tokens: Some(0), // Would need to tokenize output to get accurate count
|
||||
total_tokens: Some(0),
|
||||
};
|
||||
let usage = Usage::new(
|
||||
Some(0), // Would need to tokenize input to get accurate count
|
||||
Some(0), // Would need to tokenize output to get accurate count
|
||||
Some(0),
|
||||
);
|
||||
|
||||
// Add debug trace
|
||||
let debug_payload = serde_json::json!({
|
||||
|
||||
@@ -508,11 +508,11 @@ impl Provider for VeniceProvider {
|
||||
|
||||
// Extract usage
|
||||
let usage_data = &response_json["usage"];
|
||||
let usage = Usage {
|
||||
input_tokens: usage_data["prompt_tokens"].as_i64().map(|v| v as i32),
|
||||
output_tokens: usage_data["completion_tokens"].as_i64().map(|v| v as i32),
|
||||
total_tokens: usage_data["total_tokens"].as_i64().map(|v| v as i32),
|
||||
};
|
||||
let usage = Usage::new(
|
||||
usage_data["prompt_tokens"].as_i64().map(|v| v as i32),
|
||||
usage_data["completion_tokens"].as_i64().map(|v| v as i32),
|
||||
usage_data["total_tokens"].as_i64().map(|v| v as i32),
|
||||
);
|
||||
|
||||
Ok((
|
||||
Message::new(Role::Assistant, Utc::now().timestamp(), content),
|
||||
|
||||
Reference in New Issue
Block a user