feat: Implement async sub-agent auto-notification system (#214)

* feat: Implement async sub-agent auto-notification system

- Added async notifier functionality to handle notifications for sub-agents reaching terminal states.
- Introduced `AsyncTaskNotification` dataclass for structured notification data.
- Implemented `watch_run_and_notify` to monitor agent runs and enqueue notifications.
- Created `spawn_watcher` to manage watcher tasks and ensure proper cancellation of previous watchers.
- Developed `consume_notifications` to process notifications, deduplicate them, and format messages for LLM.
- Added tests for notification handling, including draining, deduplication, and formatting.
- Patched deepagents to integrate the new watcher functionality into start and update tools.

* Enhance async notifier with per-thread notification routing and error handling

- Introduced `origin_cli_thread_id` to `AsyncTaskNotification` for routing notifications back to the originating CLI session.
- Implemented per-thread notification queues to handle notifications based on the originating thread.
- Updated `has_pending_notifications` and `drain_notifications` to respect thread-specific queues.
- Enhanced `watch_run_and_notify` to detect in-band error events from the SSE stream and handle clean exits.
- Modified tests to verify the new notification routing behavior and ensure proper handling of notifications across threads.
- Added a fixture to restore the async watcher patch state in tests to prevent state leakage.
- Updated deepagents patching to capture the main agent's CLI thread ID for notification routing.

* feat: Enhance async notifier with thread-specific watcher management and notification filtering

* test: Enhance notification draining logic for cleaner test setup

* refactor: Remove summary field from AsyncTaskNotification and update related tests

* feat: Enhance async notification handling with target thread ID support

* Refactor async notifier and middleware for improved task management

- Removed the no-op shutdown watcher loop from async_notifier.py as it is no longer needed.
- Updated watch_run_and_notify to clarify notification handling and race conditions.
- Cleaned up shutdown handling in commands.py, interactive.py, and tui_interactive.py by removing obsolete shutdown watcher calls.
- Deleted the deepagents async watcher patch from patches.py, transitioning to a new middleware approach.
- Introduced AsyncWatcherMiddleware to handle async task notifications directly during tool calls.
- Updated tests to validate the new middleware functionality and ensure proper watcher spawning and cancellation.
- Enhanced test coverage for async watcher middleware, including edge cases and error handling.

* feat(tests): add fixture to reset notifier state before each test
This commit is contained in:
Xi Zhang
2026-05-07 17:03:02 +02:00
committed by GitHub
parent 692dc491ac
commit 80f1f4fa0f
12 changed files with 2395 additions and 36 deletions
+21 -10
View File
@@ -221,7 +221,7 @@ def _build_prompt_refs() -> dict:
}
def _maybe_swap_async_subagents(subs: list) -> list:
def _maybe_swap_async_subagents(subs: list, middleware: list | None = None) -> list:
"""Replace ``_async``-flagged sub-agents with ``AsyncSubAgent`` specs when enabled.
Reads the ``_async`` field carried through by ``utils.load_subagents._build_one``
@@ -238,6 +238,10 @@ def _maybe_swap_async_subagents(subs: list) -> list:
All return paths strip the internal ``_async`` field from sub-agent dicts
before handoff, since deepagents may schema-validate the kwarg.
When async subagents are actually swapped in and ``middleware`` is provided,
appends ``AsyncWatcherMiddleware`` so launches spawn an
``async_notifier`` watcher.
"""
cfg = _ensure_config()
if not getattr(cfg, "enable_async_subagents", False):
@@ -277,6 +281,7 @@ def _maybe_swap_async_subagents(subs: list) -> list:
port = int(getattr(cfg, "langgraph_dev_port", 6174))
out = []
agent_specs: dict[str, AsyncSubAgent] = {}
# MCP tools routed to async sub-agents (via ``expose_to: <name>`` in
# mcp.yaml) ARE delivered — the deployed factory
# ``subagents/_factory.py:build_async_subagent_graph`` loads its own MCP
@@ -285,18 +290,24 @@ def _maybe_swap_async_subagents(subs: list) -> list:
for s in subs:
name = s.get("name")
if name in async_specs:
out.append(
AsyncSubAgent(
name=name,
description=async_specs[name],
graph_id=name,
url=f"http://localhost:{port}",
)
spec = AsyncSubAgent(
name=name,
description=async_specs[name],
graph_id=name,
url=f"http://localhost:{port}",
)
agent_specs[name] = spec
out.append(spec)
else:
# Strip the internal flag before handoff to deepagents.
s.pop("_async", None)
out.append(s)
if agent_specs and middleware is not None:
from .middleware.async_watcher import AsyncWatcherMiddleware
middleware.append(AsyncWatcherMiddleware(agent_specs))
return out
@@ -316,7 +327,7 @@ def _build_base_kwargs(base_backend, base_middleware):
prompt_refs=_build_prompt_refs(),
)
_inject_subagent_middleware(subs)
subs = _maybe_swap_async_subagents(subs)
subs = _maybe_swap_async_subagents(subs, base_middleware)
return {
"name": "EvoScientist",
"model": _ensure_chat_model(),
@@ -374,7 +385,7 @@ def load_mcp_and_build_kwargs(base_backend, base_middleware, *, on_mcp_progress=
# Swap selected sub-agents to AsyncSubAgent (must happen AFTER MCP injection
# since async sub-agents are remote graphs that load their own tools).
subs = _maybe_swap_async_subagents(subs)
subs = _maybe_swap_async_subagents(subs, base_middleware)
return {
"name": "EvoScientist",
+463
View File
@@ -0,0 +1,463 @@
"""Async sub-agent auto-notification.
When a sub-agent on langgraph dev reaches a terminal state, a watcher coroutine
pushes a lightweight notification onto a thread-safe queue. The CLI loop drains
the queue, dedups against deepagents' async_tasks state, batches survivors,
and injects a synthetic user message that triggers one LLM turn.
"""
from __future__ import annotations
import asyncio
import json
import logging
import queue
import threading
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Final
TERMINAL_STATUSES: Final = frozenset({"success", "error", "timeout", "interrupted"})
"""Aligned with langgraph_sdk.schema.RunStatus terminal values.
Cancel operations transition runs into ``interrupted`` (not ``cancelled``).
"""
@dataclass(frozen=True)
class AsyncTaskNotification:
"""A completed-async-task signal pushed by a watcher."""
task_id: str
agent_name: str
status: str # one of TERMINAL_STATUSES
received_at: str # ISO-8601 UTC timestamp
prompt: str = "" # original task description sent to the sub-agent
# The CLI/main-agent thread_id under which the watcher was spawned. Used
# to route the notification back to the originating CLI session so a
# /new between launch and completion does not inject the synthetic
# message into an unrelated thread (where ``check_async_task`` cannot
# find the task_id). ``None`` means "unrouted" — the notification
# drains for any current_thread_id (back-compat for direct callers).
origin_cli_thread_id: str | None = None
# Per-thread routing: notifications with ``origin_cli_thread_id`` land in
# the matching sub-queue. Notifications without one go to ``_unrouted_queue``
# and drain regardless of current thread (back-compat for legacy callers
# and direct-put test paths).
_notifications_by_thread: dict[str, queue.Queue[AsyncTaskNotification]] = {}
_notifications_lock = threading.Lock()
_unrouted_queue: queue.Queue[AsyncTaskNotification] = queue.Queue()
# Public alias for the unrouted bucket — preserved so legacy tests and any
# external direct callers that did ``_notification_queue.put(...)`` keep
# working unchanged. New code should call ``_enqueue`` instead.
_notification_queue = _unrouted_queue
# Track active watcher tasks/futures for clean shutdown.
# dict[handle, origin_cli_thread_id] so the consumer's batching grace loop
# can filter for watchers tied to the current CLI thread (or unrouted)
# without being delayed by sibling-thread watchers.
_active_watchers: dict = {}
# Map thread_id (sub-agent thread) → current watcher handle (supports
# replacement on update_async_task). Value type widens from asyncio.Task
# to "anything with .cancel()/.done()/.add_done_callback()" so we can
# move watcher scheduling onto a background loop in a follow-up fix.
_watcher_by_thread: dict[str, object] = {}
def _has_relevant_active_watchers(current_thread_id: str | None) -> bool:
"""Are there any in-flight watchers whose notifications would drain on
a ``consume_notifications`` call for ``current_thread_id``?
A watcher is relevant if its ``origin_cli_thread_id`` matches the
current CLI thread or is ``None`` (unrouted bucket drains for any
consumer). Sibling-thread watchers are ignored.
"""
if current_thread_id is None:
return bool(_active_watchers)
return any(
origin == current_thread_id or origin is None
for origin in _active_watchers.values()
)
logger = logging.getLogger(__name__)
def _enqueue(notification: AsyncTaskNotification) -> None:
"""Route a notification to its origin-thread queue or the unrouted bucket."""
tid = notification.origin_cli_thread_id
if not tid:
_unrouted_queue.put(notification)
return
with _notifications_lock:
q = _notifications_by_thread.get(tid)
if q is None:
q = queue.Queue()
_notifications_by_thread[tid] = q
q.put(notification)
def has_pending_notifications(current_thread_id: str | None = None) -> bool:
"""Cheap predicate for poller idle paths — true iff there's anything to consume.
If ``current_thread_id`` is given, only the matching thread queue and
the unrouted bucket count. With no argument, only the unrouted bucket
counts (legacy behavior).
"""
if not _unrouted_queue.empty():
return True
if current_thread_id is None:
return False
with _notifications_lock:
q = _notifications_by_thread.get(current_thread_id)
return q is not None and not q.empty()
def pending_thread_ids() -> set[str]:
"""Return the set of thread_ids with pending routed notifications."""
with _notifications_lock:
return {tid for tid, q in _notifications_by_thread.items() if not q.empty()}
async def watch_run_and_notify(
client,
thread_id: str,
run_id: str,
agent_name: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> None:
"""Subscribe to a run's event stream; enqueue notification when it terminates.
Status detection strategy:
1. Watch for an explicit ``event="error"`` SSE part — langgraph dev
emits one when the run fails. This is authoritative, in-band, and
has no timing race against the server-side run-state writeback.
2. On clean stream exit with no error event → ``"success"``.
3. On stream exception → fall back to ``client.runs.get`` (best-effort).
Non-terminal fallback statuses (``pending`` / ``running``) are
dropped — the run is still alive, no notification is enqueued.
Reading the in-band ``event="error"`` SSE part instead of polling
``runs.get`` after every clean close avoids a race where the server-side
terminal status hasn't been written by the time the stream closes.
"""
stream_failed = False
saw_error_event = False
try:
async for chunk in client.runs.join_stream(
thread_id=thread_id, run_id=run_id, stream_mode="values"
):
ev = getattr(chunk, "event", None)
data = getattr(chunk, "data", None)
if ev == "error":
# Authoritative in-band error signal from langgraph dev.
saw_error_event = True
logger.info("Watcher saw error event for task %s: %r", thread_id, data)
except Exception:
stream_failed = True
logger.warning("Watcher stream failed for task %s", thread_id, exc_info=True)
if saw_error_event:
status = "error"
elif not stream_failed:
# Clean stream exit, no error event → trust the server-side success.
status = "success"
else:
# Stream errored without delivering a final state — fall back to
# runs.get (best-effort). REJECT non-terminal statuses (pending /
# running / unknown): the run is still alive, we shouldn't notify
# at all. Returning early without enqueueing prevents the
# "⚠ pending" notification we observed when the stream failed
# mid-flight.
try:
run = await client.runs.get(thread_id=thread_id, run_id=run_id)
raw_status = run.get("status", "")
if raw_status in TERMINAL_STATUSES:
status = raw_status
else:
logger.info(
"Watcher fallback got non-terminal status %r for task %s; "
"skipping notification (run still alive)",
raw_status,
thread_id,
)
return
except Exception:
status = "error"
notification = AsyncTaskNotification(
task_id=thread_id,
agent_name=agent_name,
status=status,
received_at=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
prompt=prompt,
origin_cli_thread_id=origin_cli_thread_id,
)
_enqueue(notification)
logger.info(
"Enqueued async notification: task=%s agent=%s status=%s origin_thread=%s",
thread_id,
agent_name,
status,
origin_cli_thread_id or "<unrouted>",
)
def spawn_watcher(
client,
thread_id: str,
run_id: str,
agent_name: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> asyncio.Task:
"""Spawn a watcher on the caller's asyncio loop.
Replacement semantics support ``update_async_task`` which creates a new
run_id on the same thread_id — we want the new watcher to take over
without the old (now obsolete) watcher firing a stale notification.
Cancellation propagates ``CancelledError`` (a BaseException), which the
watcher's ``except Exception:`` does NOT catch — so ``_enqueue(...)``
never executes for the cancelled watcher (no stale notification).
``origin_cli_thread_id`` tags the resulting notification so the consumer
only injects it back into the originating CLI session.
Caller must already be in a running asyncio event loop. Serve mode's
ephemeral per-turn loop kills watchers spawned during a turn — that
limitation is tracked separately.
"""
old_task = _watcher_by_thread.get(thread_id)
if old_task is not None and not old_task.done():
old_task.cancel()
task = asyncio.create_task(
watch_run_and_notify(
client,
thread_id,
run_id,
agent_name,
prompt,
origin_cli_thread_id=origin_cli_thread_id,
)
)
_watcher_by_thread[thread_id] = task
_active_watchers[task] = origin_cli_thread_id
def _cleanup(t: asyncio.Task) -> None:
_active_watchers.pop(t, None)
# Only remove if THIS task is still the registered one — could
# have been replaced by a newer spawn_watcher call already.
if _watcher_by_thread.get(thread_id) is t:
del _watcher_by_thread[thread_id]
task.add_done_callback(_cleanup)
return task
def _drain_one_queue(q: queue.Queue) -> list[AsyncTaskNotification]:
items: list[AsyncTaskNotification] = []
while True:
try:
items.append(q.get_nowait())
except queue.Empty:
return items
def drain_notifications(
current_thread_id: str | None = None,
) -> list[AsyncTaskNotification]:
"""Pull pending notifications off the queue (non-blocking).
With ``current_thread_id``: drains the matching per-thread queue plus
the unrouted bucket. Without it: drains EVERY queue (legacy behavior;
used by tests and diagnostics).
"""
if current_thread_id is None:
items: list[AsyncTaskNotification] = _drain_one_queue(_unrouted_queue)
with _notifications_lock:
queues = list(_notifications_by_thread.values())
for q in queues:
items.extend(_drain_one_queue(q))
return items
items = _drain_one_queue(_unrouted_queue)
with _notifications_lock:
q = _notifications_by_thread.get(current_thread_id)
if q is not None:
items.extend(_drain_one_queue(q))
return items
def dedup_notifications(
notifs: list[AsyncTaskNotification],
async_tasks: dict[str, dict] | None,
) -> list[AsyncTaskNotification]:
"""Filter notifications the agent has already 'seen' via prior check.
Logic: skip a notification if `async_tasks[task_id]` exists with a TERMINAL
status and `last_checked_at >= last_updated_at` (timestamps are ISO-8601
so lexicographic comparison is correct). Also skip if `last_checked_at`
is empty (brand-new task where agent hasn't checked yet).
"""
if not async_tasks:
return notifs
survivors: list[AsyncTaskNotification] = []
for n in notifs:
task = async_tasks.get(n.task_id)
if (
task
and task.get("status") in TERMINAL_STATUSES
and task.get("last_checked_at", "") >= task.get("last_updated_at", "")
and task.get("last_checked_at", "") != ""
):
logger.debug(
"Dedup: skipping notification for already-checked task %s", n.task_id
)
continue
survivors.append(n)
return survivors
def format_notification_lines(
notifs: list[AsyncTaskNotification],
) -> list[tuple[str, str]]:
"""Render notifications as compact tool-result-style lines for screen display.
Returns a list of (text, rich_style) tuples — one per notification.
Used by both Rich CLI (console.print) and TUI (_append_system).
The LLM still receives the full format_batch_message text; this is
purely a visual representation for the human operator.
"""
if not notifs:
return []
# Open-right compact frame: short symmetric dashes around the title.
# Bottom matches the top's width so the visual is balanced.
# ╭── ✦ Agent Teams ✦ ──
# ✔ writing Task: ... success
# ╰─────────────────────
title = " ✦ Agent Teams ✦ "
top_divider = "╭──" + title + "────" # 4 dashes on the right (2x of left)
bottom_divider = "╰" + "─" * (len(top_divider) - 1)
lines: list[tuple[str, str]] = [(top_divider, "dim")]
for n in notifs:
# Strip the "-agent" suffix so it doesn't redundantly echo the header.
# `writing-agent` → `writing`, `data-analysis-agent` → `data-analysis`.
name = n.agent_name.removesuffix("-agent")
if n.status == "success":
icon, color = "✔", "#e67e22" # carrot orange (CSS hex; Rich+Textual)
elif n.status == "error":
icon, color = "✗", "red"
else: # cancelled, timeout, interrupted
icon, color = "⚠", "yellow"
# Body format (5-space indent under "Agent:" header):
# ✔ writing Task: <prompt preview> success
# Collapse newlines, truncate prompt to 60 chars.
prompt_preview = (n.prompt or "").replace("\n", " ").strip()
if len(prompt_preview) > 60:
prompt_preview = prompt_preview[:60] + "…"
if prompt_preview:
text = f" {icon} {name:18s} Task: {prompt_preview} {n.status}"
else:
# Fallback: short task_id when no prompt is available
short_tid = (
f"{n.task_id[:8]}…{n.task_id[-4:]}"
if len(n.task_id) > 12
else n.task_id
)
text = f" {icon} {name:18s} ({short_tid}) {n.status}"
lines.append((text, color))
lines.append((bottom_divider, "dim"))
return lines
def format_batch_message(notifs: list[AsyncTaskNotification]) -> str:
"""Compose the synthetic user message that wakes the supervisor.
Each task is rendered as a compact JSON object (one per line) so the LLM
can reliably parse agent name, status, and task_id without ambiguity.
``ensure_ascii=False`` lets non-ASCII agent names pass through unchanged.
Visual decoration lives in ``format_notification_lines``.
"""
if not notifs:
return ""
lines = ["[Async tasks update]"]
for n in notifs:
lines.append(
json.dumps(
{"agent": n.agent_name, "status": n.status, "task_id": n.task_id},
ensure_ascii=False,
)
)
lines.append(
"(Signal only — fetch via check_async_task if relevant to current step, "
"else acknowledge & continue.)"
)
return "\n".join(lines)
# Brief grace window after the last drain: catch one final burst of arrivals
NOTIFICATION_BATCH_GRACE_SECONDS = 0.3
# Max time we'll wait for in-flight watchers to settle before triggering the
# agent turn — bounds latency for long-running tasks while still batching
# co-completing ones.
NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS = 3.0
async def consume_notifications(
run_message: Callable[[str, list[AsyncTaskNotification]], Awaitable[None]],
read_async_tasks_state: Callable[[], Awaitable[dict[str, dict]]],
current_thread_id: str | None = None,
) -> None:
"""Drain queue, dedup, batch, and inject as a synthetic user message.
Args:
run_message: async callable receiving (llm_text, notifs_list).
``llm_text`` is the full structured message for the LLM
(from ``format_batch_message``). ``notifs_list`` is the
survivors list so callers can render per-task visual lines
without re-parsing the text.
read_async_tasks_state: async callable returning current ``async_tasks``
from the agent's state for dedup.
current_thread_id: the active CLI thread id. When given, only
notifications whose ``origin_cli_thread_id`` matches (or that
were enqueued unrouted) are drained — notifications belonging
to other threads stay queued and naturally drain on the next
poller tick after the user ``/resume``s back into them. When
omitted (legacy callers / tests), every queue drains.
"""
notifs = drain_notifications(current_thread_id)
if not notifs:
return
# Adaptive grace: if other watchers tied to THIS thread (or unrouted) are
# still in flight, wait briefly for them to settle so co-completing tasks
# batch into a single agent turn. Sibling-thread watchers don't count —
# their notifications wouldn't drain on this tick anyway.
loop = asyncio.get_running_loop()
deadline = loop.time() + NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS
while _has_relevant_active_watchers(current_thread_id) and loop.time() < deadline:
await asyncio.sleep(0.2)
notifs.extend(drain_notifications(current_thread_id))
# Final brief grace to catch arrivals enqueued just before this tick
await asyncio.sleep(NOTIFICATION_BATCH_GRACE_SECONDS)
notifs.extend(drain_notifications(current_thread_id))
try:
async_tasks = await read_async_tasks_state()
except Exception:
logger.warning("Failed to read async_tasks state for dedup", exc_info=True)
async_tasks = {}
survivors = dedup_notifications(notifs, async_tasks)
if not survivors:
logger.info(
"All %d notifications deduped (already known to agent)", len(notifs)
)
return
text = format_batch_message(survivors)
await run_message(text, survivors)
+96 -10
View File
@@ -798,6 +798,80 @@ def _serve_process_message(
# =============================================================================
def _serve_drain_notifications(
*,
agent_holder: dict,
model: str | None,
workspace_dir: str,
show_thinking: bool,
) -> None:
"""Drain the async-task notification queue in headless serve mode.
Mirrors the Rich CLI's ``_check_channel_queue`` notification path.
Uses a dedicated event loop (same pattern as serve mode's slash dispatch).
"""
import asyncio as _aio
from EvoScientist.cli import async_notifier
from .tui_runtime import run_streaming
def _run_notification_message(text: str, notifs: list) -> None:
"""Synchronous wrapper: run the agent on the synthetic notification text."""
# Render the per-task visual frame (matches CLI/TUI aesthetic).
from EvoScientist.cli.async_notifier import format_notification_lines
for line_text, line_style in format_notification_lines(notifs):
console.print(line_text, style=line_style, markup=False)
# Use the current workspace from agent_holder (updated by /resume's
# session-rebind callback), falling back to the startup value.
runtime_workspace = agent_holder.get("workspace_dir") or workspace_dir
meta = build_metadata(runtime_workspace, model)
try:
run_streaming(
ui_backend="cli",
agent=agent_holder["agent"],
message=text,
thread_id=agent_holder["thread_id"],
show_thinking=show_thinking,
interactive=True,
metadata=meta,
)
except Exception as exc:
_serve_logger.warning("Notification agent turn failed: %s", exc)
async def _run_notification_message_async(text: str, notifs: list) -> None:
await _aio.to_thread(_run_notification_message, text, notifs)
async def _read_async_tasks() -> dict:
agent = agent_holder.get("agent")
thread_id = agent_holder.get("thread_id")
if agent is None or not thread_id:
return {}
try:
snap = await agent.aget_state({"configurable": {"thread_id": thread_id}})
return (snap.values or {}).get("async_tasks") or {}
except Exception:
return {}
async def _consume() -> None:
await async_notifier.consume_notifications(
run_message=_run_notification_message_async,
read_async_tasks_state=_read_async_tasks,
current_thread_id=agent_holder.get("thread_id"),
)
_notif_loop: _aio.AbstractEventLoop | None = None
try:
_notif_loop = _aio.new_event_loop()
_notif_loop.run_until_complete(_consume())
except Exception as exc:
_serve_logger.warning("Notification drain failed: %s", exc)
finally:
if _notif_loop is not None:
_notif_loop.close()
@app.command()
def serve(
no_thinking: bool = typer.Option(
@@ -960,23 +1034,35 @@ def serve(
try:
msg = _message_queue.get(timeout=0.5)
except queue.Empty:
continue
msg = None
if shutdown_event.is_set():
break
try:
_serve_process_message(
msg,
if msg is not None:
try:
_serve_process_message(
msg,
agent_holder=agent_holder,
model=config.model,
workspace_dir=ws,
show_thinking=effective_channel_thinking,
on_cmd_completed=_serve_on_cmd_completed,
start_new_session_cb=_serve_start_new_session_cb,
channel_runtime=channel_runtime,
)
except KeyboardInterrupt:
shutdown_event.set()
break
# Poll notification queue when idle (no channel message was pending).
from EvoScientist.cli import async_notifier
if async_notifier.has_pending_notifications(agent_holder.get("thread_id")):
_serve_drain_notifications(
agent_holder=agent_holder,
model=config.model,
workspace_dir=ws,
show_thinking=effective_channel_thinking,
on_cmd_completed=_serve_on_cmd_completed,
start_new_session_cb=_serve_start_new_session_cb,
channel_runtime=channel_runtime,
)
except KeyboardInterrupt:
shutdown_event.set()
break
except KeyboardInterrupt:
shutdown_event.set()
finally:
+94 -3
View File
@@ -967,15 +967,106 @@ def cmd_interactive(
finally:
_ch_mod._complete_channel_request(msg.msg_id)
async def _inject_notification_message(
text: str,
notifs: list,
*,
target_thread_id: str | None,
) -> None:
"""Inject a batched async-task notification as a synthetic user message.
Renders one compact tool-result-style line per task (matching the
TaskList spinner aesthetic) for the human operator. The LLM
receives the full structured ``text`` from ``format_batch_message``
unchanged — only the screen visual is simplified.
"""
from EvoScientist.cli.async_notifier import format_notification_lines
for line_text, line_style in format_notification_lines(notifs):
console.print(line_text, style=line_style, markup=False)
meta = build_metadata(state["workspace_dir"], model)
await _refresh_status_snapshot(text, reset_streaming_text=True)
run_streaming(
ui_backend=state["ui_backend"],
agent=await _await_agent_ready(),
message=text,
# Falls back to live state["thread_id"] if no override is
# passed (legacy / direct-call paths). Dedup reader has no
# fallback and returns {} for a falsey id; the asymmetry
# is intentional — we'd rather inject into the live thread
# than drop the notification entirely.
thread_id=target_thread_id or state["thread_id"],
show_thinking=show_thinking,
interactive=True,
metadata=meta,
on_stream_event=_handle_stream_status_event,
status_footer_builder=_stream_status_footer,
)
await _refresh_status_snapshot(reset_streaming_text=True)
console.print()
_print_separator()
sys.stdout.write("\033[34;1m❯\033[0m ")
sys.stdout.flush()
async def _read_current_async_tasks(
target_thread_id: str | None,
) -> dict[str, dict]:
"""Snapshot async_tasks from the active agent state for dedup.
Uses ``agent_loader.agent`` (the currently loaded agent) and
``target_thread_id`` (the thread id captured at the start of
``consume_notifications`` — frozen so a mid-consume ``/new``
cannot make us read the wrong thread's state).
"""
agent = agent_loader.agent
if agent is None or not target_thread_id:
return {}
try:
snap = await agent.aget_state(
{"configurable": {"thread_id": target_thread_id}}
)
return (snap.values or {}).get("async_tasks") or {}
except Exception:
return {}
async def _check_channel_queue() -> None:
"""Poll the channel message queue and dispatch to the agent."""
"""Poll the channel + notification queues and dispatch."""
from EvoScientist.cli import async_notifier
while True:
try:
msg = _message_queue.get_nowait()
except queue.Empty:
await asyncio.sleep(0.1)
msg = None
if msg is not None:
await _process_channel_message(msg)
continue # check queues again immediately
# Notification path (only when no channel message was pending).
# Wrap in try/except so an exception in dedup/inject can't
# kill the poller task — channel + notification dispatch
# would silently die otherwise (Fix #4).
current_tid = state.get("thread_id")
if async_notifier.has_pending_notifications(current_tid):
try:
await async_notifier.consume_notifications(
run_message=lambda text, notifs, _tid=current_tid: (
_inject_notification_message(
text, notifs, target_thread_id=_tid
)
),
read_async_tasks_state=lambda _tid=current_tid: (
_read_current_async_tasks(_tid)
),
current_thread_id=current_tid,
)
except Exception:
_channel_logger.warning(
"async-notifier consume failed", exc_info=True
)
continue
await _process_channel_message(msg)
await asyncio.sleep(0.1)
queue_task = asyncio.create_task(_check_channel_queue())
+158 -12
View File
@@ -380,6 +380,9 @@ def run_textual_interactive(
self._channel_timer: Any = None
self._started_channel_types: list[str] = []
self._busy = False
self._notification_consuming: bool = (
False # prevent overlapping consume coroutines
)
self._run_task: Any = None # asyncio.Task for current _run_turn
self._queued_messages: list[
str
@@ -785,18 +788,131 @@ def run_textual_interactive(
self._channel_timer = self.set_interval(0.1, self._poll_channel_queue)
def _poll_channel_queue(self) -> None:
"""Poll the channel message queue (called every 100ms)."""
"""Poll the channel + notification queues (every 100ms)."""
from EvoScientist.cli import async_notifier
try:
msg = _message_queue.get_nowait()
except queue.Empty:
msg = None
if msg is not None:
if self._busy:
_message_queue.put(msg)
return
self.call_later(
lambda m=msg: asyncio.ensure_future(
self._process_channel_message(m)
)
)
return
if self._busy:
_message_queue.put(msg)
return
self.call_later(
lambda m=msg: asyncio.ensure_future(self._process_channel_message(m))
# Notification path (only when idle and NOT already consuming).
# _notification_consuming is set synchronously at the schedule point
# so that the next poll tick cannot schedule a second consumer before
# the first one has a chance to run (fixes overlapping-turn bug).
if (
async_notifier.has_pending_notifications(self._conversation_tid)
and not self._busy
and not self._notification_consuming
):
self._notification_consuming = True
self.call_later(
lambda: asyncio.ensure_future(self._consume_notifications_tui())
)
async def _consume_notifications_tui(self) -> None:
"""Drain the notification queue and inject a synthetic agent turn.
Wraps the consume call in a swallowing try/except (Fix #4) so an
exception inside dedup/inject doesn't bubble out of the
``asyncio.ensure_future(...)`` scheduled by ``_poll_channel_queue``
and silently kill notification + channel dispatch.
"""
from EvoScientist.cli import async_notifier
target_tid = self._conversation_tid
try:
try:
await async_notifier.consume_notifications(
run_message=lambda text, notifs: self._inject_notification_tui(
text, notifs, target_thread_id=target_tid
),
read_async_tasks_state=lambda: self._read_async_tasks_tui(
target_tid
),
current_thread_id=target_tid,
)
except Exception:
import logging
logging.getLogger(__name__).warning(
"async-notifier consume failed (TUI)", exc_info=True
)
finally:
# Clear the guard flag regardless of success or exception so
# future notifications can schedule a new consume coroutine.
self._notification_consuming = False
async def _inject_notification_tui(
self,
text: str,
notifs: list,
*,
target_thread_id: str | None = None,
) -> None:
"""Run a synthetic user turn for the batched async-task notification.
Renders one compact tool-result-style line per task (matching the
Rich CLI aesthetic) instead of a single breadcrumb. The LLM still
receives the full ``format_batch_message`` text; only the visual
representation changes.
Args:
text: Full structured LLM message from ``format_batch_message``.
notifs: Survivor notification list for per-task visual rendering.
target_thread_id: Pinned thread id for the synthetic turn —
forwarded to ``_run_turn`` so a mid-consume ``/new`` cannot
misroute the notification into a different thread.
"""
from EvoScientist.cli.async_notifier import format_notification_lines
for line_text, line_style in format_notification_lines(notifs):
self._append_system(line_text, style=line_style)
# Fire-and-forget the turn as an INDEPENDENT task — matches the
# keyboard input path (line ~2113). Queue-triggered turns that
# `await _run_turn` from inside a nested call_later chain don't
# get viewport-follow during streaming (only after completion).
# Mark busy synchronously so the next poll tick doesn't re-enter.
self._busy = True
self._run_task = asyncio.ensure_future(
self._run_turn(
text,
skip_user_message=True,
resolve_mentions=False,
thread_id_override=target_thread_id,
)
)
async def _read_async_tasks_tui(
self, target_thread_id: str | None
) -> dict[str, dict]:
"""Read async_tasks from agent state for dedup, against a frozen tid.
``target_thread_id`` is captured by ``_consume_notifications_tui`` at
the start of the consume call so a mid-consume thread switch cannot
make us read the wrong thread's state.
"""
agent = self._agent_loader.agent
if agent is None or not target_thread_id:
return {}
try:
snap = await agent.aget_state(
{"configurable": {"thread_id": target_thread_id}}
)
return (snap.values or {}).get("async_tasks") or {}
except Exception:
return {}
async def _on_channel_cmd_completed(
self,
ctx: CommandContext,
@@ -1051,6 +1167,7 @@ def run_textual_interactive(
channel_hitl_fn: Callable[[list], list[dict] | None] | None = None,
channel_ask_user_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
thread_id_override: str | None = None,
) -> str:
"""Stream agent events and mount widgets. Returns response text.
@@ -1259,7 +1376,7 @@ def run_textual_interactive(
async for event in stream_agent_events(
self._agent_loader.agent,
_stream_input,
self._conversation_tid,
thread_id_override or self._conversation_tid,
metadata=metadata,
):
if is_stream_cancel_requested(cancel_scope):
@@ -1803,8 +1920,31 @@ def run_textual_interactive(
return response
async def _run_turn(self, user_text: str) -> None:
"""Handle a user turn: stream agent response with widgets."""
async def _run_turn(
self,
user_text: str,
*,
skip_user_message: bool = False,
resolve_mentions: bool = True,
thread_id_override: str | None = None,
) -> None:
"""Handle a user turn: stream agent response with widgets.
Args:
user_text: The user's message text.
skip_user_message: If True, suppress the UserMessage widget echo
(caller has already displayed a visual representation of the
input — e.g. async-notifier per-task lines).
resolve_mentions: If False, skip ``@file`` mention expansion.
Used by synthetic notifier turns whose payload is a fixed
JSON template — keeps the TUI path consistent with the
Rich CLI notifier path which never expands mentions.
thread_id_override: Pin the agent stream to this thread instead
of the live ``self._conversation_tid``. Used by the async
notifier path so a mid-consume ``/new`` cannot redirect a
notification meant for thread A into thread B. Falls back
to the live tid when ``None``.
"""
cancelled = False
try:
self._busy = True
@@ -1813,9 +1953,13 @@ def run_textual_interactive(
# Resolve @file mentions — inject file contents before sending to agent.
# Use self._workspace_dir (current session) not the startup-captured
# workspace_dir closure, which becomes stale after /new or /resume.
_, message_to_send, file_warnings = await asyncio.to_thread(
resolve_file_mentions, user_text, self._workspace_dir
)
if resolve_mentions:
_, message_to_send, file_warnings = await asyncio.to_thread(
resolve_file_mentions, user_text, self._workspace_dir
)
else:
message_to_send = user_text
file_warnings = []
await self._refresh_status_snapshot(message_to_send)
# Block the turn on MCP tools finishing, if still in flight.
@@ -1830,6 +1974,8 @@ def run_textual_interactive(
message_to_send,
display_text=user_text,
file_warnings=file_warnings,
skip_user_message=skip_user_message,
thread_id_override=thread_id_override,
)
except asyncio.CancelledError:
cancelled = True
+114
View File
@@ -0,0 +1,114 @@
"""Spawn async-task watchers when the agent invokes start/update_async_task.
Hooks via ``awrap_tool_call`` so it only fires on the two launch tools — every
other tool call is a no-op pass-through. The middleware does not register tools
of its own; deepagents' built-in async-subagents middleware already publishes
``start_async_task`` / ``update_async_task`` to the agent.
Stable contract this depends on:
* The two public tool names ``start_async_task`` and ``update_async_task``.
* The ``Command(update={"async_tasks": {task_id: AsyncTask}})`` state schema
returned by both tools.
* ``runtime.config["configurable"]["thread_id"]`` for capturing the originating
CLI thread.
It also imports ``_ClientCache`` from deepagents — that is a private symbol but
a stable typed class, used here only as a connection-pool helper keyed by
``(url, headers)``.
"""
from __future__ import annotations
import logging
from collections.abc import Awaitable, Callable
from typing import Any
from langchain.agents.middleware import AgentMiddleware
from langchain.agents.middleware.types import ToolCallRequest
from langchain_core.messages import ToolMessage
from langgraph.types import Command
logger = logging.getLogger(__name__)
_LAUNCH_TOOL_NAMES = ("start_async_task", "update_async_task")
class AsyncWatcherMiddleware(AgentMiddleware):
"""Spawn an ``async_notifier`` watcher whenever the agent launches or
updates an async sub-agent task.
Args:
async_agents: Mapping of subagent name → ``AsyncSubAgent`` TypedDict
(must contain at least ``url`` and ``graph_id``). Used to construct
a ``_ClientCache`` for resolving the LangGraph client per agent.
"""
def __init__(self, async_agents: dict[str, Any]) -> None:
from deepagents.middleware.async_subagents import _ClientCache
super().__init__()
self._clients = _ClientCache(async_agents)
async def awrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
) -> ToolMessage | Command:
from EvoScientist.cli import async_notifier
name = request.tool_call.get("name")
args = request.tool_call.get("args") or {}
# Pre-cancel the existing watcher BEFORE the new run interrupts the old
# one. ``update_async_task`` creates a new run on the same thread_id with
# ``multitask_strategy="interrupt"``, which closes the old run's stream
# cleanly — without pre-cancellation the old watcher would observe a
# clean exit and enqueue a stale "success" notification before the new
# spawn can replace it.
if name == "update_async_task" and (tid := args.get("task_id")):
try:
old = async_notifier._watcher_by_thread.get(tid)
if old is not None and not old.done():
old.cancel()
except Exception:
logger.warning(
"Pre-cancel of stale watcher for task %s failed; a stale "
"success notification may be enqueued",
tid,
exc_info=True,
)
result = await handler(request)
if name in _LAUNCH_TOOL_NAMES and isinstance(result, Command):
cli_thread_id = None
cfg = getattr(getattr(request, "runtime", None), "config", None)
if isinstance(cfg, dict):
cli_thread_id = cfg.get("configurable", {}).get("thread_id")
# Tool-name-gated prompt field — `start_async_task` defines
# `description`, `update_async_task` defines `message`.
prompt_field = "description" if name == "start_async_task" else "message"
prompt = args.get(prompt_field, "")
tasks_update = (result.update or {}).get("async_tasks") or {}
for task_id, task in tasks_update.items():
try:
client = self._clients.get_async(task["agent_name"])
async_notifier.spawn_watcher(
client,
task_id,
task["run_id"],
task["agent_name"],
prompt=prompt,
origin_cli_thread_id=cli_thread_id,
)
except Exception:
logger.warning(
"Failed to spawn watcher for task %s",
task_id,
exc_info=True,
)
return result
+36 -1
View File
@@ -9,7 +9,9 @@ The main agent's system prompt is assembled by :func:`get_system_prompt` from:
- :data:`REPORT_TEMPLATE` — final-report structure
- :data:`WRITING_GUIDELINES` — style rules for written output
- :data:`SHELL_GUIDELINES` — sandbox limits and `execute` tool usage
- :data:`DELEGATION_STRATEGY` — sub-agent delegation strategy
- :data:`DELEGATION_STRATEGY` — sub-agent delegation strategy (sync sub-agents)
- :data:`ASYNC_NOTIFICATIONS` — how to triage `[Async tasks update]` signals
from async sub-agents
:data:`RESEARCHER_INSTRUCTIONS` is the research-agent sub-agent prompt,
loaded via ``_build_prompt_refs`` in ``EvoScientist.py``.
@@ -315,6 +317,37 @@ After each stage, ask: "Would a critical reviewer accept this evidence?"
- Each sub-agent returns self-contained findings with concrete artifacts.
"""
# =============================================================================
# Async sub-agent notifications
# =============================================================================
ASYNC_NOTIFICATIONS = """# Async Task Notifications
A `[Async tasks update]` message is a SIGNAL of background completion, not a
new request.
## Hard rules (read these first)
NEVER:
- Switch the topic away from an ongoing user-clarification dialogue.
- Hijack a literature search or experiment step into a summary of the
unrelated finished task.
- Silently ignore — always at minimum acknowledge so the user knows the
signal was seen.
## Per-task triage
For EACH task in the batch, independently:
- Result needed for the CURRENT step → fetch the result, integrate,
continue your work in the same turn.
- Otherwise → acknowledge in ONE short line (e.g. "Noted: data-analysis-agent
finished — will fetch when relevant"), then RESUME what you were doing.
- `status="error"` → surface briefly to the user even if not currently
relevant; ask whether to retry or wait.
It is fine to fetch one task and defer another from the same batch.
"""
# =============================================================================
# Sub-agent research instructions
# =============================================================================
@@ -381,6 +414,7 @@ def get_system_prompt() -> str:
4. :data:`WRITING_GUIDELINES`
5. :data:`SHELL_GUIDELINES`
6. :data:`DELEGATION_STRATEGY`
7. :data:`ASYNC_NOTIFICATIONS`
The current date is injected per-turn by
:class:`EvoScientist.middleware.EvoMemoryMiddleware` (piggy-backing on its
@@ -397,5 +431,6 @@ def get_system_prompt() -> str:
WRITING_GUIDELINES,
SHELL_GUIDELINES,
DELEGATION_STRATEGY,
ASYNC_NOTIFICATIONS,
]
return "\n".join(sections)
+1
View File
@@ -779,6 +779,7 @@ def create_streaming_display(
elements.append(Spinner("dots", text=" Processing...", style="cyan"))
if status_footer is not None:
elements.append(status_footer)
return Group(*elements)
+6
View File
@@ -25,6 +25,12 @@ def run_async(coro):
loop.close()
@pytest.fixture(name="run_async")
def run_async_fixture():
"""Pytest fixture that exposes run_async as a callable for test functions."""
return run_async
@pytest.fixture
def sample_tool_call():
"""A minimal tool call dict."""
+859
View File
@@ -0,0 +1,859 @@
"""Tests for async sub-agent auto-notification."""
import asyncio
import queue
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
from EvoScientist.cli import async_notifier
from EvoScientist.cli.async_notifier import (
dedup_notifications,
drain_notifications,
format_batch_message,
format_notification_lines,
)
def test_notification_dataclass_fields():
n = async_notifier.AsyncTaskNotification(
task_id="tid-1",
agent_name="writing-agent",
status="success",
received_at="2026-05-06T12:00:00Z",
)
assert n.task_id == "tid-1"
assert n.status == "success"
def test_notification_queue_is_module_level_fifo():
# Drain anything left over from other tests
while True:
try:
async_notifier._notification_queue.get_nowait()
except queue.Empty:
break
n1 = async_notifier.AsyncTaskNotification("a", "x", "success", "")
n2 = async_notifier.AsyncTaskNotification("b", "x", "success", "")
async_notifier._notification_queue.put(n1)
async_notifier._notification_queue.put(n2)
assert async_notifier._notification_queue.get_nowait().task_id == "a"
assert async_notifier._notification_queue.get_nowait().task_id == "b"
def _drain_queue(q):
items = []
while True:
try:
items.append(q.get_nowait())
except queue.Empty:
return items
def test_watcher_pushes_notification_on_stream_end(run_async):
# Stream yields one "values" chunk with the final state, then closes
final_state = {
"messages": [{"type": "ai", "content": "Quantum superposition is..."}]
}
chunks = [SimpleNamespace(event="values", data=final_state)]
async def fake_stream(thread_id, run_id, stream_mode):
for c in chunks:
yield c
client = MagicMock()
client.runs.join_stream = fake_stream
# runs.get is used to fetch terminal status when stream ends
client.runs.get = AsyncMock(return_value={"status": "success"})
_drain_all(async_notifier)
run_async(
async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent")
)
notifs = _drain_queue(async_notifier._notification_queue)
assert len(notifs) == 1
assert notifs[0].task_id == "thr-1"
assert notifs[0].agent_name == "writing-agent"
assert notifs[0].status == "success"
def test_watcher_pushes_error_status_on_stream_exception(run_async):
async def fake_stream(*a, **kw):
raise RuntimeError("network broken")
yield # unreachable; makes this an async generator
client = MagicMock()
client.runs.join_stream = fake_stream
# On stream failure, watcher falls back to runs.get for terminal status
client.runs.get = AsyncMock(
return_value={"status": "error", "error": "network broken"}
)
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "error"
def test_spawn_watcher_replaces_existing_for_same_thread(run_async):
"""A second spawn_watcher with the same thread_id cancels the old watcher
and registers the new one — supports update_async_task creating a new
run_id on the same thread_id."""
spawn_starts = []
async def fake_stream_long(*a, **kw):
spawn_starts.append("started")
# Simulate a long-running stream that gets cancelled
try:
while True:
await asyncio.sleep(0.01)
yield SimpleNamespace(event="values", data={"messages": []})
except asyncio.CancelledError:
raise
client = MagicMock()
client.runs.join_stream = fake_stream_long
client.runs.get = AsyncMock(return_value={"status": "success"})
async def scenario():
# Clear all queues and the watcher registries
async_notifier._active_watchers.clear()
async_notifier._watcher_by_thread.clear()
_drain_all(async_notifier)
# First spawn for thread X, run R1
t1 = async_notifier.spawn_watcher(client, "thr-X", "R1", "agent")
assert t1 is not None
assert async_notifier._watcher_by_thread["thr-X"] is t1
await asyncio.sleep(0.02) # let it start streaming
# Second spawn for SAME thread X, NEW run R2
t2 = async_notifier.spawn_watcher(client, "thr-X", "R2", "agent")
assert t2 is not None
assert t2 is not t1
assert async_notifier._watcher_by_thread["thr-X"] is t2
# Old watcher should be cancelled
await asyncio.sleep(0.02)
assert t1.cancelled() or t1.done()
# Cleanup the new task too
t2.cancel()
try:
await t2
except asyncio.CancelledError:
pass
# Cancelled watchers don't push notifications
assert _drain_one_queue_helper(async_notifier._notification_queue) == []
assert _drain_one_queue_helper(async_notifier._unrouted_queue) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
run_async(scenario())
# ============================================================================
# Tests for drain_notifications, dedup_notifications, format_batch_message
# ============================================================================
def test_format_notification_lines_returns_decorated_block():
"""Output is: divider with 'Agent' inset, body lines, plain bottom divider."""
notifs = [
async_notifier.AsyncTaskNotification("t1", "writing-agent", "success", "", ""),
async_notifier.AsyncTaskNotification("t2", "data-agent", "error", "", ""),
async_notifier.AsyncTaskNotification("t3", "code-agent", "cancelled", "", ""),
]
lines = format_notification_lines(notifs)
# top divider (with title) + 3 body lines + bottom divider = 5
assert len(lines) == 5
# Top divider — open-right frame with ornaments: "╭── ✦ Agent ✦ ─────"
top_text, top_style = lines[0]
assert "Agent" in top_text
assert "✦" in top_text
assert top_text.startswith("╭")
assert top_text.endswith("─") # open right side
assert top_style == "dim"
# Body lines (indented, "-agent" suffix stripped)
text1, style1 = lines[1]
text2, style2 = lines[2]
text3, style3 = lines[3]
assert text1.startswith(" ")
assert "writing" in text1
assert "writing-agent" not in text1
assert "success" in text1
assert "✔" in text1
assert style1.startswith("#")
assert " data " in text2
assert "data-agent" not in text2
assert "error" in text2
assert "✗" in text2
assert style2 == "red"
assert " code " in text3
assert "code-agent" not in text3
assert "cancelled" in text3
assert "⚠" in text3
assert style3 == "yellow"
# Bottom divider — open-right frame, same width as top
bot_text, bot_style = lines[4]
assert "Agent" not in bot_text
assert bot_text.startswith("╰")
assert bot_text.endswith("─")
assert len(bot_text) == len(top_text)
assert bot_style == "dim"
def test_format_notification_lines_empty_returns_empty():
"""format_notification_lines returns an empty list for no notifications."""
lines = format_notification_lines([])
assert lines == []
def test_format_notification_lines_renders_prompt_when_provided():
"""When prompt is set, the body line shows `Task: <prompt preview>`."""
notifs = [
async_notifier.AsyncTaskNotification(
task_id="019dfe2f-aaaa",
agent_name="writing-agent",
status="success",
received_at="",
prompt="请用中文写一段关于量子叠加的简短介绍",
),
]
lines = format_notification_lines(notifs)
# top divider (with title) + 1 body + bottom divider = 3
assert len(lines) == 3
body_text, _style = lines[1]
assert "Task:" in body_text
assert "量子叠加" in body_text
assert "writing" in body_text
assert "writing-agent" not in body_text
assert "success" in body_text
def test_format_notification_lines_truncates_long_prompt():
"""Prompts longer than 60 chars get truncated with an ellipsis."""
long_prompt = "x" * 200
notifs = [
async_notifier.AsyncTaskNotification(
task_id="t1",
agent_name="agent",
status="success",
received_at="",
prompt=long_prompt,
),
]
# Body line is at index 1: [top divider with title, body, bottom divider]
body_text = format_notification_lines(notifs)[1][0]
assert "…" in body_text
assert "x" * 200 not in body_text
def test_format_notification_lines_collapses_newlines_in_prompt():
"""Multi-line prompts collapse to single line for the visual."""
notifs = [
async_notifier.AsyncTaskNotification(
task_id="t1",
agent_name="agent",
status="success",
received_at="",
prompt="line one\nline two\nline three",
),
]
body_text = format_notification_lines(notifs)[1][0]
assert "\n" not in body_text
assert "line one line two" in body_text
def test_format_notification_lines_falls_back_to_task_id_when_no_prompt():
"""Without a prompt, fall back to the short task_id."""
notifs = [
async_notifier.AsyncTaskNotification(
"019dfe2f-821a-7d43-ac5b-6bb8781be5cf",
"writing-agent",
"success",
"",
"", # no prompt
),
]
body_text = format_notification_lines(notifs)[1][0]
assert "019dfe2f" in body_text
assert "e5cf" in body_text
assert "Task:" not in body_text
def test_format_notification_lines_timeout_uses_warning_icon():
"""Timeout and interrupted statuses get the warning icon and yellow style."""
for status in ("timeout", "interrupted"):
notifs = [
async_notifier.AsyncTaskNotification("t", "some-agent", status, "", "")
]
lines = format_notification_lines(notifs)
# top divider with title + 1 body + bottom divider = 3
assert len(lines) == 3
body_text, body_style = lines[1]
assert "⚠" in body_text
assert body_style == "yellow"
assert status in body_text
def test_drain_returns_all_pending_and_empties_queue():
"""drain_notifications pulls every pending notification and empties queue."""
# Clear the queue first
while True:
try:
async_notifier._notification_queue.get_nowait()
except queue.Empty:
break
# Add three notifications
for tid in ("a", "b", "c"):
async_notifier._notification_queue.put(
async_notifier.AsyncTaskNotification(tid, "x", "success", "", "")
)
drained = drain_notifications()
assert [n.task_id for n in drained] == ["a", "b", "c"]
assert async_notifier._notification_queue.empty()
def test_dedup_skips_tasks_already_checked_after_terminal():
"""dedup_notifications skips tasks with terminal status and last_checked_at >= last_updated_at."""
async_tasks = {
"a": {
"status": "success",
"last_checked_at": "2026-05-06T12:01:00Z",
"last_updated_at": "2026-05-06T12:00:00Z",
}, # already known → skip
"b": {
"status": "success",
"last_checked_at": "2026-05-06T12:00:00Z",
"last_updated_at": "2026-05-06T12:00:30Z",
}, # checked stale → keep
"c": {
"status": "running",
"last_checked_at": "",
"last_updated_at": "",
}, # not terminal → keep
}
notifs = [
async_notifier.AsyncTaskNotification("a", "x", "success", "", ""),
async_notifier.AsyncTaskNotification("b", "x", "success", "", ""),
async_notifier.AsyncTaskNotification(
"d", "x", "success", "", ""
), # not in map → keep
]
survivors = dedup_notifications(notifs, async_tasks)
assert {n.task_id for n in survivors} == {"b", "d"}
def test_format_batch_message_single_notification():
"""format_batch_message produces compact JSON for a single notification."""
notifs = [
async_notifier.AsyncTaskNotification(
task_id="tid-1",
agent_name="writing-agent",
status="success",
received_at="2026-05-07T12:00:00Z",
prompt="Done writing.",
)
]
msg = format_batch_message(notifs)
assert msg.startswith("[Async tasks update]")
# The task line is valid JSON with the expected fields.
task_line = msg.splitlines()[1]
obj = __import__("json").loads(task_line)
assert obj["agent"] == "writing-agent"
assert obj["task_id"] == "tid-1"
assert obj["status"] == "success"
def test_format_batch_message_multiple():
"""format_batch_message handles multiple notifications as separate JSON lines."""
notifs = [
async_notifier.AsyncTaskNotification(
task_id="t1",
agent_name="writing-agent",
status="success",
received_at="2026-05-07T12:00:00Z",
prompt="A",
),
async_notifier.AsyncTaskNotification(
task_id="t2",
agent_name="data-analysis-agent",
status="error",
received_at="2026-05-07T12:00:01Z",
prompt="B",
),
]
msg = format_batch_message(notifs)
lines = msg.splitlines()
assert lines[0] == "[Async tasks update]"
obj1 = __import__("json").loads(lines[1])
obj2 = __import__("json").loads(lines[2])
assert obj1 == {"agent": "writing-agent", "status": "success", "task_id": "t1"}
assert obj2 == {"agent": "data-analysis-agent", "status": "error", "task_id": "t2"}
assert "check_async_task" in msg.lower() # hint to LLM
def test_dedup_preserves_order():
"""dedup_notifications preserves the original order of notifications."""
notifs = [
async_notifier.AsyncTaskNotification("a", "x", "success", "", ""),
async_notifier.AsyncTaskNotification("b", "x", "success", "", ""),
]
survivors = dedup_notifications(notifs, async_tasks={})
assert [n.task_id for n in survivors] == ["a", "b"]
# ============================================================================
# Tests for consume_notifications (integration path)
# ============================================================================
def test_consume_notifications_calls_runner_with_batched_message(run_async):
"""When notifications arrive and agent is idle, consume_notifications fires
the supplied async runner once with the formatted batch message and notifs list."""
from EvoScientist.cli import async_notifier as an
# Set up two pending notifications, no dedup match
while True:
try:
an._notification_queue.get_nowait()
except queue.Empty:
break
an._notification_queue.put(an.AsyncTaskNotification("t1", "wA", "success", "", ""))
an._notification_queue.put(an.AsyncTaskNotification("t2", "wB", "success", "", ""))
captured: dict = {}
async def fake_runner(text: str, notifs: list) -> None:
captured["text"] = text
captured["notifs"] = notifs
async def fake_state_reader() -> dict:
return {} # no dedup info
run_async(an.consume_notifications(fake_runner, fake_state_reader))
assert "wA" in captured["text"]
assert "wB" in captured["text"]
assert len(captured["notifs"]) == 2
def test_consume_notifications_no_op_when_queue_empty(run_async):
from EvoScientist.cli import async_notifier as an
while True:
try:
an._notification_queue.get_nowait()
except queue.Empty:
break
called = False
async def fake_runner(text: str, notifs: list):
nonlocal called
called = True
async def fake_state_reader():
return {}
run_async(an.consume_notifications(fake_runner, fake_state_reader))
assert called is False
# ============================================================================
# Tests for TUI consumer reentry guard (_notification_consuming flag)
# Exercises Fix 2: the flag prevents two overlapping consume coroutines from
# both eventually calling _inject_notification_tui (and thus _run_turn).
# ============================================================================
def test_notification_consuming_flag_prevents_reentry(run_async):
"""The _notification_consuming guard prevents two overlapping consumers.
Verifies the flag contract used by _consume_notifications_tui:
- The flag is checked before scheduling a new consumer.
- The flag is cleared in a try/finally so exceptions don't freeze it.
We model the guard using a dict (avoids nonlocal-in-nested-scope issues)
and run three scenarios sequentially:
1. Normal: flag cleared after first consume finishes → second can run.
2. Blocked: flag pre-set to True → guarded_consume bails out immediately.
3. Exception path: runner raises → flag is still cleared by finally.
"""
from EvoScientist.cli import async_notifier as an
# Clear the queue
while True:
try:
an._notification_queue.get_nowait()
except queue.Empty:
break
state = {"inject_count": 0, "consuming": False}
async def counting_runner(text: str, notifs: list) -> None:
state["inject_count"] += 1
async def fake_state_reader() -> dict:
return {}
async def guarded_consume(notif):
"""Mirror the TUI pattern: check flag, set it, run with try/finally."""
if state["consuming"]:
return # blocked
state["consuming"] = True
try:
an._notification_queue.put(notif)
await an.consume_notifications(counting_runner, fake_state_reader)
finally:
state["consuming"] = False
n1 = an.AsyncTaskNotification("g1", "writing-agent", "success", "", "")
n2 = an.AsyncTaskNotification("g2", "data-agent", "success", "", "")
async def scenario():
# Scenario 1: normal flow — flag cleared, second consumer runs fine.
await guarded_consume(n1)
assert state["inject_count"] == 1
assert state["consuming"] is False # finally ran
state["inject_count"] = 0
await guarded_consume(n2)
assert state["inject_count"] == 1
assert state["consuming"] is False
# Scenario 2: flag pre-set (first consumer in-flight) → second bails.
state["inject_count"] = 0
state["consuming"] = True # simulate first consumer running
an._notification_queue.put(n1)
await guarded_consume(n1) # should be blocked immediately
assert state["inject_count"] == 0 # runner never called
state["consuming"] = False # cleanup
# Scenario 3: exception in runner → flag still cleared by finally.
async def raising_runner(text: str, notifs: list) -> None:
raise RuntimeError("boom")
async def guarded_consume_raising(notif):
if state["consuming"]:
return
state["consuming"] = True
try:
an._notification_queue.put(notif)
await an.consume_notifications(raising_runner, fake_state_reader)
except RuntimeError:
pass
finally:
state["consuming"] = False
await guarded_consume_raising(n2)
assert state["consuming"] is False # cleared despite exception
run_async(scenario())
# ============================================================================
# Tests for Fix #3 — per-thread notification routing
# ============================================================================
def _drain_all(an_mod):
"""Drain every queue (per-thread + unrouted) so tests start clean."""
if hasattr(an_mod, "_notification_queue"):
while True:
try:
an_mod._notification_queue.get_nowait()
except queue.Empty:
break
if hasattr(an_mod, "_notifications_by_thread"):
for q in list(an_mod._notifications_by_thread.values()):
while True:
try:
q.get_nowait()
except queue.Empty:
break
if hasattr(an_mod, "_unrouted_queue"):
while True:
try:
an_mod._unrouted_queue.get_nowait()
except queue.Empty:
break
def test_consume_only_drains_matching_thread(run_async):
"""Notifications tagged with origin_cli_thread_id only drain when the
consumer is invoked with the matching current_thread_id."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
n_a = an.AsyncTaskNotification(
"tA", "writing-agent", "success", "", "", origin_cli_thread_id="threadA"
)
n_b = an.AsyncTaskNotification(
"tB", "writing-agent", "success", "", "", origin_cli_thread_id="threadB"
)
an._enqueue(n_a)
an._enqueue(n_b)
captured: dict = {"runs": []}
async def runner(text: str, notifs: list) -> None:
captured["runs"].append([n.task_id for n in notifs])
async def state_reader() -> dict:
return {}
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadA")
)
assert captured["runs"] == [["tA"]]
# B's notification should still be queued
assert an.has_pending_notifications("threadB")
_drain_all(an)
def test_unrouted_notifications_drain_on_any_thread(run_async):
"""Notifications without origin_cli_thread_id (legacy / direct put) drain
regardless of the current_thread_id arg."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
an._notification_queue.put(
an.AsyncTaskNotification("tU", "writing-agent", "success", "", "")
)
captured: dict = {}
async def runner(text: str, notifs: list) -> None:
captured["notifs"] = notifs
async def state_reader() -> dict:
return {}
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="anything")
)
assert [n.task_id for n in captured["notifs"]] == ["tU"]
_drain_all(an)
def test_thread_switch_drains_pending(run_async):
"""Pending notifications for thread B are not delivered while consumer
asks for thread A; once consumer runs with thread B they drain."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
an._enqueue(
an.AsyncTaskNotification(
"tB", "writing-agent", "success", "", "", origin_cli_thread_id="threadB"
)
)
captured: dict = {"runs": []}
async def runner(text: str, notifs: list) -> None:
captured["runs"].append([n.task_id for n in notifs])
async def state_reader() -> dict:
return {}
# First consume in thread A → no drain, B's notif still queued
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadA")
)
assert captured["runs"] == []
assert an.has_pending_notifications("threadB")
# Now switch to thread B → drains
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadB")
)
assert captured["runs"] == [["tB"]]
_drain_all(an)
def test_has_pending_notifications_respects_routing():
"""has_pending_notifications returns true only for matching or unrouted."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
# Unrouted always counts
an._notification_queue.put(
an.AsyncTaskNotification("tU", "writing-agent", "success", "", "")
)
assert an.has_pending_notifications("threadA") is True
assert an.has_pending_notifications() is True
_drain_all(an)
# Routed only counts for the matching current thread
an._enqueue(
an.AsyncTaskNotification(
"tA", "writing-agent", "success", "", "", origin_cli_thread_id="threadA"
)
)
assert an.has_pending_notifications("threadA") is True
assert an.has_pending_notifications("threadB") is False
assert an.has_pending_notifications() is False # no unrouted, no current_thread
_drain_all(an)
# ============================================================================
# Tests for Fix #1 (v2) — in-band error detection from SSE stream.
#
# We don't poll runs.get after a clean stream close (it had a server-side
# write-back race that returned "error" for successful runs). Instead we
# watch for ``event="error"`` SSE parts which langgraph dev emits when a
# run fails — that signal is authoritative and arrives in-band before the
# stream closes.
# ============================================================================
def test_watcher_reports_error_on_in_band_error_event(run_async):
"""SSE error event in the stream → notification.status == 'error'."""
async def fake_stream(*a, **kw):
yield SimpleNamespace(
event="values", data={"messages": [{"type": "ai", "content": "partial"}]}
)
yield SimpleNamespace(event="error", data={"message": "subagent crashed"})
client = MagicMock()
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(
return_value={"status": "success"}
) # would mislead — should NOT be consulted
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "error"
# We must NOT have polled runs.get — the in-band signal is authoritative.
client.runs.get.assert_not_awaited()
def test_watcher_clean_exit_without_error_event_is_success(run_async):
"""Clean stream exit with no error event → success (no runs.get poll)."""
async def fake_stream(*a, **kw):
yield SimpleNamespace(
event="values", data={"messages": [{"type": "ai", "content": "ok"}]}
)
client = MagicMock()
client.runs.join_stream = fake_stream
# If we (wrongly) poll runs.get and it returned "error", the test would
# fail — proving we no longer have the timing race.
client.runs.get = AsyncMock(return_value={"status": "error"})
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
client.runs.get.assert_not_awaited()
# ============================================================================
# Tests for Fix #4 — consume_notifications surfaces exceptions to caller
# (callers wrap the await in try/except — verify the inner contract is to
# propagate so the wrapper sees + logs).
# ============================================================================
def test_consume_notifications_propagates_inject_exception(run_async):
"""If the run_message callback raises, consume_notifications propagates
the exception to the caller — pollers wrap it in try/except so the
poller task does not die."""
import pytest
from EvoScientist.cli import async_notifier as an
_drain_all(an)
an._notification_queue.put(
an.AsyncTaskNotification("tX", "writing-agent", "success", "", "")
)
async def boom_runner(text: str, notifs: list) -> None:
raise RuntimeError("kaboom")
async def state_reader() -> dict:
return {}
with pytest.raises(RuntimeError, match="kaboom"):
run_async(an.consume_notifications(boom_runner, state_reader))
_drain_all(an)
def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(run_async):
"""When the SSE stream errors AND runs.get returns a non-terminal status
(e.g. ``pending`` because the run is still alive), the watcher must
NOT enqueue a notification — otherwise the user sees a confusing
``⚠ pending`` line for a task that's still working. This is the early-
return guard added alongside the Fix #2 revert."""
async def fake_stream(*a, **kw):
# Simulate transient transport error mid-stream.
raise RuntimeError("connection reset")
yield # unreachable; makes this an async generator
client = MagicMock()
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "pending"})
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrP", "rP", "agentP"))
# No notification should have been enqueued in any queue.
assert _drain_one_queue_helper(async_notifier._unrouted_queue) == []
assert _drain_one_queue_helper(async_notifier._notification_queue) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
def _drain_one_queue_helper(q):
items = []
while True:
try:
items.append(q.get_nowait())
except queue.Empty:
return items
def test_active_watchers_grace_filters_by_thread():
"""Verifies _has_relevant_active_watchers ignores sibling-thread watchers
(otherwise consume_notifications grace period would block thread A by up
to 3s waiting for thread B's unrelated watchers to finish)."""
async_notifier._active_watchers.clear()
# Sentinel handles — only their identity matters here, not their type
handle_a = object()
handle_b = object()
handle_unrouted = object()
async_notifier._active_watchers[handle_a] = "threadA"
async_notifier._active_watchers[handle_b] = "threadB"
async_notifier._active_watchers[handle_unrouted] = None
# Current thread A → A's own watcher + unrouted are relevant
assert async_notifier._has_relevant_active_watchers("threadA") is True
# Current thread C (no active watcher of its own) → only unrouted matters
assert async_notifier._has_relevant_active_watchers("threadC") is True
# Drop the unrouted handle → C now has nothing relevant
del async_notifier._active_watchers[handle_unrouted]
assert async_notifier._has_relevant_active_watchers("threadC") is False
# A still has its own watcher
assert async_notifier._has_relevant_active_watchers("threadA") is True
# Legacy: None argument falls back to "any active watcher counts"
assert async_notifier._has_relevant_active_watchers(None) is True
async_notifier._active_watchers.clear()
assert async_notifier._has_relevant_active_watchers("threadA") is False
assert async_notifier._has_relevant_active_watchers(None) is False
+80
View File
@@ -171,3 +171,83 @@ def test_swap_uses_configured_port():
):
out = _maybe_swap_async_subagents(subs)
assert out[0]["url"] == "http://localhost:9999"
# =============================================================================
# AsyncWatcherMiddleware appended on swap
# =============================================================================
def test_maybe_swap_appends_watcher_middleware_when_enabled():
"""When async subagents are swapped, AsyncWatcherMiddleware must be appended."""
from EvoScientist.middleware.async_watcher import AsyncWatcherMiddleware
cfg = SimpleNamespace(enable_async_subagents=True, langgraph_dev_port=6174)
middleware: list = []
with (
patch("EvoScientist.EvoScientist._ensure_config", return_value=cfg),
patch(
"EvoScientist.langgraph_dev.manager.is_async_subagents_available",
return_value=True,
),
):
subs = [_sub("writing-agent", async_flag=True)]
_maybe_swap_async_subagents(subs, middleware)
assert len(middleware) == 1
assert isinstance(middleware[0], AsyncWatcherMiddleware)
def test_maybe_swap_skips_middleware_when_no_async_flagged():
"""Without any _async-flagged subagents, no middleware is appended."""
cfg = SimpleNamespace(enable_async_subagents=True, langgraph_dev_port=6174)
middleware: list = []
with (
patch("EvoScientist.EvoScientist._ensure_config", return_value=cfg),
patch(
"EvoScientist.langgraph_dev.manager.is_async_subagents_available",
return_value=True,
),
):
subs = [_sub("planner-agent", async_flag=False)]
_maybe_swap_async_subagents(subs, middleware)
assert middleware == []
def test_maybe_swap_skips_middleware_when_langgraph_unreachable():
"""Fallback path must not append middleware (no watchers will fire)."""
cfg = SimpleNamespace(enable_async_subagents=True, langgraph_dev_port=6174)
middleware: list = []
with (
patch("EvoScientist.EvoScientist._ensure_config", return_value=cfg),
patch(
"EvoScientist.langgraph_dev.manager.is_async_subagents_available",
return_value=False,
),
):
subs = [_sub("writing-agent", async_flag=True)]
_maybe_swap_async_subagents(subs, middleware)
assert middleware == []
def test_maybe_swap_no_middleware_arg_does_not_crash():
"""Backward-compat: middleware parameter is optional."""
cfg = SimpleNamespace(enable_async_subagents=True, langgraph_dev_port=6174)
with (
patch("EvoScientist.EvoScientist._ensure_config", return_value=cfg),
patch(
"EvoScientist.langgraph_dev.manager.is_async_subagents_available",
return_value=True,
),
):
subs = [_sub("writing-agent", async_flag=True)]
out = _maybe_swap_async_subagents(subs)
assert len(out) == 1
assert out[0]["graph_id"] == "writing-agent"
+467
View File
@@ -0,0 +1,467 @@
"""Tests for ``EvoScientist.middleware.async_watcher.AsyncWatcherMiddleware``.
The middleware is the public-API replacement for the old monkey-patch on
deepagents internals. It hooks into ``awrap_tool_call`` and only fires on
``start_async_task`` / ``update_async_task`` tool invocations.
"""
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from EvoScientist.cli import async_notifier
def _drain_all_notifications():
"""Drain every routed/unrouted/legacy queue between tests."""
if hasattr(async_notifier, "_notifications_by_thread"):
for q in list(async_notifier._notifications_by_thread.values()):
while not q.empty():
try:
q.get_nowait()
except Exception:
break
for attr in ("_unrouted_queue", "_notification_queue"):
if hasattr(async_notifier, attr):
q = getattr(async_notifier, attr)
while not q.empty():
try:
q.get_nowait()
except Exception:
break
@pytest.fixture(autouse=True)
def _clean_notifier_state():
"""Reset shared module-level notifier state before and after every test.
Cleared here:
- All notification queues (per-thread, unrouted, legacy global)
- ``_watcher_by_thread`` (replacement-on-update registry)
- ``_active_watchers`` (in-flight watcher → origin-thread index used
by ``consume_notifications`` to gate the grace-window wait)
Without this, a test that touches any of these dicts/queues would
silently leak state into the next test. ``_active_watchers`` is cleared
even though current tests patch ``spawn_watcher`` (so they never insert
into it) — kept as a safeguard for future tests that exercise the real
spawn path.
"""
_drain_all_notifications()
async_notifier._watcher_by_thread.clear()
async_notifier._active_watchers.clear()
yield
_drain_all_notifications()
async_notifier._watcher_by_thread.clear()
async_notifier._active_watchers.clear()
def _build_request(tool_name: str, args: dict, *, thread_id: str | None = None):
"""Construct a minimal ToolCallRequest stand-in.
The middleware reads only ``request.tool_call`` and ``request.runtime``.
"""
runtime = SimpleNamespace(
config={"configurable": {"thread_id": thread_id}} if thread_id else {}
)
return SimpleNamespace(
tool_call={
"name": tool_name,
"args": args,
"id": "call-1",
"type": "tool_call",
},
runtime=runtime,
state={},
tool=None,
)
def _make_middleware():
"""Build an AsyncWatcherMiddleware with a stubbed ``_ClientCache``."""
from EvoScientist.middleware.async_watcher import AsyncWatcherMiddleware
fake_client = MagicMock(name="LangGraphClient")
fake_cache = MagicMock(name="ClientCache")
fake_cache.get_async.return_value = fake_client
with patch(
"deepagents.middleware.async_subagents._ClientCache",
return_value=fake_cache,
):
mw = AsyncWatcherMiddleware(
{
"writing-agent": {
"name": "writing-agent",
"url": "http://x",
"graph_id": "writing-agent",
}
}
)
return mw, fake_client
def test_middleware_spawns_watcher_on_start_async_task():
"""A successful start_async_task tool call must spawn one watcher per task."""
from langgraph.types import Command
mw, _fake_client = _make_middleware()
spawn_calls = []
def fake_spawn(
client, thread_id, run_id, agent_name, prompt="", origin_cli_thread_id=None
):
spawn_calls.append(
(thread_id, run_id, agent_name, prompt, origin_cli_thread_id)
)
request = _build_request(
"start_async_task",
{"description": "do thing", "subagent_type": "writing-agent"},
thread_id="cli-thread-A",
)
async def fake_handler(req):
return Command(
update={
"async_tasks": {
"task-1": {
"task_id": "task-1",
"agent_name": "writing-agent",
"run_id": "run-1",
"thread_id": "task-1",
"status": "running",
}
}
}
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert isinstance(result, Command)
assert spawn_calls == [
("task-1", "run-1", "writing-agent", "do thing", "cli-thread-A")
]
def test_middleware_spawns_watcher_on_update_async_task():
"""A successful update_async_task call must also spawn a (replacement) watcher."""
from langgraph.types import Command
mw, _ = _make_middleware()
spawn_calls = []
def fake_spawn(*args, **kwargs):
spawn_calls.append((args, kwargs))
request = _build_request(
"update_async_task",
{"task_id": "task-1", "message": "do more"},
thread_id="cli-thread-A",
)
async def fake_handler(req):
return Command(
update={
"async_tasks": {
"task-1": {
"task_id": "task-1",
"agent_name": "writing-agent",
"run_id": "run-2",
"thread_id": "task-1",
"status": "running",
}
}
}
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert len(spawn_calls) == 1
args, kwargs = spawn_calls[0]
# spawn_watcher(client, task_id, run_id, agent_name, prompt=..., origin_cli_thread_id=...)
assert args[1] == "task-1"
assert args[2] == "run-2"
assert args[3] == "writing-agent"
assert kwargs["prompt"] == "do more"
assert kwargs["origin_cli_thread_id"] == "cli-thread-A"
def test_middleware_pre_cancels_old_watcher_on_update():
"""update_async_task must cancel the existing watcher BEFORE invoking the handler.
Otherwise the new run interrupts the old run's stream, which closes
cleanly, and the old watcher would enqueue a stale "success" notification.
"""
from langgraph.types import Command
mw, _ = _make_middleware()
old_watcher = MagicMock()
old_watcher.done.return_value = False
async_notifier._watcher_by_thread["task-1"] = old_watcher
cancel_observed_before_handler = {"value": False}
async def fake_handler(req):
cancel_observed_before_handler["value"] = old_watcher.cancel.called
return Command(update={"async_tasks": {}})
request = _build_request(
"update_async_task", {"task_id": "task-1", "message": "x"}, thread_id="t"
)
try:
with patch.object(async_notifier, "spawn_watcher"):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
finally:
async_notifier._watcher_by_thread.pop("task-1", None)
assert cancel_observed_before_handler["value"] is True
def test_middleware_passes_through_unrelated_tools():
"""A non-launch tool call must not spawn any watcher and must return result unchanged."""
mw, _ = _make_middleware()
sentinel = object()
async def fake_handler(req):
return sentinel
request = _build_request("ls", {"path": "/"}, thread_id="t")
with patch.object(async_notifier, "spawn_watcher") as mock_spawn:
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert result is sentinel
assert mock_spawn.call_count == 0
def test_middleware_handles_non_command_results_gracefully():
"""If the launch tool returns a string (validation error), no watcher is spawned."""
mw, _ = _make_middleware()
async def fake_handler(req):
return "Unknown async subagent type `bogus`"
request = _build_request(
"start_async_task",
{"description": "x", "subagent_type": "bogus"},
thread_id="t",
)
with patch.object(async_notifier, "spawn_watcher") as mock_spawn:
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert result == "Unknown async subagent type `bogus`"
assert mock_spawn.call_count == 0
def test_middleware_origin_thread_id_is_none_when_runtime_config_missing():
"""When runtime.config is empty, origin_cli_thread_id must be None (not crash)."""
from langgraph.types import Command
mw, _ = _make_middleware()
captured = {}
def fake_spawn(*args, **kwargs):
captured.update(kwargs)
request = _build_request(
"start_async_task",
{"description": "x", "subagent_type": "writing-agent"},
thread_id=None,
)
async def fake_handler(req):
return Command(
update={
"async_tasks": {
"t1": {
"task_id": "t1",
"agent_name": "writing-agent",
"run_id": "r1",
"thread_id": "t1",
"status": "running",
}
}
}
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert captured.get("origin_cli_thread_id") is None
def test_middleware_swallows_spawn_exceptions():
"""spawn_watcher errors must not propagate up — middleware logs and continues."""
from langgraph.types import Command
mw, _ = _make_middleware()
def boom(*a, **kw):
raise RuntimeError("intentional")
request = _build_request(
"start_async_task",
{"description": "x", "subagent_type": "writing-agent"},
thread_id="t",
)
async def fake_handler(req):
return Command(
update={
"async_tasks": {
"t1": {
"task_id": "t1",
"agent_name": "writing-agent",
"run_id": "r1",
"thread_id": "t1",
"status": "running",
}
}
}
)
with patch.object(async_notifier, "spawn_watcher", side_effect=boom):
# Should not raise.
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert isinstance(result, Command)
@pytest.mark.parametrize(
("tool_name", "args", "prompt_field"),
[
(
"start_async_task",
{"description": "from start", "subagent_type": "writing-agent"},
"from start",
),
(
"update_async_task",
{"task_id": "t1", "message": "from update"},
"from update",
),
],
)
def test_middleware_picks_correct_prompt_field_per_tool(tool_name, args, prompt_field):
"""start_async_task uses 'description'; update_async_task uses 'message'."""
from langgraph.types import Command
mw, _ = _make_middleware()
captured_prompt = {}
def fake_spawn(*a, prompt="", **kw):
captured_prompt["value"] = prompt
async def fake_handler(req):
return Command(
update={
"async_tasks": {
"t1": {
"task_id": "t1",
"agent_name": "writing-agent",
"run_id": "r1",
"thread_id": "t1",
"status": "running",
}
}
}
)
request = _build_request(tool_name, args, thread_id="t")
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert captured_prompt["value"] == prompt_field
def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained():
"""update_async_task with extra `description` arg must still use `message`.
Guards against the previous `args.get('description') or args.get('message')`
chained-fallback shape, which would have picked `description` for an update
call that happened to carry both fields.
"""
from langgraph.types import Command
mw, _ = _make_middleware()
captured_prompt = {}
def fake_spawn(*a, prompt="", **kw):
captured_prompt["value"] = prompt
async def fake_handler(req):
return Command(
update={
"async_tasks": {
"t1": {
"task_id": "t1",
"agent_name": "writing-agent",
"run_id": "r1",
"thread_id": "t1",
"status": "running",
}
}
}
)
request = _build_request(
"update_async_task",
{
"task_id": "t1",
"message": "use this",
"description": "do NOT use this",
},
thread_id="t",
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert captured_prompt["value"] == "use this"
def test_middleware_pre_cancel_swallows_unexpected_errors():
"""A faulty old-watcher handle must not block the handler from running."""
from langgraph.types import Command
mw, _ = _make_middleware()
bad_watcher = MagicMock()
bad_watcher.done.side_effect = RuntimeError("watcher state corrupted")
async_notifier._watcher_by_thread["t1"] = bad_watcher
handler_called = {"value": False}
async def fake_handler(req):
handler_called["value"] = True
return Command(update={"async_tasks": {}})
request = _build_request(
"update_async_task", {"task_id": "t1", "message": "x"}, thread_id="t"
)
try:
with patch.object(async_notifier, "spawn_watcher"):
# Should not raise.
asyncio.run(mw.awrap_tool_call(request, fake_handler))
finally:
async_notifier._watcher_by_thread.pop("t1", None)
assert handler_called["value"] is True