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::{
|
use goose::conversation::message::{
|
||||||
FrontendToolRequest, Message, MessageContent, MessageMetadata, RedactedThinkingContent,
|
FrontendToolRequest, Message, MessageContent, MessageMetadata, RedactedThinkingContent,
|
||||||
SystemNotificationContent, SystemNotificationType, ThinkingContent, ToolConfirmationRequest,
|
SystemNotificationContent, SystemNotificationType, ThinkingContent, TokenState,
|
||||||
ToolRequest, ToolResponse,
|
ToolConfirmationRequest, ToolRequest, ToolResponse,
|
||||||
};
|
};
|
||||||
|
|
||||||
use crate::routes::reply::MessageEvent;
|
use crate::routes::reply::MessageEvent;
|
||||||
@@ -404,6 +404,7 @@ derive_utoipa!(Icon as IconSchema);
|
|||||||
Message,
|
Message,
|
||||||
MessageContent,
|
MessageContent,
|
||||||
MessageMetadata,
|
MessageMetadata,
|
||||||
|
TokenState,
|
||||||
ContentSchema,
|
ContentSchema,
|
||||||
EmbeddedResourceSchema,
|
EmbeddedResourceSchema,
|
||||||
ImageContentSchema,
|
ImageContentSchema,
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ use axum::{
|
|||||||
};
|
};
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use futures::{stream::StreamExt, Stream};
|
use futures::{stream::StreamExt, Stream};
|
||||||
use goose::conversation::message::{Message, MessageContent};
|
use goose::conversation::message::{Message, MessageContent, TokenState};
|
||||||
use goose::conversation::Conversation;
|
use goose::conversation::Conversation;
|
||||||
use goose::permission::{Permission, PermissionConfirmation};
|
use goose::permission::{Permission, PermissionConfirmation};
|
||||||
use goose::session::SessionManager;
|
use goose::session::SessionManager;
|
||||||
@@ -126,6 +126,7 @@ impl IntoResponse for SseResponse {
|
|||||||
pub enum MessageEvent {
|
pub enum MessageEvent {
|
||||||
Message {
|
Message {
|
||||||
message: Message,
|
message: Message,
|
||||||
|
token_state: TokenState,
|
||||||
},
|
},
|
||||||
Error {
|
Error {
|
||||||
error: String,
|
error: String,
|
||||||
@@ -159,6 +160,7 @@ async fn stream_event(
|
|||||||
e
|
e
|
||||||
)
|
)
|
||||||
});
|
});
|
||||||
|
|
||||||
if tx.send(format!("data: {}\n\n", json)).await.is_err() {
|
if tx.send(format!("data: {}\n\n", json)).await.is_err() {
|
||||||
tracing::info!("client hung up");
|
tracing::info!("client hung up");
|
||||||
cancel_token.cancel();
|
cancel_token.cancel();
|
||||||
@@ -305,7 +307,32 @@ pub async fn reply(
|
|||||||
}
|
}
|
||||||
|
|
||||||
all_messages.push(message.clone());
|
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)))) => {
|
Ok(Some(Ok(AgentEvent::HistoryReplaced(new_messages)))) => {
|
||||||
all_messages = new_messages.clone();
|
all_messages = new_messages.clone();
|
||||||
|
|||||||
@@ -825,9 +825,11 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
yield AgentEvent::Message(Message::assistant().with_text(
|
yield AgentEvent::Message(
|
||||||
format!("Ran into this error trying to compact: {e}.\n\nPlease try again or create a new session")
|
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 let Some(final_output_tool) = self.final_output_tool.lock().await.as_ref() {
|
||||||
if final_output_tool.final_output.is_some() {
|
if final_output_tool.final_output.is_some() {
|
||||||
let final_event = AgentEvent::Message(
|
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;
|
yield final_event;
|
||||||
break;
|
break;
|
||||||
@@ -926,9 +928,11 @@ impl Agent {
|
|||||||
|
|
||||||
turns_taken += 1;
|
turns_taken += 1;
|
||||||
if turns_taken > max_turns {
|
if turns_taken > max_turns {
|
||||||
yield AgentEvent::Message(Message::assistant().with_text(
|
yield AgentEvent::Message(
|
||||||
"I've reached the maximum number of actions I can do without user input. Would you like me to continue?"
|
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;
|
break;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1178,18 +1182,22 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
error!("Error: {}", 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.")
|
format!("Ran into this error trying to compact: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||||
));
|
)
|
||||||
|
);
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
error!("Error: {}", 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.")
|
format!("Ran into this error: {e}.\n\nPlease retry if you think this is a transient or recoverable error.")
|
||||||
));
|
)
|
||||||
|
);
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1224,9 +1232,11 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
error!("Retry logic failed: {}", e);
|
error!("Retry logic failed: {}", e);
|
||||||
yield AgentEvent::Message(Message::assistant().with_text(
|
yield AgentEvent::Message(
|
||||||
format!("Retry logic encountered an error: {}", e)
|
Message::assistant().with_text(
|
||||||
));
|
format!("Retry logic encountered an error: {}", e)
|
||||||
|
)
|
||||||
|
);
|
||||||
exit_chat = true;
|
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)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use crate::conversation::message::{Message, MessageContent, MessageMetadata};
|
use crate::conversation::message::{Message, MessageContent, MessageMetadata};
|
||||||
|
|||||||
@@ -278,11 +278,11 @@ impl Add for Usage {
|
|||||||
type Output = Self;
|
type Output = Self;
|
||||||
|
|
||||||
fn add(self, other: Self) -> Self {
|
fn add(self, other: Self) -> Self {
|
||||||
Self {
|
Self::new(
|
||||||
input_tokens: sum_optionals(self.input_tokens, other.input_tokens),
|
sum_optionals(self.input_tokens, other.input_tokens),
|
||||||
output_tokens: sum_optionals(self.output_tokens, other.output_tokens),
|
sum_optionals(self.output_tokens, other.output_tokens),
|
||||||
total_tokens: sum_optionals(self.total_tokens, other.total_tokens),
|
sum_optionals(self.total_tokens, other.total_tokens),
|
||||||
}
|
)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -298,10 +298,21 @@ impl Usage {
|
|||||||
output_tokens: Option<i32>,
|
output_tokens: Option<i32>,
|
||||||
total_tokens: Option<i32>,
|
total_tokens: Option<i32>,
|
||||||
) -> Self {
|
) -> 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 {
|
Self {
|
||||||
input_tokens,
|
input_tokens,
|
||||||
output_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 {
|
pub fn from_bedrock_usage(usage: &bedrock::TokenUsage) -> Usage {
|
||||||
Usage {
|
Usage::new(
|
||||||
input_tokens: Some(usage.input_tokens),
|
Some(usage.input_tokens),
|
||||||
output_tokens: Some(usage.output_tokens),
|
Some(usage.output_tokens),
|
||||||
total_tokens: Some(usage.total_tokens),
|
Some(usage.total_tokens),
|
||||||
}
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn from_bedrock_json(document: &Document) -> Result<Value> {
|
pub fn from_bedrock_json(document: &Document) -> Result<Value> {
|
||||||
|
|||||||
@@ -307,11 +307,11 @@ impl Provider for SageMakerTgiProvider {
|
|||||||
let message = self.parse_tgi_response(response)?;
|
let message = self.parse_tgi_response(response)?;
|
||||||
|
|
||||||
// TGI doesn't provide usage statistics, so we estimate
|
// TGI doesn't provide usage statistics, so we estimate
|
||||||
let usage = Usage {
|
let usage = Usage::new(
|
||||||
input_tokens: Some(0), // Would need to tokenize input to get accurate count
|
Some(0), // Would need to tokenize input to get accurate count
|
||||||
output_tokens: Some(0), // Would need to tokenize output to get accurate count
|
Some(0), // Would need to tokenize output to get accurate count
|
||||||
total_tokens: Some(0),
|
Some(0),
|
||||||
};
|
);
|
||||||
|
|
||||||
// Add debug trace
|
// Add debug trace
|
||||||
let debug_payload = serde_json::json!({
|
let debug_payload = serde_json::json!({
|
||||||
|
|||||||
@@ -508,11 +508,11 @@ impl Provider for VeniceProvider {
|
|||||||
|
|
||||||
// Extract usage
|
// Extract usage
|
||||||
let usage_data = &response_json["usage"];
|
let usage_data = &response_json["usage"];
|
||||||
let usage = Usage {
|
let usage = Usage::new(
|
||||||
input_tokens: usage_data["prompt_tokens"].as_i64().map(|v| v as i32),
|
usage_data["prompt_tokens"].as_i64().map(|v| v as i32),
|
||||||
output_tokens: usage_data["completion_tokens"].as_i64().map(|v| v as i32),
|
usage_data["completion_tokens"].as_i64().map(|v| v as i32),
|
||||||
total_tokens: usage_data["total_tokens"].as_i64().map(|v| v as i32),
|
usage_data["total_tokens"].as_i64().map(|v| v as i32),
|
||||||
};
|
);
|
||||||
|
|
||||||
Ok((
|
Ok((
|
||||||
Message::new(Role::Assistant, Utc::now().timestamp(), content),
|
Message::new(Role::Assistant, Utc::now().timestamp(), content),
|
||||||
|
|||||||
@@ -3278,12 +3278,16 @@
|
|||||||
"type": "object",
|
"type": "object",
|
||||||
"required": [
|
"required": [
|
||||||
"message",
|
"message",
|
||||||
|
"token_state",
|
||||||
"type"
|
"type"
|
||||||
],
|
],
|
||||||
"properties": {
|
"properties": {
|
||||||
"message": {
|
"message": {
|
||||||
"$ref": "#/components/schemas/Message"
|
"$ref": "#/components/schemas/Message"
|
||||||
},
|
},
|
||||||
|
"token_state": {
|
||||||
|
"$ref": "#/components/schemas/TokenState"
|
||||||
|
},
|
||||||
"type": {
|
"type": {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"enum": [
|
"enum": [
|
||||||
@@ -4526,6 +4530,43 @@
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
"TokenState": {
|
||||||
|
"type": "object",
|
||||||
|
"required": [
|
||||||
|
"inputTokens",
|
||||||
|
"outputTokens",
|
||||||
|
"totalTokens",
|
||||||
|
"accumulatedInputTokens",
|
||||||
|
"accumulatedOutputTokens",
|
||||||
|
"accumulatedTotalTokens"
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"accumulatedInputTokens": {
|
||||||
|
"type": "integer",
|
||||||
|
"format": "int32"
|
||||||
|
},
|
||||||
|
"accumulatedOutputTokens": {
|
||||||
|
"type": "integer",
|
||||||
|
"format": "int32"
|
||||||
|
},
|
||||||
|
"accumulatedTotalTokens": {
|
||||||
|
"type": "integer",
|
||||||
|
"format": "int32"
|
||||||
|
},
|
||||||
|
"inputTokens": {
|
||||||
|
"type": "integer",
|
||||||
|
"format": "int32"
|
||||||
|
},
|
||||||
|
"outputTokens": {
|
||||||
|
"type": "integer",
|
||||||
|
"format": "int32"
|
||||||
|
},
|
||||||
|
"totalTokens": {
|
||||||
|
"type": "integer",
|
||||||
|
"format": "int32"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
"Tool": {
|
"Tool": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"required": [
|
"required": [
|
||||||
|
|||||||
@@ -367,6 +367,7 @@ export type MessageContent = (TextContent & {
|
|||||||
|
|
||||||
export type MessageEvent = {
|
export type MessageEvent = {
|
||||||
message: Message;
|
message: Message;
|
||||||
|
token_state: TokenState;
|
||||||
type: 'Message';
|
type: 'Message';
|
||||||
} | {
|
} | {
|
||||||
error: string;
|
error: string;
|
||||||
@@ -789,6 +790,15 @@ export type ThinkingContent = {
|
|||||||
thinking: string;
|
thinking: string;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
export type TokenState = {
|
||||||
|
accumulatedInputTokens: number;
|
||||||
|
accumulatedOutputTokens: number;
|
||||||
|
accumulatedTotalTokens: number;
|
||||||
|
inputTokens: number;
|
||||||
|
outputTokens: number;
|
||||||
|
totalTokens: number;
|
||||||
|
};
|
||||||
|
|
||||||
export type Tool = {
|
export type Tool = {
|
||||||
annotations?: ToolAnnotations | {
|
annotations?: ToolAnnotations | {
|
||||||
[key: string]: unknown;
|
[key: string]: unknown;
|
||||||
|
|||||||
@@ -132,6 +132,7 @@ function BaseChatContent({
|
|||||||
sessionOutputTokens,
|
sessionOutputTokens,
|
||||||
localInputTokens,
|
localInputTokens,
|
||||||
localOutputTokens,
|
localOutputTokens,
|
||||||
|
tokenState,
|
||||||
commandHistory,
|
commandHistory,
|
||||||
toolCallNotifications,
|
toolCallNotifications,
|
||||||
sessionMetadata,
|
sessionMetadata,
|
||||||
@@ -442,9 +443,13 @@ function BaseChatContent({
|
|||||||
commandHistory={commandHistory}
|
commandHistory={commandHistory}
|
||||||
initialValue={input || ''}
|
initialValue={input || ''}
|
||||||
setView={setView}
|
setView={setView}
|
||||||
numTokens={sessionTokenCount}
|
totalTokens={tokenState?.totalTokens ?? sessionTokenCount}
|
||||||
inputTokens={sessionInputTokens || localInputTokens}
|
accumulatedInputTokens={
|
||||||
outputTokens={sessionOutputTokens || localOutputTokens}
|
tokenState?.accumulatedInputTokens ?? sessionInputTokens ?? localInputTokens
|
||||||
|
}
|
||||||
|
accumulatedOutputTokens={
|
||||||
|
tokenState?.accumulatedOutputTokens ?? sessionOutputTokens ?? localOutputTokens
|
||||||
|
}
|
||||||
droppedFiles={droppedFiles}
|
droppedFiles={droppedFiles}
|
||||||
onFilesProcessed={() => setDroppedFiles([])} // Clear dropped files after processing
|
onFilesProcessed={() => setDroppedFiles([])} // Clear dropped files after processing
|
||||||
messages={messages}
|
messages={messages}
|
||||||
|
|||||||
@@ -72,6 +72,7 @@ function BaseChatContent({
|
|||||||
stopStreaming,
|
stopStreaming,
|
||||||
sessionLoadError,
|
sessionLoadError,
|
||||||
setRecipeUserParams,
|
setRecipeUserParams,
|
||||||
|
tokenState,
|
||||||
} = useChatStream({
|
} = useChatStream({
|
||||||
sessionId,
|
sessionId,
|
||||||
onStreamFinish,
|
onStreamFinish,
|
||||||
@@ -281,9 +282,13 @@ function BaseChatContent({
|
|||||||
//commandHistory={commandHistory}
|
//commandHistory={commandHistory}
|
||||||
initialValue={initialPrompt}
|
initialValue={initialPrompt}
|
||||||
setView={setView}
|
setView={setView}
|
||||||
numTokens={session?.total_tokens || undefined}
|
totalTokens={tokenState?.totalTokens ?? session?.total_tokens ?? undefined}
|
||||||
inputTokens={session?.input_tokens || undefined}
|
accumulatedInputTokens={
|
||||||
outputTokens={session?.output_tokens || undefined}
|
tokenState?.accumulatedInputTokens ?? session?.accumulated_input_tokens ?? undefined
|
||||||
|
}
|
||||||
|
accumulatedOutputTokens={
|
||||||
|
tokenState?.accumulatedOutputTokens ?? session?.accumulated_output_tokens ?? undefined
|
||||||
|
}
|
||||||
droppedFiles={droppedFiles}
|
droppedFiles={droppedFiles}
|
||||||
onFilesProcessed={() => setDroppedFiles([])} // Clear dropped files after processing
|
onFilesProcessed={() => setDroppedFiles([])} // Clear dropped files after processing
|
||||||
messages={messages}
|
messages={messages}
|
||||||
|
|||||||
@@ -70,9 +70,9 @@ interface ChatInputProps {
|
|||||||
droppedFiles?: DroppedFile[];
|
droppedFiles?: DroppedFile[];
|
||||||
onFilesProcessed?: () => void; // Callback to clear dropped files after processing
|
onFilesProcessed?: () => void; // Callback to clear dropped files after processing
|
||||||
setView: (view: View) => void;
|
setView: (view: View) => void;
|
||||||
numTokens?: number;
|
totalTokens?: number;
|
||||||
inputTokens?: number;
|
accumulatedInputTokens?: number;
|
||||||
outputTokens?: number;
|
accumulatedOutputTokens?: number;
|
||||||
messages?: Message[];
|
messages?: Message[];
|
||||||
sessionCosts?: {
|
sessionCosts?: {
|
||||||
[key: string]: {
|
[key: string]: {
|
||||||
@@ -103,9 +103,9 @@ export default function ChatInput({
|
|||||||
droppedFiles = [],
|
droppedFiles = [],
|
||||||
onFilesProcessed,
|
onFilesProcessed,
|
||||||
setView,
|
setView,
|
||||||
numTokens,
|
totalTokens,
|
||||||
inputTokens,
|
accumulatedInputTokens,
|
||||||
outputTokens,
|
accumulatedOutputTokens,
|
||||||
messages = [],
|
messages = [],
|
||||||
disableAnimation = false,
|
disableAnimation = false,
|
||||||
sessionCosts,
|
sessionCosts,
|
||||||
@@ -505,16 +505,16 @@ export default function ChatInput({
|
|||||||
clearAlerts();
|
clearAlerts();
|
||||||
|
|
||||||
// Show alert when either there is registered token usage, or we know the limit
|
// Show alert when either there is registered token usage, or we know the limit
|
||||||
if ((numTokens && numTokens > 0) || (isTokenLimitLoaded && tokenLimit)) {
|
if ((totalTokens && totalTokens > 0) || (isTokenLimitLoaded && tokenLimit)) {
|
||||||
addAlert({
|
addAlert({
|
||||||
type: AlertType.Info,
|
type: AlertType.Info,
|
||||||
message: 'Context window',
|
message: 'Context window',
|
||||||
progress: {
|
progress: {
|
||||||
current: numTokens || 0,
|
current: totalTokens || 0,
|
||||||
total: tokenLimit,
|
total: tokenLimit,
|
||||||
},
|
},
|
||||||
showCompactButton: true,
|
showCompactButton: true,
|
||||||
compactButtonDisabled: !numTokens,
|
compactButtonDisabled: !totalTokens,
|
||||||
onCompact: () => {
|
onCompact: () => {
|
||||||
window.dispatchEvent(new CustomEvent('hide-alert-popover'));
|
window.dispatchEvent(new CustomEvent('hide-alert-popover'));
|
||||||
|
|
||||||
@@ -542,7 +542,7 @@ export default function ChatInput({
|
|||||||
}
|
}
|
||||||
// We intentionally omit setView as it shouldn't trigger a re-render of alerts
|
// We intentionally omit setView as it shouldn't trigger a re-render of alerts
|
||||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||||
}, [numTokens, toolCount, tokenLimit, isTokenLimitLoaded, addAlert, clearAlerts]);
|
}, [totalTokens, toolCount, tokenLimit, isTokenLimitLoaded, addAlert, clearAlerts]);
|
||||||
|
|
||||||
// Cleanup effect for component unmount - prevent memory leaks
|
// Cleanup effect for component unmount - prevent memory leaks
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
@@ -1540,8 +1540,8 @@ export default function ChatInput({
|
|||||||
<>
|
<>
|
||||||
<div className="flex items-center h-full ml-1 mr-1">
|
<div className="flex items-center h-full ml-1 mr-1">
|
||||||
<CostTracker
|
<CostTracker
|
||||||
inputTokens={inputTokens}
|
inputTokens={accumulatedInputTokens}
|
||||||
outputTokens={outputTokens}
|
outputTokens={accumulatedOutputTokens}
|
||||||
sessionCosts={sessionCosts}
|
sessionCosts={sessionCosts}
|
||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
|
|||||||
@@ -78,9 +78,9 @@ export default function Hub({
|
|||||||
commandHistory={[]}
|
commandHistory={[]}
|
||||||
initialValue=""
|
initialValue=""
|
||||||
setView={setView}
|
setView={setView}
|
||||||
numTokens={0}
|
totalTokens={0}
|
||||||
inputTokens={0}
|
accumulatedInputTokens={0}
|
||||||
outputTokens={0}
|
accumulatedOutputTokens={0}
|
||||||
droppedFiles={[]}
|
droppedFiles={[]}
|
||||||
onFilesProcessed={() => {}}
|
onFilesProcessed={() => {}}
|
||||||
messages={[]}
|
messages={[]}
|
||||||
|
|||||||
@@ -77,6 +77,7 @@ export const useChatEngine = ({
|
|||||||
notifications,
|
notifications,
|
||||||
session,
|
session,
|
||||||
setError,
|
setError,
|
||||||
|
tokenState,
|
||||||
} = useMessageStream({
|
} = useMessageStream({
|
||||||
api: getApiUrl('/reply'),
|
api: getApiUrl('/reply'),
|
||||||
id: chat.sessionId,
|
id: chat.sessionId,
|
||||||
@@ -451,6 +452,7 @@ export const useChatEngine = ({
|
|||||||
sessionOutputTokens,
|
sessionOutputTokens,
|
||||||
localInputTokens,
|
localInputTokens,
|
||||||
localOutputTokens,
|
localOutputTokens,
|
||||||
|
tokenState,
|
||||||
|
|
||||||
// UI helpers
|
// UI helpers
|
||||||
commandHistory,
|
commandHistory,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import {
|
|||||||
reply,
|
reply,
|
||||||
resumeAgent,
|
resumeAgent,
|
||||||
Session,
|
Session,
|
||||||
|
TokenState,
|
||||||
updateFromSession,
|
updateFromSession,
|
||||||
updateSessionUserRecipeValues,
|
updateSessionUserRecipeValues,
|
||||||
} from '../api';
|
} from '../api';
|
||||||
@@ -60,6 +61,7 @@ interface UseChatStreamReturn {
|
|||||||
setRecipeUserParams: (values: Record<string, string>) => Promise<void>;
|
setRecipeUserParams: (values: Record<string, string>) => Promise<void>;
|
||||||
stopStreaming: () => void;
|
stopStreaming: () => void;
|
||||||
sessionLoadError?: string;
|
sessionLoadError?: string;
|
||||||
|
tokenState: TokenState;
|
||||||
}
|
}
|
||||||
|
|
||||||
function pushMessage(currentMessages: Message[], incomingMsg: Message): Message[] {
|
function pushMessage(currentMessages: Message[], incomingMsg: Message): Message[] {
|
||||||
@@ -88,6 +90,7 @@ async function streamFromResponse(
|
|||||||
stream: AsyncIterable<MessageEvent>,
|
stream: AsyncIterable<MessageEvent>,
|
||||||
initialMessages: Message[],
|
initialMessages: Message[],
|
||||||
updateMessages: (messages: Message[]) => void,
|
updateMessages: (messages: Message[]) => void,
|
||||||
|
updateTokenState: (tokenState: TokenState) => void,
|
||||||
updateChatState: (state: ChatState) => void,
|
updateChatState: (state: ChatState) => void,
|
||||||
onFinish: (error?: string) => void
|
onFinish: (error?: string) => void
|
||||||
): Promise<void> {
|
): Promise<void> {
|
||||||
@@ -119,6 +122,8 @@ async function streamFromResponse(
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
updateTokenState(event.token_state);
|
||||||
|
|
||||||
updateMessages(currentMessages);
|
updateMessages(currentMessages);
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
@@ -171,6 +176,14 @@ export function useChatStream({
|
|||||||
const [session, setSession] = useState<Session>();
|
const [session, setSession] = useState<Session>();
|
||||||
const [sessionLoadError, setSessionLoadError] = useState<string>();
|
const [sessionLoadError, setSessionLoadError] = useState<string>();
|
||||||
const [chatState, setChatState] = useState<ChatState>(ChatState.Idle);
|
const [chatState, setChatState] = useState<ChatState>(ChatState.Idle);
|
||||||
|
const [tokenState, setTokenState] = useState<TokenState>({
|
||||||
|
inputTokens: 0,
|
||||||
|
outputTokens: 0,
|
||||||
|
totalTokens: 0,
|
||||||
|
accumulatedInputTokens: 0,
|
||||||
|
accumulatedOutputTokens: 0,
|
||||||
|
accumulatedTotalTokens: 0,
|
||||||
|
});
|
||||||
const abortControllerRef = useRef<AbortController | null>(null);
|
const abortControllerRef = useRef<AbortController | null>(null);
|
||||||
|
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
@@ -288,6 +301,7 @@ export function useChatStream({
|
|||||||
stream,
|
stream,
|
||||||
currentMessages,
|
currentMessages,
|
||||||
(messages: Message[]) => setMessagesAndLog(messages, 'streaming'),
|
(messages: Message[]) => setMessagesAndLog(messages, 'streaming'),
|
||||||
|
setTokenState,
|
||||||
setChatState,
|
setChatState,
|
||||||
onFinish
|
onFinish
|
||||||
);
|
);
|
||||||
@@ -373,5 +387,6 @@ export function useChatStream({
|
|||||||
handleSubmit,
|
handleSubmit,
|
||||||
stopStreaming,
|
stopStreaming,
|
||||||
setRecipeUserParams,
|
setRecipeUserParams,
|
||||||
|
tokenState,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import {
|
|||||||
getCompactingMessage,
|
getCompactingMessage,
|
||||||
hasCompletedToolCalls,
|
hasCompletedToolCalls,
|
||||||
} from '../types/message';
|
} from '../types/message';
|
||||||
import { Conversation, Message, Role } from '../api';
|
import { Conversation, Message, Role, TokenState } from '../api';
|
||||||
|
|
||||||
import { getSession, Session } from '../api';
|
import { getSession, Session } from '../api';
|
||||||
import { ChatState } from '../types/chatState';
|
import { ChatState } from '../types/chatState';
|
||||||
@@ -35,7 +35,7 @@ export interface NotificationEvent {
|
|||||||
|
|
||||||
// Event types for SSE stream
|
// Event types for SSE stream
|
||||||
type MessageEvent =
|
type MessageEvent =
|
||||||
| { type: 'Message'; message: Message }
|
| { type: 'Message'; message: Message; token_state: TokenState }
|
||||||
| { type: 'Error'; error: string }
|
| { type: 'Error'; error: string }
|
||||||
| { type: 'Finish'; reason: string }
|
| { type: 'Finish'; reason: string }
|
||||||
| { type: 'ModelChange'; model: string; mode: string }
|
| { type: 'ModelChange'; model: string; mode: string }
|
||||||
@@ -165,6 +165,9 @@ export interface UseMessageStreamHelpers {
|
|||||||
|
|
||||||
/** Clear error state */
|
/** Clear error state */
|
||||||
setError: (error: Error | undefined) => void;
|
setError: (error: Error | undefined) => void;
|
||||||
|
|
||||||
|
/** Real-time token state from server */
|
||||||
|
tokenState: TokenState;
|
||||||
}
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
@@ -197,6 +200,14 @@ export function useMessageStream({
|
|||||||
null
|
null
|
||||||
);
|
);
|
||||||
const [session, setSession] = useState<Session | null>(null);
|
const [session, setSession] = useState<Session | null>(null);
|
||||||
|
const [tokenState, setTokenState] = useState<TokenState>({
|
||||||
|
inputTokens: 0,
|
||||||
|
outputTokens: 0,
|
||||||
|
totalTokens: 0,
|
||||||
|
accumulatedInputTokens: 0,
|
||||||
|
accumulatedOutputTokens: 0,
|
||||||
|
accumulatedTotalTokens: 0,
|
||||||
|
});
|
||||||
|
|
||||||
// expose a way to update the body so we can update the session id when CLE occurs
|
// expose a way to update the body so we can update the session id when CLE occurs
|
||||||
const updateMessageStreamBody = useCallback((newBody: object) => {
|
const updateMessageStreamBody = useCallback((newBody: object) => {
|
||||||
@@ -280,6 +291,8 @@ export function useMessageStream({
|
|||||||
// Transition from waiting to streaming on first message
|
// Transition from waiting to streaming on first message
|
||||||
mutateChatState(ChatState.Streaming);
|
mutateChatState(ChatState.Streaming);
|
||||||
|
|
||||||
|
setTokenState(parsedEvent.token_state);
|
||||||
|
|
||||||
// Create a new message object with the properties preserved or defaulted
|
// Create a new message object with the properties preserved or defaulted
|
||||||
const newMessage: Message = {
|
const newMessage: Message = {
|
||||||
...parsedEvent.message,
|
...parsedEvent.message,
|
||||||
@@ -341,7 +354,6 @@ export function useMessageStream({
|
|||||||
}
|
}
|
||||||
|
|
||||||
case 'UpdateConversation': {
|
case 'UpdateConversation': {
|
||||||
currentMessages = parsedEvent.conversation;
|
|
||||||
setMessages(parsedEvent.conversation);
|
setMessages(parsedEvent.conversation);
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
@@ -650,5 +662,6 @@ export function useMessageStream({
|
|||||||
currentModelInfo,
|
currentModelInfo,
|
||||||
session,
|
session,
|
||||||
setError,
|
setError,
|
||||||
|
tokenState,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user