feat(token-usage): implement token usage tracking and display widget

This commit is contained in:
X-iZhang
2026-03-07 15:16:29 +00:00
parent cdb5507896
commit 250c41c561
9 changed files with 283 additions and 8 deletions
+6
View File
@@ -189,6 +189,7 @@ def run_textual_interactive(
TodoWidget,
UserMessage,
SystemMessage,
UsageWidget,
)
except Exception as e: # pragma: no cover - runtime fallback path
raise RuntimeError(
@@ -790,6 +791,11 @@ def run_textual_interactive(
lambda: container.scroll_end(animate=False),
),
)
# Mount token usage stats
if state.total_input_tokens or state.total_output_tokens:
await container.mount(
UsageWidget(state.total_input_tokens, state.total_output_tokens)
)
elif event_type == "error":
error_msg = event.get("message", "Unknown error")
+2
View File
@@ -8,6 +8,7 @@ from .subagent_widget import SubAgentWidget
from .todo_widget import TodoWidget
from .user_message import UserMessage
from .system_message import SystemMessage
from .usage_widget import UsageWidget
__all__ = [
"LoadingWidget",
@@ -18,4 +19,5 @@ __all__ = [
"TodoWidget",
"UserMessage",
"SystemMessage",
"UsageWidget",
]
+29
View File
@@ -0,0 +1,29 @@
"""Token usage statistics widget."""
from __future__ import annotations
from rich.text import Text
from textual.widgets import Static
class UsageWidget(Static):
"""Displays token usage stats, right-aligned with styled numbers."""
DEFAULT_CSS = """
UsageWidget {
height: auto;
text-align: right;
}
"""
def __init__(self, input_tokens: int, output_tokens: int) -> None:
stats = Text(justify="right")
stats.append("[", style="dim italic")
stats.append("Usage: ", style="dim italic")
stats.append(f"{input_tokens:,}", style="cyan italic")
stats.append(" in · ", style="dim italic")
stats.append(f"{output_tokens:,}", style="green italic")
stats.append(" out", style="dim italic")
stats.append("]", style="dim italic")
super().__init__(stats)
+26
View File
@@ -353,6 +353,8 @@ def create_streaming_display(
final_show_thinking: bool = False,
final_thinking_max_length: int = DisplayLimits.THINKING_FINAL,
response_markdown: Any = None,
total_input_tokens: int = 0,
total_output_tokens: int = 0,
) -> Any:
"""Create Rich display layout for streaming output.
@@ -524,6 +526,18 @@ def create_streaming_display(
if clean_response:
elements.append(Text("")) # blank separator
elements.append(response_markdown or Markdown(clean_response))
# Token usage stats (right-aligned)
if total_input_tokens or total_output_tokens:
stats = Text(justify="right")
stats.append("[", style="dim italic")
stats.append("Usage: ", style="dim italic")
stats.append(f"{total_input_tokens:,}", style="cyan italic")
stats.append(" in · ", style="dim italic")
stats.append(f"{total_output_tokens:,}", style="green italic")
stats.append(" out", style="dim italic")
stats.append("]", style="dim italic")
elements.append(stats)
else:
# Intermediate narration (tools still running) -- dim italic above Task List
if latest_text and has_used_tools and not all_done:
@@ -648,6 +662,18 @@ def display_final_results(
console.print()
console.print(Markdown(clean_response or state.response_text))
# Token usage stats (right-aligned)
if state.total_input_tokens or state.total_output_tokens:
stats = Text(justify="right")
stats.append("[", style="dim italic")
stats.append("Usage: ", style="dim italic")
stats.append(f"{state.total_input_tokens:,}", style="cyan italic")
stats.append(" in · ", style="dim italic")
stats.append(f"{state.total_output_tokens:,}", style="green italic")
stats.append(" out", style="dim italic")
stats.append("]", style="dim italic")
console.print(stats)
# ---------------------------------------------------------------------------
# Async-to-sync bridge
+9
View File
@@ -88,6 +88,15 @@ class StreamEventEmitter:
"""Done event."""
return StreamEvent("done", {"type": "done", "content": response, "response": response})
@staticmethod
def usage_stats(input_tokens: int, output_tokens: int) -> StreamEvent:
"""Token usage statistics event."""
return StreamEvent("usage_stats", {
"type": "usage_stats",
"input_tokens": input_tokens,
"output_tokens": output_tokens,
})
@staticmethod
def error(message: str) -> StreamEvent:
"""Error event."""
+36 -7
View File
@@ -293,22 +293,38 @@ async def stream_agent_events(
async for chunk in agent.astream(
{"messages": [{"role": "user", "content": user_content}]},
config=config,
stream_mode="messages",
stream_mode=["messages", "updates"],
subgraphs=True,
):
# With subgraphs=True, event is (namespace, (message, metadata))
namespace: tuple = ()
data: Any = chunk
# Multi-mode + subgraphs: 3-tuple (namespace, mode, data)
# Single-mode + subgraphs: 2-tuple (namespace, data) — fallback
if not isinstance(chunk, tuple):
continue
if isinstance(chunk, tuple) and len(chunk) >= 2:
namespace: tuple = ()
data: Any
mode_str: str
if len(chunk) == 3:
namespace, mode_str, data = chunk
if not isinstance(namespace, tuple):
namespace = ()
elif len(chunk) == 2:
first = chunk[0]
if isinstance(first, tuple):
# (namespace_tuple, (message, metadata))
namespace = first
data = chunk[1]
else:
# (message, metadata) -- no namespace
data = chunk
mode_str = "messages"
else:
continue
# Skip non-messages modes (updates, etc.)
if mode_str == "updates":
continue
if mode_str != "messages":
continue
# Unpack message + metadata from data
msg: Any
@@ -319,12 +335,25 @@ async def stream_agent_events(
else:
msg = data
# Filter summarization middleware synthetic messages
if isinstance(metadata, dict) and metadata.get("lc_source") == "summarization":
continue
subagent = _get_subagent_name(namespace, metadata)
subagent_tracker = None
if subagent:
tracker_key = _get_subagent_key(namespace, metadata) or str(namespace)
subagent_tracker = _subagent_trackers.setdefault(tracker_key, ToolCallTracker())
# Extract token usage from main-agent AIMessages
if isinstance(msg, (AIMessageChunk, AIMessage)) and not subagent:
usage = getattr(msg, "usage_metadata", None)
if usage:
inp = usage.get("input_tokens", 0) if isinstance(usage, dict) else getattr(usage, "input_tokens", 0)
out = usage.get("output_tokens", 0) if isinstance(usage, dict) else getattr(usage, "output_tokens", 0)
if inp or out:
yield emitter.usage_stats(inp, out).data
# Process AIMessageChunk / AIMessage
if isinstance(msg, (AIMessageChunk, AIMessage)):
if subagent:
+9
View File
@@ -92,6 +92,9 @@ class StreamState:
self.todo_items: list[dict] = []
# Latest text segment (reset on each tool_call)
self.latest_text = ""
# Token usage tracking
self.total_input_tokens = 0
self.total_output_tokens = 0
# Cached Markdown object for Rich CLI display (avoids O(n²) re-parsing)
self._cached_md_text: str = ""
self._cached_md: object | None = None
@@ -251,6 +254,10 @@ class StreamState:
sa.is_active = False
break
elif event_type == "usage_stats":
self.total_input_tokens += event.get("input_tokens", 0)
self.total_output_tokens += event.get("output_tokens", 0)
elif event_type == "done":
self.is_processing = False
if not self.response_text:
@@ -278,6 +285,8 @@ class StreamState:
"is_processing": self.is_processing,
"subagents": self.subagents,
"todo_items": self.todo_items,
"total_input_tokens": self.total_input_tokens,
"total_output_tokens": self.total_output_tokens,
}
+129 -1
View File
@@ -1,8 +1,12 @@
"""Tests for EvoScientist/stream/events.py helpers."""
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
from EvoScientist.stream.events import _extract_tool_content
from langchain_core.messages import AIMessageChunk
from EvoScientist.stream.events import _extract_tool_content, stream_agent_events
class TestExtractToolContent:
@@ -85,3 +89,127 @@ class TestExtractToolContent:
content, is_image = _extract_tool_content(msg)
assert is_image is False
assert content == "some result"
# =============================================================================
# Multi-mode streaming chunk unpacking
# =============================================================================
def _make_ai_chunk(content: str = "hello", **kwargs):
"""Create a minimal AIMessageChunk for testing."""
return AIMessageChunk(content=content, **kwargs)
def _collect_events(agent, message="hi", thread_id="t1"):
"""Collect all events from stream_agent_events synchronously."""
async def _run():
events = []
async for ev in stream_agent_events(agent, message, thread_id):
events.append(ev)
return events
loop = asyncio.new_event_loop()
try:
return loop.run_until_complete(_run())
finally:
loop.close()
async def _async_iter(items):
"""Create an async iterator from a list."""
for item in items:
yield item
class TestMultiModeChunkUnpacking:
"""Test 3-tuple (multi-mode) and 2-tuple (single-mode) chunk handling."""
def test_3tuple_chunk_unpacking(self):
"""Multi-mode yields 3-tuples (namespace, mode, data); messages are processed."""
chunk = _make_ai_chunk("hello world")
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), "messages", (chunk, {})),
]))
events = _collect_events(mock_agent)
text_events = [e for e in events if e.get("type") == "text"]
assert len(text_events) == 1
assert text_events[0]["content"] == "hello world"
def test_2tuple_fallback(self):
"""Single-mode yields 2-tuples; should still work."""
chunk = _make_ai_chunk("fallback")
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), (chunk, {})),
]))
events = _collect_events(mock_agent)
text_events = [e for e in events if e.get("type") == "text"]
assert len(text_events) == 1
assert text_events[0]["content"] == "fallback"
def test_updates_mode_graceful_skip(self):
"""Updates mode chunks are skipped without error."""
chunk = _make_ai_chunk("should appear")
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), "updates", {"some": "state"}),
((), "messages", (chunk, {})),
]))
events = _collect_events(mock_agent)
text_events = [e for e in events if e.get("type") == "text"]
assert len(text_events) == 1
assert text_events[0]["content"] == "should appear"
def test_summarization_filtered(self):
"""Chunks with lc_source=summarization metadata are filtered out."""
chunk_real = _make_ai_chunk("real content")
chunk_synth = _make_ai_chunk("synthetic summary")
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), "messages", (chunk_synth, {"lc_source": "summarization"})),
((), "messages", (chunk_real, {})),
]))
events = _collect_events(mock_agent)
text_events = [e for e in events if e.get("type") == "text"]
assert len(text_events) == 1
assert text_events[0]["content"] == "real content"
class TestUsageStatsExtraction:
"""Test token usage extraction from AIMessageChunk."""
def test_usage_metadata_emitted(self):
"""AIMessageChunk with usage_metadata emits usage_stats event."""
chunk = _make_ai_chunk("hi", usage_metadata={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150})
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), "messages", (chunk, {})),
]))
events = _collect_events(mock_agent)
usage_events = [e for e in events if e.get("type") == "usage_stats"]
assert len(usage_events) == 1
assert usage_events[0]["input_tokens"] == 100
assert usage_events[0]["output_tokens"] == 50
def test_no_usage_metadata_no_event(self):
"""AIMessageChunk without usage_metadata does not emit usage_stats."""
chunk = _make_ai_chunk("hi")
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), "messages", (chunk, {})),
]))
events = _collect_events(mock_agent)
usage_events = [e for e in events if e.get("type") == "usage_stats"]
assert len(usage_events) == 0
def test_zero_tokens_not_emitted(self):
"""Zero input and output tokens should not emit usage_stats."""
chunk = _make_ai_chunk("hi", usage_metadata={"input_tokens": 0, "output_tokens": 0, "total_tokens": 0})
mock_agent = AsyncMock()
mock_agent.astream = MagicMock(return_value=_async_iter([
((), "messages", (chunk, {})),
]))
events = _collect_events(mock_agent)
usage_events = [e for e in events if e.get("type") == "usage_stats"]
assert len(usage_events) == 0
+37
View File
@@ -553,6 +553,43 @@ class TestParseTodoItemsAdvanced:
assert _parse_todo_items('["a", "b"]') is None
# =============================================================================
# Token usage tracking
# =============================================================================
class TestUsageStatsAccumulated:
def test_single_usage_event(self):
state = StreamState()
state.handle_event({"type": "usage_stats", "input_tokens": 100, "output_tokens": 50})
assert state.total_input_tokens == 100
assert state.total_output_tokens == 50
def test_multiple_usage_events_accumulate(self):
state = StreamState()
state.handle_event({"type": "usage_stats", "input_tokens": 100, "output_tokens": 50})
state.handle_event({"type": "usage_stats", "input_tokens": 200, "output_tokens": 80})
assert state.total_input_tokens == 300
assert state.total_output_tokens == 130
def test_usage_stats_in_display_args(self):
state = StreamState()
state.handle_event({"type": "usage_stats", "input_tokens": 500, "output_tokens": 200})
args = state.get_display_args()
assert args["total_input_tokens"] == 500
assert args["total_output_tokens"] == 200
def test_default_zero_tokens(self):
state = StreamState()
args = state.get_display_args()
assert args["total_input_tokens"] == 0
assert args["total_output_tokens"] == 0
def test_usage_stats_returns_event_type(self):
state = StreamState()
result = state.handle_event({"type": "usage_stats", "input_tokens": 10, "output_tokens": 5})
assert result == "usage_stats"
# =============================================================================
# ChannelState queue mechanism (removed — replaced by bus mode in channel.py)
# =============================================================================