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:
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user