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:
@@ -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"
|
||||
):
|
||||
...
|
||||
"""
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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,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",
|
||||
|
||||
@@ -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", "")):
|
||||
|
||||
@@ -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:
|
||||
|
||||
+672
-850
File diff suppressed because it is too large
Load Diff
+70
-121
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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 []
|
||||
@@ -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()
|
||||
@@ -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"):
|
||||
|
||||
@@ -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
@@ -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."},
|
||||
]
|
||||
|
||||
|
||||
@@ -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": []}
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
|
||||
@@ -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"),
|
||||
|
||||
+867
-293
File diff suppressed because it is too large
Load Diff
+161
-232
@@ -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())
|
||||
|
||||
|
||||
@@ -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() == []
|
||||
@@ -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"})
|
||||
|
||||
@@ -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
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user