diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index fece8b1..ebce27c 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -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: `` 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", diff --git a/EvoScientist/cli/async_notifier.py b/EvoScientist/cli/async_notifier.py new file mode 100644 index 0000000..16cec23 --- /dev/null +++ b/EvoScientist/cli/async_notifier.py @@ -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 "", + ) + + +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: 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) diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index eb9ee62..20ac814 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -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: diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index 3ccc9be..dc293db 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -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()) diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 158fe97..2e7b390 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -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 diff --git a/EvoScientist/middleware/async_watcher.py b/EvoScientist/middleware/async_watcher.py new file mode 100644 index 0000000..bb18871 --- /dev/null +++ b/EvoScientist/middleware/async_watcher.py @@ -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 diff --git a/EvoScientist/prompts.py b/EvoScientist/prompts.py index 9b0de90..52c2257 100644 --- a/EvoScientist/prompts.py +++ b/EvoScientist/prompts.py @@ -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) diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index 1c93971..97c9f56 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -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) diff --git a/tests/conftest.py b/tests/conftest.py index c1f7c20..495d8a7 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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.""" diff --git a/tests/test_async_notifier.py b/tests/test_async_notifier.py new file mode 100644 index 0000000..c0f194e --- /dev/null +++ b/tests/test_async_notifier.py @@ -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: `.""" + 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 diff --git a/tests/test_async_subagent_swap.py b/tests/test_async_subagent_swap.py index 382fc77..88d36ff 100644 --- a/tests/test_async_subagent_swap.py +++ b/tests/test_async_subagent_swap.py @@ -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" diff --git a/tests/test_async_watcher_middleware.py b/tests/test_async_watcher_middleware.py new file mode 100644 index 0000000..ae81e9d --- /dev/null +++ b/tests/test_async_watcher_middleware.py @@ -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