Skip to content

Commit 802f95b

Browse files
authored
fix(ui): correctly display token usage across sessions (kagent-dev#1509)
When testing kagent-dev#1500, I struggled to visually demonstrate the affect of the duplicated session history (aka context bloat) when chatting with agents. The current implementation of "Usage" within the UI (introduced in kagent-dev#679) uses `Math.max` independently per field (`total`, `prompt`, `completion`) across all tasks in a session. This produces incoherent values where `total` does not necessarily match `prompt` + `completion`, since each field[^1] (most notably `prompt`) can peak during differing tasks. And, it doesn't correctly update to reflect usage across multiple API calls in a multi-turn conversation. All we see is which API call has the greatest of each type of token usage. This PR addresses those shortcomings by: - Updating the existing token usage display to reflect session total across multiple turns - Show per-call token usage tooltip on tool call cards and text messages [^1]: not 100% certain here but I think that fields other than input could also peak at different levels within cached prompt scenarios. Signed-off-by: Brian Fox <878612+onematchfox@users.noreply.github.com>
1 parent b1cf420 commit 802f95b

13 files changed

Lines changed: 572 additions & 156 deletions

File tree

python/packages/kagent-adk/src/kagent/adk/_agent_executor.py

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -53,7 +53,7 @@
5353
)
5454

5555
from ._mcp_toolset import is_anyio_cross_task_cancel_scope_error
56-
from .converters.event_converter import convert_event_to_a2a_events
56+
from .converters.event_converter import convert_event_to_a2a_events, serialize_metadata_value
5757
from .converters.part_converter import convert_a2a_part_to_genai_part, convert_genai_part_to_a2a_part
5858
from .converters.request_converter import convert_a2a_request_to_adk_run_args
5959

