refactor: route middleware display events through an injected event sink (#343)

* chore: add pytest-asyncio in auto mode

* test: migrate channel and stream tests to native async

Convert run_async() wrapper tests to plain 'async def test_*' under
pytest-asyncio auto mode. collect_events() in stream_v3_fakes becomes a
coroutine awaited at every call site.

* test: migrate command and model/middleware tests to native async

Convert run_async() wrappers (import, alias, and fixture forms) to plain
'async def test_*'. Multi-call tests merge onto one loop as sequential
awaits; none asserted on loop identity.

* test: migrate TUI, notifier, gateway, and session tests to native async

TUI/notifier/gateway files convert run_async wrappers to plain async
tests. test_sessions.py's unittest.TestCase classes move to
unittest.IsolatedAsyncioTestCase (pytest-asyncio does not await async
methods on plain TestCase; converting blindly would have made ~70 tests
silently vacuous). Its setUpClass keeps a one-shot asyncio.run() since
IsolatedAsyncioTestCase has no async class-level hook. TestLoadingWidget
in test_tui_widgets.py drops its TestCase base for the same reason.

* test: replace direct asyncio.run() calls with native async tests

Convert tests that called asyncio.run() (directly or via a local _run
helper) to plain 'async def test_*'; delete the local helpers.

* test: drop undeclared anyio markers and delete run_async helper

The @pytest.mark.anyio tests relied on anyio being a transitive dep of
httpx; auto-mode pytest-asyncio collects them natively. run_async() and
its fixture are unreferenced after the migration, so remove them —
pytest-asyncio's per-test loop teardown covers the pending-task
cancellation the helper existed for (verified: full suite runs with no
'Event loop is closed' errors or destroyed-task warnings).

* test: add autouse fixture for watcher cleanup

* refactor: remove redundant hasattr calls

* refactor: add typed middleware event sink and thread through assembly

Add MiddlewareEventSink protocol + NoOpSink in middleware/events.py
with a documented any-thread non-blocking contract (contract test uses a
deliberately-slow fake sink). Thread an optional `events` parameter
through create_cli_agent -> _get_default_middleware -> tool selector /
model fallback constructors; subagent stacks are always forced to
NoOpSink.

* refactor: inject a notifier port into async-watcher and background middleware

Add public pre_cancel_watcher() and enqueue_task_notification() to
cli/async_notifier.py and a small NotifierPort protocol
(middleware/notifier.py) that the module satisfies structurally.
AsyncWatcherMiddleware and BackgroundExecutionMiddleware now receive the
port by constructor injection at the composition root, deleting the lazy
'from ..cli import async_notifier' imports and the private
_watcher_by_thread / _enqueue pokes.

* refactor: invert tool-selection ownership onto a frontend event sink

The adaptive tool selector now reports on_tool_selection_started /
on_tool_selection / on_tool_selection_ended to the injected sink instead
of writing four process-global module variables. The frontend sink
(stream/sink.py FrontendEventSink) owns the selected/total/active state
with consume-once + dedup-vs-last-emitted semantics;
stream/tool_selection.py reads that sink object (a ToolSelectionView)
rather than reaching into tool_selector's globals.

Deleted: the 4 module globals, the cross-module mutations in
tool_selection.py, the track_stream_selection flag, the now-vestigial
_ToolSelectionTrackerMiddleware, reset_tool_selection_state_for_tests,
and the autouse conftest fixture. The sink is threaded from the two
interactive frontends through create_runtime_gateways ->
LocalGraphGateway (read side) and _load_agent -> create_cli_agent (write
side); subagent / headless stacks get NoOpSink.

* refactor: route model-fallback narration through the injected event sink

Delete the _ui_emit_fn / set_ui_emit module global and the
..stream.console import from model_fallback.py. The fallback middleware
now reports through its injected sink: the fallback transition via the
structured on_model_fallback (the frontend formats the '-> Falling back
to ...' line), and the surrounding narration (primary-failure header,
per-attempt outcome, exhaustion, non-fallbackable rejection) via
emit_fallback_notice, preserving the exact user-facing text. The TUI
binds its _append_system as the sink's fallback display where it used to
call set_ui_emit (cleared on exit); the Rich CLI's sink prints to the
console. _try_fallbacks / _guard_and_fallback take the sink.

* refactor: declare events on the GraphGateway protocol

Both gateway implementations now carry an explicit events attribute
(LangGraphServerGateway holds None — no frontend renders middleware
events across the HTTP boundary), so the four call sites use plain
attribute access instead of getattr probing an implicit contract.

* refactor: bind fallback display via the closure-scoped concrete sink

The App methods used gateway.events (typed as the read-side view) and
hasattr-probed for the concrete FrontendEventSink API. The enclosing
factory creates that sink two hundred lines up — close over it directly:
no probing, fully typed, and it becomes a constructor parameter
naturally when the App class is hoisted out of the factory.

* fix: end tool selection before fallback handler

* fix: keep fallback display errors non-fatal

* fix: preserve selector suppression for default streams

* fix: restore fallback notice console display

* refactor: consolidate fallback narration events

* refactor: clean middleware event sink plumbing

* fix: type gateway session events

* refactor: make all event protocols runtime-checkable

MiddlewareEventSink already carried @runtime_checkable (the stream
binding guard isinstance-checks it); ToolSelectionView and SessionEvents
now match, so mirroring that pattern against any of the three protocols
works instead of raising TypeError.

* fix(cli): close QuickJS workers after one-shot failures

* fix(cli): honor no-thinking in final output

* fix(channels): report failed startup accurately

* fix(channels): make Telegram cleanup idempotent

* fix(tui): skip command sync during exit

* fix(channels): preserve startup state during retries

* refactor(channels): share pending startup status

* refactor(cli): expose channel startup snapshot

* fix(tui): move channel startup off event loop

* test(channels): release retry gate on assertion failure

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
This commit is contained in:
dinos
2026-07-15 00:34:17 +02:00
committed by GitHub
parent db1abce8d8
commit 01845f4311
44 changed files with 1966 additions and 518 deletions
+30 -5
View File
@@ -41,6 +41,8 @@ logging.getLogger("deepagents.middleware.skills").setLevel(logging.ERROR)
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
from .middleware.events import MiddlewareEventSink
# =============================================================================
# Constants
# =============================================================================
@@ -465,9 +467,12 @@ def _maybe_swap_async_subagents(
out.append(s)
if agent_specs and middleware is not None:
from .cli import async_notifier
from .middleware.async_watcher import AsyncWatcherMiddleware
middleware.append(AsyncWatcherMiddleware(agent_specs))
# Composition root wires the concrete notifier port into the middleware;
# the middleware itself never imports the CLI layer.
middleware.append(AsyncWatcherMiddleware(agent_specs, notifier=async_notifier))
# Forward the CLI's live (model, provider) into deepagents'
# start/update_async_task tool calls so the deployed graph can
@@ -650,6 +655,7 @@ def _get_default_middleware(
cfg=None,
chat_model=None,
memory_source_agent: str = "EvoScientist",
events: "MiddlewareEventSink | None" = None,
):
"""Build the default middleware list.
@@ -669,6 +675,11 @@ def _get_default_middleware(
(avoids writing module globals on the pure path).
memory_source_agent: Attribution name for profile/observation writes.
Async sub-agent factories pass their deployed agent name here.
events: Frontend/session-supplied event sink. Middleware report
tool-selection events and model-fallback notices to it.
Defaults to the current stream run's sink for main agents; async
sub-agent stacks are always forced to ``NoOpSink`` (they must not
drive the main-agent widgets).
"""
from .middleware import (
ConfigurableModelMiddleware,
@@ -686,6 +697,13 @@ def _get_default_middleware(
default_memory_scheduler,
load_fallback_chain,
)
from .middleware.events import NO_OP_SINK, RunScopedEventSink
# Subagent stacks never drive the main-agent frontend widgets; force the
# no-op sink there regardless of what the caller passed. Main stacks built
# without an explicit frontend/session sink report into the active stream
# run's sink, preserving selector suppression for headless local runs.
events = NO_OP_SINK if for_async_subagent else (events or RunScopedEventSink())
cfg = cfg if cfg is not None else _ensure_config()
if cfg.model_fallbacks:
@@ -743,12 +761,12 @@ def _get_default_middleware(
ErrorNormalizationMiddleware(),
ConfigurableModelMiddleware(),
create_context_editing_middleware(model),
ModelFallbackMiddleware(),
ModelFallbackMiddleware(events=events),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*create_tool_selector_middleware(
model=tool_selector_model,
track_stream_selection=not for_async_subagent,
events=events,
),
# Interpreter prompt must land before runtime/memory context, so this
# middleware sits ahead of runtime_context in the stack.
@@ -783,9 +801,15 @@ def _get_default_middleware(
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
# must not spawn local OS processes.
if not for_async_subagent:
from .cli import async_notifier
from .middleware.background import BackgroundExecutionMiddleware
mw.append(BackgroundExecutionMiddleware())
# Inject the notifier port + the assembly-time dangerous-mode policy
# (agents rebuild on config change, so the captured flag never staler
# than the agent it lives on).
mw.append(
BackgroundExecutionMiddleware(async_notifier, dangerous=cfg.dangerous_mode)
)
return mw
@@ -880,6 +904,7 @@ def create_cli_agent(
chat_model=None,
*,
on_mcp_progress=None,
events: "MiddlewareEventSink | None" = None,
) -> "CompiledStateGraph":
"""Create agent with checkpointer for CLI multi-turn support.
@@ -981,7 +1006,7 @@ def create_cli_agent(
# CLI agent never drifts from the default chain. Anything CLI-specific
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
mw: list[AgentMiddleware] = _get_default_middleware(
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model, events=events
)
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
+16
View File
@@ -7,6 +7,7 @@ This module defines the Channel interface that all messaging channels
import asyncio
import logging
import re
import threading
from abc import ABC, abstractmethod
from collections import OrderedDict
from collections.abc import AsyncIterator, Awaitable, Callable
@@ -298,6 +299,8 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
maxsize=queue_maxsize
)
self._running = False
self._startup_event = threading.Event()
self._startup_error: str | None = None
# Global tracing can be enabled via shared config/env even when
# individual channel factories have not been updated yet.
@@ -1167,16 +1170,25 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
"""Run the channel with auto-reconnect (exponential backoff)."""
backoff = 1.0
max_backoff = 60.0
self._startup_event.clear()
self._startup_error = None
self._running = True
while self._running:
try:
await self.start()
self._startup_error = None
self._startup_event.set()
backoff = 1.0
async for msg in self.receive():
await self.queue_message(msg)
except asyncio.CancelledError:
if not self._startup_event.is_set():
self._startup_error = "startup cancelled"
self._startup_event.set()
break
except ChannelError as e:
self._startup_error = str(e)
self._startup_event.set()
self._trace_event(
"channel_fatal_error",
error_type=type(e).__name__,
@@ -1204,6 +1216,10 @@ class Channel(TraceMixin, ChannelPlugin, ABC):
await asyncio.sleep(backoff)
backoff = min(backoff * 2, max_backoff)
if not self._startup_event.is_set():
self._startup_error = "channel stopped before startup completed"
self._startup_event.set()
# ── Channel allow-list check ─────────────────────────────────────
def is_channel_allowed(self, channel_id: str) -> bool:
+27
View File
@@ -29,6 +29,8 @@ from .plugin import ChannelPlugin
logger = logging.getLogger(__name__)
CHANNEL_STARTUP_PENDING_DETAIL = "starting (bus)"
# ═════════════════════════════════════════════════════════════════════
# Account management (formerly account.py)
@@ -978,6 +980,31 @@ class ChannelManager:
"""Return names of currently running channels."""
return [name for name, ch in self._channels.items() if ch._running]
def startup_results(self, *, timeout: float = 0.0) -> list[tuple[str, bool, str]]:
"""Return each channel's initial connection result.
The optional timeout is shared across all channels, which start
concurrently. Channels still connecting when it expires are reported
as starting rather than connected.
"""
deadline = time.monotonic() + max(timeout, 0.0)
for channel in self._channels.values():
remaining = deadline - time.monotonic()
if remaining > 0 and not channel._startup_event.is_set():
channel._startup_event.wait(remaining)
results: list[tuple[str, bool, str]] = []
for name, channel in self._channels.items():
if not channel._startup_event.is_set():
results.append((name, False, CHANNEL_STARTUP_PENDING_DETAIL))
elif channel._startup_error:
results.append((name, False, f"failed: {channel._startup_error}"))
elif channel._running:
results.append((name, True, "connected (bus)"))
else:
results.append((name, False, "stopped during startup"))
return results
def get_stats(self) -> dict:
"""Return summary stats for all channels."""
return {
+11 -6
View File
@@ -94,12 +94,17 @@ class TelegramChannel(Channel):
logger.info("Telegram channel started (polling)")
async def _cleanup(self) -> None:
if self._app:
if self._app.updater and self._app.updater.running:
await self._app.updater.stop()
await self._app.stop()
await self._app.shutdown()
logger.info("Telegram channel stopped")
app = self._app
self._app = None
if app is None:
return
if app.updater and app.updater.running:
await app.updater.stop()
if app.running:
await app.stop()
await app.shutdown()
logger.info("Telegram channel stopped")
# ── Typing indicator (override base) ────────────────────────────
+2
View File
@@ -69,6 +69,7 @@ def _load_agent(
chat_model=None,
*,
on_mcp_progress=None,
events=None,
) -> "CompiledStateGraph":
"""Load the CLI agent with optional persistent checkpointer.
@@ -92,4 +93,5 @@ def _load_agent(
config=config,
chat_model=chat_model,
on_mcp_progress=on_mcp_progress,
events=events,
)
+65
View File
@@ -117,6 +117,71 @@ def _enqueue(notification: AsyncTaskNotification) -> None:
q.put(notification)
def enqueue_task_notification(notification: AsyncTaskNotification) -> None:
"""Public :class:`~EvoScientist.middleware.notifier.NotifierPort` entry point.
Route a completed-task notification onto the consumer queue. Thin wrapper
over :func:`_enqueue` so middleware can enqueue without reaching into the
module's private symbols.
"""
_enqueue(notification)
def enqueue_bg_process_notification(
*,
task_id: str,
agent_name: str,
status: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> None:
"""Build and enqueue a background-process completion notification.
:class:`~EvoScientist.middleware.notifier.NotifierPort` entry point used by
the background middleware so it never constructs the CLI-owned
:class:`AsyncTaskNotification` itself — the ``kind="bg-process"`` tag and the
UTC ``received_at`` timestamp are filled in here.
"""
_enqueue(
AsyncTaskNotification(
task_id=task_id,
agent_name=agent_name,
status=status,
received_at=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
prompt=prompt,
kind="bg-process",
origin_cli_thread_id=origin_cli_thread_id,
)
)
def pre_cancel_watcher(task_id: str) -> None:
"""Cancel a stale watcher for ``task_id`` before a new run replaces it.
``update_async_task`` starts 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 that clean
exit and enqueue a stale "success" notification before the new spawn can
replace it. Cancellation propagates ``CancelledError`` (a ``BaseException``)
which the watcher's ``except Exception:`` does not catch, so ``_enqueue``
never runs for the cancelled watcher.
No-op when there is no live watcher; swallows any error (a failed
pre-cancel only risks one stale notification, never a crashed tool call).
"""
try:
old = _watcher_by_thread.get(task_id)
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",
task_id,
exc_info=True,
)
def has_pending_notifications(current_thread_id: str | None = None) -> bool:
"""Cheap predicate for poller idle paths — true iff there's anything to consume.
+25 -10
View File
@@ -824,6 +824,11 @@ _bus_loop: asyncio.AbstractEventLoop | None = None
_bus_thread: threading.Thread | None = None
def get_channel_startup_results() -> list[tuple[str, bool, str]]:
"""Return the current channel startup snapshot without waiting."""
return _manager.startup_results() if _manager is not None else []
def _channels_is_running(channel_type: str | None = None) -> bool:
"""Check whether channels are running."""
if _manager is None:
@@ -854,7 +859,7 @@ def _channels_stop(
if channel_type is None:
# Stop everything
if _bus_loop and _manager:
if _bus_loop and _manager and not _bus_loop.is_closed():
try:
future = asyncio.run_coroutine_threadsafe(
_manager.stop_all(),
@@ -893,7 +898,7 @@ def _start_channels_bus_mode(
thread_id: str,
*,
send_thinking: bool | None = None,
) -> None:
) -> list[tuple[str, bool, str]]:
"""Start all channels in bus mode with MessageBus + ChannelManager.
Creates a single event loop in a daemon thread running the bus,
@@ -926,6 +931,10 @@ def _start_channels_bus_mode(
try:
await mgr.start_all()
finally:
# ``start_all`` returns when all channel tasks terminate. This
# includes immediate fatal startup failures, so tear down the
# dispatcher and health server before closing the bus loop.
await mgr.stop_all()
consumer.cancel()
try:
await consumer
@@ -952,6 +961,8 @@ def _start_channels_bus_mode(
break
time.sleep(0.1)
return mgr.startup_results(timeout=2.0)
def _add_channel_to_running_bus(
channel_type: str,
@@ -1208,7 +1219,7 @@ def _auto_start_channel(
*,
send_thinking: bool | None = None,
runtime: ChannelRuntime | None = None,
) -> None:
) -> list[tuple[str, bool, str]]:
"""Start channels automatically from config (bus mode).
Args:
@@ -1220,18 +1231,22 @@ def _auto_start_channel(
is accepted for callers that don't yet pass one.
"""
if not config.channel_enabled:
return
return []
_start_channels_bus_mode(
results = _start_channels_bus_mode(
config,
agent,
thread_id,
send_thinking=send_thinking,
)
# Bind only after startup succeeds; a failure above must not leave
# a stale runtime binding pointing at channels that never started.
if runtime is not None:
# A channel that is still starting may connect later and needs the runtime
# binding. Immediate failures must not leave a stale binding behind.
from ..channels.channel_manager import CHANNEL_STARTUP_PENDING_DETAIL
has_active_channel = any(
ok or detail == CHANNEL_STARTUP_PENDING_DETAIL for _, ok, detail in results
)
if runtime is not None and has_active_channel:
runtime.bind(agent, thread_id)
types = [t.strip() for t in config.channel_enabled.split(",") if t.strip()]
results = [(ct, True, "connected (bus)") for ct in types]
_print_channel_panel(results)
return results
+9
View File
@@ -2393,6 +2393,15 @@ def _main_callback(
runtime_gateways=runtime_gateways,
)
finally:
# Model failures can bypass middleware ``after_agent``
# hooks. Close any remaining QuickJS workers while this
# event loop is still available; their synchronous GC
# fallback can deadlock during interpreter shutdown.
from ..middleware.code_interpreter import (
aclose_code_interpreters,
)
await aclose_code_interpreters()
try:
print_resume_hint(tid, console=console)
except Exception:
+19 -2
View File
@@ -448,7 +448,17 @@ def cmd_interactive(
on_progress=_on_mcp_progress,
)
runtime_gateways = create_runtime_gateways()
# One frontend event sink for the whole session — injected into the agent's
# middleware (write side) and the local gateway's streaming path (read side)
# so both share one owner. It survives agent rebuilds (/model, /new, MCP
# reload) because the session, not the agent, holds it.
from ..stream.sink import SessionEventSink
event_sink = SessionEventSink(
fallback_display=lambda text, style: console.print(text, style=style)
)
runtime_gateways = create_runtime_gateways(events=event_sink)
graph_gateway = runtime_gateways.graph_gateway
requested_thread_id = thread_id
@@ -486,6 +496,7 @@ def cmd_interactive(
workspace_dir=state["workspace_dir"],
checkpointer=checkpointer,
config=config,
events=event_sink,
)
async def _await_agent_ready() -> "CompiledStateGraph":
@@ -1527,7 +1538,13 @@ def cmd_run(
raise typer.Exit(1) from e
else:
console.print(f"[red]Error: {e}[/red]")
raise
# This is the process boundary for single-shot text mode. Letting
# provider exceptions escape makes Typer/Rich render the complete
# async exception chain after we already printed a concise error;
# large OpenAI/httpx chains can keep the CLI busy well after the
# resume hint is shown. Convert the failure to Click's controlled
# exit signal while preserving the cause for programmatic callers.
raise typer.Exit(1) from e
def _wait_for_memory_workers_before_exit(
+103 -28
View File
@@ -11,6 +11,7 @@ import logging
import queue
import random
import sys
import threading
from collections.abc import Callable
from dataclasses import dataclass
from datetime import datetime
@@ -53,11 +54,11 @@ from .channel import (
ChannelMessage,
_auto_start_channel,
_channels_is_running,
_channels_running_list,
_channels_stop,
_message_queue,
_set_channel_response,
dispatch_channel_slash_command,
get_channel_startup_results,
)
from .file_mentions import complete_file_mention, resolve_file_mentions
from .history_suggester import HistorySuggester
@@ -92,6 +93,45 @@ def _shorten_path(path: str) -> str:
return _sp(path)
async def _auto_start_channel_in_worker(
agent: Any,
thread_id: str,
config: Any,
*,
send_thinking: bool,
runtime: Any,
stop_requested: threading.Event,
) -> list[tuple[str, bool, str]]:
"""Run blocking channel startup without occupying the TUI event loop."""
def _start() -> list[tuple[str, bool, str]]:
try:
return _auto_start_channel(
agent,
thread_id,
config,
send_thinking=send_thinking,
runtime=runtime,
)
finally:
if stop_requested.is_set():
_channels_stop(runtime=runtime)
worker = asyncio.create_task(asyncio.to_thread(_start))
try:
return await asyncio.shield(worker)
except asyncio.CancelledError:
stop_requested.set()
try:
await worker
except Exception:
_channel_logger.debug(
"Channel startup worker failed during cancellation",
exc_info=True,
)
raise
def _build_welcome_banner(
*,
thread_id: str,
@@ -220,6 +260,9 @@ async def _sync_tui_command_completion(
cmd: Command,
) -> None:
"""Adopt successful command-side state changes back into the TUI app."""
if app._exiting:
return
agent_swapped = ctx.agent is not None and ctx.agent is not original_agent
if agent_swapped:
from ..EvoScientist import _ensure_config
@@ -294,7 +337,15 @@ def run_textual_interactive(
config = get_effective_config()
runtime_gateways = create_runtime_gateways()
# One frontend event sink for the whole TUI session — injected into the
# agent's middleware (write side) and the local gateway's streaming path
# (read side). The fallback-notice display is bound to the App's
# _append_system once the App exists (on_mount); tool-selection needs no
# display hook (its widget is mounted from the stream event).
from ..stream.sink import SessionEventSink
event_sink = SessionEventSink()
runtime_gateways = create_runtime_gateways(events=event_sink)
graph_gateway = runtime_gateways.graph_gateway
try:
@@ -438,7 +489,8 @@ def run_textual_interactive(
self._resumed = resumed
self._resume_warning = resume_warning
self._channel_timer: Any = None
self._started_channel_types: list[str] = []
self._channel_start_results: list[tuple[str, bool, str]] = []
self._channel_start_stop = threading.Event()
self._busy = False
self._notification_consuming: bool = (
False # prevent overlapping consume coroutines
@@ -465,6 +517,7 @@ def run_textual_interactive(
self._channel_runtime = ChannelRuntime()
self._quit_pending: bool = False
self._exiting: bool = False
self._current_model: str | None = model
self._current_provider: str | None = provider
self._status_started_at = datetime.now()
@@ -516,6 +569,7 @@ def run_textual_interactive(
self._agent_loader.start(
workspace_dir=workspace,
checkpointer=self._checkpointer,
events=self._runtime_gateways.graph_gateway.events,
)
def _mount_mcp_loader_widget(self) -> None:
@@ -794,11 +848,13 @@ def run_textual_interactive(
yield Static("", id="status")
def on_mount(self) -> None:
# Register fallback middleware UI callback so messages appear
# as SystemMessage widgets in the chat container.
from ..middleware.model_fallback import set_ui_emit
set_ui_emit(lambda text, style: self._append_system(text, style))
# Bind the session sink's fallback-notice display so model-fallback
# messages appear as SystemMessage widgets in the chat container.
# ``event_sink`` is the concrete SessionEventSink created by the
# enclosing factory — the same instance the gateway carries.
event_sink.set_fallback_display(
lambda text, style: self._append_system(text, style)
)
self._render_welcome()
self._render_status()
@@ -850,7 +906,7 @@ def run_textual_interactive(
exc_info=True,
)
return
self._start_channels()
await self._start_channels()
ch_task = asyncio.create_task(_deferred_start_channels())
self._background_tasks.add(ch_task)
@@ -877,28 +933,46 @@ def run_textual_interactive(
# ── Channel integration ────────────────────────────────
def _start_channels(self) -> None:
async def _start_channels(self) -> None:
"""Auto-start channels if enabled in config."""
try:
from ..config import load_config
cfg = load_config()
cfg = await asyncio.to_thread(load_config)
if cfg and cfg.channel_enabled and not _channels_is_running():
_auto_start_channel(
results = await _auto_start_channel_in_worker(
self._agent_loader.agent,
self._conversation_tid,
cfg,
send_thinking=self._channel_send_thinking,
runtime=self._channel_runtime,
stop_requested=self._channel_start_stop,
)
types = [
t.strip() for t in cfg.channel_enabled.split(",") if t.strip()
]
self._started_channel_types = types
if self._exiting:
return
current_agent = self._agent_loader.agent
if current_agent is not None and _channels_is_running():
self._channel_runtime.bind(
current_agent,
self._conversation_tid,
)
self._channel_start_results = results
self._render_welcome()
except asyncio.CancelledError:
self._channel_start_stop.set()
raise
except Exception as e:
_channel_logger.debug(f"Channel auto-start failed: {e}")
self._channel_timer = self.set_interval(0.1, self._poll_channel_queue)
finally:
if (
not self._exiting
and not self._channel_start_stop.is_set()
and self._channel_timer is None
):
self._channel_timer = self.set_interval(
0.1,
self._poll_channel_queue,
)
def _poll_channel_queue(self) -> None:
"""Poll the channel + notification queues (every 100ms)."""
@@ -2864,8 +2938,9 @@ def run_textual_interactive(
self._render_status()
finally:
self._busy = False
prompt_widget.disabled = False
prompt_widget.focus()
if not self._exiting:
prompt_widget.disabled = False
prompt_widget.focus()
async def _render_history(self, thread_id_value: str) -> None:
"""Render conversation history from a saved thread.
@@ -2970,13 +3045,13 @@ def run_textual_interactive(
def _do_exit(self) -> None:
"""Clean up channels, unregister callbacks, and exit."""
from ..middleware.model_fallback import set_ui_emit
set_ui_emit(None)
self._exiting = True
self._channel_start_stop.set()
event_sink.set_fallback_display(None)
if self._channel_timer is not None:
self._channel_timer.stop()
self._channel_timer = None
self._started_channel_types.clear()
self._channel_start_results.clear()
if _channels_is_running():
try:
_channels_stop(runtime=self._channel_runtime)
@@ -3105,11 +3180,11 @@ def run_textual_interactive(
def _render_welcome(self) -> None:
channels_info: list[tuple[str, bool, str]] | None = None
try:
running = _channels_running_list()
started = self._started_channel_types
if running or started:
all_types = list(dict.fromkeys(running + started))
channels_info = [(ct, True, "connected (bus)") for ct in all_types]
current = get_channel_startup_results()
if current:
self._channel_start_results = current
if self._channel_start_results:
channels_info = self._channel_start_results
else:
from ..config import load_config
@@ -151,6 +151,11 @@ class ModelCommand(Command):
temp_cfg.model = model_name
temp_cfg.provider = provider
# Re-thread the session's frontend event sink so the rebuilt agent's
# middleware keeps driving the tool-selection widget / fallback notices
# after a /model switch (the sink lives on the gateway, not the agent).
events = ctx.graph_gateway.events
try:
new_chat_model = _build_chat_model(temp_cfg)
new_agent = _load_agent(
@@ -158,6 +163,7 @@ class ModelCommand(Command):
checkpointer=ctx.checkpointer,
config=temp_cfg,
chat_model=new_chat_model,
events=events,
)
except Exception as e:
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
+11 -1
View File
@@ -19,6 +19,8 @@ from .types import (
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
from ..middleware.events import SessionEvents
@dataclass(frozen=True, slots=True)
class LocalThreadStore:
@@ -61,9 +63,16 @@ class LocalThreadStore:
@dataclass(slots=True)
class LocalGraphGateway:
"""Gateway backed by the current in-process graph and session helpers."""
"""Gateway backed by the current in-process graph and session helpers.
``events`` is the frontend/session event sink for this runtime — normally
the same instance injected into the agent's middleware. If it is ``None``,
``stream_agent_events`` creates a per-run session sink and binds it for
default main-agent middleware via ``RunScopedEventSink``.
"""
thread_store: ThreadStore = field(default_factory=LocalThreadStore)
events: SessionEvents | None = None
async def create_thread(
self,
@@ -155,6 +164,7 @@ class LocalGraphGateway:
request.thread_id,
metadata=request.metadata,
media=request.media,
events=self.events,
)
try:
async for event in inner:
+12 -3
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal
from typing import TYPE_CHECKING, Literal
from langgraph_sdk import get_client
from langgraph_sdk.client import LangGraphClient
@@ -14,6 +14,9 @@ from .server import (
)
from .types import GraphGateway, ThreadStore
if TYPE_CHECKING:
from ..middleware.events import SessionEvents
RuntimeGatewayBackend = Literal["local", "langgraph_server"]
@@ -32,8 +35,14 @@ def create_runtime_gateways(
graph_id: str = DEFAULT_GRAPH_ID,
headers: dict[str, str] | None = None,
langgraph_client: LangGraphClient | None = None,
events: SessionEvents | None = None,
) -> RuntimeGateways:
"""Create gateway handles for CLI/TUI/serve execution."""
"""Create gateway handles for CLI/TUI/serve execution.
``events`` is the frontend event sink; it is attached to the local gateway
so the streaming path shares the same sink instance the frontend injects
into the agent's middleware. Server backends ignore it (headless).
"""
if backend == "langgraph_server":
if base_url is None and langgraph_client is None:
raise ValueError("base_url is required for langgraph_server gateways")
@@ -59,5 +68,5 @@ def create_runtime_gateways(
return RuntimeGateways(
thread_store=local_thread_store,
graph_gateway=LocalGraphGateway(thread_store=local_thread_store),
graph_gateway=LocalGraphGateway(thread_store=local_thread_store, events=events),
)
+5 -1
View File
@@ -7,7 +7,10 @@ import uuid
from collections.abc import AsyncIterator, Mapping
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import Any
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from ..middleware.events import SessionEvents
from langchain_core.messages import BaseMessage, convert_to_messages, messages_from_dict
from langgraph.types import Command
@@ -448,6 +451,7 @@ class LangGraphServerGateway:
thread_store: LangGraphServerThreadStore
graph_id: str = DEFAULT_GRAPH_ID
interrupt_wait_seconds: float = 5.0
events: SessionEvents | None = None
def _target_graph_id(self, target: GraphTarget | None = None) -> str:
return target.graph_id if target is not None else self.graph_id
+4
View File
@@ -11,6 +11,8 @@ from langgraph.types import Command
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
from ..middleware.events import SessionEvents
GraphEvent: TypeAlias = dict[str, Any]
GraphRunInput: TypeAlias = str | Command
GraphStateValues: TypeAlias = dict[str, Any]
@@ -94,6 +96,8 @@ class ThreadStore(Protocol):
class GraphGateway(Protocol):
"""One authority for graph runs and thread lifecycle operations."""
events: SessionEvents | None
async def create_thread(
self,
target: GraphTarget | None = None,
+12 -21
View File
@@ -22,13 +22,16 @@ from __future__ import annotations
import logging
from collections.abc import Awaitable, Callable
from typing import Any
from typing import TYPE_CHECKING, 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
if TYPE_CHECKING:
from .notifier import NotifierPort
logger = logging.getLogger(__name__)
_LAUNCH_TOOL_NAMES = ("start_async_task", "update_async_task")
@@ -42,42 +45,30 @@ class AsyncWatcherMiddleware(AgentMiddleware):
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.
notifier: Injected :class:`~EvoScientist.middleware.notifier.NotifierPort`
used to pre-cancel stale watchers and spawn new ones. The composition
root supplies ``EvoScientist.cli.async_notifier``.
"""
def __init__(self, async_agents: dict[str, Any]) -> None:
def __init__(self, async_agents: dict[str, Any], notifier: NotifierPort) -> None:
from deepagents.middleware.async_subagents import _ClientCache
super().__init__()
self._clients = _ClientCache(async_agents)
self._notifier = notifier
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.
# one (see NotifierPort.pre_cancel_watcher for the full rationale).
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,
)
self._notifier.pre_cancel_watcher(tid)
result = await handler(request)
@@ -96,7 +87,7 @@ class AsyncWatcherMiddleware(AgentMiddleware):
for task_id, task in tasks_update.items():
try:
client = self._clients.get_async(task["agent_name"])
async_notifier.spawn_watcher(
self._notifier.spawn_watcher(
client,
task_id,
task["run_id"],
+76 -64
View File
@@ -11,7 +11,7 @@ sub-agents are *tasks*, future cron is *schedules*).
from __future__ import annotations
from datetime import UTC, datetime
from typing import TYPE_CHECKING
from langchain.agents.middleware import AgentMiddleware
from langchain.tools import ToolRuntime
@@ -20,6 +20,9 @@ from langchain_core.tools import tool
from .. import background, paths
from ..backends import prepare_sandbox_command
if TYPE_CHECKING:
from .notifier import NotifierPort
def _origin_thread_id(runtime: ToolRuntime | None) -> str | None:
"""Best-effort current CLI thread_id, used to route the completion notification."""
@@ -29,11 +32,15 @@ def _origin_thread_id(runtime: ToolRuntime | None) -> str | None:
return None
def _notify_done(proc: background.BgProcess, origin_thread_id: str | None) -> None:
"""Watcher ``on_exit`` hook: enqueue a completion notification (reuses async_notifier).
def _notify_done(
proc: background.BgProcess,
origin_thread_id: str | None,
notifier: NotifierPort,
) -> None:
"""Watcher ``on_exit`` hook: enqueue a completion notification via the port.
Skipped for user-stopped processes (the user already knows). The notifier is imported
lazily to keep this module free of a load-time dependency on the CLI layer.
Skipped for user-stopped processes (the user already knows). The notifier
port owns the notification type, so this module never imports the CLI layer.
"""
if proc.stopped:
return
@@ -44,69 +51,69 @@ def _notify_done(proc: background.BgProcess, origin_thread_id: str | None) -> No
status = "interrupted" # terminated by a signal
else:
status = "error"
from ..cli import async_notifier
async_notifier._enqueue(
async_notifier.AsyncTaskNotification(
task_id=proc.process_id,
agent_name=proc.name,
status=status,
received_at=datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"),
prompt=proc.command,
kind="bg-process",
origin_cli_thread_id=origin_thread_id,
)
notifier.enqueue_bg_process_notification(
task_id=proc.process_id,
agent_name=proc.name,
status=status,
prompt=proc.command,
origin_cli_thread_id=origin_thread_id,
)
@tool(parse_docstring=True)
def run_in_background(
command: str, name: str | None = None, runtime: ToolRuntime = None
) -> str:
"""Launch a long-running shell command in the background and return immediately.
def _make_run_in_background(notifier: NotifierPort, dangerous: bool):
"""Build the ``run_in_background`` tool bound to an injected notifier + policy.
Use for unbounded or very long tasks (model training, large downloads, servers)
that should not block the conversation. Output streams to a log file; poll it with
check_process and stop it with stop_process. For a bounded command that just needs
more time, prefer execute(..., timeout=N) instead of backgrounding.
Args:
command: The shell command to run in the background.
name: Optional short label to recognize the process later.
``dangerous`` is captured from ``cfg.dangerous_mode`` at assembly (the agent
is rebuilt when config changes, so the captured value never goes stale), and
the notifier is the injected port used for the completion notification.
"""
cwd = str(paths.resolve_virtual_path("/"))
# Honor dangerous mode so background commands match `execute`'s policy
# (real-filesystem access, no virtual-path rewriting). Read the env flag that
# apply_config_to_env round-trips at startup (and the subprocess inherits) —
# cheaper than reloading the full config from disk on every launch, and uses
# the same truthy parsing as every other bool env flag.
from ..llm.models import _env_flag_enabled
dangerous = _env_flag_enabled("EVOSCIENTIST_DANGEROUS_MODE")
# Same path-rewriting + validation as execute (shared helper) so virtual paths
# resolve to the workspace and the command can't bypass the sandbox checks.
command, error = prepare_sandbox_command(
command, cwd, virtual_mode=not dangerous, dangerous=dangerous
)
if error:
return error
tid = _origin_thread_id(runtime)
process_id = background.launch(
command, cwd, name, origin_thread_id=tid, on_exit=lambda p: _notify_done(p, tid)
)
label = f" (name={name!r})" if name else ""
# In dangerous mode `/` is the real root, so advertise the real log path;
# in virtual mode `/.bg_processes/...` correctly maps to the workspace.
log_path = (
f"{cwd}/.bg_processes/{process_id}.log"
if dangerous
else f"/.bg_processes/{process_id}.log"
)
return (
f"Started background process {process_id}{label}. "
f"Output -> {log_path}. "
f"Poll with check_process('{process_id}'), stop with stop_process('{process_id}')."
)
@tool(parse_docstring=True)
def run_in_background(
command: str, name: str | None = None, runtime: ToolRuntime = None
) -> str:
"""Launch a long-running shell command in the background and return immediately.
Use for unbounded or very long tasks (model training, large downloads, servers)
that should not block the conversation. Output streams to a log file; poll it with
check_process and stop it with stop_process. For a bounded command that just needs
more time, prefer execute(..., timeout=N) instead of backgrounding.
Args:
command: The shell command to run in the background.
name: Optional short label to recognize the process later.
"""
cwd = str(paths.resolve_virtual_path("/"))
# Same path-rewriting + validation as execute (shared helper) so virtual paths
# resolve to the workspace and the command can't bypass the sandbox checks.
command, error = prepare_sandbox_command(
command, cwd, virtual_mode=not dangerous, dangerous=dangerous
)
if error:
return error
tid = _origin_thread_id(runtime)
process_id = background.launch(
command,
cwd,
name,
origin_thread_id=tid,
on_exit=lambda p: _notify_done(p, tid, notifier),
)
label = f" (name={name!r})" if name else ""
# In dangerous mode `/` is the real root, so advertise the real log path;
# in virtual mode `/.bg_processes/...` correctly maps to the workspace.
log_path = (
f"{cwd}/.bg_processes/{process_id}.log"
if dangerous
else f"/.bg_processes/{process_id}.log"
)
return (
f"Started background process {process_id}{label}. "
f"Output -> {log_path}. "
f"Poll with check_process('{process_id}'), stop with stop_process('{process_id}')."
)
return run_in_background
@tool(parse_docstring=True)
@@ -146,6 +153,11 @@ class BackgroundExecutionMiddleware(AgentMiddleware):
Attached to the main agent only (async sub-agents must not spawn local processes).
"""
def __init__(self) -> None:
def __init__(self, notifier: NotifierPort, *, dangerous: bool = False) -> None:
super().__init__()
self.tools = [run_in_background, check_process, stop_process, list_processes]
self.tools = [
_make_run_in_background(notifier, dangerous),
check_process,
stop_process,
list_processes,
]
+27 -1
View File
@@ -29,6 +29,9 @@ Usage::
from __future__ import annotations
import contextlib
import weakref
from langchain.agents.middleware.types import ModelRequest
from langchain_quickjs import CodeInterpreterMiddleware
@@ -43,6 +46,8 @@ _MEMORY_FIRST_INTERPRETER_PROMPT = (
"them before `code_interpreter` for workspace inspection or implementation work."
)
_live_interpreters: weakref.WeakSet[EvoCodeInterpreterMiddleware]
class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
"""Code interpreter middleware with EvoScientist's memory preflight hint.
@@ -64,6 +69,25 @@ class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
def _prepare_for_call(self, request: ModelRequest) -> str:
return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT
async def aclose(self) -> None:
"""Evict active REPLs on their worker loops before event-loop shutdown."""
registry = self._registry
with registry._lock:
thread_ids = tuple(registry._slots)
for thread_id in thread_ids:
with contextlib.suppress(Exception):
await registry.aevict(thread_id)
self._ptc_tools_by_thread.clear()
_live_interpreters = weakref.WeakSet()
async def aclose_code_interpreters() -> None:
"""Close all live EvoScientist QuickJS middleware instances."""
for middleware in tuple(_live_interpreters):
await middleware.aclose()
# Read-only, batchable tools that benefit from being callable inside JS.
# Multi-agent orchestration is the killer use case: ``Promise.all`` over
@@ -109,9 +133,11 @@ def create_code_interpreter_middleware(
Configured ``CodeInterpreterMiddleware`` ready to append to an agent's
middleware stack.
"""
return EvoCodeInterpreterMiddleware(
middleware = EvoCodeInterpreterMiddleware(
ptc=_DEFAULT_PTC_ALLOWLIST,
timeout=timeout,
max_result_chars=max_result_chars,
tool_name="code_interpreter",
)
_live_interpreters.add(middleware)
return middleware
+191
View File
@@ -0,0 +1,191 @@
"""Typed middleware → frontend event sink.
This module lives deliberately inside ``middleware/`` so the dependency
direction is always **frontends → middleware**, never the reverse. Middleware
reports facts about what happened during a model call; a frontend supplies a
sink implementation that owns its own display state and renders (or ignores)
those facts.
Two families of events are evidenced today and modelled here:
* **Tool selection** — the adaptive ``LLMToolSelectorMiddleware`` wrapper
reports when a selection LLM call starts, which tools survived filtering,
and when it ends.
* **Model fallback** — the fallback middleware reports lifecycle narration for
failed primary calls and fallback attempts.
Tool-selection events are structured. Fallback notices are pre-formatted
narration plus a style, because the middleware owns the wording and the sink
only decides where to display it.
Threading / blocking contract
-----------------------------
Sink methods may be called **from any thread**: synchronous middleware hooks
run on LangChain worker threads, async hooks on whichever loop runs the graph.
A sink implementation therefore MUST be:
* **thread-safe** — any state it mutates is touched under its own lock, and
* **non-blocking** — it marshals to its UI itself (Textual:
``call_from_thread`` / ``post_message``; Rich: the console's internal lock)
and returns promptly.
A sink that blocks stalls the model call that emitted the event — the emitting
worker thread is held until the sink returns. Nothing in the framework isolates
a slow sink from the run.
"""
from __future__ import annotations
from contextvars import ContextVar, Token
from typing import Protocol, runtime_checkable
@runtime_checkable
class MiddlewareEventSink(Protocol):
"""Structured display events emitted by middleware hooks.
Implementations are supplied by frontends/sessions and injected at the
agent composition root (see ``EvoScientist.EvoScientist``). Main agents
built without an explicit sink use :class:`RunScopedEventSink`; subagent
stacks use :class:`NoOpSink`.
All methods must honour the module-level threading/blocking contract:
callable from any thread, thread-safe, and non-blocking.
"""
def on_tool_selection_started(self, total_tools: int) -> None:
"""A tool-selection LLM call has begun over ``total_tools`` tools."""
...
def on_tool_selection(self, selected: list[str], total_tools: int) -> None:
"""The selection kept ``selected`` out of ``total_tools`` tools."""
...
def on_tool_selection_ended(self) -> None:
"""The tool-selection LLM call has finished (or failed)."""
...
def emit_fallback_notice(self, text: str, style: str = "yellow") -> None:
"""Render a pre-formatted fallback lifecycle line."""
...
@runtime_checkable
class ToolSelectionView(Protocol):
"""Read side of the tool-selection state the stream suppressor consumes.
A frontend sink both *records* tool-selection facts (via the
:class:`MiddlewareEventSink` write side) and *exposes* them here so
``stream/tool_selection.py`` can decide whether to suppress selector chatter
and when to surface the selection widget. Ownership lives in the frontend;
the stream layer only reads. :class:`NoOpSink` implements this as "never
active, nothing pending" so headless stacks render no widget.
"""
@property
def tool_selection_active(self) -> bool:
"""Whether a selection LLM call is currently in flight."""
...
def tool_selection_pending(self) -> bool:
"""Whether an unconsumed selection result is waiting to render."""
...
def consume_tool_selection(self) -> tuple[bool, list[str] | None]:
"""Consume the pending selection once, applying dedup-vs-last-emitted.
Returns ``(had_pending, render)``:
* ``had_pending`` — a pending selection existed and was consumed.
* ``render`` — the tool list to display, or ``None`` when the selection
should not render (it kept every tool, or duplicates the last one
shown). ``None`` with ``had_pending=True`` still counts as consumed.
"""
...
@runtime_checkable
class SessionEvents(MiddlewareEventSink, ToolSelectionView, Protocol):
"""Gateway-carried session sink for both middleware writes and stream reads."""
class NoOpSink:
"""Default sink: drops every event and never renders a selection.
Used for headless / gateway / deploy paths and for every subagent stack,
where there is no frontend to render middleware events. Trivially
thread-safe and non-blocking. Implements both the write
(:class:`MiddlewareEventSink`) and read (:class:`ToolSelectionView`) sides.
"""
__slots__ = ()
def on_tool_selection_started(self, total_tools: int) -> None:
pass
def on_tool_selection(self, selected: list[str], total_tools: int) -> None:
pass
def on_tool_selection_ended(self) -> None:
pass
def emit_fallback_notice(self, text: str, style: str = "yellow") -> None:
pass
# --- ToolSelectionView (read side) -----------------------------------
@property
def tool_selection_active(self) -> bool:
return False
def tool_selection_pending(self) -> bool:
return False
def consume_tool_selection(self) -> tuple[bool, list[str] | None]:
return (False, None)
NO_OP_SINK = NoOpSink()
_current_run_event_sink: ContextVar[MiddlewareEventSink | None] = ContextVar(
"evoscientist_current_run_event_sink", default=None
)
def bind_run_event_sink(
events: MiddlewareEventSink,
) -> Token[MiddlewareEventSink | None]:
"""Bind middleware events to the sink for the current streamed run."""
return _current_run_event_sink.set(events)
def reset_run_event_sink(token: Token[MiddlewareEventSink | None]) -> None:
"""Restore the previous run-scoped event sink binding."""
_current_run_event_sink.reset(token)
class RunScopedEventSink:
"""Proxy sink for default main agents.
A main agent can be constructed before the frontend or local gateway exists.
This proxy lets that agent report middleware events to whichever sink the
active ``stream_agent_events`` call bound for the current run. If the agent
is invoked outside that streaming path, events are dropped.
"""
__slots__ = ()
def _sink(self) -> MiddlewareEventSink:
return _current_run_event_sink.get() or NO_OP_SINK
def on_tool_selection_started(self, total_tools: int) -> None:
self._sink().on_tool_selection_started(total_tools)
def on_tool_selection(self, selected: list[str], total_tools: int) -> None:
self._sink().on_tool_selection(selected, total_tools)
def on_tool_selection_ended(self) -> None:
self._sink().on_tool_selection_ended()
def emit_fallback_notice(self, text: str, style: str = "yellow") -> None:
self._sink().emit_fallback_notice(text, style)
+38 -56
View File
@@ -3,7 +3,8 @@
Uses LangChain's AgentMiddleware to intercept model calls. When the primary
model raises an exception, the middleware walks the configured fallback chain,
trying each alternative model in order. Every fallback attempt and its
outcome is surfaced to the user via the registered UI callback.
outcome is reported to the injected event sink as fallback narration, and the
frontend sink renders it.
Errors that indicate a client-side bug (malformed request / HTTP 400) or a
context-length breach are not eligible for fallback and are re-raised
@@ -16,6 +17,7 @@ from __future__ import annotations
import logging
import threading
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING
from langchain.agents.middleware.types import (
AgentMiddleware,
@@ -23,10 +25,10 @@ from langchain.agents.middleware.types import (
ModelResponse,
)
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from .events import MiddlewareEventSink
_ui_emit_fn: Callable[[str, str], None] | None = None
"""UI callback registered by the CLI/TUI entrypoint. ``None`` until set."""
logger = logging.getLogger(__name__)
_fallback_chain_lock = threading.Lock()
_fallback_chain: list[tuple[str, str]] = []
@@ -62,40 +64,6 @@ These are intentionally *not* treated as non-fallbackable because a different
provider in the chain may have valid credentials."""
def set_ui_emit(fn: Callable[[str, str], None] | None) -> None:
"""Register (or clear) the UI callback for fallback status messages.
Args:
fn: Callable with signature ``fn(text, style)`` where *style* is a
Rich style string (``"yellow"``, ``"red"``, ``"green"``).
Pass ``None`` to unregister.
"""
global _ui_emit_fn
_ui_emit_fn = fn
def _emit(text: str, style: str = "yellow") -> None:
"""Surface a fallback status message to the user.
Dispatches to the registered UI callback when available (TUI mode),
otherwise falls back to the shared Rich console on stdout (CLI mode).
Args:
text: Plain-text message to display.
style: Rich style string applied to the message.
"""
if _ui_emit_fn is not None:
try:
_ui_emit_fn(text, style)
return
except Exception:
pass
from ..stream.console import console
console.print(text, style=style)
def get_fallback_chain() -> list[tuple[str, str]]:
"""Return a snapshot of the current fallback chain.
@@ -234,6 +202,7 @@ async def _try_fallbacks(
request: ModelRequest,
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
primary_exc: Exception,
events: MiddlewareEventSink,
) -> ModelResponse:
"""Walk the fallback chain, trying each model until one succeeds.
@@ -246,6 +215,7 @@ async def _try_fallbacks(
request: The original model request.
invoke: Async callable that invokes the handler on a request.
primary_exc: The exception raised by the primary model.
events: Injected event sink for fallback narration.
Returns:
The ``ModelResponse`` from the first successful fallback.
@@ -255,9 +225,9 @@ async def _try_fallbacks(
"""
from ..llm.models import get_chat_model
_emit(
events.emit_fallback_notice(
f"Primary model failed: {type(primary_exc).__name__}: {primary_exc}",
style="yellow",
"yellow",
)
logger.warning(
"Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc
@@ -273,35 +243,35 @@ async def _try_fallbacks(
last_failing_request = request
for model_name, provider in get_fallback_chain():
_emit(
f" -> Falling back to {model_name} ({provider}) "
f"due to: {type(last_exc).__name__}: {last_exc}",
style="yellow",
events.emit_fallback_notice(
f" -> Falling back to {model_name} ({provider}) due to: "
f"{type(last_exc).__name__}: {last_exc}",
"yellow",
)
try:
fallback_model = get_chat_model(model=model_name, provider=provider)
fb_request = request.override(model=fallback_model)
result = await invoke(fb_request)
_emit(
events.emit_fallback_notice(
f" Fallback to {model_name} ({provider}) succeeded",
style="green",
"green",
)
logger.info("Fallback to %s (%s) succeeded", model_name, provider)
return result
except Exception as fb_exc:
reason = _is_non_fallbackable(fb_exc)
if reason is not None:
_emit(
events.emit_fallback_notice(
f" {model_name} hit non-fallbackable error ({reason}) "
f"-- aborting fallback chain",
style="red",
"red",
)
_raise_normalized(fb_request, fb_exc)
last_exc = fb_exc
last_failing_request = fb_request
_emit(
events.emit_fallback_notice(
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
style="red",
"red",
)
logger.warning(
"Fallback %s (provider=%s) failed: %s: %s",
@@ -311,7 +281,9 @@ async def _try_fallbacks(
fb_exc,
)
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
events.emit_fallback_notice(
" All fallbacks exhausted -- re-raising last error", "red"
)
_raise_normalized(last_failing_request, last_exc)
@@ -336,6 +308,7 @@ def _guard_and_fallback(
primary_exc: Exception,
request: ModelRequest,
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
events: MiddlewareEventSink,
) -> Awaitable[ModelResponse]:
"""Check non-fallbackable conditions, then delegate to ``_try_fallbacks``.
@@ -343,6 +316,7 @@ def _guard_and_fallback(
primary_exc: The exception raised by the primary model.
request: The original model request.
invoke: Async callable that invokes the handler on a request.
events: Injected event sink for fallback narration.
Returns:
Coroutine that resolves to the fallback ``ModelResponse``.
@@ -352,12 +326,12 @@ def _guard_and_fallback(
"""
reason = _is_non_fallbackable(primary_exc)
if reason is not None:
_emit(
events.emit_fallback_notice(
f"Model error ({reason}) -- not eligible for fallback, re-raising",
style="red",
"red",
)
_raise_normalized(request, primary_exc)
return _try_fallbacks(request, invoke, primary_exc)
return _try_fallbacks(request, invoke, primary_exc, events)
class ModelFallbackMiddleware(AgentMiddleware):
@@ -373,6 +347,12 @@ class ModelFallbackMiddleware(AgentMiddleware):
name = "model_fallback"
def __init__(self, events: MiddlewareEventSink | None = None) -> None:
super().__init__()
from .events import NO_OP_SINK
self._events = events or NO_OP_SINK
def wrap_model_call(
self,
request: ModelRequest,
@@ -389,7 +369,9 @@ class ModelFallbackMiddleware(AgentMiddleware):
import asyncio
return asyncio.run(_guard_and_fallback(exc, request, _sync_invoke))
return asyncio.run(
_guard_and_fallback(exc, request, _sync_invoke, self._events)
)
async def awrap_model_call(
self,
@@ -401,4 +383,4 @@ class ModelFallbackMiddleware(AgentMiddleware):
try:
return await handler(request)
except Exception as exc:
return await _guard_and_fallback(exc, request, handler)
return await _guard_and_fallback(exc, request, handler, self._events)
+68
View File
@@ -0,0 +1,68 @@
"""Notifier port for async-task / background-process notifications.
The ``async_watcher`` and ``background`` middleware need to (a) pre-cancel a
stale watcher, (b) spawn a watcher, and (c) enqueue a completion notification.
Those are infrastructure calls with behaviour, not display events — so they do
not belong on the :mod:`~EvoScientist.middleware.events` display sink.
Instead the composition root injects a :class:`NotifierPort`: a small,
structural interface implemented by ``EvoScientist.cli.async_notifier`` (the
module itself satisfies it — its public functions match these methods). The
middleware depends only on this port, never on ``EvoScientist.cli``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Protocol
if TYPE_CHECKING:
import asyncio
class NotifierPort(Protocol):
"""Behaviour the notifier layer exposes to middleware.
``EvoScientist.cli.async_notifier`` implements this structurally; the
composition root passes that module in as the port.
"""
def pre_cancel_watcher(self, task_id: str) -> None:
"""Cancel any in-flight watcher registered for ``task_id``.
No-op when there is no live watcher. Swallows cancellation errors —
a failed pre-cancel only risks a stale success notification, never a
crash of the launching tool call.
"""
...
def spawn_watcher(
self,
client: Any,
thread_id: str,
run_id: str,
agent_name: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> asyncio.Task[None]:
"""Spawn a run watcher on the caller's asyncio loop."""
...
def enqueue_task_notification(self, notification: Any) -> None:
"""Route a completed-task notification onto the consumer queue."""
...
def enqueue_bg_process_notification(
self,
*,
task_id: str,
agent_name: str,
status: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> None:
"""Build and enqueue a background-process completion notification.
The notifier owns the notification type, so the background middleware
never constructs it (and never imports the CLI layer).
"""
...
+62 -103
View File
@@ -1,16 +1,18 @@
"""LLMToolSelectorMiddleware configuration for EvoScientist.
Wraps LangChain's built-in ``LLMToolSelectorMiddleware`` with project-specific
defaults and an optional stream tracker that captures which tools were selected.
defaults. The wrapper reports what it did through an injected
:class:`~EvoScientist.middleware.events.MiddlewareEventSink`; the frontend sink
owns any display state (there are no process-global variables here).
The selector only activates when the agent has more than ``threshold`` tools
(default 20). Below that, the extra LLM call isn't worth the token savings.
(default 26). Below that, the extra LLM call isn't worth the token savings.
Usage::
from EvoScientist.middleware import create_tool_selector_middleware
middleware = create_tool_selector_middleware() # returns [selector, tracker]
middleware = create_tool_selector_middleware(events=sink)
"""
from __future__ import annotations
@@ -29,14 +31,9 @@ from langchain.agents.middleware.types import (
from langchain_core.language_models import BaseChatModel
from langchain_core.tools import BaseTool
logger = logging.getLogger(__name__)
from .events import NO_OP_SINK, MiddlewareEventSink
# Module-level storage for main-agent tool-selection UI state.
# Updated only when stream tracking is enabled; read by stream/events.py.
_current_selected_tools: list[str] = []
_last_emitted_tools: list[str] = [] # last selection shown to user
_total_tools_count: int = 0 # total tools before selection
_selector_active: bool = False
logger = logging.getLogger(__name__)
# Default threshold: only run tool selection when tools exceed this count.
# Base tools are ~14; selector activates when MCP tools push count above 26.
@@ -74,8 +71,12 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
Skips the selection LLM call when ``len(request.tools) <= threshold``,
avoiding unnecessary overhead for agents with few tools.
When stream tracking is enabled, sets ``_selector_active`` during the
selector's internal LLM call so the streaming layer can suppress its output.
When selection runs, reports the lifecycle to the injected sink:
``on_tool_selection_started`` before the selector call, ``on_tool_selection``
with the surviving tools once the selector hands off the filtered request,
and ``on_tool_selection_ended`` when the call finishes (or fails). The sink
(a frontend one, or :class:`NoOpSink` for subagent / headless stacks) owns
all display state.
"""
name = "conditional_tool_selector"
@@ -86,13 +87,13 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
threshold: int = DEFAULT_TOOL_THRESHOLD,
*,
always_include: frozenset[str] | None = None,
track_stream_selection: bool = True,
events: MiddlewareEventSink | None = None,
):
super().__init__()
self._selector_factory = selector_factory
self._threshold = threshold
self._always_include = always_include or frozenset()
self._track_stream_selection = track_stream_selection
self._events = events or NO_OP_SINK
# Agent tools are fixed after graph construction, so the filtered
# always-include set is stable for this middleware instance.
self._selector: AgentMiddleware | None = None
@@ -103,6 +104,10 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
self._selector = self._selector_factory(names)
return self._selector
@staticmethod
def _selected_names(request: ModelRequest) -> list[str]:
return [name for tool in request.tools if (name := _tool_name(tool))]
def wrap_model_call(
self,
request: ModelRequest,
@@ -111,21 +116,29 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
if len(request.tools) <= self._threshold:
return handler(request)
if self._track_stream_selection:
global _selector_active, _total_tools_count
_selector_active = True
_total_tools_count = len(request.tools)
total = len(request.tools)
self._events.on_tool_selection_started(total)
# Track whether handler was called — if so, any exception is from
# the downstream model, not the selector, and must propagate.
_handler_called = False
_selection_open = True
def _end_selection() -> None:
nonlocal _selection_open
if _selection_open:
self._events.on_tool_selection_ended()
_selection_open = False
def _handler_after_selection(req: ModelRequest) -> ModelResponse:
nonlocal _handler_called
_handler_called = True
if self._track_stream_selection:
global _selector_active
_selector_active = False
# ``req.tools`` is the selector-filtered set here.
selected = self._selected_names(req)
self._events.on_tool_selection(selected, total)
if selected:
logger.debug("Selected tools: %s", selected)
_end_selection()
return handler(req)
try:
@@ -148,12 +161,10 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
# Structured-output shape / config failure — gracefully
# degrade to using all tools.
logger.debug("Tool selector failed, using all tools", exc_info=True)
if self._track_stream_selection:
_selector_active = False
_end_selection()
return handler(request)
finally:
if self._track_stream_selection:
_selector_active = False
_end_selection()
async def awrap_model_call(
self,
@@ -163,19 +174,26 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
if len(request.tools) <= self._threshold:
return await handler(request)
if self._track_stream_selection:
global _selector_active, _total_tools_count
_selector_active = True
_total_tools_count = len(request.tools)
total = len(request.tools)
self._events.on_tool_selection_started(total)
_handler_called = False
_selection_open = True
def _end_selection() -> None:
nonlocal _selection_open
if _selection_open:
self._events.on_tool_selection_ended()
_selection_open = False
async def _handler_after_selection(req: ModelRequest) -> ModelResponse:
nonlocal _handler_called
_handler_called = True
if self._track_stream_selection:
global _selector_active
_selector_active = False
selected = self._selected_names(req)
self._events.on_tool_selection(selected, total)
if selected:
logger.debug("Selected tools: %s", selected)
_end_selection()
return await handler(req)
try:
@@ -193,71 +211,32 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
# on shape / config failures.
raise
logger.debug("Tool selector failed, using all tools", exc_info=True)
if self._track_stream_selection:
_selector_active = False
_end_selection()
return await handler(request)
finally:
if self._track_stream_selection:
_selector_active = False
class _ToolSelectionTrackerMiddleware(AgentMiddleware):
"""Captures which tools the model actually receives after filtering.
Sits right AFTER the selector in the middleware chain (more inner),
so ``request.tools`` already contains only the selected tools when
this middleware's ``wrap_model_call`` runs.
"""
name = "tool_selection_tracker"
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
global _current_selected_tools
tools = [name for tool in request.tools if (name := _tool_name(tool))]
_current_selected_tools = tools
if tools:
logger.debug("Selected tools: %s", tools)
return handler(request)
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
global _current_selected_tools
tools = [name for tool in request.tools if (name := _tool_name(tool))]
_current_selected_tools = tools
if tools:
logger.debug("Selected tools: %s", tools)
return await handler(request)
_end_selection()
def create_tool_selector_middleware(
threshold: int = DEFAULT_TOOL_THRESHOLD,
*,
model: BaseChatModel | None = None,
track_stream_selection: bool = True,
events: MiddlewareEventSink | None = None,
):
"""Build LLMToolSelectorMiddleware + tracker with EvoScientist defaults.
"""Build the conditional ``LLMToolSelectorMiddleware`` wrapper.
Returns middleware for adaptive tool selection:
1. Conditional wrapper around ``LLMToolSelectorMiddleware`` — only
activates when ``len(tools) > threshold``
2. Optional ``_ToolSelectionTrackerMiddleware`` — captures selected tool
names for the main-agent stream UI when ``track_stream_selection`` is true
Returns a single-element middleware list (kept as a list so the assembly
site can splat it) that adaptively selects tools only when
``len(tools) > threshold``. The wrapper reports the selection lifecycle to
``events``; pass a frontend sink for the main agent, or omit it (subagent /
headless stacks) to get the silent :class:`NoOpSink`.
Args:
model: Chat model for tool selection. If *None*, the default
model is resolved via ``_ensure_chat_model()``.
threshold: Minimum number of tools to trigger selection.
Default 26. Set to 0 to always run selection.
track_stream_selection: Whether to update process-global stream/UI
state. Disable for async sub-agents that should still select tools
but should not drive the main-agent tool-selection widget.
events: Frontend event sink to report selection to.
``think_tool``, ``task``, and memory tools are always included because:
@@ -293,31 +272,11 @@ def create_tool_selector_middleware(
always_include=always_include,
)
middleware: list[AgentMiddleware] = [
return [
_ConditionalToolSelectorMiddleware(
selector_factory=selector_factory,
threshold=threshold,
always_include=DEFAULT_ALWAYS_INCLUDE_TOOLS,
track_stream_selection=track_stream_selection,
events=events,
),
]
if track_stream_selection:
middleware.append(_ToolSelectionTrackerMiddleware())
return middleware
def reset_tool_selection_state_for_tests() -> None:
"""Reset the process-global tool-selection state.
The selector/tracker record the last selected tools and the selector-active
flag in module globals that ``stream/tool_selection.py`` reads to suppress
selector chatter. Tests that drive the selector must not leak that state
into later tests; an autouse fixture resets it around every test.
"""
global _current_selected_tools, _last_emitted_tools
global _total_tools_count, _selector_active
_current_selected_tools = []
_last_emitted_tools = []
_total_tools_count = 0
_selector_active = False
+3 -1
View File
@@ -558,7 +558,9 @@ def create_streaming_display(
return Group(*elements)
# Thinking panel
_show_thinking = final_show_thinking if is_final else show_thinking
# ``final_show_thinking`` controls the final-frame layout, but it must
# never override the caller's user-level visibility setting.
_show_thinking = show_thinking and (final_show_thinking if is_final else True)
if _show_thinking and thinking_text:
thinking_title = "Thinking"
display_thinking = thinking_text.rstrip()
+27 -2
View File
@@ -10,7 +10,7 @@ import mimetypes
import os
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
from dataclasses import dataclass
from typing import Any, TypeAlias
from typing import TYPE_CHECKING, Any, TypeAlias
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
from langgraph.graph import END
@@ -26,6 +26,9 @@ from .summarization import (
from .tool_results import _extract_command_tool_content, _extract_tool_content
from .tool_selection import _ToolSelectionSuppressor
from .utils import DisplayLimits, is_success
if TYPE_CHECKING:
from ..middleware.events import ToolSelectionView
from .v3_payloads import (
RawMap,
_as_raw_map,
@@ -195,7 +198,10 @@ class _V3EventProcessor:
existing_summarization_event: Mapping[str, object] | None,
existing_messages: object = None,
process_value_messages: bool = False,
events: "ToolSelectionView | None" = None,
) -> None:
from ..middleware.events import NO_OP_SINK
self.emitter = emitter
self.subagents = subagents
self._suppressed_summarization_signature = _summarization_event_signature(
@@ -210,7 +216,7 @@ class _V3EventProcessor:
] = {}
self._emitted_tool_calls: set[tuple[tuple[str, ...], str]] = set()
self._emitted_interrupts: set[str] = set()
self._selector = _ToolSelectionSuppressor(emitter)
self._selector = _ToolSelectionSuppressor(emitter, events or NO_OP_SINK)
@staticmethod
def _tool_scope(
@@ -797,6 +803,7 @@ async def stream_agent_events(
thread_id: str,
metadata: dict[str, Any] | None = None,
media: list[str] | None = None,
events: "ToolSelectionView | None" = None,
) -> AsyncGenerator[dict[str, Any], None]:
"""Stream events from a DeepAgents/LangGraph v3 run.
@@ -818,6 +825,11 @@ async def stream_agent_events(
subagent_start, subagent_tool_call, subagent_tool_result, subagent_end,
done, error
"""
if events is None:
from .sink import SessionEventSink
events = SessionEventSink()
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
if metadata:
config["metadata"] = metadata
@@ -837,9 +849,19 @@ async def stream_agent_events(
stream: Any | None = None
producers: list[asyncio.Task[Any]] = []
_run_raised: bool = False
event_sink_token = None
try:
from langgraph.stream.transformers import UpdatesTransformer
from ..middleware.events import (
MiddlewareEventSink,
bind_run_event_sink,
reset_run_event_sink,
)
if isinstance(events, MiddlewareEventSink):
event_sink_token = bind_run_event_sink(events)
try:
stream_result = agent.astream_events(
astream_input,
@@ -858,6 +880,7 @@ async def stream_agent_events(
emitter,
subagents,
existing_summarization_event,
events=events,
)
queue: asyncio.Queue[Any] = asyncio.Queue()
producer_done = object()
@@ -964,6 +987,8 @@ async def stream_agent_events(
task.cancel()
if producers:
await asyncio.gather(*producers, return_exceptions=True)
if event_sink_token is not None:
reset_run_event_sink(event_sink_token)
# When the run ended with an exception the LangGraph checkpoint may be
# left interrupted (``next`` non-empty). Clear it — unless it's a real
# human-in-the-loop pause — so the next user message starts a fresh turn
+114
View File
@@ -0,0 +1,114 @@
"""Session-owned event sink.
A single sink instance, created and owned by a frontend or local stream run,
that both **records** tool-selection state for the stream suppressor and
optionally **renders** model-fallback notices. It is injected into the agent's
middleware (write side) and read by ``stream/tool_selection.py`` (read side).
The tool-selection state machine — active / total / pending / last-emitted with
consume-once + dedup semantics — lives here, in the session sink, replacing the
process-global module variables that used to live in ``middleware/tool_selector``.
The state is guarded by a lock because middleware hooks fire on worker threads
while the stream suppressor reads on the runtime thread (see the threading
contract in :mod:`EvoScientist.middleware.events`).
"""
from __future__ import annotations
import logging
import threading
from collections.abc import Callable
from .console import console
logger = logging.getLogger(__name__)
class SessionEventSink:
"""Tool-selection state holder + model-fallback renderer for a session.
Args:
fallback_display: Optional ``(text, style)`` callback the frontend
supplies to render a model-fallback notice (Rich: ``console.print``;
TUI: append a system message). ``None`` renders to the shared
Rich console.
"""
def __init__(
self, fallback_display: Callable[[str, str], None] | None = None
) -> None:
self._lock = threading.Lock()
self._active = False
self._total = 0
self._pending: list[str] | None = None
self._last_emitted: list[str] = []
self._fallback_display = fallback_display
def set_fallback_display(
self, fallback_display: Callable[[str, str], None] | None
) -> None:
"""(Re)bind the fallback-notice display sink.
Frontends whose display target only exists after construction (the TUI
binds its ``_append_system`` on mount, clears it on exit) use this
instead of the constructor argument.
"""
with self._lock:
self._fallback_display = fallback_display
def _display_fallback(self, text: str, style: str) -> None:
with self._lock:
fallback_display = self._fallback_display
if fallback_display is None:
console.print(text, style=style)
return
try:
fallback_display(text, style)
except Exception:
logger.warning("Fallback display callback failed", exc_info=True)
console.print(text, style=style)
# --- MiddlewareEventSink write side (any thread) ---------------------
def on_tool_selection_started(self, total_tools: int) -> None:
with self._lock:
self._active = True
self._total = total_tools
def on_tool_selection(self, selected: list[str], total_tools: int) -> None:
with self._lock:
self._pending = list(selected)
self._total = total_tools
def on_tool_selection_ended(self) -> None:
with self._lock:
self._active = False
def emit_fallback_notice(self, text: str, style: str = "yellow") -> None:
"""Render a fallback lifecycle line verbatim."""
self._display_fallback(text, style)
# --- ToolSelectionView read side (runtime thread) -------------------
@property
def tool_selection_active(self) -> bool:
with self._lock:
return self._active
def tool_selection_pending(self) -> bool:
with self._lock:
return bool(self._pending)
def consume_tool_selection(self) -> tuple[bool, list[str] | None]:
with self._lock:
pending = self._pending
if not pending:
return (False, None)
# Consume-once: clear before deciding whether to render.
self._pending = None
# Only surface when the selection actually filtered tools and it
# differs from the last selection already shown to the user.
if len(pending) < self._total and sorted(pending) != sorted(
self._last_emitted
):
self._last_emitted = list(pending)
return (True, list(pending))
return (True, None)
+24 -23
View File
@@ -1,15 +1,26 @@
"""Tool-selection stream suppression helpers."""
from typing import Any
from typing import TYPE_CHECKING, Any
from .emitter import StreamEventEmitter
if TYPE_CHECKING:
from ..middleware.events import ToolSelectionView
class _ToolSelectionSuppressor:
"""Suppress selector model JSON while preserving the UI selection event."""
"""Suppress selector model JSON while preserving the UI selection event.
def __init__(self, emitter: StreamEventEmitter) -> None:
Reads its selection state from the injected frontend sink (a
:class:`~EvoScientist.middleware.events.ToolSelectionView`) rather than
process globals: ``tool_selection_active`` gates chatter suppression and
``consume_tool_selection`` yields the deduped list to surface. The sink
owns the state; this class only observes the stream and asks the sink.
"""
def __init__(self, emitter: StreamEventEmitter, sink: "ToolSelectionView") -> None:
self._emitter = emitter
self._sink = sink
self._buffering = False
self._buffer = ""
self._was_active = False
@@ -93,17 +104,11 @@ class _ToolSelectionSuppressor:
return True
return self._selector_call_active() or self._selection_pending()
@staticmethod
def _selector_call_active() -> bool:
import EvoScientist.middleware.tool_selector as selector_mod
def _selector_call_active(self) -> bool:
return bool(self._sink.tool_selection_active)
return bool(selector_mod._selector_active)
@staticmethod
def _selection_pending() -> bool:
import EvoScientist.middleware.tool_selector as selector_mod
return bool(selector_mod._current_selected_tools)
def _selection_pending(self) -> bool:
return bool(self._sink.tool_selection_pending())
def flush_selection(self) -> list[dict[str, Any]]:
return self._emit_selection_if_ready("")
@@ -119,17 +124,13 @@ class _ToolSelectionSuppressor:
def _emit_selection_if_ready(self, text: str) -> list[dict[str, Any]]:
if not self._was_active:
return []
import EvoScientist.middleware.tool_selector as selector_mod
if selector_mod._current_selected_tools:
had_pending, render = self._sink.consume_tool_selection()
if had_pending:
# Pending consumed (once) — clear our observation flag regardless of
# whether the sink chose to render it (dedup / kept-all cases).
self._was_active = False
selected = selector_mod._current_selected_tools
selector_mod._current_selected_tools = []
if len(selected) < selector_mod._total_tools_count and sorted(
selected
) != sorted(selector_mod._last_emitted_tools):
selector_mod._last_emitted_tools = list(selected)
return [self._emitter.tool_selection(list(selected)).data]
if render is not None:
return [self._emitter.tool_selection(render).data]
elif text:
self._was_active = False
return []
-20
View File
@@ -7,26 +7,6 @@ import pytest
_NONEXISTENT_DOTENV = str(Path(__file__).with_name(".pytest-dotenv-does-not-exist"))
@pytest.fixture(autouse=True)
def _reset_tool_selection_state():
"""Isolate the process-global tool-selection state around every test.
``middleware.tool_selector`` records the last selected tools and the
selector-active flag in module globals that ``stream/tool_selection.py``
reads to decide whether to suppress selector output. A test that drives the
selector or tracker would otherwise leave those globals set and silently
flip unrelated streaming tests later in the same process. Reset on both ends
so order and worker sharding can't reintroduce the leak.
"""
from EvoScientist.middleware.tool_selector import (
reset_tool_selection_state_for_tests,
)
reset_tool_selection_state_for_tests()
yield
reset_tool_selection_state_for_tests()
@pytest.fixture
def sample_tool_call():
"""A minimal tool call dict."""
+11 -4
View File
@@ -20,16 +20,23 @@ async def collect_events(
agent,
message: str = "hi",
thread_id: str = "t1",
*,
events=None,
):
"""Collect stream_agent_events output for tests."""
events = []
"""Collect stream_agent_events output for tests.
``events`` is the frontend tool-selection sink to drive suppression /
selection rendering (defaults to the silent NoOpSink inside the stream).
"""
collected = []
async for ev in stream_agent_events(
agent,
message,
thread_id,
events=events,
):
events.append(ev)
return events
collected.append(ev)
return collected
def protocol_event(
+4 -1
View File
@@ -308,8 +308,11 @@ def test_async_subagent_disables_tool_selector_stream_tracking(
mock_chat.return_value = MagicMock(profile={"max_input_tokens": 200_000})
from EvoScientist.EvoScientist import _get_default_middleware
from EvoScientist.middleware.events import NoOpSink
_get_default_middleware(for_async_subagent=True)
# Async subagents still select tools, but are wired to the silent NoOpSink
# so they never drive the main-agent tool-selection widget.
mock_tool_selector.assert_called_once()
assert mock_tool_selector.call_args.kwargs["track_stream_selection"] is False
assert isinstance(mock_tool_selector.call_args.kwargs["events"], NoOpSink)
+2 -1
View File
@@ -99,7 +99,8 @@ def _make_middleware():
"url": "http://x",
"graph_id": "writing-agent",
}
}
},
notifier=async_notifier,
)
return mw, fake_client
+18 -23
View File
@@ -6,15 +6,21 @@ import time
import pytest
from EvoScientist import background as bg
from EvoScientist.cli import async_notifier
from EvoScientist.middleware.background import (
BackgroundExecutionMiddleware,
_make_run_in_background,
check_process,
list_processes,
run_in_background,
stop_process,
)
def _run_bg(*, dangerous: bool = False, notifier=async_notifier):
"""Build the injected ``run_in_background`` tool for direct-invoke tests."""
return _make_run_in_background(notifier, dangerous)
def _sleep_cmd(seconds: int) -> str:
"""Cross-platform command that sleeps for *seconds* and exits 0."""
if sys.platform == "win32":
@@ -56,7 +62,7 @@ def _clean_registry():
def test_middleware_registers_four_tools():
mw = BackgroundExecutionMiddleware()
mw = BackgroundExecutionMiddleware(async_notifier)
names = {t.name for t in mw.tools}
assert names == {
"run_in_background",
@@ -68,7 +74,7 @@ def test_middleware_registers_four_tools():
def test_no_job_in_tool_names():
"""Naming ADR: the word 'job' must not appear in the tool surface."""
mw = BackgroundExecutionMiddleware()
mw = BackgroundExecutionMiddleware(async_notifier)
assert not any("job" in t.name.lower() for t in mw.tools)
@@ -80,7 +86,7 @@ def test_run_rejects_dangerous_command_without_launching(monkeypatch):
return "should-not-happen"
monkeypatch.setattr(bg, "launch", _spy)
out = run_in_background.invoke({"command": "sudo rm -rf /"})
out = _run_bg().invoke({"command": "sudo rm -rf /"})
assert launched["called"] is False
assert "blocked" in out.lower()
@@ -88,7 +94,7 @@ def test_run_rejects_dangerous_command_without_launching(monkeypatch):
def test_run_launches_valid_command(tmp_path, monkeypatch):
# Pin the workspace cwd to a temp dir so the launch is isolated.
monkeypatch.setattr("EvoScientist.paths.resolve_virtual_path", lambda _vp: tmp_path)
out = run_in_background.invoke({"command": "echo ok", "name": "demo"})
out = _run_bg().invoke({"command": "echo ok", "name": "demo"})
assert "Started background process" in out
assert "check_process" in out
assert len(bg._PROCESSES) == 1
@@ -104,24 +110,14 @@ def test_run_applies_virtual_path_rewriting(tmp_path, monkeypatch):
return "pidX"
monkeypatch.setattr(bg, "launch", _spy)
run_in_background.invoke({"command": "python /train.py"})
_run_bg().invoke({"command": "python /train.py"})
# virtual absolute path -> workspace-relative, same as execute would produce
assert captured["command"] == "python ./train.py"
def _force_dangerous(monkeypatch, value=True):
"""Make run_in_background see dangerous mode via the env flag it reads.
monkeypatch.setenv tracks the change and restores it on teardown, so this
cannot leak EVOSCIENTIST_DANGEROUS_MODE into other tests.
"""
monkeypatch.setenv("EVOSCIENTIST_DANGEROUS_MODE", "true" if value else "false")
def test_run_dangerous_allows_real_path_no_rewrite(tmp_path, monkeypatch):
"""In dangerous mode, background commands keep real absolute paths (parity with execute)."""
monkeypatch.setattr("EvoScientist.paths.resolve_virtual_path", lambda _vp: tmp_path)
_force_dangerous(monkeypatch)
captured = {}
def _spy(command, cwd, name=None, *, origin_thread_id=None, on_exit=None):
@@ -130,7 +126,7 @@ def test_run_dangerous_allows_real_path_no_rewrite(tmp_path, monkeypatch):
monkeypatch.setattr(bg, "launch", _spy)
# Absolute path + traversal would be BLOCKED in normal mode; allowed here.
out = run_in_background.invoke({"command": "cat /etc/hosts && cat ../x"})
out = _run_bg(dangerous=True).invoke({"command": "cat /etc/hosts && cat ../x"})
assert "blocked" not in out.lower()
assert captured["command"] == "cat /etc/hosts && cat ../x" # no ./ rewrite
# Advertised log path is the real path, not the virtual /.bg_processes/.
@@ -141,7 +137,6 @@ def test_run_dangerous_allows_real_path_no_rewrite(tmp_path, monkeypatch):
def test_run_dangerous_still_blocks_privileged_command(tmp_path, monkeypatch):
"""Dangerous mode must NOT relax the privileged-command blocklist."""
monkeypatch.setattr("EvoScientist.paths.resolve_virtual_path", lambda _vp: tmp_path)
_force_dangerous(monkeypatch)
launched = {"called": False}
def _spy(*args, **kwargs):
@@ -149,7 +144,7 @@ def test_run_dangerous_still_blocks_privileged_command(tmp_path, monkeypatch):
return "should-not-happen"
monkeypatch.setattr(bg, "launch", _spy)
out = run_in_background.invoke({"command": "sudo rm x"})
out = _run_bg(dangerous=True).invoke({"command": "sudo rm x"})
assert launched["called"] is False
assert "blocked" in out.lower()
@@ -159,7 +154,7 @@ def test_run_enqueues_completion_notification(tmp_path, monkeypatch):
from EvoScientist.cli import async_notifier
monkeypatch.setattr("EvoScientist.paths.resolve_virtual_path", lambda _vp: tmp_path)
run_in_background.invoke({"command": _true_cmd(), "name": "quick"})
_run_bg().invoke({"command": _true_cmd(), "name": "quick"})
# drain consumes, so accumulate across polls until the watcher's on_exit enqueues.
notifs = []
deadline = time.time() + 4.0
@@ -189,7 +184,7 @@ def test_notify_done_routes_to_origin_thread(tmp_path):
pid = bg.launch(_true_cmd(), str(tmp_path)) # no on_exit -> no auto-notify here
assert _wait_until(lambda: bg._PROCESSES[pid].finished_ts is not None)
_notify_done(bg._PROCESSES[pid], "T-123")
_notify_done(bg._PROCESSES[pid], "T-123", async_notifier)
routed = async_notifier.drain_notifications("T-123")
assert any(n.task_id == pid and n.origin_cli_thread_id == "T-123" for n in routed)
@@ -199,7 +194,7 @@ def test_stopped_process_suppresses_notification(tmp_path, monkeypatch):
from EvoScientist.cli import async_notifier
monkeypatch.setattr("EvoScientist.paths.resolve_virtual_path", lambda _vp: tmp_path)
run_in_background.invoke({"command": _sleep_cmd(600)})
_run_bg().invoke({"command": _sleep_cmd(600)})
(pid,) = list(bg._PROCESSES.keys())
stop_process.invoke({"process_id": pid})
# Wait until the watcher observed the exit — it would have enqueued here if the
@@ -310,7 +305,7 @@ def test_shell_notification_hints_check_process():
def test_check_and_list_route_to_manager(tmp_path, monkeypatch):
monkeypatch.setattr("EvoScientist.paths.resolve_virtual_path", lambda _vp: tmp_path)
run_in_background.invoke({"command": _sleep_cmd(1)})
_run_bg().invoke({"command": _sleep_cmd(1)})
(pid,) = bg._PROCESSES.keys()
assert pid in check_process.invoke({"process_id": pid})
assert pid in list_processes.invoke({})
+90 -2
View File
@@ -847,11 +847,16 @@ class TestChannelTyping:
class TestChannelReconnect:
async def test_run_reconnects_on_error(self):
async def test_run_reconnects_on_error(self, monkeypatch):
"""Channel.run() should reconnect with backoff on transient errors."""
ch = StubChannel()
mgr = ChannelManager(MessageBus())
mgr.register(ch)
start_count = 0
sleep_count = 0
first_retry_waiting = asyncio.Event()
allow_retries = asyncio.Event()
original_start = ch.start
async def flaky_start():
@@ -863,9 +868,73 @@ class TestChannelReconnect:
# Stop after successful start to end the test
ch._running = False
async def controlled_sleep(_delay):
nonlocal sleep_count
sleep_count += 1
if sleep_count == 1:
first_retry_waiting.set()
await allow_retries.wait()
ch.start = flaky_start
await ch.run()
monkeypatch.setattr(asyncio, "sleep", controlled_sleep)
run_task = asyncio.create_task(ch.run())
await first_retry_waiting.wait()
try:
assert ch._startup_event.is_set() is False
assert ch._startup_error is None
assert mgr.startup_results() == [("stub", False, "starting (bus)")]
finally:
allow_retries.set()
await run_task
assert start_count == 3
assert ch._startup_event.is_set()
assert ch._startup_error is None
async def test_runtime_error_does_not_overwrite_successful_startup(
self, monkeypatch
):
"""A receive failure should preserve the completed startup result."""
ch = StubChannel()
start_count = 0
receive_count = 0
state_before_reconnect: list[tuple[bool, str | None]] = []
original_start = ch.start
async def tracking_start():
nonlocal start_count
start_count += 1
if start_count == 2:
state_before_reconnect.append(
(ch._startup_event.is_set(), ch._startup_error)
)
await original_start()
if start_count == 2:
ch._running = False
async def flaky_receive():
nonlocal receive_count
receive_count += 1
if receive_count == 1:
raise ConnectionError("receive transient")
if False: # pragma: no cover - marks this as an async generator
yield None
async def no_sleep(_delay):
return None
ch.start = tracking_start
ch.receive = flaky_receive
monkeypatch.setattr(asyncio, "sleep", no_sleep)
await ch.run()
assert state_before_reconnect == [(True, None)]
assert ch._startup_event.is_set()
assert ch._startup_error is None
async def test_run_stops_on_channel_error(self):
"""ChannelError should stop the channel permanently."""
@@ -878,6 +947,8 @@ class TestChannelReconnect:
ch.start = fatal_start
await ch.run()
assert ch._running is False
assert ch._startup_event.is_set()
assert ch._startup_error == "fatal"
class TestExtractRetryAfter:
@@ -1274,6 +1345,23 @@ class TestChannelManagerStatus:
ch._running = True
assert mgr.running_channels() == ["stub"]
def test_startup_results_report_fatal_error(self):
bus = MessageBus()
mgr = ChannelManager(bus)
ch = StubChannel()
mgr.register(ch)
ch._startup_error = "dependency missing"
ch._startup_event.set()
assert mgr.startup_results() == [("stub", False, "failed: dependency missing")]
def test_startup_results_do_not_assume_pending_channel_is_connected(self):
bus = MessageBus()
mgr = ChannelManager(bus)
mgr.register(StubChannel())
assert mgr.startup_results() == [("stub", False, "starting (bus)")]
def test_get_stats(self):
bus = MessageBus()
mgr = ChannelManager(bus)
+71
View File
@@ -32,6 +32,7 @@ def test_auto_start_channel_passes_send_thinking(monkeypatch):
captured["send_thinking"] = send_thinking
captured["thread_id"] = thread_id
captured["agent"] = agent
return [("telegram", True, "connected (bus)")]
monkeypatch.setattr(channel_cli, "_start_channels_bus_mode", _fake_start)
monkeypatch.setattr(channel_cli, "_print_channel_panel", lambda _rows: None)
@@ -52,3 +53,73 @@ def test_auto_start_channel_passes_send_thinking(monkeypatch):
assert captured["agent"] is agent
assert runtime.agent is agent
assert runtime.thread_id == "thread-1"
def test_auto_start_channel_reports_startup_failure(monkeypatch):
from EvoScientist.commands.base import ChannelRuntime
rows = [("telegram", False, "failed: dependency missing")]
rendered = []
monkeypatch.setattr(
channel_cli,
"_start_channels_bus_mode",
lambda *_args, **_kwargs: rows,
)
monkeypatch.setattr(channel_cli, "_print_channel_panel", rendered.append)
runtime = ChannelRuntime()
result = channel_cli._auto_start_channel(
object(),
"thread-1",
SimpleNamespace(channel_enabled="telegram"),
runtime=runtime,
)
assert result == rows
assert rendered == [rows]
assert runtime.agent is None
assert runtime.thread_id is None
def test_auto_start_channel_binds_runtime_while_starting(monkeypatch):
from EvoScientist.channels.channel_manager import CHANNEL_STARTUP_PENDING_DETAIL
from EvoScientist.commands.base import ChannelRuntime
rows = [("telegram", False, CHANNEL_STARTUP_PENDING_DETAIL)]
monkeypatch.setattr(
channel_cli,
"_start_channels_bus_mode",
lambda *_args, **_kwargs: rows,
)
monkeypatch.setattr(channel_cli, "_print_channel_panel", lambda _rows: None)
agent = object()
runtime = ChannelRuntime()
result = channel_cli._auto_start_channel(
agent,
"thread-1",
SimpleNamespace(channel_enabled="telegram"),
runtime=runtime,
)
assert result == rows
assert runtime.agent is agent
assert runtime.thread_id == "thread-1"
def test_get_channel_startup_results_without_manager():
channel_cli._manager = None
assert channel_cli.get_channel_startup_results() == []
def test_get_channel_startup_results_uses_manager_snapshot():
rows = [("telegram", True, "connected (bus)")]
class Manager:
def startup_results(self):
return rows
channel_cli._manager = Manager()
assert channel_cli.get_channel_startup_results() is rows
+13 -1
View File
@@ -9,13 +9,14 @@ allowlist (``task()`` stays reachable as the REPL global, with responseSchema).
from __future__ import annotations
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from EvoScientist.middleware.code_interpreter import (
_DEFAULT_PTC_ALLOWLIST,
aclose_code_interpreters,
create_code_interpreter_middleware,
)
@@ -51,6 +52,17 @@ def test_create_code_interpreter_middleware_builds():
assert create_code_interpreter_middleware() is not None
@pytest.mark.asyncio
async def test_aclose_code_interpreters_closes_registered_instances(monkeypatch):
middleware = create_code_interpreter_middleware()
close = AsyncMock()
monkeypatch.setattr(middleware, "aclose", close)
await aclose_code_interpreters()
close.assert_awaited_once_with()
def test_middleware_uses_thread_mode():
"""Upstream ``mode="thread"`` (the default) preserves cross-turn REPL
state as ``langchain-ai/deepagents#3064`` shipped it. The wire-cost
+28
View File
@@ -7,6 +7,7 @@ from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import typer
from langchain_core.messages import AIMessage, HumanMessage
from EvoScientist.gateway import (
@@ -282,6 +283,33 @@ def test_cmd_run_passes_local_graph_gateway(monkeypatch):
assert seen["gateway"].thread_store is thread_store
def test_cmd_run_converts_stream_failure_to_controlled_exit(monkeypatch):
from EvoScientist.cli import interactive
runtime_gateways = RuntimeGateways(
thread_store=FakeThreadStore(),
graph_gateway=FakeGraphGateway(),
)
provider_error = RuntimeError("provider unavailable")
monkeypatch.setattr(
interactive,
"run_streaming",
MagicMock(side_effect=provider_error),
)
with pytest.raises(typer.Exit) as exc_info:
interactive.cmd_run(
MagicMock(),
"hello",
thread_id="failed-thread",
show_thinking=False,
runtime_gateways=runtime_gateways,
)
assert exc_info.value.exit_code == 1
assert exc_info.value.__cause__ is provider_error
async def test_langgraph_server_thread_store_delegates_to_sdk_threads():
threads = FakeLangGraphThreadsClient(
threads=[
+111
View File
@@ -0,0 +1,111 @@
"""Contract tests for the middleware event sink.
Pins two things:
1. The protocol / :class:`NoOpSink` shape is stable and structural.
2. The threading/blocking contract: sinks may be called from any thread, and
a sink that blocks stalls its caller (nothing isolates a slow sink from the
emitting thread). The deliberately-slow fake sink documents this.
"""
from __future__ import annotations
import threading
import time
from EvoScientist.middleware.events import MiddlewareEventSink, NoOpSink
class _SlowSink:
"""A deliberately-slow, thread-safe sink used to exercise the contract.
Every event method sleeps ``delay`` seconds under a lock and records the
thread it was called on. A real frontend must NOT do this — it exists only
to demonstrate that a blocking sink holds the emitting thread.
"""
def __init__(self, delay: float) -> None:
self._delay = delay
self._lock = threading.Lock()
self.calls: list[tuple[str, threading.Thread]] = []
def _record(self, name: str) -> None:
time.sleep(self._delay)
with self._lock:
self.calls.append((name, threading.current_thread()))
def on_tool_selection_started(self, total_tools: int) -> None:
self._record("started")
def on_tool_selection(self, selected: list[str], total_tools: int) -> None:
self._record("selection")
def on_tool_selection_ended(self) -> None:
self._record("ended")
def emit_fallback_notice(self, text: str, style: str = "yellow") -> None:
self._record("notice")
def test_noopsink_satisfies_protocol():
sink = NoOpSink()
assert isinstance(sink, MiddlewareEventSink)
# Every event is a no-op and returns None regardless of arguments.
assert sink.on_tool_selection_started(10) is None
assert sink.on_tool_selection(["a", "b"], 10) is None
assert sink.on_tool_selection_ended() is None
assert sink.emit_fallback_notice("fallback notice") is None
def test_slow_sink_satisfies_protocol():
assert isinstance(_SlowSink(0.0), MiddlewareEventSink)
def test_sink_is_callable_from_any_thread():
"""Sink methods may be invoked from worker threads (the sync-hook world)."""
sink = _SlowSink(0.0)
main = threading.current_thread()
def _worker() -> None:
sink.on_tool_selection_started(3)
sink.on_tool_selection(["think_tool"], 3)
sink.on_tool_selection_ended()
t = threading.Thread(target=_worker)
t.start()
t.join()
names = [name for name, _ in sink.calls]
assert names == ["started", "selection", "ended"]
# All calls landed on the worker thread, not the caller's thread.
assert all(thread is not main for _, thread in sink.calls)
assert all(thread is t for _, thread in sink.calls)
def test_blocking_sink_stalls_the_emitting_thread():
"""A slow sink holds its caller: the contract requires non-blocking sinks.
This is the negative proof — nothing in the framework isolates the caller
from a blocking sink, so the emitting thread waits the full delay.
"""
delay = 0.2
sink = _SlowSink(delay)
start = time.perf_counter()
sink.emit_fallback_notice("fallback notice")
elapsed = time.perf_counter() - start
# The caller was blocked for at least the sink's delay.
assert elapsed >= delay
assert [name for name, _ in sink.calls] == ["notice"]
def test_noopsink_never_blocks():
sink = NoOpSink()
start = time.perf_counter()
for _ in range(10_000):
sink.on_tool_selection_started(50)
sink.emit_fallback_notice("fallback notice")
elapsed = time.perf_counter() - start
# 20k no-op calls are effectively free.
assert elapsed < 0.5
+3
View File
@@ -467,6 +467,7 @@ class TestApplyModelIntegration:
chat_model=None,
*,
on_mcp_progress=None,
events=None,
):
# The pure path threads the freshly built chat model in; bind it
# directly instead of re-deriving via _ensure_chat_model.
@@ -557,6 +558,7 @@ class TestApplyModelPreservesConfigByReference:
chat_model=None,
*,
on_mcp_progress=None,
events=None,
):
return MagicMock(name="fake-agent")
@@ -690,6 +692,7 @@ class TestApplyModelLoadAgentFailureTransactional:
chat_model=None,
*,
on_mcp_progress=None,
events=None,
):
# The pure path writes no globals; mimic a failure partway through
# agent wiring (middleware build, deepagents, MCP reconnect, ...).
+102 -24
View File
@@ -13,17 +13,21 @@ import pytest
from langchain_core.exceptions import ContextOverflowError
from langchain_core.messages import AIMessage, HumanMessage
from EvoScientist.middleware.events import NoOpSink
from EvoScientist.middleware.model_fallback import (
_guard_and_fallback,
_is_non_fallbackable,
_try_fallbacks,
add_fallback,
clear_fallbacks,
set_ui_emit,
)
from EvoScientist.stream.sink import SessionEventSink
# ── Helpers ──────────────────────────────────────────────────────
# Silent sink for tests that don't assert on the fallback narration.
_SINK = NoOpSink()
def _fake_request():
"""Build a minimal ModelRequest stub with an .override() method."""
@@ -38,12 +42,10 @@ AI_RESPONSE = AIMessage(content="ok")
@pytest.fixture(autouse=True)
def _clean_chain():
"""Ensure a clean fallback chain and no UI callback for every test."""
"""Ensure a clean fallback chain for every test."""
clear_fallbacks()
set_ui_emit(None)
yield
clear_fallbacks()
set_ui_emit(None)
# ═════════════════════════════════════════════════════════════════
@@ -154,7 +156,7 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, invoke, Exception("503 boom"))
result = await _try_fallbacks(req, invoke, Exception("503 boom"), _SINK)
assert result is AI_RESPONSE
invoke.assert_awaited_once()
@@ -177,7 +179,7 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, _invoke, Exception("503 boom"))
result = await _try_fallbacks(req, _invoke, Exception("503 boom"), _SINK)
assert result is AI_RESPONSE
assert call_count == 2
@@ -202,7 +204,7 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
with pytest.raises(Exception, match="429 from fb-b") as exc_info:
await _try_fallbacks(req, _invoke, Exception("503 primary"))
await _try_fallbacks(req, _invoke, Exception("503 primary"), _SINK)
assert exc_info.value is last_error
@@ -218,7 +220,7 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
with pytest.raises(Exception, match="context_length_exceeded"):
await _try_fallbacks(req, _invoke, Exception("503 primary"))
await _try_fallbacks(req, _invoke, Exception("503 primary"), _SINK)
# get_chat_model should only have been called once (for fb-a),
# fb-b should never be reached.
@@ -264,7 +266,12 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = fallback_model
with pytest.raises(ProviderStreamError) as exc_info:
await _try_fallbacks(req, _invoke, Exception("openai primary failed"))
await _try_fallbacks(
req,
_invoke,
Exception("openai primary failed"),
_SINK,
)
# Attribution flipped to moonshot (the failing fallback), not
# openai (the original request's model).
@@ -304,7 +311,12 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = model
with pytest.raises(InvalidUpdateError) as exc_info:
await _try_fallbacks(req, _invoke, Exception("primary failed"))
await _try_fallbacks(
req,
_invoke,
Exception("primary failed"),
_SINK,
)
assert exc_info.value is raised
@@ -322,7 +334,9 @@ class TestGuardAndFallback:
invoke = AsyncMock()
with pytest.raises(ContextOverflowError):
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
await _guard_and_fallback(
ContextOverflowError("overflow"), req, invoke, _SINK
)
invoke.assert_not_awaited()
@@ -349,7 +363,7 @@ class TestGuardAndFallback:
raised = ContextOverflowError("context length exceeded")
with pytest.raises(ContextOverflowError) as exc_info:
await _guard_and_fallback(raised, req, invoke)
await _guard_and_fallback(raised, req, invoke, _SINK)
assert exc_info.value is raised
invoke.assert_not_awaited()
@@ -361,7 +375,7 @@ class TestGuardAndFallback:
with pytest.raises(Exception, match="invalid_request_error"):
await _guard_and_fallback(
Exception("400: invalid_request_error"), req, invoke
Exception("400: invalid_request_error"), req, invoke, _SINK
)
invoke.assert_not_awaited()
@@ -373,7 +387,9 @@ class TestGuardAndFallback:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _guard_and_fallback(Exception("503 overloaded"), req, invoke)
result = await _guard_and_fallback(
Exception("503 overloaded"), req, invoke, _SINK
)
assert result is AI_RESPONSE
invoke.assert_awaited_once()
@@ -387,7 +403,7 @@ class TestGuardAndFallback:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _guard_and_fallback(
Exception("400 Bad Request: invalid_api_key"), req, invoke
Exception("400 Bad Request: invalid_api_key"), req, invoke, _SINK
)
assert result is AI_RESPONSE
@@ -400,35 +416,97 @@ class TestGuardAndFallback:
class TestUiEmit:
"""Verify that fallback events are surfaced via the registered callback."""
"""Verify that fallback narration reaches the injected frontend sink.
The fallback middleware sends its narration lines through the same
``fallback_display`` callback the frontend supplies, so capturing that
callback exercises the exact user-facing text.
"""
def _capturing_sink(self):
messages: list[tuple[str, str]] = []
sink = SessionEventSink(
fallback_display=lambda text, style: messages.append((text, style))
)
return sink, messages
async def test_emit_captures_messages(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
messages: list[tuple[str, str]] = []
set_ui_emit(lambda text, style: messages.append((text, style)))
sink, messages = self._capturing_sink()
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
await _try_fallbacks(req, invoke, Exception("503 down"))
await _try_fallbacks(req, invoke, Exception("503 down"), sink)
texts = [t for t, _ in messages]
assert any("Primary model failed" in t for t in texts)
assert any("Falling back to fb (prov)" in t for t in texts)
assert any("succeeded" in t for t in texts)
async def test_default_sink_prints_to_console(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
sink = SessionEventSink()
with (
patch("EvoScientist.llm.models.get_chat_model") as mock_gcm,
patch("EvoScientist.stream.sink.console.print") as mock_print,
):
mock_gcm.return_value = MagicMock()
await _try_fallbacks(req, invoke, Exception("503 down"), sink)
texts = [call.args[0] for call in mock_print.call_args_list]
assert any("Primary model failed" in t for t in texts)
assert any("Falling back to fb (prov)" in t for t in texts)
assert any("succeeded" in t for t in texts)
assert all(
call.kwargs == {"style": "yellow"} for call in mock_print.call_args_list[:2]
)
async def test_display_failure_does_not_abort_fallback(self, caplog):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
sink = SessionEventSink(
fallback_display=MagicMock(side_effect=RuntimeError("ui unavailable"))
)
with (
patch("EvoScientist.llm.models.get_chat_model") as mock_gcm,
patch("EvoScientist.stream.sink.console.print") as mock_print,
):
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, invoke, Exception("503 down"), sink)
assert result is AI_RESPONSE
invoke.assert_awaited_once()
assert "Fallback display callback failed" in caplog.text
texts = [call.args[0] for call in mock_print.call_args_list]
assert any("Primary model failed" in t for t in texts)
assert any("Falling back to fb (prov)" in t for t in texts)
assert any("succeeded" in t for t in texts)
def test_noopsink_keeps_fallback_notices_silent(self):
sink = NoOpSink()
with patch("EvoScientist.stream.sink.console.print") as mock_print:
sink.emit_fallback_notice("hidden")
mock_print.assert_not_called()
async def test_emit_shows_non_fallbackable_rejection(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
messages: list[tuple[str, str]] = []
set_ui_emit(lambda text, style: messages.append((text, style)))
sink, messages = self._capturing_sink()
with pytest.raises(ContextOverflowError):
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
await _guard_and_fallback(
ContextOverflowError("overflow"), req, invoke, sink
)
texts = [t for t, _ in messages]
assert any("not eligible for fallback" in t for t in texts)
+15
View File
@@ -29,6 +29,21 @@ def test_resolve_final_status_footer_keeps_footer_for_noninteractive():
assert resolve_final_status_footer(False, lambda: "footer") == "footer"
def test_final_display_respects_disabled_thinking():
"""A final-frame preference must not override the user visibility flag."""
renderable = create_streaming_display(
thinking_text="private reasoning",
show_thinking=False,
is_final=True,
final_show_thinking=True,
)
rendered = _render_text(renderable)
assert "private reasoning" not in rendered
assert "Thinking" not in rendered
def test_streaming_display_keeps_narration_visible_with_pending_memory_tool():
"""Profile-memory reads still block, while lead-in text remains visible."""
narration = "Here is the answer."
+181 -51
View File
@@ -1,6 +1,7 @@
"""Tests for EvoScientist/stream/events.py helpers."""
import asyncio
from types import SimpleNamespace
import pytest
from deepagents import create_deep_agent
@@ -26,6 +27,7 @@ from tests.stream_v3_fakes import (
FakeV3Agent,
HangingV3Agent,
SubscriptionSensitiveV3Agent,
async_iter,
collect_events,
message_delta,
message_finish,
@@ -410,14 +412,53 @@ class TestV3ProtocolStreaming:
async def test_tool_selector_reasoning_delta_is_suppressed(self):
"""Selector reasoning must not appear as main-agent thinking."""
import EvoScientist.middleware.tool_selector as selector_mod
from EvoScientist.stream.sink import SessionEventSink
original_active = selector_mod._selector_active
selector_mod._selector_active = True
try:
agent = FakeV3Agent(
[
protocol_event(
sink = SessionEventSink()
sink.on_tool_selection_started(30) # selector call in flight
agent = FakeV3Agent(
[
protocol_event(
"messages",
(
{
"event": "content-block-delta",
"index": 0,
"delta": {
"type": "reasoning-delta",
"reasoning": "selector-only thought",
},
},
{},
),
)
]
)
events = await collect_events(agent, events=sink)
assert not any(
e.get("type") == "thinking" and e.get("content") == "selector-only thought"
for e in events
)
async def test_default_run_scoped_sink_suppresses_selector_reasoning(self):
"""Default main-agent middleware reports into the current stream sink."""
from EvoScientist.middleware.events import RunScopedEventSink
middleware_events = RunScopedEventSink()
class Run:
def __init__(self):
self.subagents = async_iter([])
self.aborted = False
def __aiter__(self):
return self._iter_events()
async def _iter_events(self):
middleware_events.on_tool_selection_started(30)
try:
yield protocol_event(
"messages",
(
{
@@ -425,38 +466,134 @@ class TestV3ProtocolStreaming:
"index": 0,
"delta": {
"type": "reasoning-delta",
"reasoning": "selector-only thought",
"reasoning": "default selector thought",
},
},
{},
),
)
]
)
events = await collect_events(agent)
finally:
selector_mod._selector_active = original_active
finally:
middleware_events.on_tool_selection_ended()
async def abort(self):
self.aborted = True
class Agent:
async def aget_state(self, _config):
return SimpleNamespace(values={})
def astream_events(self, *_args, **_kwargs):
return Run()
events = await collect_events(Agent())
assert not any(
e.get("type") == "thinking" and e.get("content") == "selector-only thought"
e.get("type") == "thinking"
and e.get("content") == "default selector thought"
for e in events
)
@pytest.mark.filterwarnings(
"ignore:The v3 streaming protocol on Pregel is experimental"
)
async def test_sync_middleware_event_reaches_bound_stream_sink_via_executor(self):
"""LangChain executor context carries the active stream binding."""
from langchain.agents.middleware.types import AgentMiddleware
from langchain_core.runnables.config import run_in_executor
from EvoScientist.middleware.events import RunScopedEventSink
from EvoScientist.stream.sink import SessionEventSink
class SyncSelectionProbeMiddleware(AgentMiddleware):
name = "sync_selection_probe"
def __init__(self):
super().__init__()
self.called = False
self.events = RunScopedEventSink()
def wrap_model_call(self, request, handler):
self.called = True
self.events.on_tool_selection_started(2)
self.events.on_tool_selection(["probe_tool"], 2)
try:
return handler(request)
finally:
self.events.on_tool_selection_ended()
middleware = SyncSelectionProbeMiddleware()
sink = SessionEventSink()
inner_agent = create_deep_agent(
model=_ToolCallingFakeModel(responses=[AIMessage(content="inner answer")]),
tools=[],
system_prompt="Answer directly.",
middleware=[middleware],
)
class ExecutorBackedAgent:
async def aget_state(self, _config):
return SimpleNamespace(values={})
def astream_events(self, astream_input, config, **_kwargs):
return ExecutorBackedRun(astream_input, config)
class ExecutorBackedRun:
def __init__(self, astream_input, config):
self._astream_input = astream_input
self._config = config
self.subagents = async_iter([])
def __aiter__(self):
return self._iter_events()
async def _iter_events(self):
await run_in_executor(
None,
lambda: inner_agent.invoke(
self._astream_input,
config=self._config,
),
)
yield protocol_event(
"messages", (AIMessage(content="final answer"), {})
)
async def abort(self):
pass
agent = ExecutorBackedAgent()
events = [
event
async for event in stream_agent_events(
agent,
"answer",
"live-deepagents-sync-contextvar",
events=sink,
)
]
assert middleware.called is True
assert any(
event.get("type") == "done" and event.get("content") == "final answer"
for event in events
)
assert sink.tool_selection_active is False
assert sink.tool_selection_pending() is True
assert sink.consume_tool_selection() == (True, ["probe_tool"])
async def test_tool_selector_whole_message_reasoning_is_suppressed(self):
"""Selector reasoning in whole-message payloads is also hidden."""
import EvoScientist.middleware.tool_selector as selector_mod
from EvoScientist.stream.sink import SessionEventSink
original_active = selector_mod._selector_active
selector_mod._selector_active = True
try:
message = AIMessage(
additional_kwargs={"reasoning_content": "selector whole thought"},
content="",
)
agent = FakeV3Agent([protocol_event("messages", (message, {}))])
events = await collect_events(agent)
finally:
selector_mod._selector_active = original_active
sink = SessionEventSink()
sink.on_tool_selection_started(30)
message = AIMessage(
additional_kwargs={"reasoning_content": "selector whole thought"},
content="",
)
agent = FakeV3Agent([protocol_event("messages", (message, {}))])
events = await collect_events(agent, events=sink)
assert not any(
e.get("type") == "thinking" and e.get("content") == "selector whole thought"
@@ -779,32 +916,25 @@ class TestV3ProtocolStreaming:
async def test_tool_selection_flushes_before_tool_only_step(self):
"""Selector UI event is emitted even when selection is followed only by a tool."""
import EvoScientist.middleware.tool_selector as selector_mod
from EvoScientist.stream.sink import SessionEventSink
original_selected = selector_mod._current_selected_tools
original_total = selector_mod._total_tools_count
original_last = selector_mod._last_emitted_tools
selector_mod._current_selected_tools = ["read_file"]
selector_mod._total_tools_count = 3
selector_mod._last_emitted_tools = []
try:
output = ToolMessage(
content="File content",
name="read_file",
tool_call_id="tc1",
)
agent = FakeV3Agent(
[
message_delta('{"tools":["read_file"]}'),
tool_started("read_file", {"path": "notes.txt"}),
tool_finished(output),
]
)
events = await collect_events(agent)
finally:
selector_mod._current_selected_tools = original_selected
selector_mod._total_tools_count = original_total
selector_mod._last_emitted_tools = original_last
# The frontend sink holds a pending selection (1 of 3 tools) — the
# suppressor must surface it before the tool-only step.
sink = SessionEventSink()
sink.on_tool_selection(["read_file"], 3)
output = ToolMessage(
content="File content",
name="read_file",
tool_call_id="tc1",
)
agent = FakeV3Agent(
[
message_delta('{"tools":["read_file"]}'),
tool_started("read_file", {"path": "notes.txt"}),
tool_finished(output),
]
)
events = await collect_events(agent, events=sink)
event_types = [e["type"] for e in events]
assert event_types.index("tool_selection") < event_types.index("tool_call")
+44
View File
@@ -1,5 +1,8 @@
"""Tests for Telegram channel implementation."""
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from EvoScientist.channels.base import ChannelError
@@ -42,6 +45,47 @@ class TestTelegramChannel:
channel = TelegramChannel(config)
await channel.stop()
async def test_cleanup_is_idempotent(self):
channel = TelegramChannel(TelegramConfig(bot_token="test"))
app = SimpleNamespace(
updater=SimpleNamespace(running=True, stop=AsyncMock()),
running=True,
shutdown=AsyncMock(),
)
async def stop_once():
if not app.running:
raise RuntimeError("This Application is not running!")
app.running = False
app.stop = AsyncMock(side_effect=stop_once)
channel._app = app
await channel._cleanup()
await channel._cleanup()
app.updater.stop.assert_awaited_once()
app.stop.assert_awaited_once()
app.shutdown.assert_awaited_once()
assert channel._app is None
async def test_cleanup_partially_initialized_application(self):
channel = TelegramChannel(TelegramConfig(bot_token="test"))
app = SimpleNamespace(
updater=SimpleNamespace(running=False, stop=AsyncMock()),
running=False,
stop=AsyncMock(),
shutdown=AsyncMock(),
)
channel._app = app
await channel._cleanup()
app.updater.stop.assert_not_awaited()
app.stop.assert_not_awaited()
app.shutdown.assert_awaited_once()
assert channel._app is None
async def test_send_returns_false_without_app(self):
from EvoScientist.channels.base import OutboundMessage
+177 -64
View File
@@ -1,16 +1,45 @@
"""Tests for LLMToolSelectorMiddleware integration."""
"""Tests for LLMToolSelectorMiddleware integration and the event-sink handoff."""
from typing import Any
from unittest.mock import MagicMock, patch
import pytest
from langchain.agents.middleware.types import ModelRequest
from langchain_core.tools import BaseTool, StructuredTool
from EvoScientist.middleware.tool_selector import (
_ConditionalToolSelectorMiddleware,
_ToolSelectionTrackerMiddleware,
create_tool_selector_middleware,
)
from EvoScientist.stream.emitter import StreamEventEmitter
from EvoScientist.stream.sink import SessionEventSink
from EvoScientist.stream.tool_selection import _ToolSelectionSuppressor
class _RecordingSink:
"""Records selection lifecycle calls for assertions."""
def __init__(self) -> None:
self.calls: list[tuple] = []
self.active = False
def on_tool_selection_started(self, total_tools: int) -> None:
self.active = True
self.calls.append(("started", total_tools))
def on_tool_selection(self, selected: list[str], total_tools: int) -> None:
self.calls.append(("selection", list(selected), total_tools))
def on_tool_selection_ended(self) -> None:
self.active = False
self.calls.append(("ended",))
def emit_fallback_notice(self, text: str, style: str = "yellow") -> None:
pass
@property
def tool_selection_active(self) -> bool:
return self.active
def _tool(name: str) -> BaseTool:
@@ -36,17 +65,6 @@ def _mock_model():
return m
def _patched_create():
"""Create tool selector middleware without real LLM init."""
return [
_ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=MagicMock()),
threshold=20,
),
_ToolSelectionTrackerMiddleware(),
]
# Helper: patches needed to call create_tool_selector_middleware without LLM
def _factory_patches():
return (
@@ -68,12 +86,13 @@ def _factory_patches():
# ---------------------------------------------------------------------------
def test_create_tool_selector_returns_list():
def test_create_tool_selector_returns_single_middleware():
p1, p2, p3 = _factory_patches()
with p1, p2, p3:
result = create_tool_selector_middleware()
assert isinstance(result, list)
assert len(result) == 2
assert len(result) == 1
assert type(result[0]).__name__ == "_ConditionalToolSelectorMiddleware"
def test_create_tool_selector_always_include():
@@ -106,17 +125,19 @@ def test_custom_threshold():
# ---------------------------------------------------------------------------
# Conditional + tracker unit tests
# Conditional selector unit tests
# ---------------------------------------------------------------------------
def test_conditional_skips_below_threshold():
"""When tools <= threshold, selector is skipped."""
"""When tools <= threshold, selector is skipped and nothing is reported."""
mock_selector = MagicMock()
selector_factory = MagicMock(return_value=mock_selector)
sink = _RecordingSink()
cond = _ConditionalToolSelectorMiddleware(
selector_factory=selector_factory,
threshold=10,
events=sink,
)
request = MagicMock()
@@ -127,6 +148,7 @@ def test_conditional_skips_below_threshold():
handler.assert_called_once_with(request)
selector_factory.assert_not_called()
mock_selector.wrap_model_call.assert_not_called()
assert sink.calls == [] # no selection ran → no events
def test_conditional_runs_above_threshold():
@@ -148,58 +170,107 @@ def test_conditional_runs_above_threshold():
handler.assert_not_called()
def test_selector_active_flag():
"""_selector_active flag is True during selection, False after."""
import EvoScientist.middleware.tool_selector as ts_mod
mock_selector = MagicMock()
def test_selection_lifecycle_reported_to_sink():
"""started(total) → selection(selected, total) → ended, reported to the sink."""
# The fake selector filters the request down to two named tools before
# calling the downstream handler.
filtered = _request([_tool("read_file"), _tool("think_tool")])
def fake_selector_call(request, handler):
assert ts_mod._selector_active is True
return handler(request)
return handler(filtered)
mock_selector = MagicMock()
mock_selector.wrap_model_call.side_effect = fake_selector_call
sink = _RecordingSink()
cond = _ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=mock_selector),
threshold=5,
events=sink,
)
request = MagicMock()
request.tools = [MagicMock() for _ in range(10)]
handler = MagicMock()
request = _request([_tool(f"t{i}") for i in range(10)])
cond.wrap_model_call(request, MagicMock())
cond.wrap_model_call(request, handler)
assert ts_mod._selector_active is False
assert sink.calls == [
("started", 10),
("selection", ["read_file", "think_tool"], 10),
("ended",),
]
def test_selector_can_disable_stream_tracking():
"""Selection can run without touching the main-agent stream/UI globals."""
import EvoScientist.middleware.tool_selector as ts_mod
def test_selector_failure_reports_ended_without_selection():
"""A selector that raises before the handler surfaces no selection event."""
mock_selector = MagicMock()
def fake_selector_call(request, handler):
assert ts_mod._selector_active is False
return handler(request)
mock_selector.wrap_model_call.side_effect = fake_selector_call
mock_selector.wrap_model_call.side_effect = RuntimeError("no structured output")
sink = _RecordingSink()
cond = _ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=mock_selector),
threshold=5,
track_stream_selection=False,
events=sink,
)
ts_mod._total_tools_count = 99
request = MagicMock()
request.tools = [MagicMock() for _ in range(10)]
request = _request([_tool(f"t{i}") for i in range(10)])
handler = MagicMock()
cond.wrap_model_call(request, handler)
# Falls back to all tools; only started/ended reported, no selection.
handler.assert_called_once_with(request)
assert ("started", 10) in sink.calls
assert not any(c[0] == "selection" for c in sink.calls)
assert sink.calls[-1] == ("ended",)
def test_selector_failure_ends_before_sync_fallback_handler():
"""All-tools fallback must not run while selector suppression is active."""
mock_selector = MagicMock()
mock_selector.wrap_model_call.side_effect = RuntimeError("no structured output")
sink = _RecordingSink()
cond = _ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=mock_selector),
threshold=5,
events=sink,
)
request = _request([_tool(f"t{i}") for i in range(10)])
def handler(req):
sink.calls.append(("handler", sink.tool_selection_active))
return MagicMock()
cond.wrap_model_call(request, handler)
mock_selector.wrap_model_call.assert_called_once()
handler.assert_called_once()
assert ts_mod._selector_active is False
assert ts_mod._total_tools_count == 99
assert sink.calls == [
("started", 10),
("ended",),
("handler", False),
]
@pytest.mark.asyncio
async def test_selector_failure_ends_before_async_fallback_handler():
"""Async all-tools fallback must see selection already closed."""
mock_selector = MagicMock()
mock_selector.awrap_model_call.side_effect = RuntimeError("no structured output")
sink = _RecordingSink()
cond = _ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=mock_selector),
threshold=5,
events=sink,
)
request = _request([_tool(f"t{i}") for i in range(10)])
async def handler(req):
sink.calls.append(("handler", sink.tool_selection_active))
return MagicMock()
await cond.awrap_model_call(request, handler)
assert sink.calls == [
("started", 10),
("ended",),
("handler", False),
]
def test_selector_always_includes_available_memory_tools():
@@ -278,22 +349,56 @@ def test_selector_resolved_once_across_repeated_requests():
assert mock_selector.wrap_model_call.call_count == 3
def test_tracker_captures_tools():
"""Tracker middleware captures tool names from request."""
tracker = _ToolSelectionTrackerMiddleware()
tool1 = _tool("read_file")
tool2 = _tool("execute")
# ---------------------------------------------------------------------------
# R1: consume-once + dedup render sequences (sink + suppressor)
# ---------------------------------------------------------------------------
request = MagicMock()
request.tools = [tool1, tool2]
handler = MagicMock()
tracker.wrap_model_call(request, handler)
handler.assert_called_once_with(request)
def _drive_selection(sink, suppressor, selected, total):
"""Mimic one selection turn: sink records it, the suppressor observes the
selector JSON block, then a flush surfaces (or not) the UI event."""
sink.on_tool_selection_started(total)
sink.on_tool_selection(selected, total)
sink.on_tool_selection_ended()
# Suppressor observes the selector's structured-output tool block.
suppressor.observe_tool_block("ToolSelectionResponse")
return suppressor.flush_selection()
import EvoScientist.middleware.tool_selector as ts_mod
assert ts_mod._current_selected_tools == ["read_file", "execute"]
def test_render_sequences_table():
"""select → render; same selection again → no repeat; new selection → render."""
cases = [
# (label, selected, total, expect_render)
("first selection renders", ["read_file", "think_tool"], 5, True),
("same selection again does not repeat", ["read_file", "think_tool"], 5, False),
("new selection renders", ["execute", "think_tool"], 5, True),
("kept-all selection does not render", ["a", "b", "c"], 3, False),
]
sink = SessionEventSink()
suppressor = _ToolSelectionSuppressor(StreamEventEmitter(), sink)
for label, selected, total, expect_render in cases:
events = _drive_selection(sink, suppressor, selected, total)
rendered = [e for e in events if e.get("type") == "tool_selection"]
if expect_render:
assert rendered, f"{label}: expected a tool_selection event"
assert rendered[0]["tools"] == selected, label
else:
assert not rendered, f"{label}: expected no tool_selection event"
def test_consume_is_once_only():
"""A pending selection renders once; a second flush yields nothing."""
sink = SessionEventSink()
suppressor = _ToolSelectionSuppressor(StreamEventEmitter(), sink)
first = _drive_selection(sink, suppressor, ["read_file"], 3)
assert any(e.get("type") == "tool_selection" for e in first)
# No new selection recorded; the observation flag was consumed.
suppressor.observe_tool_block("ToolSelectionResponse")
second = suppressor.flush_selection()
assert not any(e.get("type") == "tool_selection" for e in second)
# ---------------------------------------------------------------------------
@@ -303,7 +408,12 @@ def test_tracker_captures_tools():
@patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=lambda *a, **kw: _patched_create(),
side_effect=lambda *a, **kw: [
_ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=MagicMock()),
threshold=20,
)
],
)
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
@@ -321,7 +431,6 @@ def test_default_middleware_includes_tool_selector(mock_config, mock_model, mock
mw = _get_default_middleware()
type_names = [type(m).__name__ for m in mw]
assert "_ConditionalToolSelectorMiddleware" in type_names
assert "_ToolSelectionTrackerMiddleware" in type_names
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@@ -339,7 +448,12 @@ def test_subagent_no_tool_selector(mock_model):
@patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=lambda *a, **kw: _patched_create(),
side_effect=lambda *a, **kw: [
_ConditionalToolSelectorMiddleware(
selector_factory=MagicMock(return_value=MagicMock()),
threshold=20,
)
],
)
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
@@ -359,7 +473,6 @@ def test_tool_selector_ordering(mock_config, mock_model, mock_ts):
type_names = [type(m).__name__ for m in mw]
ts_idx = type_names.index("_ConditionalToolSelectorMiddleware")
tracker_idx = type_names.index("_ToolSelectionTrackerMiddleware")
te_idx = type_names.index("ToolErrorHandlerMiddleware")
mem_idx = type_names.index("EvoMemoryMiddleware")
assert te_idx < ts_idx < tracker_idx < mem_idx
assert te_idx < ts_idx < mem_idx
+85
View File
@@ -0,0 +1,85 @@
from __future__ import annotations
import asyncio
import threading
from types import SimpleNamespace
import pytest
pytest.importorskip("textual")
from EvoScientist.cli import tui_interactive as tui_mod
from EvoScientist.commands.base import ChannelRuntime
async def test_channel_startup_worker_keeps_event_loop_responsive(monkeypatch):
started = threading.Event()
release = threading.Event()
finished = threading.Event()
worker_thread: list[int] = []
rows = [("telegram", True, "connected (bus)")]
def blocking_start(*_args, **_kwargs):
worker_thread.append(threading.get_ident())
started.set()
release.wait(timeout=2.0)
finished.set()
return rows
monkeypatch.setattr(tui_mod, "_auto_start_channel", blocking_start)
main_thread = threading.get_ident()
task = asyncio.create_task(
tui_mod._auto_start_channel_in_worker(
object(),
"thread-1",
SimpleNamespace(channel_enabled="telegram"),
send_thinking=False,
runtime=ChannelRuntime(),
stop_requested=threading.Event(),
)
)
for _ in range(100):
if started.is_set():
break
await asyncio.sleep(0.01)
try:
assert started.is_set()
assert finished.is_set() is False
assert len(worker_thread) == 1
assert worker_thread[0] != main_thread
finally:
release.set()
assert await task == rows
assert finished.is_set()
async def test_channel_startup_worker_stops_channels_after_exit(monkeypatch):
runtime = ChannelRuntime()
stop_requested = threading.Event()
stop_requested.set()
stopped_with: list[ChannelRuntime | None] = []
monkeypatch.setattr(
tui_mod,
"_auto_start_channel",
lambda *_args, **_kwargs: [("telegram", False, "starting (bus)")],
)
monkeypatch.setattr(
tui_mod,
"_channels_stop",
lambda _channel_type=None, *, runtime=None: stopped_with.append(runtime),
)
await tui_mod._auto_start_channel_in_worker(
object(),
"thread-1",
SimpleNamespace(channel_enabled="telegram"),
send_thinking=False,
runtime=runtime,
stop_requested=stop_requested,
)
assert stopped_with == [runtime]
+24
View File
@@ -24,6 +24,7 @@ class _StubApp:
self._agent_loader = _Loader()
self._conversation_tid = "thread-1"
self._channel_runtime = ChannelRuntime()
self._exiting = False
self.model_updates: list[tuple[str, str | None]] = []
self.refresh_calls: list[bool] = []
@@ -91,6 +92,29 @@ async def test_sync_tui_command_completion_refreshes_without_agent_swap(monkeypa
assert app.refresh_calls == [True]
async def test_sync_tui_command_completion_skips_unmounted_app(monkeypatch):
import EvoScientist.cli.tui_interactive as tui_mod
app = _StubApp()
app._exiting = True
ctx = CommandContext(
agent="new-agent",
thread_id="thread-1",
ui=SimpleNamespace(),
)
cmd = SimpleNamespace(name="/exit")
monkeypatch.setattr(tui_mod, "_channels_is_running", lambda: True)
await tui_mod._sync_tui_command_completion(app, ctx, "old-agent", cmd)
assert app._agent_loader.adopt_calls == []
assert app.model_updates == []
assert app.refresh_calls == []
assert app._channel_runtime.agent is None
assert app._channel_runtime.thread_id is None
async def test_sync_tui_rebinds_runtime_on_thread_rotation_without_agent_swap(
monkeypatch,
):