diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 7239472..5ca381c 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -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 diff --git a/EvoScientist/channels/base.py b/EvoScientist/channels/base.py index bd0ad61..7478adf 100644 --- a/EvoScientist/channels/base.py +++ b/EvoScientist/channels/base.py @@ -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: diff --git a/EvoScientist/channels/channel_manager.py b/EvoScientist/channels/channel_manager.py index c548a2f..bf010c5 100644 --- a/EvoScientist/channels/channel_manager.py +++ b/EvoScientist/channels/channel_manager.py @@ -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 { diff --git a/EvoScientist/channels/telegram/channel.py b/EvoScientist/channels/telegram/channel.py index e29059a..a6869f4 100644 --- a/EvoScientist/channels/telegram/channel.py +++ b/EvoScientist/channels/telegram/channel.py @@ -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) ──────────────────────────── diff --git a/EvoScientist/cli/agent.py b/EvoScientist/cli/agent.py index ce12fd4..bbf4147 100644 --- a/EvoScientist/cli/agent.py +++ b/EvoScientist/cli/agent.py @@ -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, ) diff --git a/EvoScientist/cli/async_notifier.py b/EvoScientist/cli/async_notifier.py index cec1ef2..69fbceb 100644 --- a/EvoScientist/cli/async_notifier.py +++ b/EvoScientist/cli/async_notifier.py @@ -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. diff --git a/EvoScientist/cli/channel.py b/EvoScientist/cli/channel.py index a8dcc13..f0857a7 100644 --- a/EvoScientist/cli/channel.py +++ b/EvoScientist/cli/channel.py @@ -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 diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index 2b95042..b7fcf32 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -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: diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index a79c050..1782a49 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -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( diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index cdfac47..3523970 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -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 diff --git a/EvoScientist/commands/implementation/model.py b/EvoScientist/commands/implementation/model.py index 0172650..348d2b5 100644 --- a/EvoScientist/commands/implementation/model.py +++ b/EvoScientist/commands/implementation/model.py @@ -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") diff --git a/EvoScientist/gateway/local.py b/EvoScientist/gateway/local.py index aac0096..4473eb6 100644 --- a/EvoScientist/gateway/local.py +++ b/EvoScientist/gateway/local.py @@ -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: diff --git a/EvoScientist/gateway/runtime.py b/EvoScientist/gateway/runtime.py index 5736bc9..438a441 100644 --- a/EvoScientist/gateway/runtime.py +++ b/EvoScientist/gateway/runtime.py @@ -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), ) diff --git a/EvoScientist/gateway/server.py b/EvoScientist/gateway/server.py index 31b14df..04bae63 100644 --- a/EvoScientist/gateway/server.py +++ b/EvoScientist/gateway/server.py @@ -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 diff --git a/EvoScientist/gateway/types.py b/EvoScientist/gateway/types.py index 65c30fc..e68d853 100644 --- a/EvoScientist/gateway/types.py +++ b/EvoScientist/gateway/types.py @@ -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, diff --git a/EvoScientist/middleware/async_watcher.py b/EvoScientist/middleware/async_watcher.py index bb18871..aeb7884 100644 --- a/EvoScientist/middleware/async_watcher.py +++ b/EvoScientist/middleware/async_watcher.py @@ -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"], diff --git a/EvoScientist/middleware/background.py b/EvoScientist/middleware/background.py index d1c6ee6..ed10540 100644 --- a/EvoScientist/middleware/background.py +++ b/EvoScientist/middleware/background.py @@ -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, + ] diff --git a/EvoScientist/middleware/code_interpreter.py b/EvoScientist/middleware/code_interpreter.py index 467cb77..49b1222 100644 --- a/EvoScientist/middleware/code_interpreter.py +++ b/EvoScientist/middleware/code_interpreter.py @@ -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 diff --git a/EvoScientist/middleware/events.py b/EvoScientist/middleware/events.py new file mode 100644 index 0000000..fc47fa9 --- /dev/null +++ b/EvoScientist/middleware/events.py @@ -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) diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py index 29a0de7..4deb2ac 100644 --- a/EvoScientist/middleware/model_fallback.py +++ b/EvoScientist/middleware/model_fallback.py @@ -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) diff --git a/EvoScientist/middleware/notifier.py b/EvoScientist/middleware/notifier.py new file mode 100644 index 0000000..ff195e1 --- /dev/null +++ b/EvoScientist/middleware/notifier.py @@ -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). + """ + ... diff --git a/EvoScientist/middleware/tool_selector.py b/EvoScientist/middleware/tool_selector.py index b498d0f..8cb9298 100644 --- a/EvoScientist/middleware/tool_selector.py +++ b/EvoScientist/middleware/tool_selector.py @@ -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 diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index 164afeb..dda0c10 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -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() diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index f22c161..5143c10 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -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 diff --git a/EvoScientist/stream/sink.py b/EvoScientist/stream/sink.py new file mode 100644 index 0000000..08e24bc --- /dev/null +++ b/EvoScientist/stream/sink.py @@ -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) diff --git a/EvoScientist/stream/tool_selection.py b/EvoScientist/stream/tool_selection.py index 3029e4f..4d1a84c 100644 --- a/EvoScientist/stream/tool_selection.py +++ b/EvoScientist/stream/tool_selection.py @@ -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 [] diff --git a/tests/conftest.py b/tests/conftest.py index 1e46d61..96f7679 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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.""" diff --git a/tests/stream_v3_fakes.py b/tests/stream_v3_fakes.py index a3e747f..4269fc8 100644 --- a/tests/stream_v3_fakes.py +++ b/tests/stream_v3_fakes.py @@ -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( diff --git a/tests/test_async_subagent_factory.py b/tests/test_async_subagent_factory.py index f8e0c3c..0d00d5c 100644 --- a/tests/test_async_subagent_factory.py +++ b/tests/test_async_subagent_factory.py @@ -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) diff --git a/tests/test_async_watcher_middleware.py b/tests/test_async_watcher_middleware.py index 09295c6..49b3703 100644 --- a/tests/test_async_watcher_middleware.py +++ b/tests/test_async_watcher_middleware.py @@ -99,7 +99,8 @@ def _make_middleware(): "url": "http://x", "graph_id": "writing-agent", } - } + }, + notifier=async_notifier, ) return mw, fake_client diff --git a/tests/test_background_middleware.py b/tests/test_background_middleware.py index 38b55b6..e8db5b2 100644 --- a/tests/test_background_middleware.py +++ b/tests/test_background_middleware.py @@ -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({}) diff --git a/tests/test_channel_comprehensive.py b/tests/test_channel_comprehensive.py index 274c29a..57f0f73 100644 --- a/tests/test_channel_comprehensive.py +++ b/tests/test_channel_comprehensive.py @@ -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) diff --git a/tests/test_cli_channel_bus_mode.py b/tests/test_cli_channel_bus_mode.py index a8a35a3..7c90a20 100644 --- a/tests/test_cli_channel_bus_mode.py +++ b/tests/test_cli_channel_bus_mode.py @@ -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 diff --git a/tests/test_code_interpreter_middleware.py b/tests/test_code_interpreter_middleware.py index 8baa89d..68fddf5 100644 --- a/tests/test_code_interpreter_middleware.py +++ b/tests/test_code_interpreter_middleware.py @@ -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 diff --git a/tests/test_graph_gateway.py b/tests/test_graph_gateway.py index ebd2859..c910408 100644 --- a/tests/test_graph_gateway.py +++ b/tests/test_graph_gateway.py @@ -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=[ diff --git a/tests/test_middleware_event_sink.py b/tests/test_middleware_event_sink.py new file mode 100644 index 0000000..b54c0a6 --- /dev/null +++ b/tests/test_middleware_event_sink.py @@ -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 diff --git a/tests/test_model_command.py b/tests/test_model_command.py index 35141f7..6af2586 100644 --- a/tests/test_model_command.py +++ b/tests/test_model_command.py @@ -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, ...). diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index 30f1a03..66a1f8b 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -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) diff --git a/tests/test_stream_display.py b/tests/test_stream_display.py index b162dfb..049e370 100644 --- a/tests/test_stream_display.py +++ b/tests/test_stream_display.py @@ -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." diff --git a/tests/test_stream_events.py b/tests/test_stream_events.py index c1d3491..a2f6525 100644 --- a/tests/test_stream_events.py +++ b/tests/test_stream_events.py @@ -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") diff --git a/tests/test_telegram_channel.py b/tests/test_telegram_channel.py index af1a461..83e58a6 100644 --- a/tests/test_telegram_channel.py +++ b/tests/test_telegram_channel.py @@ -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 diff --git a/tests/test_tool_selector_middleware.py b/tests/test_tool_selector_middleware.py index 60687a3..b2496ad 100644 --- a/tests/test_tool_selector_middleware.py +++ b/tests/test_tool_selector_middleware.py @@ -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 diff --git a/tests/test_tui_channel_startup.py b/tests/test_tui_channel_startup.py new file mode 100644 index 0000000..f1fb6c0 --- /dev/null +++ b/tests/test_tui_channel_startup.py @@ -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] diff --git a/tests/test_tui_command_sync.py b/tests/test_tui_command_sync.py index 43015cd..2a42eb6 100644 --- a/tests/test_tui_command_sync.py +++ b/tests/test_tui_command_sync.py @@ -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, ):