@@ -585,6 +585,7 @@ async def _handle_request(
585585
# For streaming A2A update events, the invocation_id is added through event converter
586586
# This adds the invocation_id of the run to the metadata of the FINAL event (completed or failed)
587587
real_invocation_id: str | None = None
588+
last_usage_metadata = None
588589

589590
task_result_aggregator = TaskResultAggregator()
590591
async with Aclosing(runner.run_async(**run_args)) as agen:
@@ -595,6 +596,12 @@ async def _handle_request(
595596
real_invocation_id = event_inv_id
596597
run_metadata[get_kagent_metadata_key("invocation_id")] = real_invocation_id
597598

599+
# Track the last usage_metadata so it can be included in the final
600+
# event's run_metadata. The A2A task_manager merges run_metadata into
601+
# task.metadata, making it available to callers (e.g. KAgentRemoteA2ATool).
602+
if getattr(adk_event, "usage_metadata", None) is not None:
603+
last_usage_metadata = adk_event.usage_metadata
604+
598605
for a2a_event in convert_event_to_a2a_events(
599606
adk_event, invocation_context, context.task_id, context.context_id
600607
):
@@ -608,6 +615,11 @@ async def _handle_request(
608615
if getattr(adk_event, "long_running_tool_ids", None):
609616
break
610617

618+
# Attach the last LLM usage to run_metadata so the A2A task_manager
619+
# merges it into task.metadata on the completed Task object.
620+
if last_usage_metadata is not None:
621+
run_metadata[get_kagent_metadata_key("usage_metadata")] = serialize_metadata_value(last_usage_metadata)
622+
611623
# publish the task result event - this is final
612624
if (
613625
task_result_aggregator.task_state == TaskState.working

python/packages/kagent-adk/src/kagent/adk/_remote_a2a_tool.py

Lines changed: 27 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,21 @@ def _extract_text_from_task(task: Task) -> str:
8383
return ""
8484

8585

86+
def _extract_usage_from_task(task: Task) -> Optional[dict]:
87+
"""Extract kagent_usage_metadata from a completed task.
88+
89+
The A2A task_manager merges the final TaskStatusUpdateEvent.metadata into
90+
task.metadata. The agent executor now adds the last LLM invocation's
91+
usage_metadata to run_metadata before publishing the final event, so it
92+
is available here for non-streaming callers like KAgentRemoteA2ATool.
93+
"""
94+
if task.metadata:
95+
usage = task.metadata.get("kagent_usage_metadata")
96+
if usage and isinstance(usage, dict):
97+
return usage
98+
return None
99+
100+
86101
class KAgentRemoteA2ATool(BaseTool):
87102
"""A tool that calls a remote A2A agent and propagates HITL state."""
88103

@@ -215,8 +230,13 @@ async def _handle_first_call(self, args: dict[str, Any], tool_context: ToolConte
215230
error_text = _extract_text_from_task(task)
216231
return error_text or f"Remote agent '{self.name}' failed."
217232

218-
# completed or any other terminal state
219-
return _extract_text_from_task(task) or ""
233+
# completed — include the sub-agent's final LLM usage from task.metadata
234+
# so the parent can display it on the AgentCall card in the UI.
235+
result_text = _extract_text_from_task(task)
236+
usage = _extract_usage_from_task(task)
237+
if usage:
238+
return {"result": result_text, "kagent_usage_metadata": usage}
239+
return result_text or ""
220240

221241
def _handle_input_required(self, task: Task, tool_context: ToolContext) -> dict[str, Any]:
222242
"""Handle a subagent that returned input_required (HITL).
@@ -351,7 +371,11 @@ async def _handle_resume(self, tool_context: ToolContext) -> Any:
351371
error_text = _extract_text_from_task(task)
352372
return error_text or f"Remote agent '{subagent_name}' failed after resume."
353373

354-
return _extract_text_from_task(task) or ""
374+
result_text = _extract_text_from_task(task)
375+
usage = _extract_usage_from_task(task)
376+
if usage:
377+
return {"result": result_text, "kagent_usage_metadata": usage}
378+
return result_text or ""
355379

356380
@staticmethod
357381
def _extract_text_from_message(message: A2AMessage) -> str:

python/packages/kagent-adk/src/kagent/adk/converters/event_converter.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,7 @@
3333
logger = logging.getLogger("kagent_adk." + __name__)
3434

3535

36-
def _serialize_metadata_value(value: Any) -> str:
36+
def serialize_metadata_value(value: Any) -> str:
3737
"""Safely serializes metadata values to string format.
3838
3939
Args:
@@ -90,7 +90,7 @@ def _get_context_metadata(event: Event, invocation_context: InvocationContext) -
9090

9191
for field_name, field_value in optional_fields:
9292
if field_value is not None:
93-
metadata[get_kagent_metadata_key(field_name)] = _serialize_metadata_value(field_value)
93+
metadata[get_kagent_metadata_key(field_name)] = serialize_metadata_value(field_value)
9494

9595
return metadata
9696

ui/src/components/ToolDisplay.tsx

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,11 @@
11
import { useState } from "react";
2-
import { FunctionCall } from "@/types";
2+
import { FunctionCall, TokenStats } from "@/types";
33
import { ScrollArea } from "@radix-ui/react-scroll-area";
44
import { FunctionSquare, CheckCircle, Clock, Code, ChevronUp, ChevronDown, Loader2, Text, Check, Copy, AlertCircle, ShieldAlert } from "lucide-react";
55
import { Button } from "@/components/ui/button";
66
import { Textarea } from "@/components/ui/textarea";
77
import { Card, CardHeader, CardTitle, CardContent } from "@/components/ui/card";
8+
import TokenStatsTooltip from "@/components/chat/TokenStatsTooltip";
89
import { convertToUserFriendlyName } from "@/lib/utils";
910

1011
export type ToolCallStatus = "requested" | "executing" | "completed" | "pending_approval" | "approved" | "rejected";
@@ -23,9 +24,10 @@ interface ToolDisplayProps {
2324
subagentName?: string;
2425
onApprove?: () => void;
2526
onReject?: (reason?: string) => void;
27+
tokenStats?: TokenStats;
2628
}
2729

28-
const ToolDisplay = ({ call, result, status = "requested", isError = false, isDecided = false, subagentName, onApprove, onReject }: ToolDisplayProps) => {
30+
const ToolDisplay = ({ call, result, status = "requested", isError = false, isDecided = false, subagentName, onApprove, onReject, tokenStats }: ToolDisplayProps) => {
2931
const [areArgumentsExpanded, setAreArgumentsExpanded] = useState(status === "pending_approval");
3032
const [areResultsExpanded, setAreResultsExpanded] = useState(false);
3133
const [isCopied, setIsCopied] = useState(false);
@@ -166,7 +168,8 @@ const ToolDisplay = ({ call, result, status = "requested", isError = false, isDe
166168
)}
167169
<div className="font-light">{call.id}</div>
168170
</CardTitle>
169-
<div className="flex justify-center items-center text-xs">
171+
<div className="flex items-center gap-2 text-xs">
172+
{tokenStats && <TokenStatsTooltip stats={tokenStats} />}
170173
{getStatusDisplay()}
171174
</div>
172175
</CardHeader>

ui/src/components/chat/AgentCallDisplay.tsx

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
11
import { useMemo, useState } from "react";
2-
import { FunctionCall } from "@/types";
2+
import { FunctionCall, TokenStats } from "@/types";
33
import { Card, CardHeader, CardTitle, CardContent } from "@/components/ui/card";
44
import { convertToUserFriendlyName } from "@/lib/utils";
55
import { ChevronDown, ChevronUp, MessageSquare, Loader2, AlertCircle, CheckCircle } from "lucide-react";
66
import KagentLogo from "../kagent-logo";
7+
import TokenStatsTooltip from "@/components/chat/TokenStatsTooltip";
78

89
export type AgentCallStatus = "requested" | "executing" | "completed";
910

@@ -15,9 +16,10 @@ interface AgentCallDisplayProps {
1516
};
1617
status?: AgentCallStatus;
1718
isError?: boolean;
19+
tokenStats?: TokenStats;
1820
}
1921

20-
const AgentCallDisplay = ({ call, result, status = "requested", isError = false }: AgentCallDisplayProps) => {
22+
const AgentCallDisplay = ({ call, result, status = "requested", isError = false, tokenStats }: AgentCallDisplayProps) => {
2123
const [areInputsExpanded, setAreInputsExpanded] = useState(false);
2224
const [areResultsExpanded, setAreResultsExpanded] = useState(false);
2325

@@ -78,7 +80,8 @@ const AgentCallDisplay = ({ call, result, status = "requested", isError = false
7880
</div>
7981
<div className="font-light">{call.id}</div>
8082
</CardTitle>
81-
<div className="flex justify-center items-center text-xs">
83+
<div className="flex items-center gap-2 text-xs">
84+
{tokenStats && <TokenStatsTooltip stats={tokenStats} />}
8285
{getStatusDisplay()}
8386
</div>
8487
</CardHeader>

ui/src/components/chat/ChatInterface.tsx

Lines changed: 15 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
"use client";
22

33
import type React from "react";
4-
import { useState, useRef, useEffect } from "react";
4+
import { useState, useRef, useEffect, useMemo } from "react";
55
import { ArrowBigUp, X, Loader2, Mic, Square } from "lucide-react";
66
import { Button } from "@/components/ui/button";
77
import {
@@ -15,7 +15,7 @@ import { Textarea } from "@/components/ui/textarea";
1515
import { ScrollArea } from "@/components/ui/scroll-area";
1616
import ChatMessage from "@/components/chat/ChatMessage";
1717
import StreamingMessage from "./StreamingMessage";
18-
import TokenStatsDisplay from "./TokenStats";
18+
import SessionTokenStatsDisplay from "@/components/chat/TokenStats";
1919
import type { TokenStats, Session, ChatStatus, ToolDecision } from "@/types";
2020
import StatusDisplay from "./StatusDisplay";
2121
import { createSession, getSessionTasks, checkSessionExists } from "@/app/actions/sessions";
@@ -39,11 +39,6 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
3939
const router = useRouter();
4040
const containerRef = useRef<HTMLDivElement>(null);
4141
const [currentInputMessage, setCurrentInputMessage] = useState("");
42-
const [tokenStats, setTokenStats] = useState<TokenStats>({
43-
total: 0,
44-
input: 0,
45-
output: 0,
46-
});
4742

4843
const [chatStatus, setChatStatus] = useState<ChatStatus>("ready");
4944

@@ -58,6 +53,9 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
5853
const [sessionNotFound, setSessionNotFound] = useState<boolean>(false);
5954
const isCreatingSessionRef = useRef<boolean>(false);
6055
const [isFirstMessage, setIsFirstMessage] = useState<boolean>(!sessionId);
56+
const [sessionStats, setSessionStats] = useState<TokenStats>({ total: 0, prompt: 0, completion: 0 });
57+
// Mutable ref so pendingTurnStats survives re-renders between A2A stream events
58+
const pendingTurnStatsRef = useRef<TokenStats | undefined>(undefined);
6159
const [pendingDecisions, setPendingDecisions] = useState<Record<string, ToolDecision>>({});
6260
const pendingDecisionsRef = useRef<Record<string, ToolDecision>>({});
6361
/** Per-tool rejection reasons collected as the user rejects individual tools. */
@@ -78,25 +76,27 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
7876
},
7977
});
8078

81-
const { handleMessageEvent } = createMessageHandlers({
79+
const { handleMessageEvent } = useMemo(() => createMessageHandlers({
8280
setMessages: setStreamingMessages,
8381
setIsStreaming,
8482
setStreamingContent,
85-
setTokenStats,
8683
setChatStatus,
84+
setSessionStats,
85+
pendingTurnStats: pendingTurnStatsRef,
8786
agentContext: {
8887
namespace: selectedNamespace,
8988
agentName: selectedAgentName
9089
}
91-
});
90+
}), [selectedNamespace, selectedAgentName]);
9291

9392
useEffect(() => {
9493
async function initializeChat() {
95-
setTokenStats({ total: 0, input: 0, output: 0 });
94+
setSessionStats({ total: 0, prompt: 0, completion: 0 });
9695
setStreamingMessages([]);
9796
setPendingDecisions({});
9897
pendingDecisionsRef.current = {};
9998
pendingRejectionReasonsRef.current = {};
99+
pendingTurnStatsRef.current = undefined;
100100

101101
// Skip completely if this is a first message session creation flow
102102
if (isFirstMessage || isCreatingSessionRef.current) {
@@ -129,11 +129,11 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
129129
}
130130
if (!messagesResponse.data || messagesResponse?.data?.length === 0) {
131131
setStoredMessages([]);
132-
setTokenStats({ total: 0, input: 0, output: 0 });
132+
setSessionStats({ total: 0, prompt: 0, completion: 0 });
133133
}
134134
else {
135135
const extractedMessages = extractMessagesFromTasks(messagesResponse.data);
136-
const extractedTokenStats = extractTokenStatsFromTasks(messagesResponse.data);
136+
setSessionStats(extractTokenStatsFromTasks(messagesResponse.data));
137137

138138
// Resolved approvals are already inline in extractedMessages (with
139139
// approved/rejected badges). Only pending approvals need appending.
@@ -144,7 +144,6 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
144144
? [...extractedMessages, ...pendingApprovalMessages]
145145
: extractedMessages
146146
);
147-
setTokenStats(extractedTokenStats);
148147

149148
if (hasPendingApproval) {
150149
setChatStatus("input_required");
@@ -192,6 +191,7 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
192191
setPendingDecisions({});
193192
pendingDecisionsRef.current = {};
194193
pendingRejectionReasonsRef.current = {};
194+
pendingTurnStatsRef.current = undefined;
195195

196196
// For new sessions or when no stored messages exist, show the user message immediately
197197
const userMessage: Message = {
@@ -682,7 +682,7 @@ export default function ChatInterface({ selectedAgentName, selectedNamespace, se
682682
<div className="w-full sticky bg-secondary bottom-0 md:bottom-2 rounded-none md:rounded-lg p-4 border overflow-hidden transition-all duration-300 ease-in-out">
683683
<div className="flex items-center justify-between mb-4">
684684
<StatusDisplay chatStatus={chatStatus} />
685-
<TokenStatsDisplay stats={tokenStats} />
685+
{sessionStats.total > 0 && <SessionTokenStatsDisplay stats={sessionStats} />}
686686
</div>
687687

688688
<form onSubmit={handleSendMessage}>

ui/src/components/chat/ChatMessage.tsx

Lines changed: 25 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,8 @@ import ToolCallDisplay from "@/components/chat/ToolCallDisplay";
44
import AskUserDisplay, { AskUserQuestion } from "@/components/chat/AskUserDisplay";
55
import KagentLogo from "../kagent-logo";
66
import { ThumbsUp, ThumbsDown } from "lucide-react";
7+
import TokenStatsTooltip from "@/components/chat/TokenStatsTooltip";
8+
import type { TokenStats } from "@/types";
79
import { useState } from "react";
810
import { FeedbackDialog } from "./FeedbackDialog";
911
import { toast } from "sonner";
@@ -28,10 +30,13 @@ export default function ChatMessage({ message, allMessages, agentContext, onAppr
2830
const [feedbackDialogOpen, setFeedbackDialogOpen] = useState(false);
2931
const [isPositiveFeedback, setIsPositiveFeedback] = useState(true);
3032

33+
if (!message) return null;
34+
3135
const textParts = message.parts?.filter(part => part.kind === "text") || [];
3236
const content = textParts.map(part => (part as TextPart).text).join("");
3337

3438
const source = message.role === "user" ? "user" : "assistant";
39+
const tokenStats = (message.metadata as Record<string, unknown> | undefined)?.tokenStats as TokenStats | undefined;
3540
const messageId = message.messageId;
3641

3742
// Extract agent name from metadata for display
@@ -68,10 +73,6 @@ export default function ChatMessage({ message, allMessages, agentContext, onAppr
6873
return a & a;
6974
}, 0)) : 0;
7075

71-
if (!message) {
72-
return null;
73-
}
74-
7576
const metadata = message.metadata as ADKMetadata;
7677
const originalType = metadata?.originalType;
7778

@@ -168,22 +169,27 @@ export default function ChatMessage({ message, allMessages, agentContext, onAppr
168169
<div className="text-xs font-bold">{displayName}</div>
169170
</div> : <div className="text-xs font-bold">{displayName}</div>}
170171
<TruncatableText content={String(content)} className="break-all text-primary-foreground" />
171-
{source !== "user" && messageId !== undefined && (
172+
{source !== "user" && (
172173
<div className="flex mt-2 justify-end items-center gap-2">
173-
<button
174-
onClick={() => handleFeedback(true)}
175-
className="p-1 rounded-full hover:bg-gray-200 dark:hover:bg-gray-700 transition-colors"
176-
aria-label="Thumbs up"
177-
>
178-
<ThumbsUp className="w-4 h-4" />
179-
</button>
180-
<button
181-
onClick={() => handleFeedback(false)}
182-
className="p-1 rounded-full hover:bg-gray-200 dark:hover:bg-gray-700 transition-colors"
183-
aria-label="Thumbs down"
184-
>
185-
<ThumbsDown className="w-4 h-4" />
186-
</button>
174+
{tokenStats && <TokenStatsTooltip stats={tokenStats} />}
175+
{messageId !== undefined && (
176+
<>
177+
<button
178+
onClick={() => handleFeedback(true)}
179+
className="p-1 rounded-full hover:bg-gray-200 dark:hover:bg-gray-700 transition-colors"
180+
aria-label="Thumbs up"
181+
>
182+
<ThumbsUp className="w-4 h-4" />
183+
</button>
184+
<button
185+
onClick={() => handleFeedback(false)}
186+
className="p-1 rounded-full hover:bg-gray-200 dark:hover:bg-gray-700 transition-colors"
187+
aria-label="Thumbs down"
188+
>
189+
<ThumbsDown className="w-4 h-4" />
190+
</button>
191+
</>
192+
)}
187193
</div>
188194
)}
189195
</div>

ui/src/components/chat/TokenStats.tsx

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,23 +1,23 @@
11
import { ArrowLeft, ArrowRightFromLine } from "lucide-react";
22
import { TokenStats } from "@/types";
33

4-
interface TokenStatsDisplayProps {
4+
interface SessionTokenStatsDisplayProps {
55
stats: TokenStats;
66
}
77

8-
export default function TokenStatsDisplay({ stats }: TokenStatsDisplayProps) {
8+
export default function SessionTokenStatsDisplay({ stats }: SessionTokenStatsDisplayProps) {
99
return (
1010
<div className="flex items-center gap-2 text-xs">
1111
<span>Usage: </span>
1212
<span>{stats.total}</span>
1313
<div className="flex items-center gap-2">
1414
<div className="flex items-center gap-1">
1515
<ArrowLeft className="h-3 w-3" />
16-
<span>{stats.input}</span>
16+
<span>{stats.prompt}</span>
1717
</div>
1818
<div className="flex items-center gap-1">
1919
<ArrowRightFromLine className="h-3 w-3" />
20-
<span>{stats.output}</span>
20+
<span>{stats.completion}</span>
2121
</div>
2222
</div>
2323
</div>

0 commit comments

Comments
 (0)