From 250c41c561836947d979d19be278655800a2cc76 Mon Sep 17 00:00:00 2001 From: X-iZhang Date: Sat, 7 Mar 2026 15:16:29 +0000 Subject: [PATCH] feat(token-usage): implement token usage tracking and display widget --- EvoScientist/cli/tui_interactive.py | 6 ++ EvoScientist/cli/widgets/__init__.py | 2 + EvoScientist/cli/widgets/usage_widget.py | 29 +++++ EvoScientist/stream/display.py | 26 +++++ EvoScientist/stream/emitter.py | 9 ++ EvoScientist/stream/events.py | 43 ++++++-- EvoScientist/stream/state.py | 9 ++ tests/test_stream_events.py | 130 ++++++++++++++++++++++- tests/test_stream_state.py | 37 +++++++ 9 files changed, 283 insertions(+), 8 deletions(-) create mode 100644 EvoScientist/cli/widgets/usage_widget.py diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 7e70f5e..89b4e50 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -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") diff --git a/EvoScientist/cli/widgets/__init__.py b/EvoScientist/cli/widgets/__init__.py index 487d48f..8e0283e 100644 --- a/EvoScientist/cli/widgets/__init__.py +++ b/EvoScientist/cli/widgets/__init__.py @@ -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", ] diff --git a/EvoScientist/cli/widgets/usage_widget.py b/EvoScientist/cli/widgets/usage_widget.py new file mode 100644 index 0000000..1f301a4 --- /dev/null +++ b/EvoScientist/cli/widgets/usage_widget.py @@ -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) diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index 95228e5..b3ef324 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -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 diff --git a/EvoScientist/stream/emitter.py b/EvoScientist/stream/emitter.py index b6bb2bc..4fcfca3 100644 --- a/EvoScientist/stream/emitter.py +++ b/EvoScientist/stream/emitter.py @@ -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.""" diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index 0ed9ea4..7ccc323 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -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: diff --git a/EvoScientist/stream/state.py b/EvoScientist/stream/state.py index f8a364c..ad92550 100644 --- a/EvoScientist/stream/state.py +++ b/EvoScientist/stream/state.py @@ -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, } diff --git a/tests/test_stream_events.py b/tests/test_stream_events.py index 2acebc1..48aba12 100644 --- a/tests/test_stream_events.py +++ b/tests/test_stream_events.py @@ -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 diff --git a/tests/test_stream_state.py b/tests/test_stream_state.py index 29a3d00..2563bb9 100644 --- a/tests/test_stream_state.py +++ b/tests/test_stream_state.py @@ -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) # =============================================================================