From 8bb1d6c0e3a11af9203c2b7d4dfc42a66b975210 Mon Sep 17 00:00:00 2001 From: dinos Date: Mon, 8 Jun 2026 18:14:25 +0200 Subject: [PATCH] refactor(stream): langgraph streaming v3 (#268) * refactor(stream): langgraph streaming v3 * fix: address CR comments * chore(stream): add success field to state --------- Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com> --- EvoScientist/EvoScientist.py | 7 +- EvoScientist/channels/consumer.py | 20 +- EvoScientist/cli/tui_interactive.py | 99 +- EvoScientist/cli/widgets/approval_widget.py | 18 +- EvoScientist/stream/__init__.py | 3 - EvoScientist/stream/display.py | 114 +- EvoScientist/stream/emitter.py | 42 +- EvoScientist/stream/events.py | 1522 ++++++++----------- EvoScientist/stream/state.py | 191 +-- EvoScientist/stream/summarization.py | 94 ++ EvoScientist/stream/tool_results.py | 72 + EvoScientist/stream/tool_selection.py | 117 ++ EvoScientist/stream/tracker.py | 115 -- EvoScientist/stream/utils.py | 6 - EvoScientist/stream/v3_payloads.py | 84 + tests/conftest.py | 19 +- tests/stream_v3_fakes.py | 269 ++++ tests/test_event_loop.py | 31 + tests/test_hitl.py | 86 +- tests/test_stream_display.py | 44 +- tests/test_stream_emitter.py | 60 +- tests/test_stream_events.py | 1160 ++++++++++---- tests/test_stream_state.py | 393 ++--- tests/test_stream_tracker.py | 99 -- tests/test_stream_utils.py | 4 +- tests/test_subagent_summarize.py | 221 +-- tests/test_summarization.py | 82 +- tests/test_tui_widgets.py | 4 +- 28 files changed, 2801 insertions(+), 2175 deletions(-) create mode 100644 EvoScientist/stream/summarization.py create mode 100644 EvoScientist/stream/tool_results.py create mode 100644 EvoScientist/stream/tool_selection.py delete mode 100644 EvoScientist/stream/tracker.py create mode 100644 EvoScientist/stream/v3_payloads.py create mode 100644 tests/stream_v3_fakes.py delete mode 100644 tests/test_stream_tracker.py diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index fe5f84a..753f0c8 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -7,12 +7,11 @@ use so that importing this module is fast and non-agent CLI commands Usage: from EvoScientist import EvoScientist_agent + from EvoScientist.stream.events import stream_agent_events # Notebook / programmatic usage - for state in EvoScientist_agent.stream( - {"messages": [HumanMessage(content="your question")]}, - config={"configurable": {"thread_id": "1"}}, - stream_mode="values", + async for event in stream_agent_events( + EvoScientist_agent, "your question", thread_id="1" ): ... """ diff --git a/EvoScientist/channels/consumer.py b/EvoScientist/channels/consumer.py index cd9438d..9b83d0d 100644 --- a/EvoScientist/channels/consumer.py +++ b/EvoScientist/channels/consumer.py @@ -134,14 +134,10 @@ def _should_auto_approve(action_requests: list[dict]) -> bool: ) for req in action_requests: - name = ( - req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "") - ) + name = req.get("name", "") if name not in HITL_SHELL_TOOLS: continue - args = ( - req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {}) - ) + args = req.get("args", {}) command = args.get("command", "") if isinstance(args, dict) else "" cmd = command.strip() if not any(cmd.startswith(prefix) for prefix in shell_allow_list): @@ -159,12 +155,8 @@ def _format_approval_prompt( """ lines = ["\u26a0\ufe0f Approval Required\n"] for i, req in enumerate(action_requests, 1): - name = ( - req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "") - ) - args = ( - req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {}) - ) + name = req.get("name", "") + args = req.get("args", {}) if isinstance(args, dict): command = args.get("command", args.get("path", "")) else: @@ -555,7 +547,9 @@ class InboundConsumer: elif event_type == "subagent_text": sa_name = event.get("subagent", "unknown") - instance_id = event.get("instance_id") or sa_name + instance_id = event.get("instance_id") + if not instance_id: + continue if instance_id not in subagent_text_buffers: subagent_text_buffers[instance_id] = (sa_name, []) subagent_text_buffers[instance_id][1].append( diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 685c2ca..ba50809 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -1439,22 +1439,18 @@ def run_textual_interactive( for tw in tool_widgets.values(): tw.display = True - def _find_or_rename_sa_widget( - resolved_name: str, + def _get_sa_widget( + instance_id: str, + name: str = "", description: str = "", ) -> SubAgentWidget | None: - """Look up a sub-agent widget, renaming 'sub-agent' entry if needed.""" - if resolved_name in subagent_widgets: - w = subagent_widgets[resolved_name] + if instance_id in subagent_widgets: + w = subagent_widgets[instance_id] + if name and w._sa_name != name: + w.update_name(name, description or w._description) if description and not w._description: w.update_name(w._sa_name, description) return w - # Rename "sub-agent" → real name (mirrors state._get_or_create_subagent) - if resolved_name != "sub-agent" and "sub-agent" in subagent_widgets: - w = subagent_widgets.pop("sub-agent") - w.update_name(resolved_name, description) - subagent_widgets[resolved_name] = w - return w return None _MAX_HITL_ROUNDS = 50 @@ -1640,7 +1636,7 @@ def run_textual_interactive( elif event_type == "tool_call": tool_name = event.get("name", "unknown") - tool_id = event.get("id", "") + tool_id = event["id"] tool_args = event.get("args", {}) # Finalize thinking if still active if thinking_w is not None and thinking_w._is_active: @@ -1668,8 +1664,7 @@ def run_textual_interactive( else: w = ToolCallWidget(tool_name, tool_args, tool_id) await container.mount(w) - if tool_id: - tool_widgets[tool_id] = w + tool_widgets[tool_id] = w # Update todo widget on write_todos. # Insert before tool call widget so Task List # panel appears above the tool call. @@ -1690,38 +1685,18 @@ def run_textual_interactive( result_name = event.get("name", "unknown") result_content = event.get("content", "") result_success = event.get("success", True) - # Match via state's deduplicated tool_calls (uses tool_id) - matched = False - matched_tid = "" - result_idx = len(state.tool_results) - 1 - if 0 <= result_idx < len(state.tool_calls): - tc = state.tool_calls[result_idx] - tid = tc.get("id", "") - if tid and tid in tool_widgets: - tw = tool_widgets[tid] - if tw._status == "running": - if result_success: - tw.set_success(result_content) - else: - tw.set_error(result_content) - matched = True - matched_tid = tid - # Fallback: match first running widget with same name - if not matched: - for fid, tw in tool_widgets.items(): - if ( - tw.tool_name == result_name - and tw._status == "running" - ): - if result_success: - tw.set_success(result_content) - else: - tw.set_error(result_content) - matched = True - matched_tid = fid - break + matched_tid = event["id"] + tw = tool_widgets.get(matched_tid) + if tw is not None and tw._status == "running": + if result_success: + tw.set_success(result_content) + else: + tw.set_error(result_content) # Track completion order for collapsing - if matched_tid and matched_tid not in completed_tool_order: + if ( + tw is not None + and matched_tid not in completed_tool_order + ): completed_tool_order.append(matched_tid) await _collapse_completed_tools() # Update todo from results @@ -1746,44 +1721,44 @@ def run_textual_interactive( await container.mount(processing_w) elif event_type == "subagent_start": - sa_name = event.get("name", "sub-agent") + sa_name = event["name"] sa_desc = event.get("description", "") - existing = _find_or_rename_sa_widget(sa_name, sa_desc) + instance_id = event["instance_id"] + existing = _get_sa_widget(instance_id, sa_name, sa_desc) if existing is None: sa_w = SubAgentWidget(sa_name, sa_desc) await container.mount(sa_w) - subagent_widgets[sa_name] = sa_w + subagent_widgets[instance_id] = sa_w elif event_type == "subagent_tool_call": - sa_name = event.get("subagent", "sub-agent") - sa_name = state._resolve_subagent_name(sa_name) - sa_w = _find_or_rename_sa_widget(sa_name) + instance_id = event["instance_id"] + sa_w = _get_sa_widget( + instance_id, event.get("subagent", "") + ) if sa_w is None: - sa_w = SubAgentWidget(sa_name) - await container.mount(sa_w) - subagent_widgets[sa_name] = sa_w + continue await sa_w.add_tool_call( event.get("name", "unknown"), event.get("args", {}), - event.get("id", ""), + event["id"], ) elif event_type == "subagent_tool_result": - sa_name = event.get("subagent", "sub-agent") - sa_name = state._resolve_subagent_name(sa_name) - sa_w = _find_or_rename_sa_widget(sa_name) + instance_id = event["instance_id"] + sa_w = _get_sa_widget( + instance_id, event.get("subagent", "") + ) if sa_w is not None: sa_w.complete_tool( event.get("name", "unknown"), event.get("content", ""), event.get("success", True), - event.get("id", ""), + event["id"], ) elif event_type == "subagent_end": - sa_name = event.get("name", "sub-agent") - sa_name = state._resolve_subagent_name(sa_name) - sa_w = _find_or_rename_sa_widget(sa_name) + instance_id = event["instance_id"] + sa_w = _get_sa_widget(instance_id, event.get("name", "")) if sa_w is not None: sa_w.finalize() diff --git a/EvoScientist/cli/widgets/approval_widget.py b/EvoScientist/cli/widgets/approval_widget.py index b7e574e..53f0fc9 100644 --- a/EvoScientist/cli/widgets/approval_widget.py +++ b/EvoScientist/cli/widgets/approval_widget.py @@ -105,11 +105,7 @@ class ApprovalWidget(Widget): self._option_widgets = [] count = len(self._action_requests) if count == 1: - name = ( - self._action_requests[0].get("name", "") - if isinstance(self._action_requests[0], dict) - else getattr(self._action_requests[0], "name", "") - ) + name = self._action_requests[0].get("name", "") title = f">>> {name} Requires Approval <<<" else: title = f">>> {count} Tool Calls Require Approval <<<" @@ -117,16 +113,8 @@ class ApprovalWidget(Widget): # Show each action request as a compact line for req in self._action_requests: - name = ( - req.get("name", "") - if isinstance(req, dict) - else getattr(req, "name", "") - ) - args = ( - req.get("args", {}) - if isinstance(req, dict) - else getattr(req, "args", {}) - ) + name = req.get("name", "") + args = req.get("args", {}) if isinstance(args, dict): command = args.get("command", args.get("path", "")) else: diff --git a/EvoScientist/stream/__init__.py b/EvoScientist/stream/__init__.py index fd9cd8f..4d1e6dc 100644 --- a/EvoScientist/stream/__init__.py +++ b/EvoScientist/stream/__init__.py @@ -3,7 +3,6 @@ Stream module - streaming event processing for CLI display. Provides: - StreamEventEmitter: Standardized event creation -- ToolCallTracker: Incremental JSON parsing for tool parameters - ToolResultFormatter: Content-aware result formatting with Rich - Utility functions and constants - SubAgentState / StreamState: Stream state tracking @@ -28,7 +27,6 @@ __getattr__, __dir__, __all__ = _lazy.attach( "events", "formatter", "state", - "tracker", "utils", ], submod_attrs={ @@ -50,7 +48,6 @@ __getattr__, __dir__, __all__ = _lazy.attach( "_build_todo_stats", "_parse_todo_items", ], - "tracker": ["ToolCallInfo", "ToolCallTracker"], "utils": [ "FAILURE_PREFIX", "SUCCESS_PREFIX", diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index 2882bcd..c205e60 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -311,6 +311,18 @@ def _render_tool_call_line(tc: dict, tr: dict | None) -> Text: return tool_text +def _tool_result_for_call( + tool_results: list, + tool_call: dict, +) -> dict | None: + """Match root tool results by DeepAgents tool_call_id.""" + tool_id = tool_call["id"] + return next( + (result for result in tool_results if result["id"] == tool_id), + None, + ) + + # --------------------------------------------------------------------------- # Sub-agent section rendering # --------------------------------------------------------------------------- @@ -626,24 +638,13 @@ def create_streaming_display( elements.append(Text("")) # blank separator elements.append(narration_markdown) - def _find_task_subagent(tc: dict, shown_sa_names: set[str]) -> SubAgentState | None: - sa_name = tc.get("args", {}).get("subagent_type", "") - task_desc = tc.get("args", {}).get("description", "") + def _find_task_subagent(tc: dict, shown_sa_ids: set[str]) -> SubAgentState | None: + tool_id = tc["id"] for sa in subagents: - if sa.name in shown_sa_names: + if sa.instance_id in shown_sa_ids: continue - if sa.name == sa_name or ( - task_desc and task_desc in (sa.description or "") - ): + if sa.parent_tool_call_id == tool_id: return sa - - candidates = [ - sa - for sa in subagents - if sa.name not in shown_sa_names and (sa.tool_calls or sa.is_active) - ] - if len(candidates) == 1: - return candidates[0] return None def _append_task_entry( @@ -651,14 +652,14 @@ def create_streaming_display( tc: dict, tr: dict | None, *, - shown_sa_names: set[str], + shown_sa_ids: set[str], compact: bool, ) -> None: _append_narration_before_tool(tool_index) elements.append(_render_tool_call_line(tc, tr)) - matched_sa = _find_task_subagent(tc, shown_sa_names) + matched_sa = _find_task_subagent(tc, shown_sa_ids) if matched_sa is not None: - shown_sa_names.add(matched_sa.name) + shown_sa_ids.add(matched_sa.instance_id) elements.extend(_render_subagent_section(matched_sa, compact=compact)) # Tool calls and results paired display @@ -666,7 +667,7 @@ def create_streaming_display( # Task tool calls are ALWAYS visible (they represent sub-agent delegations) MAX_VISIBLE_TOOLS = 4 MAX_VISIBLE_RUNNING = 3 - shown_sa_names: set[str] = set() + shown_sa_ids: set[str] = set() if tool_calls: # Split into categories @@ -675,8 +676,8 @@ def create_streaming_display( running_regular = [] # running non-task tools for i, tc in enumerate(tool_calls): - has_result = i < len(tool_results) - tr = tool_results[i] if has_result else None + tr = _tool_result_for_call(tool_results, tc) + has_result = tr is not None is_task = tc.get("name") == "task" if is_task: @@ -697,7 +698,7 @@ def create_streaming_display( tool_index, tc, tr, - shown_sa_names=shown_sa_names, + shown_sa_ids=shown_sa_ids, compact=True, ) continue @@ -716,7 +717,9 @@ def create_streaming_display( # Render any sub-agents not already shown via task tool calls for sa in subagents: - if sa.name not in shown_sa_names and (sa.tool_calls or sa.is_active): + if sa.instance_id not in shown_sa_ids and ( + sa.tool_calls or sa.is_active + ): elements.extend(_render_subagent_section(sa, compact=True)) else: @@ -759,12 +762,12 @@ def create_streaming_display( key=lambda item: item[0], ): if tc.get("name") == "task": - matched_sa = _find_task_subagent(tc, shown_sa_names) + matched_sa = _find_task_subagent(tc, shown_sa_ids) _append_task_entry( tool_index, tc, tr, - shown_sa_names=shown_sa_names, + shown_sa_ids=shown_sa_ids, compact=not matched_sa.is_active if matched_sa else True, ) continue @@ -826,7 +829,7 @@ def create_streaming_display( # Sub-agent activity sections # Active: full bordered view; Completed: compact 1-line summary for sa in subagents: - if sa.name not in shown_sa_names and (sa.tool_calls or sa.is_active): + if sa.instance_id not in shown_sa_ids and (sa.tool_calls or sa.is_active): elements.extend(_render_subagent_section(sa, compact=not sa.is_active)) # Processing state after tool execution @@ -914,11 +917,11 @@ def display_final_results( ) if show_tools and state.tool_calls: - shown_sa_names: set[str] = set() + shown_sa_ids: set[str] = set() - for i, tc in enumerate(state.tool_calls): - has_result = i < len(state.tool_results) - tr = state.tool_results[i] if has_result else None + for tc in state.tool_calls: + tr = _tool_result_for_call(state.tool_results, tc) + has_result = tr is not None content = tr.get("content", "") if tr is not None else "" tool_name = tc.get("name", "") is_task = tool_name.lower() == "task" @@ -926,17 +929,13 @@ def display_final_results( # Task tools: show delegation line + compact sub-agent summary if is_task: console.print(_render_tool_call_line(tc, tr)) - sa_name = tc.get("args", {}).get("subagent_type", "") - task_desc = tc.get("args", {}).get("description", "") matched_sa = None for sa in state.subagents: - if sa.name == sa_name or ( - task_desc and task_desc in (sa.description or "") - ): + if sa.parent_tool_call_id == tc["id"]: matched_sa = sa break if matched_sa: - shown_sa_names.add(matched_sa.name) + shown_sa_ids.add(matched_sa.instance_id) for elem in _render_subagent_section(matched_sa, compact=True): console.print(elem) continue @@ -955,7 +954,7 @@ def display_final_results( # Render any sub-agents not already shown via task tool calls for sa in state.subagents: - if sa.name not in shown_sa_names and (sa.tool_calls or sa.is_active): + if sa.instance_id not in shown_sa_ids and (sa.tool_calls or sa.is_active): for elem in _render_subagent_section(sa, compact=True): console.print(elem) @@ -1047,12 +1046,8 @@ def _resolve_hitl_approval( needs_prompt = False for req in action_requests: - name = ( - req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "") - ) - args = ( - req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {}) - ) + name = req.get("name", "") + args = req.get("args", {}) if name not in HITL_SHELL_TOOLS: continue # Only shell-running tools need manual approval @@ -1091,12 +1086,8 @@ def _prompt_hitl_approval(action_requests: list) -> list[dict] | None: console.print() panel_text = Text() for i, req in enumerate(action_requests): - name = ( - req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "") - ) - args = ( - req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {}) - ) + name = req.get("name", "") + args = req.get("args", {}) desc = format_tool_compact(name, args if isinstance(args, dict) else {}) if panel_text.plain: panel_text.append("\n") @@ -1669,18 +1660,15 @@ async def _astream_to_console( # Only show subagent starts as real-time progress. # Full results rendered by display_final_results() after streaming. if etype == "subagent_start": - name = event.get("name", "sub-agent") - # Skip generic "sub-agent" — real name arrives later; - # static prints can't be overwritten like Live display. - if name and name != "sub-agent": - desc = event.get("description", "") - line = Text() - line.append("\u25b6 ", style="cyan bold") - line.append(f"Cooking with {name}", style="cyan bold") - if desc: - short = desc[:50] + "\u2026" if len(desc) > 50 else desc - line.append(f" \u2014 {short}", style="dim") - console.print(line) + name = event["name"] + desc = event.get("description", "") + line = Text() + line.append("\u25b6 ", style="cyan bold") + line.append(f"Cooking with {name}", style="cyan bold") + if desc: + short = desc[:50] + "\u2026" if len(desc) > 50 else desc + line.append(f" \u2014 {short}", style="dim") + console.print(line) # Final output (streaming layout: tools → Task List → subagents → response) @@ -1707,10 +1695,10 @@ async def _astream_to_console( ) # 1) Regular (non-task) tools — above Task List - for i, tc in enumerate(state.tool_calls): + for tc in state.tool_calls: if tc.get("name", "").lower() == "task": continue - tr = state.tool_results[i] if i < len(state.tool_results) else None + tr = _tool_result_for_call(state.tool_results, tc) console.print(_render_tool_call_line(tc, tr)) if tr and not is_success(tr.get("content", "")): for elem in format_tool_result_compact(tr["name"], tr.get("content", "")): diff --git a/EvoScientist/stream/emitter.py b/EvoScientist/stream/emitter.py index 1198ca3..cfa268a 100644 --- a/EvoScientist/stream/emitter.py +++ b/EvoScientist/stream/emitter.py @@ -32,7 +32,7 @@ class StreamEventEmitter: return StreamEvent("text", {"type": "text", "content": content}) @staticmethod - def tool_call(name: str, args: dict[str, Any], tool_id: str = "") -> StreamEvent: + def tool_call(name: str, args: dict[str, Any], tool_id: str) -> StreamEvent: """Tool call event.""" return StreamEvent( "tool_call", @@ -40,7 +40,12 @@ class StreamEventEmitter: ) @staticmethod - def tool_result(name: str, content: str, success: bool = True) -> StreamEvent: + def tool_result( + name: str, + content: str, + success: bool, + tool_call_id: str, + ) -> StreamEvent: """Tool result event.""" return StreamEvent( "tool_result", @@ -49,11 +54,14 @@ class StreamEventEmitter: "name": name, "content": content, "success": success, + "id": tool_call_id, }, ) @staticmethod - def subagent_start(name: str, description: str) -> StreamEvent: + def subagent_start( + name: str, description: str, instance_id: str, tool_call_id: str + ) -> StreamEvent: """Sub-agent delegation started.""" return StreamEvent( "subagent_start", @@ -61,12 +69,18 @@ class StreamEventEmitter: "type": "subagent_start", "name": name, "description": description, + "instance_id": instance_id, + "tool_call_id": tool_call_id, }, ) @staticmethod def subagent_tool_call( - subagent: str, name: str, args: dict[str, Any], tool_id: str = "" + subagent: str, + name: str, + args: dict[str, Any], + tool_id: str, + instance_id: str, ) -> StreamEvent: """Tool call from inside a sub-agent.""" return StreamEvent( @@ -77,6 +91,7 @@ class StreamEventEmitter: "name": name, "args": args, "id": tool_id, + "instance_id": instance_id, }, ) @@ -85,8 +100,9 @@ class StreamEventEmitter: subagent: str, name: str, content: str, - success: bool = True, - tool_call_id: str = "", + success: bool, + tool_call_id: str, + instance_id: str, ) -> StreamEvent: """Tool result from inside a sub-agent.""" return StreamEvent( @@ -98,14 +114,13 @@ class StreamEventEmitter: "content": content, "success": success, "id": tool_call_id, + "instance_id": instance_id, }, ) @staticmethod - def subagent_text( - subagent: str, content: str, instance_id: str = "" - ) -> StreamEvent: - """Text content from a sub-agent (for fallback extraction).""" + def subagent_text(subagent: str, content: str, instance_id: str) -> StreamEvent: + """Text content from a sub-agent.""" return StreamEvent( "subagent_text", { @@ -117,9 +132,12 @@ class StreamEventEmitter: ) @staticmethod - def subagent_end(name: str) -> StreamEvent: + def subagent_end(name: str, instance_id: str) -> StreamEvent: """Sub-agent delegation completed.""" - return StreamEvent("subagent_end", {"type": "subagent_end", "name": name}) + return StreamEvent( + "subagent_end", + {"type": "subagent_end", "name": name, "instance_id": instance_id}, + ) @staticmethod def done(response: str = "") -> StreamEvent: diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index 2a96058..78f21f9 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -1,200 +1,571 @@ -"""Stream event generator and chunk processing helpers. +"""Stream event generator and v3 protocol helpers. -Async generator that streams events from an agent graph, -plus helpers for processing AI message chunks and tool results. +Async generator that streams UI events from a DeepAgents/LangGraph v3 run. """ import asyncio import base64 +import inspect import mimetypes import os -import re -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator, AsyncIterator +from dataclasses import dataclass from typing import Any -from langchain_core.messages import ( # type: ignore[import-untyped] - AIMessage, - AIMessageChunk, -) +from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage +from langgraph.types import Command, Interrupt from ..memory.worker_activity import clear_memory_worker_saved_counts from .emitter import StreamEventEmitter -from .tracker import ToolCallTracker +from .summarization import ( + _extract_summary_message_text, + _find_summarization_event_payload, + _summarization_event_signature, +) +from .tool_results import _extract_command_tool_content, _extract_tool_content +from .tool_selection import _ToolSelectionSuppressor from .utils import DisplayLimits, is_success - -# Safety net: older ccproxy versions may embed thinking as XML tags in content -# strings. Strip them so they never leak to users or channels. -_THINKING_TAG_RE = re.compile(r".*?", re.DOTALL) -_SUMMARY_TAG_RE = re.compile( - r"\s*(.*?)\s*", re.DOTALL | re.IGNORECASE +from .v3_payloads import ( + RawMap, + _as_raw_map, + _event_data, + _event_namespace, + _reasoning_from_content, + _split_message_event_data, + _strip_legacy_thinking_tags, + _text_from_content, + _usage_counts, ) -def _strip_legacy_thinking_tags(content: str) -> str: - """Remove ``...`` tags from content strings.""" - return _THINKING_TAG_RE.sub("", content) +def _is_interrupt_error_message(message: object) -> bool: + if not isinstance(message, str): + return False + stripped = message.strip() + return stripped.startswith("(Interrupt(value=") -# Image media types returned by DeepAgents read_file -_IMAGE_MEDIA_TYPES = { - "image/png", - "image/jpeg", - "image/gif", - "image/webp", - "image/bmp", - "image/svg+xml", -} +@dataclass(frozen=True) +class _SubagentInfo: + path: tuple[str, ...] + name: str + description: str = "" + + @property + def instance_id(self) -> str: + return ":".join(self.path) -def _extract_tool_content(msg) -> tuple[str, bool]: - """Extract display-safe content from a ToolMessage. +class _SubagentRegistry: + """Track v3 DeepAgents subagent handles by namespace path.""" - DeepAgents ``read_file`` returns image content as - ``ToolMessage(content=[ImageContentBlock])`` with - ``additional_kwargs["read_file_media_type"]`` set. - Stringifying that would dump huge base64 data into the display. + def __init__(self) -> None: + self._by_path: dict[tuple[str, ...], _SubagentInfo] = {} + self._changed = asyncio.Event() + self._closed = False - Returns: - (content_string, is_image) — a short summary for images, - or the raw string content for normal results. - """ - additional = getattr(msg, "additional_kwargs", None) or {} - media_type = additional.get("read_file_media_type", "") - if media_type and media_type in _IMAGE_MEDIA_TYPES: - # Extract path from the tool call args if available - file_path = additional.get("read_file_path", "") - if not file_path: - file_path = getattr(msg, "name", "image") - return f"[OK] Image displayed: {file_path} ({media_type})", True + def register(self, path: tuple[str, ...], name: str, description: str = "") -> None: + if not path: + return + self._by_path[path] = _SubagentInfo(path, name, description) + self._changed.set() - content = getattr(msg, "content", "") - # Guard against list-type content (image content blocks without metadata) - if isinstance(content, list): - # Check if any block looks like image data - for block in content: - if isinstance(block, dict): - if block.get("type") == "image" or "base64" in block: - return "[OK] Image displayed", True - # Non-image list content — join text blocks - parts = [] - for block in content: - if isinstance(block, dict): - text = block.get("text", "") - if text: - parts.append(text) - elif isinstance(block, str): - parts.append(block) - return "\n".join(parts) if parts else str(content), False + def close(self) -> None: + self._closed = True + self._changed.set() - return str(content), False - - -def _extract_summarization_text(msg: Any) -> str: - """Extract plain text from a summarization chunk. - - The summarization LLM streams ``AIMessageChunk`` objects whose - ``content`` may be a plain string **or** a list of content blocks - (e.g. ``[{'type': 'text', 'text': '...', 'index': 1}]``) depending - on the provider. This helper normalises both forms to a plain string. - """ - if not hasattr(msg, "content"): - return "" - content = msg.content - if isinstance(content, str): - return content - if isinstance(content, list): - parts: list[str] = [] - for block in content: - if isinstance(block, dict): - text = block.get("text") - if isinstance(text, str): - parts.append(text) - elif isinstance(block, str): - parts.append(block) - return "".join(parts) - return "" - - -def _extract_summary_message_text(summary_message: Any) -> str: - """Extract user-facing summary text from a stored summarization event. - - DeepAgents persists summary messages as ``HumanMessage`` objects with wrapper - text like ``Here is a summary of the conversation to date:`` or an XML-ish - ``...`` block. For UI display we only want the summary - body itself. - """ - text = _extract_summarization_text(summary_message) - if not text: - return "" - - match = _SUMMARY_TAG_RE.search(text) - if match: - return match.group(1).strip() - - prefix = "Here is a summary of the conversation to date:" - if text.startswith(prefix): - return text[len(prefix) :].strip() - - return text.strip() - - -def _find_summarization_event_payload(data: Any) -> dict[str, Any] | None: - """Find a `_summarization_event` dict anywhere inside an updates payload.""" - seen: set[int] = set() - stack: list[Any] = [data] - - while stack: - item = stack.pop() - item_id = id(item) - if item_id in seen: - continue - seen.add(item_id) - - if isinstance(item, dict): - event = item.get("_summarization_event") - if isinstance(event, dict): - return event - stack.extend(item.values()) - continue - - if isinstance(item, list | tuple): - stack.extend(item) - continue - - if hasattr(item, "__dict__"): - try: - stack.append(vars(item)) - except TypeError: - pass - - return None - - -def _summarization_event_signature( - event: dict[str, Any] | None, -) -> tuple[Any, ...] | None: - """Build a stable signature for a persisted summarization event.""" - if not isinstance(event, dict): + def resolve(self, namespace: tuple[str, ...]) -> _SubagentInfo | None: + for depth in range(len(namespace), 0, -1): + info = self._by_path.get(namespace[:depth]) + if info is not None: + return info return None - summary_message = event.get("summary_message") - summary_text = _extract_summary_message_text(summary_message) - return ( - event.get("cutoff_index"), - event.get("file_path"), - summary_text, - ) + + async def wait_resolve(self, namespace: tuple[str, ...]) -> _SubagentInfo | None: + if not namespace: + return None + while not self._closed: + if info := self.resolve(namespace): + return info + self._changed.clear() + if info := self.resolve(namespace): + return info + if self._closed: + break + await self._changed.wait() + return self.resolve(namespace) + + +class _V3EventProcessor: + """Translate v3 protocol channel events into EvoScientist UI events.""" + + def __init__( + self, + emitter: StreamEventEmitter, + subagents: _SubagentRegistry, + baseline_summarization_signature: tuple[object, ...] | None, + ) -> None: + self.emitter = emitter + self.subagents = subagents + self.baseline_summarization_signature = baseline_summarization_signature + self.full_response = "" + self._summarization_in_progress = False + self._tool_inputs: dict[ + tuple[tuple[str, ...], str], tuple[str, dict[str, Any]] + ] = {} + self._emitted_tool_calls: set[tuple[tuple[str, ...], str]] = set() + self._emitted_interrupts: set[str] = set() + self._selector = _ToolSelectionSuppressor(emitter) + + @staticmethod + def _tool_scope( + namespace: tuple[str, ...], subagent: _SubagentInfo | None + ) -> tuple[str, ...]: + return subagent.path if subagent is not None else namespace + + async def process(self, event: dict[str, Any]) -> list[dict[str, Any]]: + method = event.get("method") + namespace = _event_namespace(event) + subagent = None + if namespace and method in ("messages", "tools"): + subagent = await self.subagents.wait_resolve(namespace) + if subagent is None: + return [] + + if method == "messages": + return self._process_message_event(_event_data(event), subagent, namespace) + if method == "tools": + return self._process_tool_event(namespace, _event_data(event), subagent) + if method == "updates": + return self._process_update_event(_event_data(event)) + if method == "values": + params = event.get("params") or {} + interrupts = params.get("interrupts") or () + if interrupts: + return self._process_update_event({"__interrupt__": interrupts}) + return [] + + def _process_message_event( + self, + data: object, + subagent: _SubagentInfo | None, + namespace: tuple[str, ...], + ) -> list[dict[str, Any]]: + payload, metadata = _split_message_event_data(data) + payload_map = _as_raw_map(payload) + if metadata.get("lc_source") == "summarization": + if payload_map is not None: + text = self._text_from_message_payload(payload_map) + elif isinstance(payload, BaseMessage): + text = self._text_from_message_payload(payload) + else: + text = "" + return self._emit_summarization_text(text) + + if payload_map is not None and "event" in payload_map: + return self._process_protocol_message_payload( + payload_map, subagent, namespace + ) + + if isinstance(payload, AIMessage | AIMessageChunk): + return self._process_whole_message(payload, subagent, namespace) + return [] + + def _process_protocol_message_payload( + self, + payload: RawMap, + subagent: _SubagentInfo | None, + namespace: tuple[str, ...], + ) -> list[dict[str, Any]]: + event_type = payload.get("event") + if event_type == "content-block-delta": + delta = _as_raw_map(payload.get("delta")) + if delta is None: + return [] + delta_type = delta.get("type") + if delta_type == "text-delta": + text = delta.get("text") + return self._emit_text(text if isinstance(text, str) else "", subagent) + if delta_type == "reasoning-delta": + reasoning = delta.get("reasoning") + return self._emit_thinking( + reasoning if isinstance(reasoning, str) else "", subagent + ) + if delta_type == "block-delta": + fields = _as_raw_map(delta.get("fields")) + if fields is not None: + return self._process_message_tool_block(fields, subagent, namespace) + return [] + + if event_type == "content-block-finish": + content = _as_raw_map(payload.get("content")) + if content is not None: + return self._process_message_tool_block(content, subagent, namespace) + return [] + + if event_type == "message-finish": + events: list[dict[str, Any]] = [] + pending = self._selector.flush_pending_text() + if pending and not pending.isspace(): + if subagent is not None: + events.append( + self.emitter.subagent_text( + subagent.name, + pending, + instance_id=subagent.instance_id, + ).data + ) + else: + self.full_response += pending + events.append(self.emitter.text(pending).data) + if subagent is None: + usage = _as_raw_map(payload.get("usage")) + inp, out = _usage_counts(usage) if usage is not None else (0, 0) + if inp or out: + events.append(self.emitter.usage_stats(inp, out).data) + return events + return [] + + def _process_message_tool_block( + self, + block: RawMap, + subagent: _SubagentInfo | None, + namespace: tuple[str, ...], + ) -> list[dict[str, Any]]: + block_type = block.get("type") + if block_type not in ( + "tool_call", + "tool_call_chunk", + "server_tool_call", + "server_tool_call_chunk", + ): + return [] + name = block.get("name") or "" + if self._selector.observe_tool_block(str(name)): + return [] + events = self._selector.flush_selection() + tool_call = self._tool_call_from_message_block(block) + if tool_call is None: + return events + tool_name, args, tool_call_id = tool_call + events.extend( + self._emit_tool_call_once( + namespace=namespace, + subagent=subagent, + name=tool_name, + args=args, + tool_call_id=tool_call_id, + ) + ) + return events + + @staticmethod + def _tool_call_from_message_block( + block: RawMap, + ) -> tuple[str, dict[str, Any], str] | None: + tool_call_id = str(block.get("id") or block.get("tool_call_id") or "") + name = str(block.get("name") or block.get("tool_name") or "") + if not tool_call_id or not name: + return None + raw_args = block.get("args") if "args" in block else block.get("input") + args = _as_raw_map(raw_args) + if args is None: + return None + return name, dict(args), tool_call_id + + def _emit_tool_call_once( + self, + *, + namespace: tuple[str, ...], + subagent: _SubagentInfo | None, + name: str, + args: dict[str, Any], + tool_call_id: str, + ) -> list[dict[str, Any]]: + key = (self._tool_scope(namespace, subagent), tool_call_id) + if key in self._emitted_tool_calls: + return [] + self._emitted_tool_calls.add(key) + if subagent is not None: + return [ + self.emitter.subagent_tool_call( + subagent.name, + name, + args, + tool_call_id, + instance_id=subagent.instance_id, + ).data + ] + return [self.emitter.tool_call(name, args, tool_call_id).data] + + def _process_whole_message( + self, + msg: AIMessage | AIMessageChunk, + subagent: _SubagentInfo | None, + namespace: tuple[str, ...], + ) -> list[dict[str, Any]]: + events: list[dict[str, Any]] = [] + additional = msg.additional_kwargs + reasoning = additional.get("reasoning_content") + emitted_reasoning = False + if isinstance(reasoning, str): + events.extend(self._emit_thinking(reasoning, subagent)) + emitted_reasoning = bool(reasoning) + + content = msg.content + if not emitted_reasoning: + events.extend( + self._emit_thinking(_reasoning_from_content(content), subagent) + ) + events.extend(self._emit_text(_text_from_content(content), subagent)) + + for raw_tool_call_block in msg.tool_calls: + tool_call_block = _as_raw_map(raw_tool_call_block) + if tool_call_block is None: + continue + tool_call = self._tool_call_from_message_block(tool_call_block) + if tool_call is None: + continue + tool_name, args, tool_call_id = tool_call + events.extend( + self._emit_tool_call_once( + namespace=namespace, + subagent=subagent, + name=tool_name, + args=args, + tool_call_id=tool_call_id, + ) + ) + + if subagent is None: + inp, out = _usage_counts(msg.usage_metadata) + if inp or out: + events.append(self.emitter.usage_stats(inp, out).data) + return events + + def _process_tool_event( + self, + namespace: tuple[str, ...], + data: object, + subagent: _SubagentInfo | None, + ) -> list[dict[str, Any]]: + data_map = _as_raw_map(data) + if data_map is None: + return [] + event_type = data_map.get("event") + tool_call_id = str(data_map.get("tool_call_id") or "") + if event_type == "tool-started": + events = self._selector.flush_selection() + if not tool_call_id: + return events + name = str(data_map.get("tool_name") or "") + if not name: + return events + input_args = _as_raw_map(data_map.get("input")) + args = dict(input_args) if input_args is not None else {} + self._tool_inputs[(self._tool_scope(namespace, subagent), tool_call_id)] = ( + name, + args, + ) + events.extend( + self._emit_tool_call_once( + namespace=namespace, + subagent=subagent, + name=name, + args=args, + tool_call_id=tool_call_id, + ) + ) + return events + + if event_type in ("tool-finished", "tool-error"): + events = self._selector.flush_selection() + if not tool_call_id: + return events + name, _args = self._tool_inputs.pop( + (self._tool_scope(namespace, subagent), tool_call_id), + (str(data_map.get("tool_name") or "unknown"), {}), + ) + message = data_map.get("message") + # LangGraph v3 reports interrupt() as a tool-error before the + # structured __interrupt__ update; the real result arrives on resume. + if event_type == "tool-error" and _is_interrupt_error_message(message): + return events + if event_type == "tool-error": + content = str(message or "") + success = False + else: + output = data_map.get("output") + if isinstance(output, ToolMessage): + name = output.name or name + raw_content, _ = _extract_tool_content(output) + elif isinstance(output, Command): + command_content = _extract_command_tool_content( + output, tool_call_id + ) + raw_content = ( + command_content if command_content is not None else str(output) + ) + else: + raw_content = "" if output is None else str(output) + content = raw_content[: DisplayLimits.TOOL_RESULT_MAX] + if len(raw_content) > DisplayLimits.TOOL_RESULT_MAX: + content += "\n... (truncated)" + success = is_success(content) + + if subagent is not None: + events.append( + self.emitter.subagent_tool_result( + subagent.name, + name, + content, + success, + tool_call_id, + instance_id=subagent.instance_id, + ).data + ) + return events + events.append( + self.emitter.tool_result( + name, content, success, tool_call_id=tool_call_id + ).data + ) + return events + + return [] + + def _process_update_event(self, data: object) -> list[dict[str, Any]]: + events: list[dict[str, Any]] = [] + data_map = _as_raw_map(data) + if data_map is not None and "__interrupt__" in data_map: + events.extend(self._process_interrupts(data_map["__interrupt__"])) + + summarization_event = _find_summarization_event_payload(data) + if summarization_event and not self._summarization_in_progress: + signature = _summarization_event_signature(summarization_event) + if ( + signature is not None + and signature == self.baseline_summarization_signature + ): + return events + summary_message = summarization_event.get("summary_message") + summary_text = _extract_summary_message_text( + summary_message if isinstance(summary_message, BaseMessage) else None + ) + events.extend(self._emit_summarization_text(summary_text)) + return events + + def _process_interrupts(self, interrupts: object) -> list[dict[str, Any]]: + events: list[dict[str, Any]] = [] + if not isinstance(interrupts, list | tuple): + return events + + for interrupt_obj in interrupts: + if not isinstance(interrupt_obj, Interrupt): + continue + + interrupt_value = interrupt_obj.value + if not isinstance(interrupt_value, dict): + continue + + iv_type = interrupt_value.get("type") + interrupt_id = interrupt_obj.id or "default" + if iv_type == "ask_user": + questions = interrupt_value.get("questions", []) + tc_id = str(interrupt_value.get("tool_call_id", "")) + events.extend( + self._dedupe_interrupt_event( + self.emitter.ask_user_interrupt( + interrupt_id, questions, tc_id + ).data + ) + ) + continue + + action_reqs = interrupt_value.get("action_requests", []) + review_cfgs = interrupt_value.get("review_configs", []) + if action_reqs: + events.extend( + self._dedupe_interrupt_event( + self.emitter.interrupt( + interrupt_id, action_reqs, review_cfgs + ).data + ) + ) + return events + + def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]: + signature = repr(event) + if signature in self._emitted_interrupts: + return [] + self._emitted_interrupts.add(signature) + return [event] + + def _emit_text( + self, text: str, subagent: _SubagentInfo | None + ) -> list[dict[str, Any]]: + if not text: + return [] + cleaned = _strip_legacy_thinking_tags(text) + if not cleaned or cleaned.isspace(): + return [] + suppressed, events, emit_text = self._selector.process_text(cleaned) + if suppressed: + return events + if not emit_text or emit_text.isspace(): + return events + if subagent is not None: + events.append( + self.emitter.subagent_text( + subagent.name, emit_text, instance_id=subagent.instance_id + ).data + ) + return events + self.full_response += emit_text + events.append(self.emitter.text(emit_text).data) + return events + + def _emit_thinking( + self, text: str, subagent: _SubagentInfo | None + ) -> list[dict[str, Any]]: + if not text or subagent is not None: + return [] + return [self.emitter.thinking(text).data] + + def _emit_summarization_text(self, text: str) -> list[dict[str, Any]]: + if not text: + return [] + events: list[dict[str, Any]] = [] + if not self._summarization_in_progress: + events.append(self.emitter.summarization_start().data) + self._summarization_in_progress = True + events.append(self.emitter.summarization(text).data) + return events + + def _text_from_message_payload(self, payload: RawMap | BaseMessage) -> str: + if not isinstance(payload, BaseMessage): + if payload.get("event") == "content-block-delta": + delta = _as_raw_map(payload.get("delta")) + if delta is not None and delta.get("type") == "text-delta": + text = delta.get("text") + return text if isinstance(text, str) else "" + return "" + return _text_from_content(payload.content) async def stream_agent_events( agent: Any, - message: Any, + message: str | Command, thread_id: str, - metadata: dict | None = None, + metadata: dict[str, Any] | None = None, media: list[str] | None = None, -) -> AsyncIterator[dict]: - """Stream events from the agent graph using async iteration. +) -> AsyncGenerator[dict[str, Any], None]: + """Stream events from a DeepAgents/LangGraph v3 run. - Uses agent.astream() with subgraphs=True to see sub-agent activity. + The ingestion side uses ``astream_events(..., version="v3")``: + raw protocol events preserve arrival order for messages/tools/updates, + while DeepAgents' native ``stream.subagents`` projection provides the + user-facing subagent identity that raw graph namespaces intentionally hide. Args: agent: Compiled state graph from create_deep_agent() @@ -213,187 +584,16 @@ async def stream_agent_events( if metadata: config["metadata"] = metadata emitter = StreamEventEmitter() - main_tracker = ToolCallTracker() - full_response = "" - # Track sub-agent names - _key_to_name: dict[str, str] = {} # subagent_key -> display name (cache) - _announced_names: list[str] = [] # ordered queue of announced task names - _assigned_names: set[str] = set() # names already assigned to a namespace - _announced_task_ids: list[str] = [] # ordered task tool_call_ids - _task_id_to_name: dict[str, str] = {} # tool_call_id -> sub-agent name - _subagent_trackers: dict[str, ToolCallTracker] = {} # namespace_key -> tracker - - def _register_task_tool_call(tc_data: dict) -> str | None: - """Register or update a task tool call, return subagent name if started/updated.""" - tool_id = tc_data.get("id", "") - if not tool_id: - return None - args = tc_data.get("args", {}) or {} - desc = str(args.get("description", "")).strip() - sa_name = str(args.get("subagent_type", "")).strip() - if not sa_name: - # Fallback to description snippet (may be empty during streaming) - sa_name = desc.split("\n")[0].strip() - sa_name = sa_name[:30] + "\u2026" if len(sa_name) > 30 else sa_name - if not sa_name: - sa_name = "sub-agent" - - if tool_id not in _announced_task_ids: - _announced_task_ids.append(tool_id) - _announced_names.append(sa_name) - _task_id_to_name[tool_id] = sa_name - return sa_name - - # Update mapping if we learned a better name later - current = _task_id_to_name.get(tool_id, "sub-agent") - if sa_name != "sub-agent" and current != sa_name: - _task_id_to_name[tool_id] = sa_name - try: - idx = _announced_task_ids.index(tool_id) - if idx < len(_announced_names): - _announced_names[idx] = sa_name - except ValueError: - pass - return sa_name - return None - - def _extract_task_id(namespace: tuple) -> tuple[str | None, str | None]: - """Extract task tool_call_id from namespace if present. - - Returns (task_id, task_ns_element) or (None, None). - """ - for part in namespace: - part_str = str(part) - if "task:" in part_str: - tail = part_str.split("task:", 1)[1] - task_id = tail.split(":", 1)[0] if tail else "" - if task_id: - return task_id, part_str - return None, None - - def _next_announced_name() -> str | None: - """Get next announced name that hasn't been assigned yet.""" - for announced in _announced_names: - if announced not in _assigned_names: - _assigned_names.add(announced) - return announced - return None - - def _find_task_id_from_metadata(metadata: dict | None) -> str | None: - """Try to find a task tool_call_id in metadata.""" - if not metadata: - return None - candidates = ( - "tool_call_id", - "task_id", - "parent_run_id", - "root_run_id", - "run_id", - ) - for key in candidates: - val = metadata.get(key) - if val and val in _task_id_to_name: - return val - return None - - def _get_subagent_key(namespace: tuple, metadata: dict | None) -> str | None: - """Stable key for tracker/mapping per sub-agent namespace.""" - if not namespace: - return None - _task_id, task_ns = _extract_task_id(namespace) - if task_ns: - return task_ns - meta_task_id = _find_task_id_from_metadata(metadata) - if meta_task_id: - return f"task:{meta_task_id}" - if metadata: - for key in ( - "parent_run_id", - "root_run_id", - "run_id", - "graph_id", - "node_id", - ): - val = metadata.get(key) - if val: - return f"{key}:{val}" - return str(namespace) - - def _get_subagent_name(namespace: tuple, metadata: dict | None) -> str | None: - """Resolve sub-agent name from namespace, or None if main agent. - - Priority: - 0) metadata["lc_agent_name"] -- most reliable, set by DeepAgents framework. - 1) Match task_id embedded in namespace to announced tool_call_id. - 2) Use cached key mapping (only real names, never "sub-agent"). - 3) Queue-based: assign next announced name to this key. - 4) Fallback: return "sub-agent" WITHOUT caching. - """ - if not namespace: - return None - - key = _get_subagent_key(namespace, metadata) or str(namespace) - - # 0) lc_agent_name from metadata -- the REAL sub-agent name - # set by the DeepAgents framework on every namespace event. - if metadata: - lc_name = metadata.get("lc_agent_name", "") - if isinstance(lc_name, str): - lc_name = lc_name.strip() - # Filter out generic/framework names - if lc_name and lc_name not in ( - "sub-agent", - "agent", - "tools", - "EvoScientist", - "LangGraph", - "", - ): - _key_to_name[key] = lc_name - return lc_name - - # 1) Resolve by task_id if present in namespace - task_id, _task_ns = _extract_task_id(namespace) - if task_id and task_id in _task_id_to_name: - name = _task_id_to_name[task_id] - if name and name != "sub-agent": - _assigned_names.add(name) - _key_to_name[key] = name - return name - - meta_task_id = _find_task_id_from_metadata(metadata) - if meta_task_id and meta_task_id in _task_id_to_name: - name = _task_id_to_name[meta_task_id] - if name and name != "sub-agent": - _assigned_names.add(name) - _key_to_name[key] = name - return name - - # 2) Cached real name for this key (skip if it's "sub-agent") - cached = _key_to_name.get(key) - if cached and cached != "sub-agent": - return cached - - # 3) Assign next announced name from queue (skip "sub-agent" entries) - for announced in _announced_names: - if announced not in _assigned_names and announced != "sub-agent": - _assigned_names.add(announced) - _key_to_name[key] = announced - return announced - - # 4) No real names available yet -- return generic WITHOUT caching - return "sub-agent" - - # Build input for agent.astream() clear_memory_worker_saved_counts() + # Build input for agent.astream_events() if isinstance(message, str): # Build user message content: text + inline images + file path references - user_content: str | list[dict[str, Any]] = message + user_content: str | list[dict[str, object]] = message if media: _IMAGE_EXTS = frozenset({".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"}) _MAX_INLINE_SIZE = 5 * 1024 * 1024 # 5 MB - content_blocks: list[dict[str, Any]] = [] + content_blocks: list[dict[str, object]] = [] if message: content_blocks.append({"type": "text", "text": message}) @@ -432,531 +632,153 @@ async def stream_agent_events( content_blocks.append({"type": "text", "text": ref_text}) if content_blocks: user_content = content_blocks - astream_input: Any = {"messages": [{"role": "user", "content": user_content}]} + astream_input: dict[str, list[dict[str, object]]] | Command = { + "messages": [{"role": "user", "content": user_content}] + } else: # HITL resume: Command object passed directly to agent astream_input = message - _summarization_in_progress = False - _baseline_summarization_signature: tuple[Any, ...] | None = None - _tool_selection_suppressing = False # True while buffering selector JSON - _tool_selection_buffer = "" # accumulates JSON chunks for parse attempt - _tool_selection_was_active = False # True after suppression, triggers Panel - - if hasattr(agent, "aget_state"): - try: - snapshot = await agent.aget_state(config) - values = getattr(snapshot, "values", None) - if isinstance(values, dict): - baseline_event = _find_summarization_event_payload(values) - _baseline_summarization_signature = _summarization_event_signature( - baseline_event - ) - except Exception: - pass + _baseline_summarization_signature: tuple[object, ...] | None = None try: - async for chunk in agent.astream( - astream_input, - config=config, - stream_mode=["messages", "updates"], - subgraphs=True, - ): - # Multi-mode + subgraphs: 3-tuple (namespace, mode, data) - # Single-mode + subgraphs: 2-tuple (namespace, data) — fallback - if not isinstance(chunk, tuple): - continue + snapshot = await agent.aget_state(config) + values = snapshot.values + if isinstance(values, dict): + baseline_event = _find_summarization_event_payload(values) + _baseline_summarization_signature = _summarization_event_signature( + baseline_event + ) + except Exception: + pass - namespace: tuple = () - data: Any - mode_str: str + stream: Any | None = None + producers: list[asyncio.Task[Any]] = [] + try: + from langgraph.stream.transformers import UpdatesTransformer - 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 = first - data = chunk[1] - else: - data = chunk - mode_str = "messages" - else: - continue + try: + stream_result = agent.astream_events( + astream_input, + config=config, + version="v3", + transformers=[UpdatesTransformer], + ) + except AttributeError as exc: + raise RuntimeError( + "This agent does not expose astream_events(); EvoScientist requires " + "DeepAgents/LangGraph stream v3." + ) from exc - # Parse HITL / ask_user interrupts from updates mode - if mode_str == "updates": - if isinstance(data, dict) and "__interrupt__" in data: - for interrupt_obj in data["__interrupt__"]: - if isinstance(interrupt_obj, dict): - interrupt_value = interrupt_obj.get("value", {}) - else: - interrupt_value = getattr(interrupt_obj, "value", {}) + subagents = _SubagentRegistry() + processor = _V3EventProcessor( + emitter, + subagents, + _baseline_summarization_signature, + ) + queue: asyncio.Queue[Any] = asyncio.Queue() + producer_done = object() - # Discriminate ask_user vs HITL interrupts - iv_type = ( - interrupt_value.get("type") - if isinstance(interrupt_value, dict) - else getattr(interrupt_value, "type", None) + stream = ( + await stream_result if inspect.isawaitable(stream_result) else stream_result + ) + + async def _put_events(events: list[dict[str, Any]]) -> None: + for event in events: + await queue.put(event) + + async def _consume_protocol_events() -> None: + async for event in stream: + await _put_events(await processor.process(event)) + + async def _await_subagent_done( + subagent: Any, name: str, instance_id: str + ) -> None: + try: + result = subagent.output() + if inspect.isawaitable(result): + await result + finally: + await queue.put( + emitter.subagent_end(name, instance_id=instance_id).data + ) + + subagent_iter: AsyncIterator[Any] = aiter(stream.subagents) + + async def _consume_subagents() -> None: + completion_tasks: list[asyncio.Task[Any]] = [] + try: + async for subagent in subagent_iter: + path = tuple(str(part) for part in subagent.path) + name = subagent.name + if not name: + continue + name = str(name) + description = "" + cause = subagent.cause + trigger_call_id = "" + if cause and cause["type"] == "toolCall": + trigger_call_id = cause["tool_call_id"] + instance_id = ":".join(path) + await queue.put( + emitter.subagent_start( + name, + description, + instance_id=instance_id, + tool_call_id=trigger_call_id, + ).data + ) + subagents.register(path, name, description) + completion_tasks.append( + asyncio.create_task( + _await_subagent_done(subagent, name, instance_id) ) - if iv_type == "ask_user": - questions = ( - interrupt_value.get("questions", []) - if isinstance(interrupt_value, dict) - else getattr(interrupt_value, "questions", []) - ) - tc_id = ( - interrupt_value.get("tool_call_id", "") - if isinstance(interrupt_value, dict) - else getattr(interrupt_value, "tool_call_id", "") - ) - ns_parts = ( - interrupt_obj.get("ns", [""]) - if isinstance(interrupt_obj, dict) - else getattr(interrupt_obj, "ns", [""]) - ) - interrupt_id = str(ns_parts[0]) if ns_parts else "default" - yield emitter.ask_user_interrupt( - interrupt_id, questions, tc_id - ).data - continue - - # Standard HITL approval interrupt - if isinstance(interrupt_value, dict): - action_reqs = interrupt_value.get("action_requests", []) - review_cfgs = interrupt_value.get("review_configs", []) - else: - action_reqs = getattr( - interrupt_value, "action_requests", [] - ) - review_cfgs = getattr(interrupt_value, "review_configs", []) - if action_reqs: - ns_parts = ( - interrupt_obj.get("ns", [""]) - if isinstance(interrupt_obj, dict) - else getattr(interrupt_obj, "ns", [""]) - ) - interrupt_id = str(ns_parts[0]) if ns_parts else "default" - yield emitter.interrupt( - interrupt_id, action_reqs, review_cfgs - ).data - summarization_event = _find_summarization_event_payload(data) - if summarization_event and not _summarization_in_progress: - signature = _summarization_event_signature(summarization_event) - if ( - signature is not None - and signature == _baseline_summarization_signature - ): - continue - summary_text = _extract_summary_message_text( - summarization_event.get("summary_message") ) - if summary_text: - yield emitter.summarization_start().data - _summarization_in_progress = True - yield emitter.summarization(summary_text).data + if completion_tasks: + await asyncio.gather(*completion_tasks) + finally: + subagents.close() + + async def _run_producer(coro: Any) -> None: + try: + await coro + except BaseException as exc: + await queue.put(exc) + finally: + await queue.put(producer_done) + + producers = [ + asyncio.create_task(_run_producer(_consume_subagents())), + asyncio.create_task(_run_producer(_consume_protocol_events())), + ] + pending_producers = len(producers) + + while pending_producers: + item = await queue.get() + if item is producer_done: + pending_producers -= 1 continue - if mode_str != "messages": - continue - - # Unpack message + metadata from data - msg: Any - metadata: dict = {} - if isinstance(data, tuple) and len(data) >= 2: - msg = data[0] - metadata = data[1] or {} - else: - msg = data - - # Accumulate summarization middleware chunks and emit text incrementally. - # The summarization LLM streams AIMessageChunks; content may be a - # plain string or a list of content blocks (provider-dependent). - if ( - isinstance(metadata, dict) - and metadata.get("lc_source") == "summarization" - ): - chunk_text = _extract_summarization_text(msg) - if chunk_text: - if not _summarization_in_progress: - yield emitter.summarization_start().data - _summarization_in_progress = True - yield emitter.summarization(chunk_text).data - continue - - # Suppress LLMToolSelectorMiddleware streaming output. - # The selector streams JSON like '{"tools":[...]}' via ainvoke - # which gets captured by astream. We detect it by content since - # the _selector_active flag is not visible in the streaming loop. - # Uses _tool_selection_suppressing to track suppression state - # and emits a tool_selection event from the tracker ContextVar. - if isinstance(msg, AIMessageChunk | AIMessage): - _raw = msg.content - _text = ( - _raw - if isinstance(_raw, str) - else "".join( - b.get("text", "") if isinstance(b, dict) else str(b) - for b in _raw - ) - if isinstance(_raw, list) - else "" - ) - - # Suppress Anthropic's structured output tool calls. - # Anthropic implements with_structured_output via tool calls - # named "ToolSelectionResponse" (exact LangChain convention). - # After the initial tool call, Anthropic streams input_json_delta - # chunks with the JSON arguments — suppress those too. - if hasattr(msg, "tool_calls") and msg.tool_calls: - _tc_names = [tc.get("name", "") for tc in msg.tool_calls] - if any(n == "ToolSelectionResponse" for n in _tc_names): - _tool_selection_was_active = True - continue - # Skip follow-up tool call chunks with empty names - # (streaming fragments of the ToolSelectionResponse) - if _tool_selection_was_active and all(n == "" for n in _tc_names): - continue - if _tool_selection_was_active and isinstance(_raw, list): - if any( - isinstance(b, dict) and b.get("type") == "input_json_delta" - for b in _raw - ): - continue # still streaming selector tool call args - - # Universal selector JSON detection via buffering. - # Buffer chunks starting with '{', try JSON parse, - # suppress if contains "tools". If not selector JSON, - # stop buffering and let the chunk through (don't drop it). - if _tool_selection_suppressing: - _tool_selection_buffer += _text - try: - import json as _json - - _parsed = _json.loads(_tool_selection_buffer.strip()) - if isinstance(_parsed, dict) and "tools" in _parsed: - _tool_selection_was_active = True - _tool_selection_suppressing = False - _tool_selection_buffer = "" - continue - # Valid JSON but not selector — stop buffering, - # fall through so this chunk is processed normally. - _tool_selection_suppressing = False - _tool_selection_buffer = "" - except (ValueError, TypeError): - _buf = _tool_selection_buffer.strip() - # Concatenated JSONs: {"tools":[...]}{"tools":[...]} - if '"tools"' in _buf and _buf.endswith("}"): - _tool_selection_was_active = True - _tool_selection_suppressing = False - _tool_selection_buffer = "" - continue - if len(_tool_selection_buffer) > 10000: - _tool_selection_suppressing = False - _tool_selection_buffer = "" - # Fall through — don't drop content - else: - continue # keep buffering - if ( - not _tool_selection_suppressing - and _text.lstrip().startswith("{") - and ('"tools"' in _text or len(_text.strip()) <= 10) - ): - # Try immediate parse (ccproxy returns full JSON in one chunk) - _stripped_text = _text.strip() - try: - import json as _json2 - - _parsed2 = _json2.loads(_stripped_text) - if isinstance(_parsed2, dict) and "tools" in _parsed2: - _tool_selection_was_active = True - continue - except (ValueError, TypeError): - # Could be concatenated JSONs: {"tools":[...]}{"tools":[...]} - # or incomplete JSON from streamed provider. - if '"tools"' in _stripped_text and _stripped_text.endswith("}"): - _tool_selection_was_active = True - continue - # Incomplete JSON — start buffering for streamed providers - _tool_selection_suppressing = True - _tool_selection_buffer = _text - continue - - # Emit tool_selection event on first non-empty chunk after - # suppression. Empty chunks arrive before the tracker has - # captured the selected tools, so we skip them. - if _tool_selection_was_active: - import EvoScientist.middleware.tool_selector as _ts_mod - - if _ts_mod._current_selected_tools: - _tool_selection_was_active = False - selected = _ts_mod._current_selected_tools - # Only show Panel when: - # 1. Tools were actually filtered (not all selected) - # 2. Selection changed from last time - if len(selected) < _ts_mod._total_tools_count and sorted( - selected - ) != sorted(_ts_mod._last_emitted_tools): - yield emitter.tool_selection(list(selected)).data - _ts_mod._last_emitted_tools = list(selected) - _ts_mod._current_selected_tools = [] - elif _text or (hasattr(msg, "tool_calls") and msg.tool_calls): - # Non-empty content arrived but no selected tools — - # tracker didn't run (shouldn't happen). Give up. - _tool_selection_was_active = False - - 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: - # Sub-agent content -- emit sub-agent events - for ev in _process_chunk_content(msg, emitter, subagent_tracker): - if ev.type == "tool_call": - yield emitter.subagent_tool_call( - subagent, - ev.data["name"], - ev.data["args"], - ev.data.get("id", ""), - ).data - # Emit sub-agent text for fallback extraction - # (not displayed in TUI, but available to consumers) - elif ev.type == "text": - yield emitter.subagent_text( - subagent, - ev.data.get("content", ""), - instance_id=tracker_key, - ).data - - if hasattr(msg, "tool_calls") and msg.tool_calls: - for tc in msg.tool_calls: - name = tc.get("name", "") - args = tc.get("args", {}) - tool_id = tc.get("id", "") - # Skip empty-name chunks (incomplete streaming fragments) - if not name and not tool_id: - continue - yield emitter.subagent_tool_call( - subagent, - name, - args if isinstance(args, dict) else {}, - tool_id, - ).data - else: - # Main agent content - for ev in _process_chunk_content(msg, emitter, main_tracker): - if ev.type == "text": - full_response += ev.data.get("content", "") - yield ev.data - - if hasattr(msg, "tool_calls") and msg.tool_calls: - for ev in _process_tool_calls( - msg.tool_calls, emitter, main_tracker - ): - yield ev.data - # Detect task tool calls -> announce sub-agent - tc_data = ev.data - if tc_data.get("name") == "task": - started_name = _register_task_tool_call(tc_data) - if started_name: - desc = str( - tc_data.get("args", {}).get("description", "") - ).strip() - yield emitter.subagent_start( - started_name, desc - ).data - - # Process ToolMessage (tool execution result) - elif hasattr(msg, "type") and msg.type == "tool": - if subagent: - if subagent_tracker: - subagent_tracker.finalize_all() - for info in subagent_tracker.emit_all_pending(): - yield emitter.subagent_tool_call( - subagent, - info.name, - info.args, - info.id, - ).data - name = getattr(msg, "name", "unknown") - tool_call_id = getattr(msg, "tool_call_id", "") or "" - raw_content, _is_img = _extract_tool_content(msg) - content = raw_content[: DisplayLimits.TOOL_RESULT_MAX] - success = is_success(content) - yield emitter.subagent_tool_result( - subagent, name, content, success, tool_call_id - ).data - else: - for ev in _process_tool_result(msg, emitter, main_tracker): - yield ev.data - # Tool result can re-emit tool_call with full args; update task mapping - if ev.type == "tool_call" and ev.data.get("name") == "task": - started_name = _register_task_tool_call(ev.data) - if started_name: - desc = str( - ev.data.get("args", {}).get("description", "") - ).strip() - yield emitter.subagent_start(started_name, desc).data - # Check if this is a task result -> sub-agent ended - name = getattr(msg, "name", "") - if name == "task": - tool_call_id = getattr(msg, "tool_call_id", "") - # Find the sub-agent name via tool_call_id map - sa_name = _task_id_to_name.get(tool_call_id, "sub-agent") - yield emitter.subagent_end(sa_name).data - + if isinstance(item, BaseException): + for task in producers: + task.cancel() + await asyncio.gather(*producers, return_exceptions=True) + raise item + yield item except Exception as e: yield emitter.error(str(e)).data raise + finally: + if stream is not None: + try: + result = stream.abort() + if inspect.isawaitable(result): + await result + except Exception: + pass + for task in producers: + if not task.done(): + task.cancel() + if producers: + await asyncio.gather(*producers, return_exceptions=True) - yield emitter.done(full_response).data - - -def _process_chunk_content( - chunk, emitter: StreamEventEmitter, tracker: ToolCallTracker -): - """Process content blocks from an AI message chunk.""" - content = chunk.content - - # OpenRouter (langchain-openrouter) stores reasoning in - # additional_kwargs["reasoning_content"] instead of content blocks. - _additional = getattr(chunk, "additional_kwargs", None) or {} - _rc = _additional.get("reasoning_content") - _emitted_thinking = False - if _rc and isinstance(_rc, str): - yield emitter.thinking(_rc) - _emitted_thinking = True - - if isinstance(content, str): - if content: - cleaned = _strip_legacy_thinking_tags(content) - # Skip whitespace-only text (OpenRouter may send '\n\n' before - # reasoning chunks, which would prematurely trigger is_responding). - if cleaned and not cleaned.isspace(): - yield emitter.text(cleaned) - return - - blocks = None - if hasattr(chunk, "content_blocks"): - try: - blocks = chunk.content_blocks - except Exception: - blocks = None - - if blocks is None: - if isinstance(content, dict): - blocks = [content] - elif isinstance(content, list): - blocks = content - else: - return - - for raw_block in blocks: - block = raw_block - if not isinstance(block, dict): - if hasattr(block, "model_dump"): - block = block.model_dump() - elif hasattr(block, "dict"): - block = block.dict() - else: - continue - - block_type = block.get("type") - - if block_type in ("thinking", "reasoning"): - thinking_text = block.get("thinking") or block.get("reasoning") or "" - # Skip if already emitted from additional_kwargs (avoid duplicates) - if thinking_text and not _emitted_thinking: - yield emitter.thinking(thinking_text) - - elif block_type == "text": - text = block.get("text") or block.get("content") or "" - if text: - text = _strip_legacy_thinking_tags(text) - if text: - yield emitter.text(text) - - elif block_type in ("tool_use", "tool_call"): - tool_id = block.get("id", "") - name = block.get("name", "") - args = block.get("input") if block_type == "tool_use" else block.get("args") - args_payload = args if isinstance(args, dict) else {} - - if tool_id: - tracker.update(tool_id, name=name, args=args_payload) - if tracker.is_ready(tool_id): - tracker.mark_emitted(tool_id) - yield emitter.tool_call(name, args_payload, tool_id) - - elif block_type == "input_json_delta": - partial_json = block.get("partial_json", "") - if partial_json: - tracker.append_json_delta(partial_json, block.get("index", 0)) - - elif block_type == "tool_call_chunk": - tool_id = block.get("id", "") - name = block.get("name", "") - if tool_id: - tracker.update(tool_id, name=name) - partial_args = block.get("args", "") - if isinstance(partial_args, str) and partial_args: - tracker.append_json_delta(partial_args, block.get("index", 0)) - - -def _process_tool_calls( - tool_calls: list, emitter: StreamEventEmitter, tracker: ToolCallTracker -): - """Process tool_calls from chunk.tool_calls attribute.""" - for tc in tool_calls: - tool_id = tc.get("id", "") - if tool_id: - name = tc.get("name", "") - args = tc.get("args", {}) - args_payload = args if isinstance(args, dict) else {} - - tracker.update(tool_id, name=name, args=args_payload) - if tracker.is_ready(tool_id): - tracker.mark_emitted(tool_id) - yield emitter.tool_call(name, args_payload, tool_id) - - -def _process_tool_result(chunk, emitter: StreamEventEmitter, tracker: ToolCallTracker): - """Process a ToolMessage result.""" - tracker.finalize_all() - - # Re-emit all tool calls with complete args - for info in tracker.get_all(): - yield emitter.tool_call(info.name, info.args, info.id) - - name = getattr(chunk, "name", "unknown") - raw_content, _is_img = _extract_tool_content(chunk) - content = raw_content[: DisplayLimits.TOOL_RESULT_MAX] - if len(raw_content) > DisplayLimits.TOOL_RESULT_MAX: - content += "\n... (truncated)" - - success = is_success(content) - yield emitter.tool_result(name, content, success) + yield emitter.done(processor.full_response).data diff --git a/EvoScientist/stream/state.py b/EvoScientist/stream/state.py index 8258661..b3b7374 100644 --- a/EvoScientist/stream/state.py +++ b/EvoScientist/stream/state.py @@ -22,39 +22,40 @@ class ResearchPhase(StrEnum): class SubAgentState: """Tracks a single sub-agent's activity.""" - def __init__(self, name: str, description: str = ""): + def __init__( + self, + name: str, + description: str = "", + instance_id: str = "", + parent_tool_call_id: str = "", + ): self.name = name self.description = description + self.instance_id = instance_id + self.parent_tool_call_id = parent_tool_call_id self.tool_calls: list[dict] = [] self.tool_results: list[dict] = [] self._result_map: dict[str, dict] = {} # tool_call_id -> result self.is_active = True - def add_tool_call(self, name: str, args: dict, tool_id: str = ""): - # Skip empty-name calls without an id (incomplete streaming chunks) - if not name and not tool_id: + def add_tool_call(self, name: str, args: dict, tool_id: str): + if not name or not tool_id: return tc_data = {"id": tool_id, "name": name, "args": args} - if tool_id: - for i, tc in enumerate(self.tool_calls): - if tc.get("id") == tool_id: - # Merge: keep the non-empty name/args - if name: - self.tool_calls[i]["name"] = name - if args: - self.tool_calls[i]["args"] = args - return - # Skip if name is empty and we can't deduplicate by id - if not name: - return + for i, tc in enumerate(self.tool_calls): + if tc["id"] == tool_id: + self.tool_calls[i]["name"] = name + if args: + self.tool_calls[i]["args"] = args + return self.tool_calls.append(tc_data) def add_tool_result( self, name: str, content: str, - success: bool = True, - tool_call_id: str = "", + success: bool, + tool_call_id: str, ): result = { "name": name, @@ -63,39 +64,14 @@ class SubAgentState: "tool_call_id": tool_call_id, } self.tool_results.append(result) - # Preferred: exact id match (correct under concurrent same-name tools). - if tool_call_id: - for tc in self.tool_calls: - if tc.get("id") == tool_call_id: - self._result_map[tool_call_id] = result - return - # Fallback: first unmatched tool call with same name. for tc in self.tool_calls: - tc_id = tc.get("id", "") - tc_name = tc.get("name", "") - if tc_id and tc_id not in self._result_map and tc_name == name: - self._result_map[tc_id] = result - return - # Last resort: first unmatched tool call regardless of name. - for tc in self.tool_calls: - tc_id = tc.get("id", "") - if tc_id and tc_id not in self._result_map: - self._result_map[tc_id] = result + if tc["id"] == tool_call_id: + self._result_map[tool_call_id] = result return def get_result_for(self, tc: dict) -> dict | None: """Get matched result for a tool call.""" - tc_id = tc.get("id", "") - if tc_id: - return self._result_map.get(tc_id) - # Fallback: index-based matching - try: - idx = self.tool_calls.index(tc) - if idx < len(self.tool_results): - return self.tool_results[idx] - except ValueError: - pass - return None + return self._result_map.get(tc["id"]) class StreamState: @@ -113,7 +89,7 @@ class StreamState: self.is_processing = False # Sub-agent tracking self.subagents: list[SubAgentState] = [] - self._subagent_map: dict[str, SubAgentState] = {} # name -> state + self._subagent_map: dict[str, SubAgentState] = {} # Todo list tracking self.todo_items: list[dict] = [] # Latest text segment (reset on each tool_call) @@ -150,51 +126,28 @@ class StreamState: return self._cached_md def _get_or_create_subagent( - self, name: str, description: str = "" + self, + name: str, + description: str, + instance_id: str, + parent_tool_call_id: str = "", ) -> SubAgentState: - if name not in self._subagent_map: - # Case 1: real name arrives, "sub-agent" entry exists -> rename it - if name != "sub-agent" and "sub-agent" in self._subagent_map: - old_sa = self._subagent_map.pop("sub-agent") - old_sa.name = name - if description: - old_sa.description = description - self._subagent_map[name] = old_sa - return old_sa - # Case 2: "sub-agent" arrives but a pre-registered real-name entry - # exists with no tool calls -> merge into it - if name == "sub-agent": - active_named = [ - sa - for sa in self.subagents - if sa.is_active and sa.name != "sub-agent" - ] - if len(active_named) == 1 and not active_named[0].tool_calls: - self._subagent_map[name] = active_named[0] - return active_named[0] - sa = SubAgentState(name, description) + key = instance_id + if key not in self._subagent_map: + sa = SubAgentState(name, description, instance_id, parent_tool_call_id) self.subagents.append(sa) - self._subagent_map[name] = sa + self._subagent_map[key] = sa else: - existing = self._subagent_map[name] + existing = self._subagent_map[key] + if name and existing.name != name: + existing.name = name if description and not existing.description: existing.description = description - # If this entry was created as "sub-agent" placeholder and the - # actual name is different, update. - if name != "sub-agent" and existing.name == "sub-agent": - existing.name = name - return self._subagent_map[name] - - def _resolve_subagent_name(self, name: str) -> str: - """Resolve "sub-agent" to the single active named sub-agent when possible.""" - if name != "sub-agent": - return name - active_named = [ - sa.name for sa in self.subagents if sa.is_active and sa.name != "sub-agent" - ] - if len(active_named) == 1: - return active_named[0] - return name + if instance_id and not existing.instance_id: + existing.instance_id = instance_id + if parent_tool_call_id and not existing.parent_tool_call_id: + existing.parent_tool_call_id = parent_tool_call_id + return self._subagent_map[key] def handle_event(self, event: dict) -> str: """Process a single stream event, update internal state, return event type.""" @@ -220,7 +173,7 @@ class StreamState: self.is_responding = False self.is_processing = False - tool_id = event.get("id", "") + tool_id = event["id"] tool_name = event.get("name", "unknown") tool_args = event.get("args", {}) tc_data = { @@ -229,21 +182,13 @@ class StreamState: "args": tool_args, } - if tool_id: - updated = False - for i, tc in enumerate(self.tool_calls): - if tc.get("id") == tool_id: - self.tool_calls[i] = tc_data - updated = True - break - if not updated: - if self.latest_text.strip(): - self.narration_segments.append( - (len(self.tool_calls), self.latest_text) - ) - self.narrated_response_end = len(self.response_text) - self.tool_calls.append(tc_data) - else: + updated = False + for i, tc in enumerate(self.tool_calls): + if tc["id"] == tool_id: + self.tool_calls[i] = tc_data + updated = True + break + if not updated: if self.latest_text.strip(): self.narration_segments.append( (len(self.tool_calls), self.latest_text) @@ -267,49 +212,53 @@ class StreamState: { "name": result_name, "content": result_content, + "id": event["id"], + "success": event.get("success", True), } ) - # Update todo list from write_todos / read_todos results (fallback) + # Update todo list from write_todos / read_todos tool results. if result_name in ("write_todos", "read_todos"): parsed = _parse_todo_items(result_content) if parsed: self.todo_items = parsed elif event_type == "subagent_start": - name = event.get("name", "sub-agent") + name = event["name"] desc = event.get("description", "") - sa = self._get_or_create_subagent(name, desc) + instance_id = event["instance_id"] + sa = self._get_or_create_subagent( + name, desc, instance_id, event["tool_call_id"] + ) sa.is_active = True elif event_type == "subagent_tool_call": - sa_name = self._resolve_subagent_name(event.get("subagent", "sub-agent")) - sa = self._get_or_create_subagent(sa_name) + instance_id = event["instance_id"] + sa = self._subagent_map.get(instance_id) + if sa is None: + return event_type sa.add_tool_call( event.get("name", "unknown"), event.get("args", {}), - event.get("id", ""), + event["id"], ) elif event_type == "subagent_tool_result": - sa_name = self._resolve_subagent_name(event.get("subagent", "sub-agent")) - sa = self._get_or_create_subagent(sa_name) + instance_id = event["instance_id"] + sa = self._subagent_map.get(instance_id) + if sa is None: + return event_type sa.add_tool_result( event.get("name", "unknown"), event.get("content", ""), event.get("success", True), - event.get("id", ""), + event["id"], ) elif event_type == "subagent_end": - name = self._resolve_subagent_name(event.get("name", "sub-agent")) - if name in self._subagent_map: - self._subagent_map[name].is_active = False - elif name == "sub-agent": - # Couldn't resolve -- deactivate the oldest active sub-agent - for sa in self.subagents: - if sa.is_active: - sa.is_active = False - break + instance_id = event["instance_id"] + key = instance_id + if key in self._subagent_map: + self._subagent_map[key].is_active = False elif event_type == "interrupt": self.pending_interrupt = event diff --git a/EvoScientist/stream/summarization.py b/EvoScientist/stream/summarization.py new file mode 100644 index 0000000..bfb6a5b --- /dev/null +++ b/EvoScientist/stream/summarization.py @@ -0,0 +1,94 @@ +"""Summarization event helpers for stream v3 updates.""" + +import re +from collections.abc import Mapping + +from langchain_core.messages import BaseMessage + +from .v3_payloads import RawMap, _as_raw_map + +_SUMMARY_TAG_RE = re.compile( + r"\s*(.*?)\s*", re.DOTALL | re.IGNORECASE +) + + +def _extract_summarization_text(msg: BaseMessage) -> str: + """Extract plain text from a summarization message or chunk.""" + content = msg.content + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for block in content: + if isinstance(block, str): + parts.append(block) + continue + block_map = _as_raw_map(block) + if block_map is not None: + text = block_map.get("text") + if isinstance(text, str): + parts.append(text) + return "".join(parts) + return "" + + +def _extract_summary_message_text(summary_message: BaseMessage | None) -> str: + """Extract user-facing summary text from a stored summarization event.""" + if summary_message is None: + return "" + text = _extract_summarization_text(summary_message) + if not text: + return "" + + match = _SUMMARY_TAG_RE.search(text) + if match: + return match.group(1).strip() + + prefix = "Here is a summary of the conversation to date:" + if text.startswith(prefix): + return text[len(prefix) :].strip() + + return text.strip() + + +def _find_summarization_event_payload(data: object) -> RawMap | None: + """Find a `_summarization_event` dict anywhere inside an updates payload.""" + seen: set[int] = set() + stack: list[object] = [data] + + while stack: + item = stack.pop() + item_id = id(item) + if item_id in seen: + continue + seen.add(item_id) + + item_map = _as_raw_map(item) + if item_map is not None: + event = _as_raw_map(item_map.get("_summarization_event")) + if event is not None: + return event + stack.extend(item_map.values()) + continue + + if isinstance(item, list | tuple): + stack.extend(item) + + return None + + +def _summarization_event_signature( + event: Mapping[str, object] | None, +) -> tuple[object, ...] | None: + """Build a stable signature for a persisted summarization event.""" + if event is None: + return None + summary_message = event.get("summary_message") + summary_text = _extract_summary_message_text( + summary_message if isinstance(summary_message, BaseMessage) else None + ) + return ( + event.get("cutoff_index"), + event.get("file_path"), + summary_text, + ) diff --git a/EvoScientist/stream/tool_results.py b/EvoScientist/stream/tool_results.py new file mode 100644 index 0000000..8417773 --- /dev/null +++ b/EvoScientist/stream/tool_results.py @@ -0,0 +1,72 @@ +"""Tool result extraction helpers for stream events.""" + +from langchain_core.messages import ToolMessage +from langgraph.types import Command + +from .v3_payloads import _as_raw_map + +_IMAGE_MEDIA_TYPES = { + "image/png", + "image/jpeg", + "image/gif", + "image/webp", + "image/bmp", + "image/svg+xml", +} + + +def _extract_tool_content(msg: ToolMessage) -> tuple[str, bool]: + """Extract display-safe content from a ToolMessage. + + DeepAgents ``read_file`` returns image content as + ``ToolMessage(content=[ImageContentBlock])`` with + ``additional_kwargs["read_file_media_type"]`` set. Stringifying that would + dump huge base64 data into the display. + """ + additional = msg.additional_kwargs + media_type = additional.get("read_file_media_type", "") + if media_type and media_type in _IMAGE_MEDIA_TYPES: + file_path = additional.get("read_file_path", "") + if not file_path: + file_path = msg.name or "image" + return f"[OK] Image displayed: {file_path} ({media_type})", True + + content = msg.content + if isinstance(content, list): + for block in content: + block_map = _as_raw_map(block) + if block_map is not None: + if block_map.get("type") == "image" or "base64" in block_map: + return "[OK] Image displayed", True + + parts: list[str] = [] + for block in content: + block_map = _as_raw_map(block) + if block_map is not None: + text = block_map.get("text") + if text: + parts.append(str(text)) + elif isinstance(block, str): + parts.append(block) + return "\n".join(parts) if parts else str(content), False + + return str(content), False + + +def _extract_command_tool_content(output: Command, tool_call_id: str) -> str | None: + if not tool_call_id: + return None + update = _as_raw_map(output.update) + if update is None: + return None + messages = update.get("messages") + if not isinstance(messages, list): + return None + for msg in messages: + if not isinstance(msg, ToolMessage): + continue + if msg.tool_call_id != tool_call_id: + continue + content, _ = _extract_tool_content(msg) + return content + return None diff --git a/EvoScientist/stream/tool_selection.py b/EvoScientist/stream/tool_selection.py new file mode 100644 index 0000000..fd9eef5 --- /dev/null +++ b/EvoScientist/stream/tool_selection.py @@ -0,0 +1,117 @@ +"""Tool-selection stream suppression helpers.""" + +from typing import Any + +from .emitter import StreamEventEmitter + + +class _ToolSelectionSuppressor: + """Suppress selector model JSON while preserving the UI selection event.""" + + def __init__(self, emitter: StreamEventEmitter) -> None: + self._emitter = emitter + self._buffering = False + self._buffer = "" + self._was_active = False + + def observe_tool_block(self, name: str) -> bool: + if name == "ToolSelectionResponse": + self._was_active = True + return True + return False + + def process_text(self, text: str) -> tuple[bool, list[dict[str, Any]], str]: + events = self._emit_selection_if_ready(text) + if not text: + return False, events, "" + + if self._buffering: + self._buffer += text + json_kind = self._json_buffer_kind(self._buffer) + if json_kind == "selector" and self._selector_context_active(): + self._was_active = True + self._buffering = False + self._buffer = "" + return True, events, "" + if json_kind == "complete": + replay = self._buffer + self._buffering = False + self._buffer = "" + return False, events, replay + if len(self._buffer) <= 10000: + return True, events, "" + replay = self._buffer + self._buffering = False + self._buffer = "" + return False, events, replay + + stripped = text.strip() + if ( + self._selector_context_active() + and stripped.startswith("{") + and ('"tools"' in stripped or len(stripped) <= 10) + ): + json_kind = self._json_buffer_kind(stripped) + if json_kind == "selector": + self._was_active = True + return True, events, "" + if json_kind == "complete": + return False, events, text + self._buffering = True + self._buffer = text + return True, events, "" + + return False, events, text + + @staticmethod + def _json_buffer_kind(text: str) -> str: + try: + import json + + parsed = json.loads(text.strip()) + except (TypeError, ValueError): + stripped = text.strip() + if '"tools"' in stripped and stripped.endswith("}"): + return "selector" + return "incomplete" + if isinstance(parsed, dict) and "tools" in parsed: + return "selector" + return "complete" + + def _selector_context_active(self) -> bool: + if self._was_active: + return True + import EvoScientist.middleware.tool_selector as selector_mod + + return bool( + selector_mod._selector_active or selector_mod._current_selected_tools + ) + + def flush_selection(self) -> list[dict[str, Any]]: + return self._emit_selection_if_ready("") + + def flush_pending_text(self) -> str: + if not self._buffering: + return "" + replay = self._buffer + self._buffering = False + self._buffer = "" + return replay + + def _emit_selection_if_ready(self, text: str) -> list[dict[str, Any]]: + if not self._was_active: + return [] + import EvoScientist.middleware.tool_selector as selector_mod + + if selector_mod._current_selected_tools: + self._was_active = False + selected = selector_mod._current_selected_tools + selector_mod._current_selected_tools = [] + if len(selected) < selector_mod._total_tools_count and sorted( + selected + ) != sorted(selector_mod._last_emitted_tools): + selector_mod._last_emitted_tools = list(selected) + return [self._emitter.tool_selection(list(selected)).data] + elif text: + self._was_active = False + return [] diff --git a/EvoScientist/stream/tracker.py b/EvoScientist/stream/tracker.py deleted file mode 100644 index 821613c..0000000 --- a/EvoScientist/stream/tracker.py +++ /dev/null @@ -1,115 +0,0 @@ -""" -ToolCallTracker - manages incremental JSON parsing for tool parameters. - -Handles tool_use blocks where arguments arrive in fragments via input_json_delta. -""" - -import json -from dataclasses import dataclass, field - - -@dataclass -class ToolCallInfo: - """Tool call information.""" - - id: str - name: str - args: dict = field(default_factory=dict) - emitted: bool = False - args_complete: bool = False - _json_buffer: str = "" - - -class ToolCallTracker: - """Tool call tracker for incremental argument parsing. - - Usage: - tracker = ToolCallTracker() - tracker.update(tool_id, name="execute") - tracker.append_json_delta('{"command') - tracker.append_json_delta('": "ls"}') - tracker.finalize_all() - info = tracker.get(tool_id) - yield emitter.tool_call(info.name, info.args) - """ - - def __init__(self): - self._calls: dict[str, ToolCallInfo] = {} - self._last_tool_id: str | None = None - - def update( - self, - tool_id: str, - name: str | None = None, - args: dict | None = None, - args_complete: bool = False, - ) -> None: - """Update tool call info (accumulative).""" - if tool_id not in self._calls: - self._calls[tool_id] = ToolCallInfo( - id=tool_id, - name=name or "", - args=args or {}, - args_complete=args_complete, - ) - self._last_tool_id = tool_id - else: - info = self._calls[tool_id] - if name: - info.name = name - if args: - info.args = args - if args_complete: - info.args_complete = True - - def append_json_delta(self, partial_json: str, index: int = 0) -> None: - """Accumulate input_json_delta fragment.""" - tool_id = self._last_tool_id - if tool_id and tool_id in self._calls: - self._calls[tool_id]._json_buffer += partial_json - - def finalize_all(self) -> None: - """Finalize all tool calls: parse accumulated JSON and mark complete.""" - for info in self._calls.values(): - if info._json_buffer: - try: - info.args = json.loads(info._json_buffer) - except json.JSONDecodeError: - pass - info._json_buffer = "" - info.args_complete = True - - def is_ready(self, tool_id: str) -> bool: - """Check if a tool call is ready to emit (has name and not yet emitted).""" - if tool_id not in self._calls: - return False - info = self._calls[tool_id] - return bool(info.name) and not info.emitted - - def get_all(self) -> list[ToolCallInfo]: - """Get all tool calls.""" - return list(self._calls.values()) - - def mark_emitted(self, tool_id: str) -> None: - """Mark a tool call as emitted.""" - if tool_id in self._calls: - self._calls[tool_id].emitted = True - - def get(self, tool_id: str) -> ToolCallInfo | None: - """Get tool call info by ID.""" - return self._calls.get(tool_id) - - def get_pending(self) -> list[ToolCallInfo]: - """Get all unemitted tool calls.""" - return [info for info in self._calls.values() if not info.emitted] - - def emit_all_pending(self) -> list[ToolCallInfo]: - """Emit all pending tool calls and mark them.""" - pending = self.get_pending() - for info in pending: - info.emitted = True - return pending - - def clear(self) -> None: - """Clear all tracked tool calls.""" - self._calls.clear() diff --git a/EvoScientist/stream/utils.py b/EvoScientist/stream/utils.py index 830bdab..f01f12b 100644 --- a/EvoScientist/stream/utils.py +++ b/EvoScientist/stream/utils.py @@ -215,12 +215,6 @@ def format_tool_compact(name: str, args: dict | None) -> str: task_desc = task_desc[:47] + "\u2026" return f"Cooking with {sa_type} — {task_desc}" return f"Cooking with {sa_type}" - # Fallback if no subagent_type - if task_desc: - if len(task_desc) > 50: - task_desc = task_desc[:47] + "\u2026" - return f"Cooking with sub-agent — {task_desc}" - return "Cooking with sub-agent" # Web search if name_lower in ("tavily_search", "internet_search"): diff --git a/EvoScientist/stream/v3_payloads.py b/EvoScientist/stream/v3_payloads.py new file mode 100644 index 0000000..a62a667 --- /dev/null +++ b/EvoScientist/stream/v3_payloads.py @@ -0,0 +1,84 @@ +"""Small helpers for raw LangGraph v3 stream payloads.""" + +import re +from collections.abc import Mapping +from typing import cast + +RawMap = Mapping[str, object] + +_THINKING_TAG_RE = re.compile(r".*?", re.DOTALL) + + +def _as_raw_map(value: object) -> RawMap | None: + if not isinstance(value, Mapping): + return None + if not all(isinstance(key, str) for key in value): + return None + return cast("RawMap", value) + + +def _strip_legacy_thinking_tags(content: str) -> str: + """Remove ``...`` tags from content strings.""" + return _THINKING_TAG_RE.sub("", content) + + +def _event_namespace(event: Mapping[str, object]) -> tuple[str, ...]: + params = _as_raw_map(event.get("params")) or {} + namespace = params.get("namespace") + if not isinstance(namespace, list | tuple): + return () + return tuple(str(part) for part in namespace) + + +def _event_data(event: Mapping[str, object]) -> object: + params = _as_raw_map(event.get("params")) or {} + return params.get("data") + + +def _split_message_event_data(data: object) -> tuple[object, RawMap]: + if isinstance(data, tuple) and len(data) >= 2: + metadata = _as_raw_map(data[1]) or {} + return data[0], metadata + return data, {} + + +def _usage_counts(usage: Mapping[str, object] | None) -> tuple[int, int]: + if not usage: + return 0, 0 + input_tokens = usage.get("input_tokens") + output_tokens = usage.get("output_tokens") + return ( + input_tokens if isinstance(input_tokens, int) else 0, + output_tokens if isinstance(output_tokens, int) else 0, + ) + + +def _text_from_content(content: object) -> str: + if isinstance(content, str): + return content + if isinstance(content, list): + parts: list[str] = [] + for block in content: + if isinstance(block, str): + parts.append(block) + continue + block_map = _as_raw_map(block) + if block_map is not None: + text = block_map.get("text") + if isinstance(text, str): + parts.append(text) + return "".join(parts) + return "" + + +def _reasoning_from_content(content: object) -> str: + if not isinstance(content, list): + return "" + parts: list[str] = [] + for block in content: + block_map = _as_raw_map(block) + if block_map is not None: + reasoning = block_map.get("reasoning") or block_map.get("thinking") + if isinstance(reasoning, str): + parts.append(reasoning) + return "".join(parts) diff --git a/tests/conftest.py b/tests/conftest.py index 9994834..a3d8dab 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -40,7 +40,12 @@ def sample_tool_call(): @pytest.fixture def sample_tool_result(): """A minimal tool result dict.""" - return {"name": "execute", "content": "[OK] file1.py file2.py", "success": True} + return { + "id": "tc_001", + "name": "execute", + "content": "[OK] file1.py file2.py", + "success": True, + } @pytest.fixture @@ -57,6 +62,7 @@ def sample_events(): }, { "type": "tool_result", + "id": "tc_001", "name": "execute", "content": "[OK] done", "success": True, @@ -65,10 +71,13 @@ def sample_events(): "type": "subagent_start", "name": "research-agent", "description": "Find papers", + "instance_id": "task:research", + "tool_call_id": "tc_task_001", }, { "type": "subagent_tool_call", "subagent": "research-agent", + "instance_id": "task:research", "name": "tavily_search", "args": {"query": "test"}, "id": "tc_sa_001", @@ -76,11 +85,17 @@ def sample_events(): { "type": "subagent_tool_result", "subagent": "research-agent", + "instance_id": "task:research", "name": "tavily_search", "content": "Results...", "success": True, + "id": "tc_sa_001", + }, + { + "type": "subagent_end", + "name": "research-agent", + "instance_id": "task:research", }, - {"type": "subagent_end", "name": "research-agent"}, {"type": "done", "response": "Here is the answer."}, ] diff --git a/tests/stream_v3_fakes.py b/tests/stream_v3_fakes.py new file mode 100644 index 0000000..b8316ab --- /dev/null +++ b/tests/stream_v3_fakes.py @@ -0,0 +1,269 @@ +"""Shared DeepAgents v3 protocol fakes for stream tests.""" + +from __future__ import annotations + +import asyncio +from collections.abc import AsyncIterator, Iterable +from dataclasses import dataclass, field +from typing import Any +from unittest.mock import MagicMock + +from EvoScientist.stream.events import stream_agent_events +from tests.conftest import run_async + + +async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]: + for item in items: + yield item + + +def collect_events(agent, message: str = "hi", thread_id: str = "t1"): + """Collect stream_agent_events output for synchronous tests.""" + + async def _run(): + events = [] + async for ev in stream_agent_events(agent, message, thread_id): + events.append(ev) + return events + + return run_async(_run()) + + +def protocol_event( + method: str, + data, + namespace: Iterable[Any] = (), + **params, +) -> dict[str, Any]: + """Build a minimal DeepAgents v3 protocol event.""" + return { + "type": "event", + "method": method, + "params": { + "namespace": list(namespace), + "timestamp": 0, + "data": data, + **params, + }, + } + + +def message_delta( + text: str, + metadata: dict[str, Any] | None = None, + namespace: Iterable[Any] = (), +) -> dict[str, Any]: + return protocol_event( + "messages", + ( + { + "event": "content-block-delta", + "index": 0, + "delta": {"type": "text-delta", "text": text}, + }, + metadata or {}, + ), + namespace, + ) + + +def message_finish( + usage: dict[str, int] | None = None, + metadata: dict[str, Any] | None = None, + namespace: Iterable[Any] = (), +) -> dict[str, Any]: + payload: dict[str, Any] = {"event": "message-finish"} + if usage is not None: + payload["usage"] = usage + return protocol_event("messages", (payload, metadata or {}), namespace) + + +def message_tool_call_block( + name: str, + args: dict[str, Any] | None = None, + *, + tool_call_id: str = "tc1", + namespace: Iterable[Any] = (), +) -> dict[str, Any]: + return protocol_event( + "messages", + ( + { + "event": "content-block-finish", + "content": { + "type": "tool_call", + "id": tool_call_id, + "name": name, + "args": args or {}, + }, + }, + {}, + ), + namespace, + ) + + +def tool_started( + name: str, + args: dict[str, Any] | None = None, + *, + tool_call_id: str = "tc1", + namespace: Iterable[Any] = (), +) -> dict[str, Any]: + return protocol_event( + "tools", + { + "event": "tool-started", + "tool_name": name, + "input": args or {}, + "tool_call_id": tool_call_id, + }, + namespace, + ) + + +def tool_finished( + output, + *, + tool_call_id: str = "tc1", + namespace: Iterable[Any] = (), +) -> dict[str, Any]: + return protocol_event( + "tools", + { + "event": "tool-finished", + "output": output, + "tool_call_id": tool_call_id, + }, + namespace, + ) + + +@dataclass +class FakeStateSnapshot: + values: dict[str, Any] = field(default_factory=dict) + + +class FakeV3Run: + def __init__(self, events: Iterable[Any], subagents: Iterable[Any] | None = None): + self._events = list(events) + self.subagents = async_iter(list(subagents or [])) + self.aborted = False + + def __aiter__(self): + return async_iter(self._events) + + async def abort(self) -> None: + self.aborted = True + + +class FakeV3Agent: + def __init__( + self, + events: Iterable[Any], + *, + state_values: dict[str, Any] | None = None, + subagents: Iterable[Any] | None = None, + ): + self._run = FakeV3Run(events, subagents=subagents) + self.astream_events = MagicMock(return_value=self._run) + self._state_values = state_values if state_values is not None else {} + + async def aget_state(self, _config): + return FakeStateSnapshot(values=self._state_values) + + +class ErroringV3Agent: + def __init__(self, exc: Exception): + self.exc = exc + + def astream_events(self, *_args, **_kwargs): + raise self.exc + + async def aget_state(self, _config): + return FakeStateSnapshot() + + +class HangingV3Run: + def __init__(self, events: Iterable[Any]): + self._events = list(events) + self.subagents = async_iter([]) + self.aborted = False + + async def abort(self) -> None: + self.aborted = True + + def __aiter__(self): + return self._iter_events() + + async def _iter_events(self): + for event in self._events: + yield event + await asyncio.Event().wait() + + +class HangingV3Agent: + def __init__(self, events: Iterable[Any]): + self._run = HangingV3Run(events) + self.astream_events = MagicMock(return_value=self._run) + + @property + def aborted(self) -> bool: + return self._run.aborted + + async def aget_state(self, _config): + return FakeStateSnapshot() + + +class LazySubagentChannel: + """Projection fake that only queues handles after subscription.""" + + def __init__(self, subagents: Iterable[Any]): + self._subagents = list(subagents) + self.subscribed = False + + def __aiter__(self): + self.subscribed = True + return async_iter(self._subagents) + + def drop_if_unsubscribed(self) -> None: + if not self.subscribed: + self._subagents = [] + + +class SubscriptionSensitiveV3Run: + def __init__(self, events: Iterable[Any], subagents: Iterable[Any]): + self._events = list(events) + self.subagents = LazySubagentChannel(subagents) + + def __aiter__(self): + self.subagents.drop_if_unsubscribed() + return async_iter(self._events) + + +class SubscriptionSensitiveV3Agent: + def __init__(self, events: Iterable[Any], subagents: Iterable[Any]): + self._run = SubscriptionSensitiveV3Run(events, subagents) + self.astream_events = MagicMock(return_value=self._run) + + async def aget_state(self, _config): + return FakeStateSnapshot() + + +@dataclass +class FakeSubagent: + path: Iterable[Any] + name: str = "research-agent" + tool_call_id: str = "" + + def __post_init__(self): + self.path = tuple(self.path) + if not self.tool_call_id: + self.tool_call_id = "call_" + "_".join(str(p) for p in self.path) + + @property + def cause(self) -> dict[str, str]: + return {"type": "toolCall", "tool_call_id": self.tool_call_id} + + async def output(self): + return {"messages": []} diff --git a/tests/test_event_loop.py b/tests/test_event_loop.py index 244f247..4f06e5f 100644 --- a/tests/test_event_loop.py +++ b/tests/test_event_loop.py @@ -3,9 +3,40 @@ import asyncio from unittest.mock import Mock, patch +import pytest + from EvoScientist.stream.display import _create_event_loop, _get_event_loop +class _TrackingEventLoopPolicy(asyncio.DefaultEventLoopPolicy): + """Event loop policy that records loops created by one test.""" + + def __init__(self): + super().__init__() + self.created_loops: list[asyncio.AbstractEventLoop] = [] + + def new_event_loop(self) -> asyncio.AbstractEventLoop: + loop = super().new_event_loop() + self.created_loops.append(loop) + return loop + + +@pytest.fixture(autouse=True) +def isolated_event_loop_policy(): + previous_policy = asyncio.get_event_loop_policy() + test_policy = _TrackingEventLoopPolicy() + asyncio.set_event_loop_policy(test_policy) + try: + yield + finally: + try: + for loop in test_policy.created_loops: + if not loop.is_closed(): + loop.close() + finally: + asyncio.set_event_loop_policy(previous_policy) + + class TestCreateEventLoop: """Tests for _create_event_loop helper.""" diff --git a/tests/test_hitl.py b/tests/test_hitl.py index 8f80df0..d235815 100644 --- a/tests/test_hitl.py +++ b/tests/test_hitl.py @@ -1,10 +1,17 @@ """Tests for HITL (Human-in-the-Loop) approval mechanism.""" -import asyncio from unittest.mock import MagicMock, patch +from langgraph.types import Interrupt + from EvoScientist.stream.emitter import StreamEvent, StreamEventEmitter from EvoScientist.stream.state import StreamState +from tests.stream_v3_fakes import ( + FakeV3Agent, + collect_events, + message_delta, + protocol_event, +) # ============================================================================= # StreamEventEmitter.interrupt() @@ -371,28 +378,12 @@ class TestHitlConfig: class TestInterruptEventParsing: - def _run_async(self, coro): - """Run async code with a fresh event loop.""" - loop = asyncio.new_event_loop() - try: - return loop.run_until_complete(coro) - finally: - loop.close() - def test_interrupt_from_updates_mode(self): """__interrupt__ in updates mode yields interrupt event.""" - from langchain_core.messages import AIMessageChunk - - from EvoScientist.stream.events import stream_agent_events - - mock_agent = MagicMock() - - ai_chunk = AIMessageChunk(content="thinking...", id="msg1") - interrupt_data = { "__interrupt__": [ - { - "value": { + Interrupt( + value={ "action_requests": [ {"name": "execute", "args": {"command": "ls"}, "id": "tc1"} ], @@ -403,30 +394,18 @@ class TestInterruptEventParsing: } ], }, - "ns": ["main"], - "resumable": True, - } + id="main", + ) ] } - chunks = [ - ((), "messages", (ai_chunk, {})), - ((), "updates", interrupt_data), - ] - - async def fake_astream(*a, **kw): - for c in chunks: - yield c - - mock_agent.astream = fake_astream - - events = [] - - async def collect(): - async for ev in stream_agent_events(mock_agent, "test", "thread-1"): - events.append(ev) - - self._run_async(collect()) + agent = FakeV3Agent( + [ + message_delta("thinking..."), + protocol_event("updates", interrupt_data), + ] + ) + events = collect_events(agent, message="test", thread_id="thread-1") types = [e["type"] for e in events] assert "interrupt" in types @@ -438,27 +417,12 @@ class TestInterruptEventParsing: def test_updates_without_interrupt_skipped(self): """Regular updates mode data is skipped as before.""" - from EvoScientist.stream.events import stream_agent_events - - mock_agent = MagicMock() - - chunks = [ - ((), "updates", {"some_node": {"key": "value"}}), - ] - - async def fake_astream(*a, **kw): - for c in chunks: - yield c - - mock_agent.astream = fake_astream - - events = [] - - async def collect(): - async for ev in stream_agent_events(mock_agent, "test", "thread-1"): - events.append(ev) - - self._run_async(collect()) + agent = FakeV3Agent( + [ + protocol_event("updates", {"some_node": {"key": "value"}}), + ] + ) + events = collect_events(agent, message="test", thread_id="thread-1") types = [e["type"] for e in events] assert "interrupt" not in types diff --git a/tests/test_stream_display.py b/tests/test_stream_display.py index 12dcb69..b162dfb 100644 --- a/tests/test_stream_display.py +++ b/tests/test_stream_display.py @@ -70,6 +70,7 @@ def test_streaming_display_keeps_narration_visible_while_processing_tool_result( ], tool_results=[ { + "id": "tc1", "name": "read_file", "content": "# User profile\n\n- Likes concise updates.", } @@ -85,6 +86,38 @@ def test_streaming_display_keeps_narration_visible_while_processing_tool_result( assert rendered.index("Here is the answer.") < rendered.index("Reading memory") +def test_streaming_display_pairs_root_tool_results_by_id(): + """Out-of-order same-name root tool results should not pair by list index.""" + renderable = create_streaming_display( + tool_calls=[ + { + "id": "tc-a", + "name": "execute", + "args": {"command": "python a.py"}, + }, + { + "id": "tc-b", + "name": "execute", + "args": {"command": "python b.py"}, + }, + ], + tool_results=[ + { + "id": "tc-b", + "name": "execute", + "content": "Error: result from b", + } + ], + ) + + rendered = _render_text(renderable) + + assert "execute(python a.py)" in rendered + assert "execute(python b.py)" in rendered + assert rendered.index("execute(python a.py)") < rendered.index("Running") + assert rendered.index("execute(python b.py)") < rendered.index("result from b") + + def test_streaming_display_keeps_narration_visible_with_pending_normal_tool(): """Ordinary tools use the same pending-tool behavior as memory reads.""" narration = "I will inspect the files first." @@ -130,6 +163,7 @@ def test_streaming_display_keeps_narration_separate_when_answer_streams(): ], tool_results=[ { + "id": "tc1", "name": "execute", "content": "check complete", } @@ -168,6 +202,7 @@ def test_streaming_display_keeps_narration_separate_in_final_frame(): ], tool_results=[ { + "id": "tc1", "name": "execute", "content": "check complete", } @@ -247,10 +282,12 @@ def test_streaming_display_interleaves_multiple_narration_segments(): ], tool_results=[ { + "id": "tc1", "name": "execute", "content": "check.py", }, { + "id": "tc2", "name": "execute", "content": "check complete", }, @@ -287,6 +324,7 @@ def test_streaming_display_preserves_narration_for_collapsed_completed_tool(): ] tool_results = [ { + "id": f"tc{i}", "name": "execute", "content": f"step {i} complete", } @@ -344,7 +382,7 @@ def test_streaming_display_preserves_narration_for_collapsed_running_tool(): def test_streaming_display_preserves_task_narration_while_subagent_runs(): """Narration before a task call should stay attached to the task section.""" narration = "I'll ask a specialist to inspect this.\n" - subagent = SubAgentState("code-agent", "inspect this") + subagent = SubAgentState("code-agent", "inspect this", "task:code", "task1") subagent.is_active = True subagent.add_tool_call("execute", {"command": "rg -n TODO ."}, "sa1") @@ -381,7 +419,7 @@ def test_streaming_display_orders_final_task_narration_by_tool_index(): first = "I'll ask a specialist to inspect this.\n" second = "Now I will run the result locally.\n" answer = "The local run passed." - subagent = SubAgentState("code-agent", "inspect this") + subagent = SubAgentState("code-agent", "inspect this", "task:code", "task1") subagent.is_active = False subagent.add_tool_call("execute", {"command": "rg -n TODO ."}, "sa1") subagent.add_tool_result("execute", "todo.py", True, "sa1") @@ -411,10 +449,12 @@ def test_streaming_display_orders_final_task_narration_by_tool_index(): ], tool_results=[ { + "id": "task1", "name": "task", "content": "todo.py", }, { + "id": "tc2", "name": "execute", "content": "passed", }, diff --git a/tests/test_stream_emitter.py b/tests/test_stream_emitter.py index 60722e9..c3baf16 100644 --- a/tests/test_stream_emitter.py +++ b/tests/test_stream_emitter.py @@ -24,48 +24,66 @@ class TestStreamEventEmitter: assert ev.data["id"] == "tc1" def test_tool_result(self): - ev = StreamEventEmitter.tool_result("execute", "[OK] done", success=True) + ev = StreamEventEmitter.tool_result( + "execute", "[OK] done", success=True, tool_call_id="tc1" + ) assert ev.type == "tool_result" assert ev.data["name"] == "execute" assert ev.data["content"] == "[OK] done" assert ev.data["success"] is True + assert ev.data["id"] == "tc1" def test_tool_result_failure(self): - ev = StreamEventEmitter.tool_result("execute", "Error: fail", success=False) + ev = StreamEventEmitter.tool_result( + "execute", "Error: fail", success=False, tool_call_id="tc1" + ) assert ev.data["success"] is False def test_subagent_start(self): - ev = StreamEventEmitter.subagent_start("research-agent", "Find papers") + ev = StreamEventEmitter.subagent_start( + "research-agent", + "Find papers", + instance_id="task:abc", + tool_call_id="call_task_1", + ) assert ev.type == "subagent_start" assert ev.data["name"] == "research-agent" assert ev.data["description"] == "Find papers" + assert ev.data["instance_id"] == "task:abc" + assert ev.data["tool_call_id"] == "call_task_1" def test_subagent_tool_call(self): ev = StreamEventEmitter.subagent_tool_call( - "research-agent", "tavily_search", {"query": "q"}, "tc2" + "research-agent", + "tavily_search", + {"query": "q"}, + "tc2", + "task:abc", ) assert ev.type == "subagent_tool_call" assert ev.data["subagent"] == "research-agent" assert ev.data["name"] == "tavily_search" + assert ev.data["instance_id"] == "task:abc" def test_subagent_tool_result(self): ev = StreamEventEmitter.subagent_tool_result( - "research-agent", "tavily_search", "results", True + "research-agent", + "tavily_search", + "results", + True, + tool_call_id="tc2", + instance_id="task:abc", ) assert ev.type == "subagent_tool_result" assert ev.data["subagent"] == "research-agent" - assert ev.data["id"] == "" # default when caller omits tool_call_id - - def test_subagent_tool_result_with_id(self): - ev = StreamEventEmitter.subagent_tool_result( - "research-agent", "execute", "ok", True, tool_call_id="tc_xyz" - ) - assert ev.data["id"] == "tc_xyz" + assert ev.data["id"] == "tc2" + assert ev.data["instance_id"] == "task:abc" def test_subagent_end(self): - ev = StreamEventEmitter.subagent_end("research-agent") + ev = StreamEventEmitter.subagent_end("research-agent", "task:abc") assert ev.type == "subagent_end" assert ev.data["name"] == "research-agent" + assert ev.data["instance_id"] == "task:abc" def test_done(self): ev = StreamEventEmitter.done("final answer") @@ -101,13 +119,15 @@ class TestStreamEventEmitter: events = [ StreamEventEmitter.thinking("x"), StreamEventEmitter.text("x"), - StreamEventEmitter.tool_call("t", {}), - StreamEventEmitter.tool_result("t", "x"), - StreamEventEmitter.subagent_start("s", "d"), - StreamEventEmitter.subagent_tool_call("s", "t", {}), - StreamEventEmitter.subagent_tool_result("s", "t", "x"), - StreamEventEmitter.subagent_text("s", "c"), - StreamEventEmitter.subagent_end("s"), + StreamEventEmitter.tool_call("t", {}, "tc1"), + StreamEventEmitter.tool_result("t", "x", True, "tc1"), + StreamEventEmitter.subagent_start("s", "d", "task:1", "call_task_1"), + StreamEventEmitter.subagent_tool_call("s", "t", {}, "tc2", "task:1"), + StreamEventEmitter.subagent_tool_result( + "s", "t", "x", True, "tc2", "task:1" + ), + StreamEventEmitter.subagent_text("s", "c", "task:1"), + StreamEventEmitter.subagent_end("s", "task:1"), StreamEventEmitter.interrupt("i", []), StreamEventEmitter.done(), StreamEventEmitter.error("e"), diff --git a/tests/test_stream_events.py b/tests/test_stream_events.py index 2d7aa17..3ce871a 100644 --- a/tests/test_stream_events.py +++ b/tests/test_stream_events.py @@ -1,19 +1,45 @@ """Tests for EvoScientist/stream/events.py helpers.""" import asyncio -from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock -from langchain_core.messages import AIMessageChunk -from langgraph.types import Command +import pytest +from deepagents import create_deep_agent +from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel +from langchain_core.messages import AIMessage, HumanMessage, ToolMessage +from langchain_core.tools import tool +from langgraph.checkpoint.memory import InMemorySaver +from langgraph.types import Command, Interrupt -from EvoScientist.stream.events import ( +from EvoScientist.middleware.ask_user import AskUserMiddleware +from EvoScientist.stream.events import stream_agent_events +from EvoScientist.stream.summarization import ( _extract_summary_message_text, - _extract_tool_content, _find_summarization_event_payload, - _process_chunk_content, - stream_agent_events, ) +from EvoScientist.stream.tool_results import ( + _extract_command_tool_content, + _extract_tool_content, +) +from tests.conftest import run_async +from tests.stream_v3_fakes import ( + ErroringV3Agent, + FakeSubagent, + FakeV3Agent, + HangingV3Agent, + SubscriptionSensitiveV3Agent, + collect_events, + message_delta, + message_finish, + message_tool_call_block, + protocol_event, + tool_finished, + tool_started, +) + + +class _ToolCallingFakeModel(FakeMessagesListChatModel): + def bind_tools(self, tools, *, tool_choice=None, **kwargs): + return self class TestExtractToolContent: @@ -21,13 +47,14 @@ class TestExtractToolContent: def test_image_via_additional_kwargs(self): """Image ToolMessages with read_file_media_type return summary.""" - msg = SimpleNamespace( + msg = ToolMessage( content=[{"type": "image", "base64": "abc123..."}], + name="read_file", + tool_call_id="tc-image", additional_kwargs={ "read_file_media_type": "image/png", "read_file_path": "/chart.png", }, - name="read_file", ) content, is_image = _extract_tool_content(msg) assert is_image is True @@ -38,13 +65,13 @@ class TestExtractToolContent: def test_image_via_list_content_blocks(self): """Image content blocks without metadata are still detected.""" - msg = SimpleNamespace( + msg = ToolMessage( content=[ {"type": "text", "text": "Image: chart.png"}, {"type": "image", "base64": "iVBORw0KGgo..."}, ], - additional_kwargs={}, name="read_file", + tool_call_id="tc-image", ) content, is_image = _extract_tool_content(msg) assert is_image is True @@ -52,10 +79,10 @@ class TestExtractToolContent: def test_normal_text_passthrough(self): """Normal text content passes through unchanged.""" - msg = SimpleNamespace( + msg = ToolMessage( content="File written successfully to /output.txt", - additional_kwargs={}, name="write_file", + tool_call_id="tc-write", ) content, is_image = _extract_tool_content(msg) assert is_image is False @@ -63,10 +90,10 @@ class TestExtractToolContent: def test_empty_content(self): """Empty content returns empty string.""" - msg = SimpleNamespace( + msg = ToolMessage( content="", - additional_kwargs={}, name="read_file", + tool_call_id="tc-empty", ) content, is_image = _extract_tool_content(msg) assert is_image is False @@ -74,164 +101,127 @@ class TestExtractToolContent: def test_list_text_blocks(self): """List of text blocks are joined.""" - msg = SimpleNamespace( + msg = ToolMessage( content=[ {"type": "text", "text": "Line 1"}, {"type": "text", "text": "Line 2"}, ], - additional_kwargs={}, name="read_file", + tool_call_id="tc-list", ) content, is_image = _extract_tool_content(msg) assert is_image is False assert "Line 1" in content assert "Line 2" in content - def test_no_additional_kwargs_attr(self): - """Messages without additional_kwargs attribute are handled.""" - msg = SimpleNamespace( - content="some result", - name="execute", - ) - content, is_image = _extract_tool_content(msg) - assert is_image is False - assert content == "some result" - - -# ============================================================================= -# _process_chunk_content — string content passthrough -# ============================================================================= - - -class TestProcessChunkContentStrings: - """Verify _process_chunk_content handles string content correctly. - - After removing strip_thinking_tags (ccproxy >=0.2.7 no longer embeds - tags), string content is emitted verbatim. These tests - serve as a regression baseline: if a future ccproxy version re-introduces - tags, the raw tags will be visible and these tests will document that. - """ - - def _emit(self, content: str) -> list: - from EvoScientist.stream.emitter import StreamEventEmitter - from EvoScientist.stream.tracker import ToolCallTracker - - emitter = StreamEventEmitter() - tracker = ToolCallTracker() - chunk = AIMessageChunk(content=content) - return list(_process_chunk_content(chunk, emitter, tracker)) - - def test_plain_text_passthrough(self): - events = self._emit("Hello world") - assert len(events) == 1 - assert events[0].type == "text" - assert events[0].data["content"] == "Hello world" - - def test_thinking_tags_stripped(self): - """Legacy tags from older ccproxy are stripped.""" - raw = "some reasoningThe answer is 42." - events = self._emit(raw) - assert len(events) == 1 - assert "" not in events[0].data["content"] - assert events[0].data["content"] == "The answer is 42." - - def test_thinking_tags_only_yields_nothing(self): - """Content that is only a thinking block yields no events.""" - events = self._emit("just reasoning") - assert events == [] - - def test_thinking_tags_preserve_surrounding_whitespace(self): - """Stripping tags does not swallow adjacent spaces.""" - raw = "before x after" - events = self._emit(raw) - assert len(events) == 1 - assert events[0].data["content"] == "before after" - - def test_empty_string_no_events(self): - events = self._emit("") - assert events == [] - - -# ============================================================================= -# 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, {})), + def test_command_tool_content_scans_multiple_messages(self): + """Command updates may contain multiple messages; match by tool_call_id.""" + output = Command( + update={ + "messages": [ + ToolMessage( + content="Ignore me", + name="read_file", + tool_call_id="other", + ), + ToolMessage( + content=[{"type": "image", "base64": "iVBORw0KGgo..."}], + name="read_file", + tool_call_id="target", + ), ] - ) + } ) - events = _collect_events(mock_agent) + + assert _extract_command_tool_content(output, "target") == "[OK] Image displayed" + + +# ============================================================================= +# v3 protocol streaming +# ============================================================================= + + +class TestV3ProtocolStreaming: + """Test stream_agent_events against v3 protocol events.""" + + def test_message_delta_emits_text(self): + """v3 content-block text deltas are processed.""" + agent = FakeV3Agent([message_delta("hello world")]) + events = collect_events(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" + _, kwargs = agent.astream_events.call_args + assert kwargs["version"] == "v3" + assert "stream_mode" not in kwargs + assert "subgraphs" not in kwargs - 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, {})), - ] - ) + def test_streamed_non_selector_json_is_replayed(self): + """Normal JSON answers are not swallowed by selector JSON buffering.""" + agent = FakeV3Agent( + [ + message_delta("{"), + message_delta('"answer"'), + message_delta(": 1}"), + ] ) - events = _collect_events(mock_agent) + events = collect_events(agent) + text_events = [e for e in events if e.get("type") == "text"] + assert "".join(e["content"] for e in text_events) == '{"answer": 1}' + assert events[-1]["type"] == "done" + assert events[-1]["response"] == '{"answer": 1}' + + def test_incomplete_non_selector_json_flushes_on_message_finish(self): + """Buffered non-selector text is not lost if the message ends mid-object.""" + agent = FakeV3Agent( + [ + message_delta("{"), + message_delta('"answer":'), + message_finish(), + ] + ) + events = collect_events(agent) + text_events = [e for e in events if e.get("type") == "text"] + assert "".join(e["content"] for e in text_events) == '{"answer":' + assert events[-1]["response"] == '{"answer":' + + def test_json_answer_with_tools_key_is_replayed_without_selector_context(self): + """Normal answers may legitimately contain a top-level tools key.""" + agent = FakeV3Agent( + [ + message_delta('{"tools":["hammer"],"answer":"use safely"}'), + ] + ) + events = collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert len(text_events) == 1 - assert text_events[0]["content"] == "fallback" + assert text_events[0]["content"] == '{"tools":["hammer"],"answer":"use safely"}' + assert events[-1]["response"] == '{"tools":["hammer"],"answer":"use safely"}' - 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, {})), - ] - ) + def test_text_delta_strips_legacy_thinking_tags(self): + """Legacy tags are still removed on the v3 text path.""" + agent = FakeV3Agent( + [message_delta("some reasoningThe answer is 42.")] ) - events = _collect_events(mock_agent) + events = collect_events(agent) + text_events = [e for e in events if e.get("type") == "text"] + assert len(text_events) == 1 + assert text_events[0]["content"] == "The answer is 42." + + def test_text_delta_with_only_legacy_thinking_tags_is_skipped(self): + agent = FakeV3Agent([message_delta("just reasoning")]) + events = collect_events(agent) + assert [e for e in events if e.get("type") == "text"] == [] + + def test_updates_event_without_summary_is_skipped(self): + """Non-summary updates are skipped without error.""" + agent = FakeV3Agent( + [ + protocol_event("updates", {"some": "state"}), + message_delta("should appear"), + ] + ) + events = collect_events(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" @@ -242,10 +232,9 @@ class TestMultiModeChunkUnpacking: "EvoScientist.stream.events.clear_memory_worker_saved_counts", lambda: calls.append(True), ) - mock_agent = AsyncMock() - mock_agent.astream = MagicMock(return_value=_async_iter([])) + agent = FakeV3Agent([]) - _collect_events(mock_agent, message="new user turn") + collect_events(agent, message="new user turn") assert calls == [True] @@ -255,29 +244,23 @@ class TestMultiModeChunkUnpacking: "EvoScientist.stream.events.clear_memory_worker_saved_counts", lambda: calls.append(True), ) - mock_agent = AsyncMock() - mock_agent.astream = MagicMock(return_value=_async_iter([])) + agent = FakeV3Agent([]) resume_command = Command(resume={"decisions": [{"type": "approve"}]}) - _collect_events(mock_agent, message=resume_command) + collect_events(agent, message=resume_command) assert calls == [True] - assert mock_agent.astream.call_args.args[0] is resume_command + assert agent.astream_events.call_args.args[0] is resume_command 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, {})), - ] - ) + """v3 messages with lc_source=summarization emit summarization events.""" + agent = FakeV3Agent( + [ + message_delta("synthetic summary", {"lc_source": "summarization"}), + message_delta("real content"), + ] ) - events = _collect_events(mock_agent) + events = collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -291,32 +274,27 @@ class TestMultiModeChunkUnpacking: def test_updates_mode_summarization_event_emitted(self): """_summarization_event updates should emit a summarization event.""" - summary_message = SimpleNamespace( + summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts", ) - chunk_real = _make_ai_chunk("real content") - mock_agent = AsyncMock() - mock_agent.astream = MagicMock( - return_value=_async_iter( - [ - ( - (), - "updates", - { - "agent": { - "_summarization_event": { - "summary_message": summary_message, - "cutoff_index": 12, - "file_path": None, - } + agent = FakeV3Agent( + [ + protocol_event( + "updates", + { + "agent": { + "_summarization_event": { + "summary_message": summary_message, + "cutoff_index": 12, + "file_path": None, } - }, - ), - ((), "messages", (chunk_real, {})), - ] - ) + } + }, + ), + message_delta("real content"), + ] ) - events = _collect_events(mock_agent) + events = collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -327,32 +305,26 @@ class TestMultiModeChunkUnpacking: def test_updates_mode_does_not_duplicate_streamed_summarization(self): """If streamed summarization already emitted, updates fallback should not duplicate it.""" - chunk_synth = _make_ai_chunk("synthetic summary") - summary_message = SimpleNamespace( + summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts" ) - chunk_real = _make_ai_chunk("real content") - mock_agent = AsyncMock() - mock_agent.astream = MagicMock( - return_value=_async_iter( - [ - ((), "messages", (chunk_synth, {"lc_source": "summarization"})), - ( - (), - "updates", - { - "_summarization_event": { - "summary_message": summary_message, - "cutoff_index": 12, - "file_path": None, - } - }, - ), - ((), "messages", (chunk_real, {})), - ] - ) + agent = FakeV3Agent( + [ + message_delta("synthetic summary", {"lc_source": "summarization"}), + protocol_event( + "updates", + { + "_summarization_event": { + "summary_message": summary_message, + "cutoff_index": 12, + "file_path": None, + } + }, + ), + message_delta("real content"), + ] ) - events = _collect_events(mock_agent) + events = collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -363,41 +335,24 @@ class TestMultiModeChunkUnpacking: def test_updates_mode_does_not_reemit_existing_summarization_event(self): """Persisted _summarization_event from a prior turn should not be replayed.""" - summary_message = SimpleNamespace( + summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts", ) - chunk_real = _make_ai_chunk("real content") - mock_agent = AsyncMock() - mock_agent.aget_state = AsyncMock( - return_value=SimpleNamespace( - values={ - "_summarization_event": { - "summary_message": summary_message, - "cutoff_index": 12, - "file_path": None, - } - } - ) + summary_event = { + "_summarization_event": { + "summary_message": summary_message, + "cutoff_index": 12, + "file_path": None, + } + } + agent = FakeV3Agent( + [ + protocol_event("updates", summary_event), + message_delta("real content"), + ], + state_values=summary_event, ) - mock_agent.astream = MagicMock( - return_value=_async_iter( - [ - ( - (), - "updates", - { - "_summarization_event": { - "summary_message": summary_message, - "cutoff_index": 12, - "file_path": None, - } - }, - ), - ((), "messages", (chunk_real, {})), - ] - ) - ) - events = _collect_events(mock_agent) + events = collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -405,46 +360,667 @@ class TestMultiModeChunkUnpacking: summary_events = [e for e in events if e.get("type") == "summarization"] assert summary_events == [] - -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, - }, + def test_whole_message_reasoning_is_not_duplicated(self): + """Providers can expose the same reasoning in kwargs and content blocks.""" + message = AIMessage( + additional_kwargs={"reasoning_content": "Think once."}, + content=[{"type": "reasoning", "reasoning": "Think once."}], ) - mock_agent = AsyncMock() - mock_agent.astream = MagicMock( - return_value=_async_iter( + agent = FakeV3Agent([protocol_event("messages", (message, {}))]) + events = collect_events(agent) + thinking_events = [e for e in events if e.get("type") == "thinking"] + assert len(thinking_events) == 1 + assert thinking_events[0]["content"] == "Think once." + + def test_tool_events_emit_call_and_result(self): + """v3 tool projection events become UI tool call/result events.""" + output = ToolMessage( + name="read_file", + content="File content", + tool_call_id="tc1", + ) + agent = FakeV3Agent( + [ + tool_started("read_file", {"path": "notes.txt"}), + tool_finished(output), + ] + ) + events = collect_events(agent) + tool_call = next(e for e in events if e.get("type") == "tool_call") + tool_result = next(e for e in events if e.get("type") == "tool_result") + assert tool_call["name"] == "read_file" + assert tool_call["args"] == {"path": "notes.txt"} + assert tool_call["id"] == "tc1" + assert tool_result["name"] == "read_file" + assert tool_result["content"] == "File content" + assert tool_result["success"] is True + assert tool_result["id"] == "tc1" + + @pytest.mark.filterwarnings( + "ignore:The v3 streaming protocol on Pregel is experimental" + ) + def test_live_deepagents_v3_tool_result_preserves_tool_call_id(self): + """DeepAgents v3 emits tool_call_id on started and finished tool events.""" + + @tool + def probe(value: str) -> str: + """Return a deterministic probe result.""" + return f"probe:{value}" + + model = _ToolCallingFakeModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + { + "name": "probe", + "args": {"value": "ok"}, + "id": "call_probe_1", + } + ], + ), + AIMessage(content="final answer"), + ] + ) + agent = create_deep_agent( + model=model, + tools=[probe], + system_prompt="Use tools when requested.", + ) + + async def _collect_events(): + return [ + event + async for event in stream_agent_events( + agent, "run probe", "live-deepagents-tool-id" + ) + ] + + events = run_async(_collect_events()) + + tool_call = next(e for e in events if e.get("type") == "tool_call") + tool_result = next(e for e in events if e.get("type") == "tool_result") + done = next(e for e in events if e.get("type") == "done") + assert tool_call == { + "type": "tool_call", + "name": "probe", + "args": {"value": "ok"}, + "id": "call_probe_1", + } + assert tool_result == { + "type": "tool_result", + "name": "probe", + "content": "probe:ok", + "success": True, + "id": "call_probe_1", + } + assert done["content"] == "final answer" + + @pytest.mark.filterwarnings( + "ignore:The v3 streaming protocol on Pregel is experimental" + ) + def test_live_deepagents_v3_hitl_emits_tool_call_and_single_interrupt(self): + """Live HITL streams the model tool call once before one interrupt.""" + + @tool + def echo_tool(value: str) -> str: + """Echo a deterministic value.""" + return f"echo:{value}" + + model = _ToolCallingFakeModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + { + "name": "echo_tool", + "args": {"value": "ok"}, + "id": "call_echo_1", + } + ], + ) + ] + ) + agent = create_deep_agent( + model=model, + tools=[echo_tool], + system_prompt="Use tools when requested.", + interrupt_on={"echo_tool": True}, + checkpointer=InMemorySaver(), + ) + + async def _collect_events(): + return [ + event + async for event in stream_agent_events( + agent, "run echo", "live-deepagents-hitl" + ) + ] + + events = run_async(_collect_events()) + + tool_calls = [e for e in events if e.get("type") == "tool_call"] + interrupts = [e for e in events if e.get("type") == "interrupt"] + assert tool_calls == [ + { + "type": "tool_call", + "name": "echo_tool", + "args": {"value": "ok"}, + "id": "call_echo_1", + } + ] + assert len(interrupts) == 1 + assert events.index(tool_calls[0]) < events.index(interrupts[0]) + assert interrupts[0]["action_requests"][0]["name"] == "echo_tool" + assert interrupts[0]["action_requests"][0]["args"] == {"value": "ok"} + + @pytest.mark.filterwarnings( + "ignore:The v3 streaming protocol on Pregel is experimental" + ) + def test_live_deepagents_v3_ask_user_suppresses_interrupt_tool_result(self): + """ask_user pause markers are not displayed as failed tool results.""" + + model = _ToolCallingFakeModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + { + "name": "ask_user", + "args": { + "questions": [ + { + "question": "What dataset?", + "type": "text", + } + ] + }, + "id": "call_ask_1", + } + ], + ), + AIMessage(content="final after ask"), + ] + ) + agent = create_deep_agent( + model=model, + tools=[], + system_prompt="Use ask_user when requested.", + middleware=[AskUserMiddleware()], + checkpointer=InMemorySaver(), + ) + + async def _collect(message): + return [ + event + async for event in stream_agent_events( + agent, message, "live-deepagents-ask-user" + ) + ] + + first_events = run_async(_collect("ask")) + first_types = [event.get("type") for event in first_events] + assert first_types == ["tool_call", "ask_user", "done"] + ask_event = next(e for e in first_events if e.get("type") == "ask_user") + assert ask_event["tool_call_id"] == "call_ask_1" + assert ask_event["questions"] == [{"question": "What dataset?", "type": "text"}] + + resumed_events = run_async( + _collect(Command(resume={"answers": ["CIFAR-10"], "status": "answered"})) + ) + tool_result = next(e for e in resumed_events if e.get("type") == "tool_result") + assert tool_result == { + "type": "tool_result", + "name": "ask_user", + "content": "Q: What dataset?\nA: CIFAR-10", + "success": True, + "id": "call_ask_1", + } + done = next(e for e in resumed_events if e.get("type") == "done") + assert done["content"] == "final after ask" + + @pytest.mark.filterwarnings( + "ignore:The v3 streaming protocol on Pregel is experimental" + ) + def test_live_deepagents_v3_task_result_uses_subagent_tool_message_content(self): + """Live task results should display the subagent ToolMessage content.""" + + root_model = _ToolCallingFakeModel( + responses=[ + AIMessage( + content="", + tool_calls=[ + { + "name": "task", + "args": { + "subagent_type": "researcher", + "description": "find answer", + }, + "id": "call_task_1", + } + ], + ), + AIMessage(content="root final"), + ] + ) + subagent_model = _ToolCallingFakeModel( + responses=[AIMessage(content="subagent final")] + ) + agent = create_deep_agent( + model=root_model, + tools=[], + system_prompt="Delegate when requested.", + subagents=[ + { + "name": "researcher", + "description": "Finds answers", + "system_prompt": "Answer directly.", + "model": subagent_model, + "tools": [], + } + ], + ) + + async def _collect_events(): + return [ + event + async for event in stream_agent_events( + agent, "delegate", "live-deepagents-subagent" + ) + ] + + events = run_async(_collect_events()) + + subagent_start = next(e for e in events if e.get("type") == "subagent_start") + subagent_end = next(e for e in events if e.get("type") == "subagent_end") + task_result = next( + e + for e in events + if e.get("type") == "tool_result" and e.get("name") == "task" + ) + assert subagent_start["name"] == "researcher" + assert subagent_start["description"] == "" + assert subagent_start["instance_id"] + assert subagent_start["tool_call_id"] == "call_task_1" + assert subagent_end["instance_id"] == subagent_start["instance_id"] + assert task_result["id"] == "call_task_1" + assert task_result["content"] == "subagent final" + assert "Command(" not in task_result["content"] + + def test_message_tool_call_block_emits_pre_execution_tool_call(self): + """Model-declared tool calls remain visible before execution starts.""" + agent = FakeV3Agent( + [ + message_tool_call_block( + "execute", + {"command": "ls"}, + tool_call_id="tc-msg", + ), + protocol_event( + "updates", + { + "__interrupt__": [ + Interrupt( + value={ + "action_requests": [ + { + "name": "execute", + "args": {"command": "ls"}, + "id": "tc-msg", + } + ], + "review_configs": [], + }, + id="main", + ) + ] + }, + ), + ] + ) + events = collect_events(agent) + event_types = [e["type"] for e in events] + assert event_types.index("tool_call") < event_types.index("interrupt") + tool_call = next(e for e in events if e.get("type") == "tool_call") + assert tool_call["id"] == "tc-msg" + assert tool_call["args"] == {"command": "ls"} + + def test_tool_selection_flushes_before_tool_only_step(self): + """Selector UI event is emitted even when selection is followed only by a tool.""" + import EvoScientist.middleware.tool_selector as selector_mod + + original_selected = selector_mod._current_selected_tools + original_total = selector_mod._total_tools_count + original_last = selector_mod._last_emitted_tools + selector_mod._current_selected_tools = ["read_file"] + selector_mod._total_tools_count = 3 + selector_mod._last_emitted_tools = [] + try: + output = ToolMessage( + content="File content", + name="read_file", + tool_call_id="tc1", + ) + agent = FakeV3Agent( [ - ((), "messages", (chunk, {})), + message_delta('{"tools":["read_file"]}'), + tool_started("read_file", {"path": "notes.txt"}), + tool_finished(output), ] ) + events = collect_events(agent) + finally: + selector_mod._current_selected_tools = original_selected + selector_mod._total_tools_count = original_total + selector_mod._last_emitted_tools = original_last + + event_types = [e["type"] for e in events] + assert event_types.index("tool_selection") < event_types.index("tool_call") + selection = next(e for e in events if e.get("type") == "tool_selection") + assert selection["tools"] == ["read_file"] + + def test_subagent_projection_routes_namespaced_events(self): + """DeepAgents subagent projection supplies identity for namespaced events.""" + namespace = ("task", "abc") + output = ToolMessage( + content="Found result", + name="search", + tool_call_id="sa-tc", ) - events = _collect_events(mock_agent) + agent = FakeV3Agent( + [ + message_delta("Sub-agent finding.", namespace=namespace), + tool_started( + "search", + {"query": "papers"}, + tool_call_id="sa-tc", + namespace=namespace, + ), + tool_finished(output, tool_call_id="sa-tc", namespace=namespace), + ], + subagents=[FakeSubagent(namespace, "research-agent")], + ) + events = collect_events(agent) + assert any(e.get("type") == "subagent_start" for e in events) + assert any(e.get("type") == "subagent_end" for e in events) + + text = next(e for e in events if e.get("type") == "subagent_text") + tool_call = next(e for e in events if e.get("type") == "subagent_tool_call") + tool_result = next(e for e in events if e.get("type") == "subagent_tool_result") + + assert text["subagent"] == "research-agent" + assert text["content"] == "Sub-agent finding." + assert text["instance_id"] == "task:abc" + start = next(e for e in events if e.get("type") == "subagent_start") + assert start["tool_call_id"] == "call_task_abc" + assert tool_call["instance_id"] == "task:abc" + assert tool_call["subagent"] == "research-agent" + assert tool_call["name"] == "search" + assert tool_call["args"] == {"query": "papers"} + assert tool_result["instance_id"] == "task:abc" + assert tool_result["subagent"] == "research-agent" + assert tool_result["content"] == "Found result" + event_types = [e["type"] for e in events] + assert event_types.index("subagent_start") < event_types.index("subagent_text") + assert event_types.index("subagent_tool_result") < event_types.index( + "subagent_end" + ) + assert event_types.index("subagent_end") < event_types.index("done") + + def test_namespaced_events_wait_for_delayed_subagent_registration(self): + """Subagent events are not dropped if protocol events arrive first.""" + namespace = ("task", "late") + + class DelayedSubagentRun: + def __init__(self): + self.subagents = self._subagent_iter() + + async def _subagent_iter(self): + await asyncio.sleep(0) + yield FakeSubagent(namespace, "research-agent") + + def __aiter__(self): + return self._events() + + async def _events(self): + yield message_delta("Sub-agent finding.", namespace=namespace) + + async def abort(self): + pass + + class Agent: + def __init__(self): + self._run = DelayedSubagentRun() + + def astream_events(self, *_args, **_kwargs): + return self._run + + async def aget_state(self, _config): + class Snapshot: + def __init__(self): + self.values = {} + + return Snapshot() + + events = collect_events(Agent()) + event_types = [e["type"] for e in events] + text = next(e for e in events if e.get("type") == "subagent_text") + + assert text["content"] == "Sub-agent finding." + assert text["instance_id"] == "task:late" + assert event_types.index("subagent_start") < event_types.index("subagent_text") + + def test_subagent_tool_dedupe_uses_resolved_path(self): + """Tool call/result events can arrive on namespace suffixes for one subagent.""" + subagent_path = ("task", "abc") + call_namespace = (*subagent_path, "agent") + tool_namespace = (*subagent_path, "tools") + output = ToolMessage( + content="Found result", + name="search", + tool_call_id="sa-tc", + ) + agent = FakeV3Agent( + [ + message_tool_call_block( + "search", + {"query": "papers"}, + tool_call_id="sa-tc", + namespace=call_namespace, + ), + tool_started( + "search", + {"query": "papers"}, + tool_call_id="sa-tc", + namespace=tool_namespace, + ), + tool_finished(output, tool_call_id="sa-tc", namespace=tool_namespace), + ], + subagents=[FakeSubagent(subagent_path, "research-agent")], + ) + events = collect_events(agent) + + calls = [e for e in events if e.get("type") == "subagent_tool_call"] + results = [e for e in events if e.get("type") == "subagent_tool_result"] + + assert len(calls) == 1 + assert calls[0]["instance_id"] == "task:abc" + assert calls[0]["id"] == "sa-tc" + assert len(results) == 1 + assert results[0]["instance_id"] == "task:abc" + assert results[0]["id"] == "sa-tc" + + def test_subagent_end_is_emitted_before_later_root_text(self): + """Finished subagents stop showing as active while root streaming continues.""" + output_returned = asyncio.Event() + + class CompletingSubagent: + path = ("task", "done-first") + name = "research-agent" + + @property + def cause(self) -> dict[str, str]: + return {"type": "toolCall", "tool_call_id": "call_done_first"} + + async def output(self): + output_returned.set() + return {"messages": []} + + class RootTextAfterSubagentDoneRun: + def __init__(self): + self.subagents = self._subagent_iter() + + async def _subagent_iter(self): + yield CompletingSubagent() + + def __aiter__(self): + return self._events() + + async def _events(self): + await output_returned.wait() + await asyncio.sleep(0) + yield message_delta("root answer") + + async def abort(self): + pass + + class Agent: + def __init__(self): + self._run = RootTextAfterSubagentDoneRun() + + def astream_events(self, *_args, **_kwargs): + return self._run + + async def aget_state(self, _config): + class Snapshot: + def __init__(self): + self.values = {} + + return Snapshot() + + events = collect_events(Agent()) + event_types = [e["type"] for e in events] + + assert event_types.index("subagent_end") < event_types.index("text") + + def test_subagent_projection_is_subscribed_before_protocol_pump(self): + """Subagent handles are not dropped by lazy projection subscription.""" + namespace = ("task", "early") + agent = SubscriptionSensitiveV3Agent( + [message_delta("Sub-agent finding.", namespace=namespace)], + [FakeSubagent(namespace, "research-agent")], + ) + events = collect_events(agent) + assert any(e.get("type") == "subagent_start" for e in events) + assert any(e.get("type") == "subagent_end" for e in events) + assert [e for e in events if e.get("type") == "text"] == [] + + text = next(e for e in events if e.get("type") == "subagent_text") + assert text["subagent"] == "research-agent" + assert text["content"] == "Sub-agent finding." + assert text["instance_id"] == "task:early" + + def test_parallel_same_name_subagent_events_carry_instance_ids(self): + """Lifecycle and tool events distinguish same-name parallel subagents.""" + ns1 = ("task", "one") + ns2 = ("task", "two") + output1 = ToolMessage( + content="Found one", + name="search", + tool_call_id="tc1", + ) + output2 = ToolMessage( + content="Found two", + name="search", + tool_call_id="tc2", + ) + agent = FakeV3Agent( + [ + tool_started( + "search", {"query": "a"}, tool_call_id="tc1", namespace=ns1 + ), + tool_started( + "search", {"query": "b"}, tool_call_id="tc2", namespace=ns2 + ), + tool_finished(output2, tool_call_id="tc2", namespace=ns2), + tool_finished(output1, tool_call_id="tc1", namespace=ns1), + ], + subagents=[ + FakeSubagent(ns1, "research-agent"), + FakeSubagent(ns2, "research-agent"), + ], + ) + events = collect_events(agent) + + starts = [e for e in events if e.get("type") == "subagent_start"] + calls = [e for e in events if e.get("type") == "subagent_tool_call"] + results = [e for e in events if e.get("type") == "subagent_tool_result"] + ends = [e for e in events if e.get("type") == "subagent_end"] + + assert {e["instance_id"] for e in starts} == {"task:one", "task:two"} + assert {e["instance_id"] for e in calls} == {"task:one", "task:two"} + assert {e["instance_id"] for e in results} == {"task:one", "task:two"} + assert {e["instance_id"] for e in ends} == {"task:one", "task:two"} + + def test_stream_construction_error_emits_error_before_reraising(self): + """astream_events construction failures preserve the UI error event contract.""" + events = [] + + async def collect(): + async for ev in stream_agent_events( + ErroringV3Agent(RuntimeError("boom")), + "hi", + "t1", + ): + events.append(ev) + + with pytest.raises(RuntimeError, match="boom"): + run_async(collect()) + assert events == [{"type": "error", "message": "boom"}] + + def test_generator_close_aborts_underlying_v3_stream(self): + """Early consumer exit should abort the caller-driven v3 run.""" + + async def consume_one_and_close(): + agent = HangingV3Agent([message_delta("hi")]) + stream = stream_agent_events(agent, "hi", "t1") + first = await stream.__anext__() + await stream.aclose() + return first, agent.aborted + + first, aborted = run_async(consume_one_and_close()) + assert first["type"] == "text" + assert first["content"] == "hi" + assert aborted is True + + +class TestUsageStatsExtraction: + """Test token usage extraction from v3 message-finish events.""" + + def test_usage_metadata_emitted(self): + """v3 message-finish usage emits usage_stats event.""" + agent = FakeV3Agent( + [ + message_delta("hi"), + message_finish( + { + "input_tokens": 100, + "output_tokens": 50, + "total_tokens": 150, + } + ), + ] + ) + events = collect_events(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) + """message-finish without usage does not emit usage_stats.""" + agent = FakeV3Agent([message_delta("hi"), message_finish()]) + events = collect_events(agent) usage_events = [e for e in events if e.get("type") == "usage_stats"] assert len(usage_events) == 0 @@ -453,13 +1029,13 @@ class TestSummarizationHelpers: """Summarization extraction helpers.""" def test_extract_summary_message_text_from_summary_tag(self): - message = SimpleNamespace( + message = HumanMessage( content="Before\n\nImportant facts\n\nAfter", ) assert _extract_summary_message_text(message) == "Important facts" def test_extract_summary_message_text_accepts_output_text_blocks(self): - message = SimpleNamespace( + message = HumanMessage( content=[{"type": "output_text", "text": "Summary body"}], ) assert _extract_summary_message_text(message) == "Summary body" @@ -469,29 +1045,27 @@ class TestSummarizationHelpers: "node": { "response": { "_summarization_event": { - "summary_message": SimpleNamespace(content="Summary body"), + "summary_message": HumanMessage(content="Summary body"), } } } } event = _find_summarization_event_payload(payload) assert event is not None - assert event["summary_message"].content == "Summary body" + summary_message = event["summary_message"] + assert isinstance(summary_message, HumanMessage) + assert summary_message.content == "Summary body" 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}, + agent = FakeV3Agent( + [ + message_delta("hi"), + message_finish( + {"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) + events = collect_events(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 90f05c1..3e74d93 100644 --- a/tests/test_stream_state.py +++ b/tests/test_stream_state.py @@ -26,37 +26,16 @@ class TestSubAgentState: assert len(sa.tool_calls) == 1 assert sa.tool_calls[0]["args"]["query"] == "updated" - def test_add_tool_call_merge_name(self): - """When first call has empty name, second should fill it in.""" + def test_add_tool_call_merge_args_by_id(self): sa = SubAgentState("agent") - sa.add_tool_call("", {}, "tc1") - # Empty name + empty id → skipped entirely - # But with an id, it can be tracked: - # Actually, empty name with id is also skipped per the code (not name check) - # Let's use a named call first, then merge args - sa2 = SubAgentState("agent") - sa2.add_tool_call("search", {}, "tc1") - sa2.add_tool_call("search", {"query": "test"}, "tc1") - assert sa2.tool_calls[0]["args"] == {"query": "test"} - - def test_skip_empty_name_no_id(self): - sa = SubAgentState("agent") - sa.add_tool_call("", {}, "") - assert len(sa.tool_calls) == 0 + sa.add_tool_call("search", {}, "tc1") + sa.add_tool_call("search", {"query": "test"}, "tc1") + assert sa.tool_calls[0]["args"] == {"query": "test"} def test_add_tool_result_matched(self): sa = SubAgentState("agent") sa.add_tool_call("execute", {}, "tc1") - sa.add_tool_result("execute", "output", True) - result = sa.get_result_for(sa.tool_calls[0]) - assert result is not None - assert result["content"] == "output" - - def test_add_tool_result_fallback(self): - """When name doesn't match, falls back to first unmatched.""" - sa = SubAgentState("agent") - sa.add_tool_call("execute", {}, "tc1") - sa.add_tool_result("different_name", "output", True) + sa.add_tool_result("execute", "output", True, tool_call_id="tc1") result = sa.get_result_for(sa.tool_calls[0]) assert result is not None assert result["content"] == "output" @@ -80,15 +59,6 @@ class TestSubAgentState: tc = {"id": "tc_missing", "name": "x", "args": {}} assert sa.get_result_for(tc) is None - def test_get_result_for_index_fallback(self): - """When no id, falls back to index-based matching.""" - sa = SubAgentState("agent") - tc = {"id": "", "name": "execute", "args": {}} - sa.tool_calls.append(tc) - sa.tool_results.append({"name": "execute", "content": "ok", "success": True}) - result = sa.get_result_for(tc) - assert result is not None - # ============================================================================= # StreamState @@ -152,6 +122,7 @@ class TestStreamState: "type": "tool_result", "name": "execute", "content": "[OK] done", + "id": "tc1", } ) assert len(state.tool_results) == 1 @@ -164,6 +135,8 @@ class TestStreamState: "type": "subagent_start", "name": "research-agent", "description": "Search", + "instance_id": "task:research", + "tool_call_id": "tc_task_research", } ) assert len(state.subagents) == 1 @@ -172,10 +145,20 @@ class TestStreamState: def test_handle_subagent_tool_call(self): state = StreamState() + state.handle_event( + { + "type": "subagent_start", + "name": "research-agent", + "description": "Search", + "instance_id": "task:research", + "tool_call_id": "tc_task_research", + } + ) state.handle_event( { "type": "subagent_tool_call", "subagent": "research-agent", + "instance_id": "task:research", "name": "tavily_search", "args": {"query": "test"}, "id": "tc_sa1", @@ -186,10 +169,20 @@ class TestStreamState: def test_handle_subagent_tool_result(self): state = StreamState() + state.handle_event( + { + "type": "subagent_start", + "name": "code-agent", + "description": "Run code", + "instance_id": "task:code", + "tool_call_id": "tc_task_code", + } + ) state.handle_event( { "type": "subagent_tool_call", "subagent": "code-agent", + "instance_id": "task:code", "name": "execute", "args": {}, "id": "tc1", @@ -199,6 +192,7 @@ class TestStreamState: { "type": "subagent_tool_result", "subagent": "code-agent", + "instance_id": "task:code", "name": "execute", "content": "output", "success": True, @@ -212,11 +206,65 @@ class TestStreamState: def test_handle_subagent_end(self): state = StreamState() state.handle_event( - {"type": "subagent_start", "name": "agent-x", "description": ""} + { + "type": "subagent_start", + "name": "agent-x", + "description": "", + "instance_id": "task:x", + "tool_call_id": "tc_task_x", + } + ) + state.handle_event( + {"type": "subagent_end", "name": "agent-x", "instance_id": "task:x"} ) - state.handle_event({"type": "subagent_end", "name": "agent-x"}) assert state.subagents[0].is_active is False + def test_same_name_subagents_are_separated_by_instance_id(self): + state = StreamState() + state.handle_event( + { + "type": "subagent_start", + "name": "research-agent", + "description": "Find A", + "instance_id": "task:one", + "tool_call_id": "tc_task_1", + } + ) + state.handle_event( + { + "type": "subagent_start", + "name": "research-agent", + "description": "Find B", + "instance_id": "task:two", + "tool_call_id": "tc_task_2", + } + ) + state.handle_event( + { + "type": "subagent_tool_call", + "subagent": "research-agent", + "instance_id": "task:two", + "name": "search", + "args": {"query": "b"}, + "id": "tc2", + } + ) + state.handle_event( + { + "type": "subagent_end", + "name": "research-agent", + "instance_id": "task:one", + } + ) + + assert len(state.subagents) == 2 + first, second = state.subagents + assert first.name == "research-agent" + assert second.name == "research-agent" + assert first.is_active is False + assert second.is_active is True + assert second.tool_calls[0]["args"] == {"query": "b"} + def test_handle_done(self): state = StreamState() state.handle_event({"type": "done", "response": "Final answer"}) @@ -247,83 +295,6 @@ class TestStreamState: assert state.subagents[0].is_active is False -# ============================================================================= -# Name merging -# ============================================================================= - - -class TestNameMerging: - def test_generic_subagent_merged(self): - state = StreamState() - # First event creates generic "sub-agent" - state.handle_event( - { - "type": "subagent_tool_call", - "subagent": "sub-agent", - "name": "execute", - "args": {}, - "id": "tc1", - } - ) - assert len(state.subagents) == 1 - assert state.subagents[0].name == "sub-agent" - - # Proper name arrives, should merge - sa = state._get_or_create_subagent("code-agent", "write code") - assert len(state.subagents) == 1 - assert sa.name == "code-agent" - assert len(sa.tool_calls) == 1 # preserved from generic entry - - def test_no_merge_when_not_generic(self): - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "research-agent", "description": ""} - ) - state._get_or_create_subagent("code-agent", "code-agent") - assert len(state.subagents) == 2 - - -class TestSubagentNameResolution: - def test_subagent_tool_call_resolves_single_active(self): - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} - ) - state.handle_event( - { - "type": "subagent_tool_call", - "subagent": "sub-agent", - "name": "execute", - "args": {}, - "id": "tc1", - } - ) - assert len(state.subagents) == 1 - assert state.subagents[0].name == "code-agent" - assert len(state.subagents[0].tool_calls) == 1 - - def test_subagent_tool_call_does_not_resolve_when_multiple_active(self): - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} - ) - state.handle_event( - {"type": "subagent_start", "name": "research-agent", "description": ""} - ) - state.handle_event( - { - "type": "subagent_tool_call", - "subagent": "sub-agent", - "name": "execute", - "args": {}, - "id": "tc1", - } - ) - # Should stay as "sub-agent" because more than one named subagent is active - assert len(state.subagents) == 3 - assert state.subagents[-1].name == "sub-agent" - - # ============================================================================= # _parse_todo_items # ============================================================================= @@ -388,79 +359,6 @@ class TestBuildTodoStats: assert "0 items" in result -# ============================================================================= -# _resolve_subagent_name -# ============================================================================= - - -class TestResolveSubagentName: - def test_real_name_passthrough(self): - state = StreamState() - assert state._resolve_subagent_name("code-agent") == "code-agent" - - def test_resolve_single_active(self): - """When exactly one named active sub-agent exists, 'sub-agent' resolves to it.""" - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} - ) - assert state._resolve_subagent_name("sub-agent") == "code-agent" - - def test_no_resolve_multiple_active(self): - """With multiple active named sub-agents, 'sub-agent' stays generic.""" - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} - ) - state.handle_event( - {"type": "subagent_start", "name": "research-agent", "description": ""} - ) - assert state._resolve_subagent_name("sub-agent") == "sub-agent" - - def test_no_resolve_no_active(self): - state = StreamState() - assert state._resolve_subagent_name("sub-agent") == "sub-agent" - - -# ============================================================================= -# subagent_end fallback -# ============================================================================= - - -class TestSubagentEndFallback: - def test_end_resolves_via_name(self): - """subagent_end with exact name deactivates sub-agent.""" - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} - ) - state.handle_event({"type": "subagent_end", "name": "code-agent"}) - assert state.subagents[0].is_active is False - - def test_end_resolves_generic_to_single_active(self): - """subagent_end with 'sub-agent' resolves to the only active named sub-agent.""" - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "research-agent", "description": ""} - ) - state.handle_event({"type": "subagent_end", "name": "sub-agent"}) - assert state.subagents[0].is_active is False - - def test_end_fallback_deactivates_oldest(self): - """When name can't be resolved, deactivates oldest active sub-agent.""" - state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "a-agent", "description": ""} - ) - state.handle_event( - {"type": "subagent_start", "name": "b-agent", "description": ""} - ) - # "sub-agent" can't resolve (2 active), falls back to oldest - state.handle_event({"type": "subagent_end", "name": "sub-agent"}) - assert state.subagents[0].is_active is False # a-agent deactivated - assert state.subagents[1].is_active is True # b-agent still active - - # ============================================================================= # Todo capture from write_todos args # ============================================================================= @@ -493,6 +391,7 @@ class TestTodoCaptureFromArgs: state.handle_event( { "type": "tool_result", + "id": "tc1", "name": "write_todos", "content": json.dumps(items), } @@ -520,6 +419,7 @@ class TestTodoCaptureFromArgs: state.handle_event( { "type": "tool_result", + "id": "tc1", "name": "write_todos", "content": json.dumps(result_todos), } @@ -536,6 +436,7 @@ class TestTodoCaptureFromArgs: state.handle_event( { "type": "tool_result", + "id": "tc1", "name": "read_todos", "content": json.dumps(items), } @@ -611,47 +512,32 @@ class TestLatestTextReset: ] -# ============================================================================= -# Name merging edge cases -# ============================================================================= - - -class TestNameMergingAdvanced: - def test_subagent_merges_into_preregistered(self): - """'sub-agent' event merges into pre-registered real-name entry with no tools.""" +class TestSubagentInstanceIds: + def test_multiple_subagents_no_cross_merge(self): state = StreamState() - state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": "code"} - ) - assert len(state.subagents) == 1 - # Now sub-agent tool call arrives — should merge into code-agent state.handle_event( { - "type": "subagent_tool_call", - "subagent": "sub-agent", - "name": "execute", - "args": {}, - "id": "tc1", + "type": "subagent_start", + "name": "code-agent", + "description": "", + "instance_id": "task:code", + "tool_call_id": "tc_task_code", } ) - # _resolve_subagent_name should resolve "sub-agent" → "code-agent" (single active) - assert len(state.subagents) == 1 - assert state.subagents[0].name == "code-agent" - assert len(state.subagents[0].tool_calls) == 1 - - def test_multiple_subagents_no_cross_merge(self): - """Two named sub-agents should not merge with each other.""" - state = StreamState() state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} - ) - state.handle_event( - {"type": "subagent_start", "name": "research-agent", "description": ""} + { + "type": "subagent_start", + "name": "research-agent", + "description": "", + "instance_id": "task:research", + "tool_call_id": "tc_task_research", + } ) state.handle_event( { "type": "subagent_tool_call", "subagent": "code-agent", + "instance_id": "task:code", "name": "execute", "args": {}, "id": "tc1", @@ -661,14 +547,15 @@ class TestNameMergingAdvanced: { "type": "subagent_tool_call", "subagent": "research-agent", + "instance_id": "task:research", "name": "tavily_search", "args": {}, "id": "tc2", } ) assert len(state.subagents) == 2 - code_sa = state._subagent_map["code-agent"] - research_sa = state._subagent_map["research-agent"] + code_sa = state._subagent_map["task:code"] + research_sa = state._subagent_map["task:research"] assert len(code_sa.tool_calls) == 1 assert code_sa.tool_calls[0]["name"] == "execute" assert len(research_sa.tool_calls) == 1 @@ -788,7 +675,9 @@ class TestComputePhase: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) assert state.compute_phase() == "researching" def test_writing_after_tools_done_and_text_starts(self): @@ -796,7 +685,9 @@ class TestComputePhase: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) state.handle_event({"type": "text", "content": "Final report"}) assert state.compute_phase() == "writing" @@ -808,16 +699,30 @@ class TestComputePhase: def test_researching_with_active_subagent(self): state = StreamState() state.handle_event( - {"type": "subagent_start", "name": "research", "description": ""} + { + "type": "subagent_start", + "name": "research", + "description": "", + "instance_id": "task:research", + "tool_call_id": "tc_task_research", + } ) assert state.compute_phase() == "researching" def test_writing_after_subagent_ends(self): state = StreamState() state.handle_event( - {"type": "subagent_start", "name": "research", "description": ""} + { + "type": "subagent_start", + "name": "research", + "description": "", + "instance_id": "task:research", + "tool_call_id": "tc_task_research", + } + ) + state.handle_event( + {"type": "subagent_end", "name": "research", "instance_id": "task:research"} ) - state.handle_event({"type": "subagent_end", "name": "research"}) state.handle_event({"type": "text", "content": "Report"}) assert state.compute_phase() == "writing" @@ -827,7 +732,9 @@ class TestComputePhase: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) state.is_processing = False assert state.compute_phase() == "researching" @@ -846,7 +753,13 @@ class TestComputePhase: {"type": "tool_call", "id": "tc1", "name": "task", "args": {}} ) state.handle_event( - {"type": "subagent_start", "name": "code-agent", "description": ""} + { + "type": "subagent_start", + "name": "code-agent", + "description": "", + "instance_id": "task:code", + "tool_call_id": "tc1", + } ) assert state.compute_phase() == "researching" @@ -868,7 +781,13 @@ class TestHasPendingWork: def test_active_subagent(self): state = StreamState() state.handle_event( - {"type": "subagent_start", "name": "agent", "description": ""} + { + "type": "subagent_start", + "name": "agent", + "description": "", + "instance_id": "task:agent", + "tool_call_id": "tc_task_agent", + } ) assert state.has_pending_work() is True @@ -877,7 +796,9 @@ class TestHasPendingWork: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) assert state.is_processing is True assert state.has_pending_work() is True @@ -886,7 +807,9 @@ class TestHasPendingWork: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) state.handle_event({"type": "text", "content": "done"}) assert state.has_pending_work() is False @@ -910,7 +833,9 @@ class TestVisibleToolCounts: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) assert state.visible_tool_counts() == (1, 1) def test_mixed(self): @@ -921,7 +846,9 @@ class TestVisibleToolCounts: state.handle_event( {"type": "tool_call", "id": "tc3", "name": "search", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) # execute done, search pending assert state.visible_tool_counts() == (1, 2) @@ -971,7 +898,9 @@ class TestIsFinalResponseDelegation: state.handle_event( {"type": "tool_call", "id": "tc1", "name": "execute", "args": {}} ) - state.handle_event({"type": "tool_result", "name": "execute", "content": "ok"}) + state.handle_event( + {"type": "tool_result", "id": "tc1", "name": "execute", "content": "ok"} + ) state.handle_event({"type": "text", "content": "done"}) assert _is_final_response(state) == (not state.has_pending_work()) diff --git a/tests/test_stream_tracker.py b/tests/test_stream_tracker.py deleted file mode 100644 index 408d103..0000000 --- a/tests/test_stream_tracker.py +++ /dev/null @@ -1,99 +0,0 @@ -"""Tests for EvoScientist/stream/tracker.py.""" - -from EvoScientist.stream.tracker import ToolCallTracker - - -class TestToolCallTracker: - def test_update_and_get(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute", args={"command": "ls"}) - info = tracker.get("tc1") - assert info is not None - assert info.name == "execute" - assert info.args == {"command": "ls"} - - def test_update_merges(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.update("tc1", args={"command": "ls"}) - info = tracker.get("tc1") - assert info.name == "execute" - assert info.args == {"command": "ls"} - - def test_get_missing(self): - tracker = ToolCallTracker() - assert tracker.get("nonexistent") is None - - def test_is_ready_true(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - assert tracker.is_ready("tc1") is True - - def test_is_ready_no_name(self): - tracker = ToolCallTracker() - tracker.update("tc1") - assert tracker.is_ready("tc1") is False - - def test_is_ready_already_emitted(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.mark_emitted("tc1") - assert tracker.is_ready("tc1") is False - - def test_is_ready_missing_id(self): - tracker = ToolCallTracker() - assert tracker.is_ready("tc_missing") is False - - def test_append_json_delta_and_finalize(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.append_json_delta('{"comma') - tracker.append_json_delta('nd": "ls"}') - tracker.finalize_all() - info = tracker.get("tc1") - assert info.args == {"command": "ls"} - assert info.args_complete is True - - def test_finalize_invalid_json(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.append_json_delta("{invalid json") - tracker.finalize_all() - info = tracker.get("tc1") - # Args should remain empty since JSON is invalid - assert info.args == {} - assert info.args_complete is True - - def test_emit_all_pending(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.update("tc2", name="read_file") - pending = tracker.emit_all_pending() - assert len(pending) == 2 - # All should now be marked emitted - assert tracker.get("tc1").emitted is True - assert tracker.get("tc2").emitted is True - # Second call should return empty - assert tracker.emit_all_pending() == [] - - def test_get_pending(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.update("tc2", name="read_file") - tracker.mark_emitted("tc1") - pending = tracker.get_pending() - assert len(pending) == 1 - assert pending[0].id == "tc2" - - def test_get_all(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.update("tc2", name="read_file") - assert len(tracker.get_all()) == 2 - - def test_clear(self): - tracker = ToolCallTracker() - tracker.update("tc1", name="execute") - tracker.clear() - assert tracker.get("tc1") is None - assert tracker.get_all() == [] diff --git a/tests/test_stream_utils.py b/tests/test_stream_utils.py index 12c365d..2d3e01b 100644 --- a/tests/test_stream_utils.py +++ b/tests/test_stream_utils.py @@ -233,11 +233,11 @@ class TestFormatToolCompact: def test_task_with_desc_only(self): result = format_tool_compact("task", {"description": "do stuff"}) - assert "Cooking with sub-agent" in result + assert result == "task(description=do stuff)" def test_task_no_info(self): result = format_tool_compact("task", {"other": "value"}) - assert result == "Cooking with sub-agent" + assert result == "task(other=value)" def test_tavily_search(self): result = format_tool_compact("tavily_search", {"query": "python testing"}) diff --git a/tests/test_subagent_summarize.py b/tests/test_subagent_summarize.py index bd8891c..474ac8a 100644 --- a/tests/test_subagent_summarize.py +++ b/tests/test_subagent_summarize.py @@ -13,16 +13,19 @@ import asyncio from dataclasses import dataclass from unittest.mock import AsyncMock, MagicMock, patch -from langchain_core.messages import AIMessageChunk - from EvoScientist.channels.base import Channel from EvoScientist.channels.bus.events import InboundMessage as BusInbound from EvoScientist.channels.bus.message_bus import MessageBus from EvoScientist.channels.channel_manager import ChannelManager from EvoScientist.channels.consumer import InboundConsumer, _join_subagent_text from EvoScientist.stream.emitter import StreamEvent, StreamEventEmitter -from EvoScientist.stream.events import stream_agent_events from tests.conftest import run_async as _run +from tests.stream_v3_fakes import ( + FakeSubagent, + FakeV3Agent, + collect_events, + message_delta, +) # ═══════════════════════════════════════════════════════════════════ # Helpers @@ -39,29 +42,6 @@ class _FakeConfig: dm_policy: str = "allowlist" -def _make_ai_chunk(content: str = "", **kwargs): - return AIMessageChunk(content=content, **kwargs) - - -async def _async_iter(items): - for item in items: - yield item - - -def _collect_events(agent, message="hi", thread_id="t1"): - async def _run_inner(): - 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_inner()) - finally: - loop.close() - - # ═══════════════════════════════════════════════════════════════════ # 1. StreamEventEmitter.subagent_text # ═══════════════════════════════════════════════════════════════════ @@ -69,29 +49,29 @@ def _collect_events(agent, message="hi", thread_id="t1"): class TestSubagentTextEmitter: def test_creates_correct_event_type(self): - ev = StreamEventEmitter.subagent_text("research-agent", "Found 3 papers.") + ev = StreamEventEmitter.subagent_text( + "research-agent", "Found 3 papers.", "task:research" + ) assert isinstance(ev, StreamEvent) assert ev.type == "subagent_text" def test_data_contains_subagent_and_content(self): - ev = StreamEventEmitter.subagent_text("analyst", "Result summary") + ev = StreamEventEmitter.subagent_text( + "analyst", "Result summary", "task:analyst" + ) assert ev.data["subagent"] == "analyst" assert ev.data["content"] == "Result summary" def test_data_contains_type_key(self): """Event data dict should include 'type' matching event type (project convention).""" - ev = StreamEventEmitter.subagent_text("a", "b") + ev = StreamEventEmitter.subagent_text("a", "b", "task:a") assert ev.data["type"] == "subagent_text" def test_empty_content(self): - ev = StreamEventEmitter.subagent_text("agent", "") + ev = StreamEventEmitter.subagent_text("agent", "", "task:agent") assert ev.data["content"] == "" - def test_instance_id_defaults_to_empty(self): - ev = StreamEventEmitter.subagent_text("agent", "content") - assert ev.data["instance_id"] == "" - - def test_instance_id_included_when_provided(self): + def test_instance_id_included(self): ev = StreamEventEmitter.subagent_text( "agent", "content", instance_id="tracker:abc" ) @@ -99,7 +79,7 @@ class TestSubagentTextEmitter: def test_included_in_all_events_type_check(self): """subagent_text must pass the same invariant as other emitters.""" - ev = StreamEventEmitter.subagent_text("s", "c") + ev = StreamEventEmitter.subagent_text("s", "c", "task:s") assert "type" in ev.data assert ev.data["type"] == ev.type @@ -114,16 +94,16 @@ class TestStreamAgentEventsSubagentText: def test_subagent_text_emitted_for_subagent_chunks(self): """When a sub-agent produces text, subagent_text events should appear.""" - subagent_chunk = _make_ai_chunk("Sub-agent finding: X is significant.") - mock_agent = AsyncMock() - mock_agent.astream = MagicMock( - return_value=_async_iter( - [ - (("sub:research",), "messages", (subagent_chunk, {})), - ] - ) + namespace = ("sub", "research") + agent = FakeV3Agent( + [ + message_delta( + "Sub-agent finding: X is significant.", namespace=namespace + ) + ], + subagents=[FakeSubagent(namespace, "research-agent")], ) - events = _collect_events(mock_agent) + events = collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] assert len(sa_text) == 1 assert "Sub-agent finding" in sa_text[0]["content"] @@ -132,16 +112,8 @@ class TestStreamAgentEventsSubagentText: def test_subagent_text_not_emitted_for_main_agent(self): """Main agent text should produce 'text' events, not 'subagent_text'.""" - chunk = _make_ai_chunk("Main agent reply.") - mock_agent = AsyncMock() - mock_agent.astream = MagicMock( - return_value=_async_iter( - [ - ((), "messages", (chunk, {})), - ] - ) - ) - events = _collect_events(mock_agent) + agent = FakeV3Agent([message_delta("Main agent reply.")]) + events = collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] text_events = [e for e in events if e.get("type") == "text"] assert len(sa_text) == 0 @@ -149,14 +121,16 @@ class TestStreamAgentEventsSubagentText: def test_multiple_subagent_text_chunks_all_emitted(self): """Multiple text chunks from a sub-agent all yield subagent_text events.""" - chunks = [ - (("sub:a",), "messages", (_make_ai_chunk("Part 1."), {})), - (("sub:a",), "messages", (_make_ai_chunk("Part 2."), {})), - (("sub:a",), "messages", (_make_ai_chunk("Part 3."), {})), - ] - mock_agent = AsyncMock() - mock_agent.astream = MagicMock(return_value=_async_iter(chunks)) - events = _collect_events(mock_agent) + namespace = ("sub", "a") + agent = FakeV3Agent( + [ + message_delta("Part 1.", namespace=namespace), + message_delta("Part 2.", namespace=namespace), + message_delta("Part 3.", namespace=namespace), + ], + subagents=[FakeSubagent(namespace, "research-agent")], + ) + events = collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] assert len(sa_text) == 3 combined = "".join(e["content"] for e in sa_text) @@ -176,33 +150,21 @@ class TestStreamAgentEventsSubagentText: the two instances apart even though their 'subagent' field is identical. """ - # Two different namespaces with task: IDs, same lc_agent_name - chunks = [ - ( - ("ns:task:id1:agent",), - "messages", - ( - _make_ai_chunk("Instance-1 text."), - {"lc_agent_name": "research-agent"}, - ), - ), - ( - ("ns:task:id2:agent",), - "messages", - ( - _make_ai_chunk("Instance-2 text."), - {"lc_agent_name": "research-agent"}, - ), - ), - ( - ("ns:task:id1:agent",), - "messages", - (_make_ai_chunk(" More from 1."), {"lc_agent_name": "research-agent"}), - ), - ] - mock_agent = AsyncMock() - mock_agent.astream = MagicMock(return_value=_async_iter(chunks)) - events = _collect_events(mock_agent) + # Two different v3 namespaces, same projected subagent display name. + ns1 = ("ns", "task", "id1", "agent") + ns2 = ("ns", "task", "id2", "agent") + agent = FakeV3Agent( + [ + message_delta("Instance-1 text.", namespace=ns1), + message_delta("Instance-2 text.", namespace=ns2), + message_delta(" More from 1.", namespace=ns1), + ], + subagents=[ + FakeSubagent(ns1, "research-agent"), + FakeSubagent(ns2, "research-agent"), + ], + ) + events = collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] assert len(sa_text) == 3 @@ -287,11 +249,13 @@ class TestConsumerSubagentTextFallback: { "type": "subagent_text", "subagent": "research", + "instance_id": "task:research", "content": "Found 3 relevant papers.", }, { "type": "subagent_text", "subagent": "research", + "instance_id": "task:research", "content": " Key insight: X is Y.", }, {"type": "done", "content": ""}, @@ -330,6 +294,7 @@ class TestConsumerSubagentTextFallback: { "type": "subagent_text", "subagent": "research", + "instance_id": "task:research", "content": "Sub-agent detail.", }, {"type": "text", "content": "Here is my summary."}, @@ -543,6 +508,7 @@ class TestConsumerSubagentTextFallback: { "type": "subagent_text", "subagent": "research", + "instance_id": "task:research", "content": "Sub-agent work.", }, {"type": "done", "content": "Final summary from done event."}, @@ -656,16 +622,19 @@ class TestConsumerParallelSubagentFallback: { "type": "subagent_text", "subagent": "research", + "instance_id": "task:research", "content": "Found papers.", }, { "type": "subagent_text", "subagent": "analysis", + "instance_id": "task:analysis", "content": "Metric is high.", }, { "type": "subagent_text", "subagent": "research", + "instance_id": "task:research", "content": " Key insight.", }, {"type": "done", "content": ""}, @@ -699,7 +668,12 @@ class TestConsumerParallelSubagentFallback: def test_single_agent_no_attribution_prefix(self): """Single sub-agent fallback has no [name]: prefix.""" events = [ - {"type": "subagent_text", "subagent": "research", "content": "Only agent."}, + { + "type": "subagent_text", + "subagent": "research", + "instance_id": "task:research", + "content": "Only agent.", + }, {"type": "done", "content": ""}, ] consumer, bus, fake_stream = _make_consumer(events) @@ -728,38 +702,6 @@ class TestConsumerParallelSubagentFallback: _run(_test()) - def test_missing_subagent_field_uses_unknown(self): - """Events without 'subagent' field are grouped under 'unknown'.""" - events = [ - {"type": "subagent_text", "content": "No agent name."}, - {"type": "done", "content": ""}, - ] - consumer, bus, fake_stream = _make_consumer(events) - - async def _test(): - with patch( - "EvoScientist.stream.events.stream_agent_events", - new=fake_stream, - ): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) - - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - - # Single agent (unknown), no prefix - assert outbound.content == "No agent name." - - await consumer.stop() - await task - - _run(_test()) - class TestConsumerSameNameInterleaved: """Two instances of the same agent type with interleaved chunks.""" @@ -831,39 +773,6 @@ class TestConsumerSameNameInterleaved: _run(_test()) - def test_same_name_no_instance_id_still_concatenated(self): - """Without instance_id (legacy events), same-name chunks still merge into one buffer.""" - events = [ - {"type": "subagent_text", "subagent": "research-agent", "content": "A."}, - {"type": "subagent_text", "subagent": "research-agent", "content": " B."}, - {"type": "done", "content": ""}, - ] - consumer, bus, fake_stream = _make_consumer(events) - - async def _test(): - with patch( - "EvoScientist.stream.events.stream_agent_events", - new=fake_stream, - ): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) - - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - - # Single instance (no instance_id), no prefix - assert outbound.content == "A. B." - - await consumer.stop() - await task - - _run(_test()) - class TestDelegationPromptSummarize: def test_framework_task_tool_contains_summarize_guidance(self): diff --git a/tests/test_summarization.py b/tests/test_summarization.py index ddc5367..9dedfdb 100644 --- a/tests/test_summarization.py +++ b/tests/test_summarization.py @@ -1,8 +1,10 @@ """Tests for the summarization event pipeline and display widgets.""" +from langchain_core.messages import HumanMessage + from EvoScientist.stream.emitter import StreamEventEmitter -from EvoScientist.stream.events import _extract_summarization_text from EvoScientist.stream.state import StreamState +from EvoScientist.stream.summarization import _extract_summarization_text # --------------------------------------------------------------------------- # Emitter @@ -277,56 +279,52 @@ class TestExtractSummarizationText: """Content extraction from summarization chunks.""" def test_string_content(self): - class Msg: - content = "hello world" - - assert _extract_summarization_text(Msg()) == "hello world" + assert ( + _extract_summarization_text(HumanMessage(content="hello world")) + == "hello world" + ) def test_content_blocks(self): - class Msg: - def __init__(self): - self.content = [ - {"type": "text", "text": "part1"}, - {"type": "text", "text": "part2"}, - ] - - assert _extract_summarization_text(Msg()) == "part1part2" + assert ( + _extract_summarization_text( + HumanMessage( + content=[ + {"type": "text", "text": "part1"}, + {"type": "text", "text": "part2"}, + ] + ) + ) + == "part1part2" + ) def test_content_blocks_with_index(self): """Content blocks may include 'index' field — should still extract text.""" - class Msg: - def __init__(self): - self.content = [{"type": "text", "text": " vs", "index": 1}] - - assert _extract_summarization_text(Msg()) == " vs" + assert ( + _extract_summarization_text( + HumanMessage(content=[{"type": "text", "text": " vs", "index": 1}]) + ) + == " vs" + ) def test_empty_list(self): - class Msg: - def __init__(self): - self.content = [] - - assert _extract_summarization_text(Msg()) == "" - - def test_no_content_attr(self): - class Msg: - pass - - assert _extract_summarization_text(Msg()) == "" + assert _extract_summarization_text(HumanMessage(content=[])) == "" def test_mixed_block_types(self): - class Msg: - def __init__(self): - self.content = [ - {"type": "text", "text": "hello"}, - {"type": "image", "url": "..."}, - ] - - assert _extract_summarization_text(Msg()) == "hello" + assert ( + _extract_summarization_text( + HumanMessage( + content=[ + {"type": "text", "text": "hello"}, + {"type": "image", "url": "..."}, + ] + ) + ) + == "hello" + ) def test_string_blocks_in_list(self): - class Msg: - def __init__(self): - self.content = ["hello", "world"] - - assert _extract_summarization_text(Msg()) == "helloworld" + assert ( + _extract_summarization_text(HumanMessage(content=["hello", "world"])) + == "helloworld" + ) diff --git a/tests/test_tui_widgets.py b/tests/test_tui_widgets.py index 5c208b8..32138c8 100644 --- a/tests/test_tui_widgets.py +++ b/tests/test_tui_widgets.py @@ -58,9 +58,9 @@ class TestLoadingWidget(unittest.TestCase): w._timer_handle = timer w.remove = AsyncMock() - import asyncio + from tests.conftest import run_async - asyncio.run(w.cleanup()) + run_async(w.cleanup()) assert timer.stopped is True assert w._timer_handle is None