feat(token-usage): implement token usage tracking and display widget
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
# =============================================================================
|
||||
|
||||
Reference in New Issue
Block a user