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:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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) ────────────────────────────
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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).
|
||||
"""
|
||||
...
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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 []
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -99,7 +99,8 @@ def _make_middleware():
|
||||
"url": "http://x",
|
||||
"graph_id": "writing-agent",
|
||||
}
|
||||
}
|
||||
},
|
||||
notifier=async_notifier,
|
||||
)
|
||||
return mw, fake_client
|
||||
|
||||
|
||||
@@ -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({})
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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=[
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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,
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user