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>
This commit is contained in:
dinos
2026-06-08 18:14:25 +02:00
committed by GitHub
parent ac052bbb4c
commit 8bb1d6c0e3
28 changed files with 2801 additions and 2175 deletions
+3 -4
View File
@@ -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"
):
...
"""
+7 -13
View File
@@ -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(
+37 -62
View File
@@ -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()
+3 -15
View File
@@ -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:
-3
View File
@@ -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",
+51 -63
View File
@@ -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", "")):
+30 -12
View File
@@ -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:
File diff suppressed because it is too large Load Diff
+70 -121
View File
@@ -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
+94
View File
@@ -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"<summary>\s*(.*?)\s*</summary>", 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,
)
+72
View File
@@ -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
+117
View File
@@ -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 []
-115
View File
@@ -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()
-6
View File
@@ -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"):
+84
View File
@@ -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"<thinking>.*?</thinking>", 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 ``<thinking>...</thinking>`` 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)
+17 -2
View File
@@ -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."},
]
+269
View File
@@ -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": []}
+31
View File
@@ -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."""
+25 -61
View File
@@ -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
+42 -2
View File
@@ -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",
},
+40 -20
View File
@@ -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"),
File diff suppressed because it is too large Load Diff
+161 -232
View File
@@ -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())
-99
View File
@@ -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() == []
+2 -2
View File
@@ -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"})
+65 -156
View File
@@ -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):
+40 -42
View File
@@ -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"
)
+2 -2
View File
@@ -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