diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index d9638af..8a5fc55 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -42,6 +42,7 @@ if TYPE_CHECKING: from langgraph.graph.state import CompiledStateGraph from .middleware.events import MiddlewareEventSink + from .runtime import AsyncRuntime # ============================================================================= # Constants @@ -246,7 +247,11 @@ def _load_mcp_config_once() -> tuple[str, dict]: return sig, cfg -def _load_mcp_tools_cached(on_progress=None) -> dict[str, list]: +def _load_mcp_tools_cached( + on_progress=None, + *, + runtime: "AsyncRuntime | None" = None, +) -> dict[str, list]: """Load MCP tools with config-aware caching. Args: @@ -267,7 +272,11 @@ def _load_mcp_tools_cached(on_progress=None) -> dict[str, list]: if _MCP_TOOLS_CACHE_KEY == cfg_key and _MCP_TOOLS_CACHE_VALUE is not None: return {k: list(v) for k, v in _MCP_TOOLS_CACHE_VALUE.items()} - loaded = load_mcp_tools(config=cfg, on_progress=on_progress) + loaded = load_mcp_tools( + config=cfg, + on_progress=on_progress, + runtime=runtime, + ) _MCP_TOOLS_CACHE_KEY = cfg_key _MCP_TOOLS_CACHE_VALUE = {k: list(v) for k, v in loaded.items()} return {k: list(v) for k, v in loaded.items()} @@ -533,6 +542,7 @@ def load_mcp_and_build_kwargs( cfg=None, chat_model=None, workspace_dir=None, + runtime: "AsyncRuntime | None" = None, ): """Load MCP tools (cached by config) and build agent kwargs. @@ -551,7 +561,10 @@ def load_mcp_and_build_kwargs( from .utils import load_subagents cfg = cfg if cfg is not None else _ensure_config() - mcp_by_agent = _load_mcp_tools_cached(on_progress=on_mcp_progress) + mcp_by_agent = _load_mcp_tools_cached( + on_progress=on_mcp_progress, + runtime=runtime, + ) if not mcp_by_agent: return _build_base_kwargs( base_backend, @@ -915,6 +928,7 @@ def create_cli_agent( *, on_mcp_progress=None, events: "MiddlewareEventSink | None" = None, + runtime: "AsyncRuntime | None" = None, ) -> "CompiledStateGraph": """Create agent with checkpointer for CLI multi-turn support. @@ -941,6 +955,8 @@ def create_cli_agent( chat_model: Optional pre-built chat model. Only triggers the pure path when ``config`` is also explicit; otherwise it is ignored in favor of the ``_ensure_chat_model()`` fallback. + runtime: Optional application-scoped runtime for synchronous MCP tool + discovery. Direct callers get a scoped runtime when omitted. """ import os as _os @@ -1041,6 +1057,7 @@ def create_cli_agent( cfg=cfg, chat_model=chat_model, workspace_dir=workspace_dir, + runtime=runtime, ) return create_deep_agent( diff --git a/EvoScientist/backends.py b/EvoScientist/backends.py index 8e680b2..61902b0 100644 --- a/EvoScientist/backends.py +++ b/EvoScientist/backends.py @@ -4,7 +4,11 @@ import os import posixpath import re import shlex +import signal +import subprocess import sys +import threading +import time import uuid from pathlib import Path @@ -22,6 +26,7 @@ from deepagents.backends.protocol import ( ) from . import paths +from .cancellation import current_cancel_event # Reproduced here to dodge a circular import from .EvoScientist (the canonical # SKILLS_DIR constant). @@ -68,6 +73,87 @@ BLOCKED_COMMANDS = [ ] +_active_shell_processes_lock = threading.RLock() +_active_shell_processes: dict[threading.Event, set[subprocess.Popen[str]]] = {} +_PROCESS_DRAIN_GRACE_SECONDS = 1.0 + + +def _terminate_process_tree(process: subprocess.Popen[str]) -> None: + """Force-stop a shell and its descendants without waiting for reaping.""" + # A completed Popen has already reaped its PID, which the OS may reuse. + # Inspect the recorded state rather than calling poll(): an exited but + # unreaped shell can still have live descendants in its process group. + if process.returncode is not None: + return + + try: + if os.name == "nt": + # CREATE_NEW_PROCESS_GROUP alone does not make terminate() recursive. + # taskkill is the native way to stop the complete descendant tree. + subprocess.run( + ["taskkill", "/PID", str(process.pid), "/T", "/F"], + check=False, + capture_output=True, + timeout=5, + ) + else: + os.killpg(process.pid, signal.SIGKILL) + except (OSError, subprocess.SubprocessError): + try: + process.kill() + except OSError: + pass + + +def _stop_collecting_process_output(process: subprocess.Popen[str]) -> None: + """Close inherited pipes and reap *process* without blocking the caller.""" + for pipe in (process.stdout, process.stderr): + if pipe is not None: + try: + pipe.close() + except OSError: + pass + + if process.poll() is None: + threading.Thread(target=process.wait, daemon=True).start() + + +def cancel_active_shell_processes(event: threading.Event) -> None: + """Terminate every active shell command associated with *event*.""" + with _active_shell_processes_lock: + processes = tuple(_active_shell_processes.get(event, ())) + for process in processes: + _terminate_process_tree(process) + + +def _register_shell_process( + event: threading.Event | None, + process: subprocess.Popen[str], +) -> None: + if event is None: + return + with _active_shell_processes_lock: + _active_shell_processes.setdefault(event, set()).add(process) + cancel_now = event.is_set() + if cancel_now: + _terminate_process_tree(process) + + +def _unregister_shell_process( + event: threading.Event | None, + process: subprocess.Popen[str], +) -> None: + if event is None: + return + with _active_shell_processes_lock: + processes = _active_shell_processes.get(event) + if processes is None: + return + processes.discard(process) + if not processes: + _active_shell_processes.pop(event, None) + + def _shell_token_spans(command: str) -> list[dict[str, object]]: """Tokenize enough shell syntax to find quoted SSH remote commands. @@ -1237,16 +1323,168 @@ class CustomSandboxBackend(LocalShellBackend): - Access to paths outside workspace - Dangerous system commands - Then delegates to LocalShellBackend.execute() for actual execution. + The validated command is handed to the owned process runner so + cancelling an agent turn can terminate the complete process tree. """ + # Preserve LocalShellBackend's public validation contract. This + # override cannot delegate execution to the base implementation because + # it must retain the Popen handle for cancellation, so validate before + # command preparation and process launch instead. + if not command or not isinstance(command, str): + return ExecuteResponse( + output="Error: Command must be a non-empty string.", + exit_code=1, + truncated=False, + ) + command, error = prepare_sandbox_command( command, self.cwd, virtual_mode=self.virtual_mode, dangerous=self._dangerous ) if error: return ExecuteResponse(output=error, exit_code=1, truncated=False) - # Delegate to parent for subprocess execution - response = super().execute(command, timeout=timeout) + return self._execute_prepared_command(command, timeout=timeout) + + def _execute_prepared_command( + self, + command: str, + *, + timeout: int | None = None, + ) -> ExecuteResponse: + """Execute an already validated command in an owned process group.""" + + effective_timeout = timeout if timeout is not None else self._default_timeout + if effective_timeout <= 0: + msg = f"timeout must be positive, got {effective_timeout}" + raise ValueError(msg) + + cancel_event = current_cancel_event() + if cancel_event is not None and cancel_event.is_set(): + return ExecuteResponse( + output="Command cancelled before execution.", + exit_code=130, + truncated=False, + ) + + process: subprocess.Popen[str] | None = None + termination_reason: str | None = None + output_abandoned = False + try: + process_options: dict[str, object] = {} + if os.name == "nt": + process_options["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP + else: + process_options["start_new_session"] = True + + process = subprocess.Popen( + command, + shell=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + stdin=subprocess.DEVNULL, + text=True, + env=self._env, + cwd=str(self.cwd), + **process_options, + ) + _register_shell_process(cancel_event, process) + deadline = time.monotonic() + effective_timeout + drain_deadline: float | None = None + + while True: + now = time.monotonic() + if ( + termination_reason is None + and cancel_event is not None + and cancel_event.is_set() + ): + termination_reason = "cancelled" + _terminate_process_tree(process) + drain_deadline = now + _PROCESS_DRAIN_GRACE_SECONDS + elif termination_reason is None and now >= deadline: + termination_reason = "timed_out" + _terminate_process_tree(process) + drain_deadline = now + _PROCESS_DRAIN_GRACE_SECONDS + + if drain_deadline is not None and now >= drain_deadline: + _stop_collecting_process_output(process) + stdout = stderr = "" + output_abandoned = True + break + + communicate_deadline = ( + drain_deadline if drain_deadline is not None else deadline + ) + try: + stdout, stderr = process.communicate( + timeout=max( + 0.01, + min(0.1, communicate_deadline - time.monotonic()), + ) + ) + break + except subprocess.TimeoutExpired: + continue + + if termination_reason == "timed_out": + if timeout is not None: + timeout_output = ( + "Error: Command timed out after " + f"{effective_timeout} seconds (custom timeout). The command " + "may be stuck or require more time." + ) + else: + timeout_output = ( + f"Error: Command timed out after {effective_timeout} seconds. " + "For long-running commands, re-run using the timeout parameter." + ) + response = ExecuteResponse( + output=timeout_output, + exit_code=124, + truncated=output_abandoned, + ) + elif termination_reason == "cancelled" or ( + cancel_event is not None and cancel_event.is_set() + ): + response = ExecuteResponse( + output="Command cancelled.", + exit_code=130, + truncated=output_abandoned, + ) + else: + output_parts = [] + if stdout: + output_parts.append(stdout) + if stderr: + stderr_lines = stderr.strip().split("\n") + output_parts.extend(f"[stderr] {line}" for line in stderr_lines) + output = "\n".join(output_parts) if output_parts else "" + + truncated = False + if len(output) > self._max_output_bytes: + output = output[: self._max_output_bytes] + output += ( + f"\n\n... Output truncated at {self._max_output_bytes} bytes." + ) + truncated = True + if process.returncode != 0: + output = f"{output.rstrip()}\n\nExit code: {process.returncode}" + response = ExecuteResponse( + output=output, + exit_code=process.returncode, + truncated=truncated, + ) + except Exception as exc: + if process is not None: + _terminate_process_tree(process) + response = ExecuteResponse( + output=f"Error executing command ({type(exc).__name__}): {exc}", + exit_code=1, + truncated=False, + ) + finally: + if process is not None: + _unregister_shell_process(cancel_event, process) # Enhance timeout errors with actionable recovery guidance if response.exit_code == 124: diff --git a/EvoScientist/cancellation.py b/EvoScientist/cancellation.py new file mode 100644 index 0000000..807ff62 --- /dev/null +++ b/EvoScientist/cancellation.py @@ -0,0 +1,27 @@ +"""Cancellation context shared by streaming frontends and blocking tools.""" + +from __future__ import annotations + +import contextvars +import threading +from collections.abc import Iterator +from contextlib import contextmanager + +_current_cancel_event: contextvars.ContextVar[threading.Event | None] = ( + contextvars.ContextVar("evoscientist_cancel_event", default=None) +) + + +@contextmanager +def bind_cancel_event(event: threading.Event) -> Iterator[None]: + """Make a stream's cancellation event visible to nested sync tool calls.""" + token = _current_cancel_event.set(event) + try: + yield + finally: + _current_cancel_event.reset(token) + + +def current_cancel_event() -> threading.Event | None: + """Return the cancellation event bound to the current agent run, if any.""" + return _current_cancel_event.get() diff --git a/EvoScientist/channels/base.py b/EvoScientist/channels/base.py index 6ea5817..1b9502c 100644 --- a/EvoScientist/channels/base.py +++ b/EvoScientist/channels/base.py @@ -18,6 +18,7 @@ from pathlib import Path from typing import Any from ..paths import MEDIA_DIR +from ..runtime import AsyncRuntime from .bus.events import InboundMessage, OutboundMessage from .capabilities import ChannelCapabilities from .debug import TraceMixin, debug_trace_enabled @@ -927,34 +928,27 @@ class Channel(TraceMixin, ChannelPlugin, ABC): return None return self._raw_to_inbound(current) - def _build_inbound(self, raw: RawIncoming) -> InboundMessage | None: + def _build_inbound( + self, + raw: RawIncoming, + *, + runtime: AsyncRuntime | None = None, + ) -> InboundMessage | None: """Run *raw* through inbound middlewares and convert to InboundMessage. - Synchronous wrapper around :meth:`_build_inbound_async`. When an - event loop is already running, the coroutine is scheduled on that - loop via :func:`asyncio.run_coroutine_threadsafe` to avoid - thread-safety issues with middleware state (DedupCache, - GroupHistoryBuffer, etc.). + Compatibility wrapper for synchronous integrations. Internal channel + implementations should await :meth:`_build_inbound_async` on their + transport loop. A caller may provide its application runtime to reuse + that owner; otherwise a runtime is scoped to this call. + + This method deliberately rejects callers already running an event + loop. Blocking such a loop while scheduling the coroutine back onto it + deadlocks; async callers must await :meth:`_build_inbound_async`. """ - import asyncio - - try: - loop = asyncio.get_running_loop() - except RuntimeError: - loop = None - - if loop is not None and loop.is_running(): - future = asyncio.run_coroutine_threadsafe( - self._build_inbound_async(raw), - loop, - ) - return future.result() - else: - new_loop = asyncio.new_event_loop() - try: - return new_loop.run_until_complete(self._build_inbound_async(raw)) - finally: - new_loop.close() + if runtime is None: + with AsyncRuntime(thread_name="evosci-channel-adapter-runtime") as owned: + return self._build_inbound(raw, runtime=owned) + return runtime.run_sync(lambda: self._build_inbound_async(raw)) def _raw_to_inbound(self, raw: RawIncoming) -> InboundMessage | None: """Convert a RawIncoming to InboundMessage (pure transformation, no filtering). diff --git a/EvoScientist/channels/email/probe.py b/EvoScientist/channels/email/probe.py index c2a09d0..87f926d 100644 --- a/EvoScientist/channels/email/probe.py +++ b/EvoScientist/channels/email/probe.py @@ -25,7 +25,7 @@ async def validate_email_imap( import asyncio - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() def _check(): try: @@ -62,7 +62,7 @@ async def validate_email_smtp( import asyncio - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() def _check(): server = None diff --git a/EvoScientist/channels/imessage/rpc_client.py b/EvoScientist/channels/imessage/rpc_client.py index 9e18091..b363cf8 100644 --- a/EvoScientist/channels/imessage/rpc_client.py +++ b/EvoScientist/channels/imessage/rpc_client.py @@ -150,7 +150,7 @@ class ImsgRpcClient: "params": params or {}, } - future: asyncio.Future = asyncio.get_event_loop().create_future() + future: asyncio.Future = asyncio.get_running_loop().create_future() self._pending[request_id] = future line = json.dumps(payload) + "\n" diff --git a/EvoScientist/channels/signal/probe.py b/EvoScientist/channels/signal/probe.py index bcc73b9..2dc81df 100644 --- a/EvoScientist/channels/signal/probe.py +++ b/EvoScientist/channels/signal/probe.py @@ -22,7 +22,7 @@ async def validate_signal( return False, "phone_number is required" # Check signal-cli binary - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() def _check(): try: diff --git a/EvoScientist/channels/standalone.py b/EvoScientist/channels/standalone.py index 902bb69..ad7a5eb 100644 --- a/EvoScientist/channels/standalone.py +++ b/EvoScientist/channels/standalone.py @@ -26,6 +26,13 @@ from .debug import emit_debug_event logger = logging.getLogger(__name__) +async def _create_standalone_agent(): + """Construct the synchronous agent without blocking the channel loop.""" + from ..EvoScientist import create_cli_agent + + return await asyncio.to_thread(create_cli_agent) + + def _channel_trace_enabled(channel: Channel) -> bool: """Check if debug tracing is enabled on the channel.""" try: @@ -107,10 +114,12 @@ async def _async_main( consumer: InboundConsumer | None = None if use_agent: logger.info("Loading EvoScientist agent...") - from ..EvoScientist import create_cli_agent from ..gateway import create_runtime_gateways - agent = create_cli_agent() + # Agent construction performs synchronous MCP discovery through the + # owned-runtime bridge. Keep it off this already-running channel loop + # (and avoid blocking channel health/startup work while it loads). + agent = await _create_standalone_agent() runtime_gateways = create_runtime_gateways() logger.info("Agent loaded") @@ -151,7 +160,7 @@ async def _async_main( await channel.stop() await manager.stop_health() - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() for sig in (signal.SIGINT, signal.SIGTERM): loop.add_signal_handler( sig, diff --git a/EvoScientist/cli/agent.py b/EvoScientist/cli/agent.py index bbf4147..a2b188c 100644 --- a/EvoScientist/cli/agent.py +++ b/EvoScientist/cli/agent.py @@ -10,6 +10,8 @@ from ..paths import new_run_dir if TYPE_CHECKING: from langgraph.graph.state import CompiledStateGraph + from ..runtime import AsyncRuntime + def _shorten_path(path: str) -> str: """Shorten absolute path to relative path from current directory.""" @@ -70,6 +72,7 @@ def _load_agent( *, on_mcp_progress=None, events=None, + runtime: "AsyncRuntime | None" = None, ) -> "CompiledStateGraph": """Load the CLI agent with optional persistent checkpointer. @@ -84,6 +87,7 @@ def _load_agent( selects the pure (no module-global write) build path. on_mcp_progress: Optional per-server MCP progress callback. Signature ``(event, server_name, detail) -> None``. + runtime: Optional application-scoped runtime used for MCP discovery. """ from ..EvoScientist import create_cli_agent @@ -94,4 +98,5 @@ def _load_agent( chat_model=chat_model, on_mcp_progress=on_mcp_progress, events=events, + runtime=runtime, ) diff --git a/EvoScientist/cli/channel.py b/EvoScientist/cli/channel.py index 1e469a7..cbfd848 100644 --- a/EvoScientist/cli/channel.py +++ b/EvoScientist/cli/channel.py @@ -43,6 +43,7 @@ from ..stream.console import console if TYPE_CHECKING: from ..gateway import GraphGateway + from ..runtime import AsyncRuntime _channel_logger = logging.getLogger(__name__) @@ -281,6 +282,7 @@ async def dispatch_channel_slash_command( await_agent_ready: Callable[[], Awaitable[Any]] | None = None, on_cmd_completed: Callable[..., Awaitable[None]] | None = None, channel_runtime: ChannelRuntime | None = None, + async_runtime: AsyncRuntime | None = None, ) -> bool: """Dispatch a slash command from a channel message. @@ -347,6 +349,7 @@ async def dispatch_channel_slash_command( on_cmd_completed=on_cmd_completed, channel_runtime=channel_runtime, graph_gateway=graph_gateway, + async_runtime=async_runtime, ) except Exception as exc: # Last-ditch safety: any uncaught exception from inside the @@ -383,6 +386,7 @@ async def _dispatch_channel_slash_impl( await_agent_ready: Callable[[], Awaitable[Any]] | None, on_cmd_completed: Callable[..., Awaitable[None]] | None, channel_runtime: ChannelRuntime | None, + async_runtime: AsyncRuntime | None, ) -> bool: """Inner body of ``dispatch_channel_slash_command``. @@ -431,6 +435,7 @@ async def _dispatch_channel_slash_impl( checkpointer=checkpointer, channel_runtime=channel_runtime, graph_gateway=graph_gateway, + async_runtime=async_runtime, ) try: diff --git a/EvoScientist/cli/channel_sends.py b/EvoScientist/cli/channel_sends.py new file mode 100644 index 0000000..2dccfde --- /dev/null +++ b/EvoScientist/cli/channel_sends.py @@ -0,0 +1,112 @@ +"""Non-blocking bridge for streaming callbacks sent through a channel loop.""" + +from __future__ import annotations + +import asyncio +import concurrent.futures +import logging +import threading +from collections.abc import Coroutine +from typing import Any + + +class PendingChannelSends: + """Schedule channel I/O without blocking the caller's event loop. + + Streaming callbacks run on the owned async runtime, while channel clients + belong to the channel bus loop. Submissions therefore only enqueue work; + the frontend settles the returned futures after streaming has unwound. + """ + + def __init__( + self, + loop: asyncio.AbstractEventLoop | None, + logger: logging.Logger, + ) -> None: + self._loop = loop + self._logger = logger + self._lock = threading.Lock() + self._pending: list[tuple[concurrent.futures.Future[Any], str, int]] = [] + self._tail: concurrent.futures.Future[Any] | None = None + + @staticmethod + def _close(coro: Coroutine[Any, Any, Any]) -> None: + coro.close() + + async def _run_after( + self, + predecessor: concurrent.futures.Future[Any] | None, + coro: Coroutine[Any, Any, Any], + ) -> Any: + if predecessor is not None: + try: + await asyncio.shield(asyncio.wrap_future(predecessor)) + except asyncio.CancelledError: + task = asyncio.current_task() + if task is not None and task.cancelling(): + self._close(coro) + raise + except Exception: + pass + return await coro + + def submit( + self, + coro: Coroutine[Any, Any, Any], + label: str, + timeout: int = 15, + ) -> None: + """Schedule one send and return immediately.""" + if self._loop is None: + self._close(coro) + return + with self._lock: + ordered_coro = self._run_after(self._tail, coro) + try: + future = asyncio.run_coroutine_threadsafe(ordered_coro, self._loop) + except Exception as exc: + self._close(ordered_coro) + self._close(coro) + self._logger.debug("%s send failed: %s", label, exc) + return + self._tail = future + self._pending.append((future, label, timeout)) + + def _take_pending( + self, + ) -> list[tuple[concurrent.futures.Future[Any], str, int]]: + with self._lock: + pending = self._pending + self._pending = [] + return pending + + def settle(self) -> None: + """Wait for scheduled sends from a synchronous frontend thread.""" + for future, label, timeout in self._take_pending(): + try: + future.result(timeout=timeout) + except Exception as exc: + future.cancel() + self._logger.debug("%s send failed: %s", label, exc) + + async def settle_async(self) -> None: + """Wait for scheduled sends without blocking the frontend loop.""" + pending = self._take_pending() + try: + for future, label, timeout in pending: + try: + await asyncio.wait_for(asyncio.wrap_future(future), timeout=timeout) + except TimeoutError as exc: + future.cancel() + self._logger.debug("%s send failed: %s", label, exc) + except asyncio.CancelledError as exc: + task = asyncio.current_task() + if task is not None and task.cancelling(): + raise + self._logger.debug("%s send failed: %s", label, exc) + except Exception as exc: + self._logger.debug("%s send failed: %s", label, exc) + except asyncio.CancelledError: + for future, _label, _timeout in pending: + future.cancel() + raise diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index b7fcf32..9b6cf12 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -1,6 +1,5 @@ """Typer command registrations — onboard, config, mcp, main callback.""" -import asyncio import logging import os import queue @@ -12,6 +11,7 @@ from importlib.metadata import version as _pkg_version from pathlib import Path from typing import TYPE_CHECKING, Annotated, Any, cast +import click import typer from rich.markup import escape from rich.table import Table @@ -26,6 +26,7 @@ from ..gateway import ( ) from ..llm.context_window import DEFAULT_CONTEXT_WINDOW_FALLBACK, resolve_context_window from ..paths import ensure_dirs, set_active_workspace, set_workspace_root +from ..runtime import AsyncRuntime from ..stream.console import console from . import async_notifier from ._app import app, channel_app, config_app, configure_app, mcp_app, sessions_app @@ -53,6 +54,7 @@ from .channel import ( publish_to_channel_origin, remember_channel_origin, ) +from .channel_sends import PendingChannelSends from .mcp_ui import ( _mcp_add_server_from_kwargs, _mcp_edit_server_fields, @@ -66,6 +68,35 @@ if TYPE_CHECKING: from ..config import EvoScientistConfig + +_ASYNC_RUNTIME_META_KEY = "evoscientist.async_runtime" + + +def _close_cli_async_runtime(runtime: AsyncRuntime) -> None: + """Close the owned runtime or surface a controlled CLI shutdown failure.""" + try: + runtime.close() + except TimeoutError as exc: + click.echo( + f"Error: Async runtime shutdown did not complete: {exc}", + err=True, + ) + raise click.exceptions.Exit(1) from None + + +def _get_cli_async_runtime(ctx: typer.Context) -> AsyncRuntime: + """Return the application-scoped runtime owned by this CLI invocation.""" + root = ctx.find_root() + runtime = root.meta.get(_ASYNC_RUNTIME_META_KEY) + if runtime is None: + runtime = AsyncRuntime() + root.meta[_ASYNC_RUNTIME_META_KEY] = runtime + root.call_on_close(lambda: _close_cli_async_runtime(runtime)) + if not isinstance(runtime, AsyncRuntime): # pragma: no cover - defensive + raise RuntimeError("CLI async runtime context is invalid") + return runtime + + # ============================================================================= # Onboard command # ============================================================================= @@ -73,6 +104,7 @@ if TYPE_CHECKING: @app.command() def onboard( + ctx: typer.Context, skip_validation: bool = typer.Option( False, "--skip-validation", help="Skip API key validation during setup" ), @@ -201,7 +233,11 @@ def onboard( strict=non_interactive, ) - _run_onboard_cli(skip_validation=skip_validation, prompter=prompter) + _run_onboard_cli( + skip_validation=skip_validation, + prompter=prompter, + runtime=_get_cli_async_runtime(ctx), + ) # ============================================================================= @@ -243,11 +279,21 @@ def _run_onboard_cli(**kwargs: Any) -> None: raise typer.Exit(code=1) from exc -def _configure_section(section: str, skip_validation: bool = False) -> None: +def _configure_section( + section: str, + skip_validation: bool = False, + *, + runtime: AsyncRuntime | None = None, +) -> None: """Run a single onboarding section, reusing the wizard's step logic.""" + kwargs: dict[str, Any] = { + "skip_validation": skip_validation, + "only_sections": {section}, + } + if runtime is not None: + kwargs["runtime"] = runtime _run_onboard_cli( - skip_validation=skip_validation, - only_sections={section}, + **kwargs, ) @@ -325,9 +371,9 @@ def configure_latex(): @configure_app.command("channels") -def configure_channels(): +def configure_channels(ctx: typer.Context): """Re-run channels selection and per-channel configuration.""" - _configure_section("channels") + _configure_section("channels", runtime=_get_cli_async_runtime(ctx)) # ============================================================================= @@ -336,24 +382,17 @@ def configure_channels(): @channel_app.command("setup") -def channel_setup(): +def channel_setup(ctx: typer.Context): """Interactive channel configuration wizard. Guides you through selecting and configuring messaging channels (Telegram, Discord, or iMessage). """ - import asyncio - - try: - asyncio.get_event_loop() - except RuntimeError: - asyncio.set_event_loop(asyncio.new_event_loop()) - from ..config import load_config, save_config from ..config.onboard.channels import _step_channels config = load_config() - updates = _step_channels(config) + updates = _step_channels(config, runtime=_get_cli_async_runtime(ctx)) if updates: for key, value in updates.items(): setattr(config, key, value) @@ -875,6 +914,7 @@ class ServeRuntimeState: workspace_dir: str | None config: "EvoScientistConfig | None" runtime_gateways: RuntimeGateways + async_runtime: AsyncRuntime resume_warning_thread_id: str | None = None def set_agent( @@ -969,6 +1009,7 @@ async def _apply_serve_resume_state( _load_agent, workspace_dir=new_workspace, config=effective_config, + runtime=runtime_state.async_runtime, ) await _sync_background_agent_server_workspace( effective_config, @@ -1119,8 +1160,6 @@ def _serve_process_message( via the ``on_cmd_completed`` hook because the command mutates ``ctx.thread_id`` / ``ctx.workspace_dir`` directly. """ - import asyncio - from .channel import _bus_loop from .tui_runtime import run_streaming @@ -1139,14 +1178,10 @@ def _serve_process_message( # -- channel callback helpers (same pattern as interactive.py) -- + pending_channel_sends = PendingChannelSends(_bus_loop, _serve_logger) + def _send_to_channel(coro, label: str, timeout: int = 15) -> None: - loop = _bus_loop - if not loop: - return - try: - asyncio.run_coroutine_threadsafe(coro, loop).result(timeout=timeout) - except Exception as e: - _serve_logger.debug(f"{label} send failed: {e}") + pending_channel_sends.submit(coro, label, timeout) def _send_thinking(thinking: str) -> None: ch = msg.channel_ref @@ -1196,31 +1231,15 @@ def _serve_process_message( # commands like ``/evoskills`` actually execute in serve mode instead # of being fed to the LLM as a plain prompt. ``await_agent_ready`` is # None because the agent is always loaded before the serve loop polls. - # Uses a dedicated event loop (not ``asyncio.run``) so SIGINT handling - # installed by ``serve()`` remains authoritative — ``asyncio.run`` - # swaps ``signal.set_wakeup_fd`` and can leave it dangling on edge - # cases, which breaks Ctrl+C between messages. - # ``set_event_loop`` is needed because some downstream commands - # (e.g. ``/install-mcp``) call ``asyncio.get_event_loop()``, which - # raises ``RuntimeError`` on Python 3.12+ when the thread has no - # current loop set. The prior loop (often ``None``) is restored in - # the ``finally`` below so subsequent messages start from a clean - # slate. Loop creation lives inside the try so an exception between - # creation and ``set_event_loop`` still closes the loop. + # Slash commands run on the application-owned runtime. The main thread + # remains the signal owner while command coroutines share one stable loop. try: - _prev_loop: asyncio.AbstractEventLoop | None - try: - _prev_loop = asyncio.get_event_loop_policy().get_event_loop() - except RuntimeError: - _prev_loop = None - _slash_loop: asyncio.AbstractEventLoop | None = None _slash_handled = False _slash_error: Exception | None = None try: - _slash_loop = asyncio.new_event_loop() - asyncio.set_event_loop(_slash_loop) - _slash_handled = _slash_loop.run_until_complete( - dispatch_channel_slash_command( + async_runtime = runtime_state.async_runtime + _slash_handled = async_runtime.run_sync( + lambda: dispatch_channel_slash_command( msg, agent=runtime_state.agent, thread_id=runtime_state.thread_id, @@ -1245,15 +1264,12 @@ def _serve_process_message( ), channel_runtime=channel_runtime, graph_gateway=runtime_gateways.graph_gateway, + async_runtime=async_runtime, ) ) except Exception as exc: _slash_error = exc _serve_logger.exception("Slash dispatch failed for %s", msg.channel_type) - finally: - if _slash_loop is not None: - _slash_loop.close() - asyncio.set_event_loop(_prev_loop) if _slash_error is not None: _set_channel_response(msg.msg_id, f"Command error: {_slash_error}") @@ -1287,11 +1303,13 @@ def _serve_process_message( ask_user_prompt_fn=_ask_user_prompt, cancel_scope=_channel_message_cancel_scope(msg), gateway=runtime_gateways.graph_gateway, + runtime=runtime_state.async_runtime, ) except Exception as e: response = f"Error: {e}" console.print(f"[red]Serve error: {e}[/red]") + pending_channel_sends.settle() _set_channel_response(msg.msg_id, response) console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]") finally: @@ -1341,6 +1359,7 @@ def _serve_drain_notifications( interactive=True, metadata=meta, gateway=runtime_state.runtime_gateways.graph_gateway, + runtime=runtime_state.async_runtime, ) except Exception as exc: _serve_logger.warning("Notification agent turn failed: %s", exc) @@ -1378,19 +1397,15 @@ def _serve_drain_notifications( current_thread_id=runtime_state.thread_id, ) - _notif_loop: _aio.AbstractEventLoop | None = None try: - _notif_loop = _aio.new_event_loop() - _notif_loop.run_until_complete(_consume()) + runtime_state.async_runtime.run_sync(_consume) except Exception as exc: _serve_logger.warning("Notification drain failed: %s", exc) - finally: - if _notif_loop is not None: - _notif_loop.close() @app.command() def serve( + ctx: typer.Context, no_thinking: bool = typer.Option( False, "--no-thinking", help="Disable thinking relay to channels" ), @@ -1445,6 +1460,7 @@ def serve( cli_overrides["log_level"] = "DEBUG" cli_overrides["channel_debug_tracing"] = True config = get_effective_config(cli_overrides) + async_runtime = _get_cli_async_runtime(ctx) if debug: os.environ["EVOSCIENTIST_LOG_LEVEL"] = "DEBUG" os.environ["EVOSCIENTIST_CHANNEL_DEBUG_TRACING"] = "true" @@ -1495,11 +1511,13 @@ def serve( f"[bold red]{DANGEROUS_BANNER_MESSAGE}[/bold red]" ) console.print("[dim]Loading agent...[/dim]") - agent = _load_agent(workspace_dir=ws, config=config) + agent = _load_agent(workspace_dir=ws, config=config, runtime=async_runtime) runtime_gateways = create_runtime_gateways() - tid = asyncio.run( - runtime_gateways.graph_gateway.create_thread(GraphTarget(workspace_dir=ws)) + tid = async_runtime.run_sync( + lambda: runtime_gateways.graph_gateway.create_thread( + GraphTarget(workspace_dir=ws) + ) ) # Mutable runtime shared with _serve_process_message so channel slash @@ -1511,6 +1529,7 @@ def serve( workspace_dir=ws, config=config, runtime_gateways=runtime_gateways, + async_runtime=async_runtime, ) channel_runtime = ChannelRuntime(agent=agent, thread_id=tid) @@ -1551,9 +1570,22 @@ def serve( import threading shutdown_event = threading.Event() + no_active_cancel_scope = object() + active_cancel_scope: str | object | None = no_active_cancel_scope def _handle_shutdown(signum: int, _frame: Any) -> None: shutdown_event.set() + # Cancelling the owned asyncio task is not enough when it is awaiting a + # blocking execute call: the executor thread and its isolated process + # group keep running until the matching stream event is set. Request + # scope cancellation before KeyboardInterrupt unwinds message cleanup + # (which discards that scope). SIGTERM also needs this to unblock the + # synchronous serve call so the poll loop can observe shutdown_event. + scope = active_cancel_scope + if scope is not no_active_cancel_scope: + from ..stream.display import request_stream_cancel + + request_stream_cancel(cast(str | None, scope)) # Fall back to Python's default SIGINT behavior (raises # KeyboardInterrupt) so blocking I/O inside ``run_streaming`` # is still interrupted. For SIGTERM there's no default that @@ -1573,6 +1605,7 @@ def serve( if shutdown_event.is_set(): break if msg is not None: + active_cancel_scope = _channel_message_cancel_scope(msg) try: _serve_process_message( msg, @@ -1588,15 +1621,22 @@ def serve( except KeyboardInterrupt: shutdown_event.set() break + finally: + active_cancel_scope = no_active_cancel_scope # Poll notification queue when idle (no channel message was pending). if async_notifier.has_pending_notifications(runtime_state.thread_id): - _serve_drain_notifications( - runtime_state=runtime_state, - model=config.model, - workspace_dir=ws, - show_thinking=effective_channel_thinking, - ) + # Notification turns use the default stream cancellation scope. + active_cancel_scope = None + try: + _serve_drain_notifications( + runtime_state=runtime_state, + model=config.model, + workspace_dir=ws, + show_thinking=effective_channel_thinking, + ) + finally: + active_cancel_scope = no_active_cancel_scope except KeyboardInterrupt: shutdown_event.set() finally: @@ -1954,20 +1994,16 @@ def sessions_callback(ctx: typer.Context): so the bare command is informative rather than silent. """ if ctx.invoked_subcommand is None: - sessions_stats() + sessions_stats(ctx) @sessions_app.command("stats") -def sessions_stats(): +def sessions_stats(ctx: typer.Context): """Show DB size, thread count, total checkpoints, top heaviest threads.""" - import asyncio - from ..sessions import db_stats - try: - stats = asyncio.get_event_loop().run_until_complete(db_stats()) - except RuntimeError: - stats = asyncio.new_event_loop().run_until_complete(db_stats()) + runtime = _get_cli_async_runtime(ctx) + stats = runtime.run_sync(db_stats) table = Table(title="EvoScientist sessions DB", show_header=True) table.add_column("Metric", style="cyan") @@ -2108,6 +2144,8 @@ def _main_callback( if ctx.invoked_subcommand is not None: return + async_runtime = _get_cli_async_runtime(ctx) + # Load and apply configuration from ..config import apply_config_to_env, get_effective_config @@ -2350,10 +2388,12 @@ def _main_callback( else: tid = await graph_gateway.create_thread() console.print("[dim]Loading agent...[/dim]") - agent = _load_agent( + agent = await asyncio.to_thread( + _load_agent, workspace_dir=workspace_dir, checkpointer=checkpointer, config=config, + runtime=async_runtime, ) try: if effective_output_format == "stream-json": @@ -2382,16 +2422,31 @@ def _main_callback( # matching the text path (cmd_run does this itself). _wait_for_memory_workers_before_exit() else: - cmd_run( - agent, - prompt, - thread_id=tid, - show_thinking=show_thinking, - workspace_dir=workspace_dir, - model=config.model, - ui_backend=config.ui_backend, - runtime_gateways=runtime_gateways, + stream_worker = asyncio.create_task( + asyncio.to_thread( + cmd_run, + agent, + prompt, + thread_id=tid, + show_thinking=show_thinking, + workspace_dir=workspace_dir, + model=config.model, + ui_backend=config.ui_backend, + runtime_gateways=runtime_gateways, + async_runtime=async_runtime, + ) ) + try: + await asyncio.shield(stream_worker) + except asyncio.CancelledError: + from ..stream.display import request_stream_cancel + from .tui_runtime import settle_cancelled_worker + + await settle_cancelled_worker( + stream_worker, + on_cancel=request_stream_cancel, + ) + raise finally: # Model failures can bypass middleware ``after_agent`` # hooks. Close any remaining QuickJS workers while this @@ -2407,10 +2462,7 @@ def _main_callback( except Exception: pass - import nest_asyncio - - nest_asyncio.apply() - asyncio.get_event_loop().run_until_complete(_single_shot()) + async_runtime.run_sync(_single_shot) else: from .interactive import cmd_interactive @@ -2427,6 +2479,7 @@ def _main_callback( thread_id=thread_id, ui_backend=config.ui_backend, config=config, + async_runtime=async_runtime, ) diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index 1782a49..658c8a2 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -4,8 +4,10 @@ import asyncio import logging import queue import random +import signal import sys -from collections.abc import Callable +import threading +from collections.abc import Awaitable, Callable from dataclasses import dataclass from datetime import datetime from typing import TYPE_CHECKING, Any @@ -62,6 +64,7 @@ from .channel import ( _set_channel_response, dispatch_channel_slash_command, ) +from .channel_sends import PendingChannelSends from .file_mentions import complete_file_mention, resolve_file_mentions from .rich_command_ui import RichCLICommandUI from .status_bar import ( @@ -83,7 +86,12 @@ from .status_bar import ( make_usage_status_snapshot, ) from .tui_interactive import run_textual_interactive -from .tui_runtime import resolve_ui_backend, run_streaming +from .tui_runtime import ( + StreamCancellationTimeout, + resolve_ui_backend, + run_streaming, + run_streaming_async, +) _MEMORY_WORKER_SHUTDOWN_WAIT_SECONDS = 120.0 _MEMORY_WORKER_SHUTDOWN_POLL_SECONDS = 0.5 @@ -97,6 +105,8 @@ _background_tasks: set[asyncio.Task] = set() if TYPE_CHECKING: from langgraph.graph.state import CompiledStateGraph + from ..runtime import AsyncRuntime + @dataclass(frozen=True, slots=True) class _StartupSession: @@ -107,6 +117,15 @@ class _StartupSession: resumed: bool +async def _run_serialized_turn( + turn_lock: asyncio.Lock, + operation: Callable[[], Awaitable[Any]], +) -> Any: + """Run one session turn without overlapping another frontend source.""" + async with turn_lock: + return await operation() + + # ============================================================================= # Banner # ============================================================================= @@ -328,6 +347,47 @@ async def _resolve_startup_session( # ============================================================================= +async def _run_rich_cli_streaming_turn(**kwargs: Any) -> str: + """Run one Rich CLI turn with a fresh, turn-local SIGINT policy. + + ``asyncio.run`` installs a SIGINT handler whose interrupt count lasts for + the lifetime of the runner. The Rich CLI intentionally recovers after a + cancelled turn, so relying on that handler makes Ctrl+C on a later turn + look like the runner's second interrupt and raises ``KeyboardInterrupt``. + + While a model turn is active, route the first Ctrl+C to a child task + instead. Restoring the runner's handler after every turn keeps Ctrl+C at + the prompt unchanged and resets the force-quit boundary for the next turn. + A second Ctrl+C before the current turn settles remains a force quit. + """ + stream_task = asyncio.create_task( + run_streaming_async(**kwargs, recover_on_cancel=True) + ) + + # Interactive CLI execution belongs on the main thread, but retaining the + # ordinary await makes this helper safe in embedded/test environments where + # Python does not permit installing process signal handlers. + if threading.current_thread() is not threading.main_thread(): + return await stream_task + + previous_sigint = signal.getsignal(signal.SIGINT) + interrupted = False + + def _cancel_turn(signum: int, frame: Any) -> None: + nonlocal interrupted + if interrupted or stream_task.done(): + signal.default_int_handler(signum, frame) + return + interrupted = True + stream_task.cancel() + + signal.signal(signal.SIGINT, _cancel_turn) + try: + return await stream_task + finally: + signal.signal(signal.SIGINT, previous_sigint) + + def cmd_interactive( show_thinking: bool = True, channel_send_thinking: bool = True, @@ -340,6 +400,7 @@ def cmd_interactive( thread_id: str | None = None, ui_backend: str = "cli", config=None, + async_runtime: "AsyncRuntime | None" = None, ) -> None: """Interactive conversation mode with streaming output. @@ -358,15 +419,15 @@ def cmd_interactive( thread_id: Optional thread ID to resume a previous session ui_backend: UI backend ('cli' or 'tui') """ - import nest_asyncio - - nest_asyncio.apply() - resolved_ui_backend = resolve_ui_backend(ui_backend, warn_fallback=True) if resolved_ui_backend == "tui": from functools import partial - load_agent = partial(_load_agent, config=config) + load_agent = partial( + _load_agent, + config=config, + runtime=async_runtime, + ) run_textual_interactive( show_thinking=show_thinking, channel_send_thinking=channel_send_thinking, @@ -380,6 +441,7 @@ def cmd_interactive( load_agent=load_agent, create_session_workspace=_create_session_workspace, config=config, + async_runtime=async_runtime, ) return @@ -497,6 +559,7 @@ def cmd_interactive( checkpointer=checkpointer, config=config, events=event_sink, + runtime=async_runtime, ) async def _await_agent_ready() -> "CompiledStateGraph": @@ -868,6 +931,8 @@ def cmd_interactive( # ---- Channel queue processing (bus → main thread) ---- + turn_lock = asyncio.Lock() + async def _process_channel_message(msg: ChannelMessage) -> None: """Process a single channel message with real-time streaming. @@ -905,17 +970,12 @@ def cmd_interactive( console.print(rx) _print_separator() + pending_channel_sends = PendingChannelSends( + _ch_mod._bus_loop, _channel_logger + ) + def _send_to_channel(coro, label: str, timeout: int = 15) -> None: - """Schedule an async channel send on the bus loop.""" - loop = _ch_mod._bus_loop - if not loop: - return - try: - asyncio.run_coroutine_threadsafe(coro, loop).result( - timeout=timeout - ) - except Exception as e: - _channel_logger.debug(f"{label} send failed: {e}") + pending_channel_sends.submit(coro, label, timeout) def _send_thinking_to_channel(thinking: str) -> None: ch = msg.channel_ref @@ -1029,6 +1089,7 @@ def cmd_interactive( on_cmd_completed=_on_channel_cmd_completed, channel_runtime=channel_runtime, graph_gateway=runtime_gateways.graph_gateway, + async_runtime=async_runtime, ) if _slash_handled: # A channel-issued /new or /resume rotates the thread @@ -1047,7 +1108,7 @@ def cmd_interactive( await _refresh_status_snapshot( msg.content, reset_streaming_text=True ) - response = run_streaming( + response = await run_streaming_async( ui_backend=state["ui_backend"], agent=ready_agent, message=msg.content, @@ -1064,11 +1125,13 @@ def cmd_interactive( status_footer_builder=_stream_status_footer, cancel_scope=_ch_mod._channel_message_cancel_scope(msg), gateway=runtime_gateways.graph_gateway, + runtime=async_runtime, ) except Exception as e: response = f"Error: {e}" console.print(f"[red]Channel error: {e}[/red]") + await pending_channel_sends.settle_async() _set_channel_response(msg.msg_id, response) await _refresh_status_snapshot(reset_streaming_text=True) @@ -1105,7 +1168,7 @@ def cmd_interactive( meta = build_metadata(state["workspace_dir"], model) await _refresh_status_snapshot(text, reset_streaming_text=True) ready_agent = await _await_agent_ready() - response = run_streaming( + response = await run_streaming_async( ui_backend=state["ui_backend"], agent=ready_agent, message=text, @@ -1121,6 +1184,7 @@ def cmd_interactive( on_stream_event=_handle_stream_status_event, status_footer_builder=_stream_status_footer, gateway=runtime_gateways.graph_gateway, + runtime=async_runtime, ) _notif_tid = target_thread_id or state["thread_id"] if _ch_mod.publish_to_channel_origin(_notif_tid, response): @@ -1176,7 +1240,10 @@ def cmd_interactive( except queue.Empty: msg = None if msg is not None: - await _process_channel_message(msg) + await _run_serialized_turn( + turn_lock, + lambda _msg=msg: _process_channel_message(_msg), + ) continue # check queues again immediately # Notification path (only when no channel message was pending). @@ -1193,8 +1260,13 @@ def cmd_interactive( try: await async_notifier.consume_notifications( run_message=lambda text, notifs, _tid=current_tid: ( - _inject_notification_message( - text, notifs, target_thread_id=_tid + _run_serialized_turn( + turn_lock, + lambda: _inject_notification_message( + text, + notifs, + target_thread_id=_tid, + ), ) ), read_async_tasks_state=read_async_tasks_state, @@ -1329,6 +1401,7 @@ def cmd_interactive( input_tokens_hint=state.get("status_last_input_tokens"), channel_runtime=channel_runtime, graph_gateway=runtime_gateways.graph_gateway, + async_runtime=async_runtime, ) await cmd_manager.execute(user_input, ctx) @@ -1411,17 +1484,23 @@ def cmd_interactive( await _refresh_status_snapshot( message_to_send, reset_streaming_text=True ) - run_streaming( - ui_backend=state["ui_backend"], - agent=ready_agent, - message=message_to_send, - thread_id=state["thread_id"], - show_thinking=show_thinking, - interactive=True, - metadata=meta, - on_stream_event=_handle_stream_status_event, - status_footer_builder=_stream_status_footer, - gateway=runtime_gateways.graph_gateway, + await _run_serialized_turn( + turn_lock, + lambda _agent=ready_agent, _message=message_to_send, _thread_id=state["thread_id"], _meta=meta: ( + _run_rich_cli_streaming_turn( + ui_backend=state["ui_backend"], + agent=_agent, + message=_message, + thread_id=_thread_id, + show_thinking=show_thinking, + interactive=True, + metadata=_meta, + on_stream_event=_handle_stream_status_event, + status_footer_builder=_stream_status_footer, + gateway=runtime_gateways.graph_gateway, + runtime=async_runtime, + ) + ), ) await _refresh_status_snapshot(reset_streaming_text=True) console.print() @@ -1436,6 +1515,14 @@ def cmd_interactive( console.print() state["running"] = False break + except StreamCancellationTimeout as e: + console.print(f"[red]{escape(str(e))}[/red]") + console.print( + "[dim]Exiting because the active turn could not be " + "stopped safely.[/dim]" + ) + state["running"] = False + break except Exception as e: error_msg = str(e) if ( @@ -1456,6 +1543,17 @@ def cmd_interactive( await queue_task except asyncio.CancelledError: pass + try: + from ..middleware.code_interpreter import ( + aclose_code_interpreters, + ) + + await aclose_code_interpreters() + except Exception: + _channel_logger.debug( + "code interpreter cleanup failed", + exc_info=True, + ) # Best-effort: guard so a DB lookup failure here can't # shadow the original exception exiting _async_main_loop. current_tid = state.get("thread_id") @@ -1493,6 +1591,7 @@ def cmd_run( ui_backend: str = "cli", *, runtime_gateways: RuntimeGateways, + async_runtime: "AsyncRuntime | None" = None, ) -> None: """Single-shot execution with streaming display. @@ -1526,6 +1625,7 @@ def cmd_run( interactive=False, metadata=meta, gateway=runtime_gateways.graph_gateway, + runtime=async_runtime, ) _wait_for_memory_workers_before_exit() except Exception as e: diff --git a/EvoScientist/cli/tui_backends.py b/EvoScientist/cli/tui_backends.py index e196a69..e5cb84c 100644 --- a/EvoScientist/cli/tui_backends.py +++ b/EvoScientist/cli/tui_backends.py @@ -4,11 +4,14 @@ from __future__ import annotations from collections.abc import Callable from dataclasses import dataclass -from typing import Any, Protocol +from typing import TYPE_CHECKING, Any, Protocol from ..gateway import GraphGateway from ..stream.display import _run_streaming +if TYPE_CHECKING: + from ..runtime import AsyncRuntime + class StreamingTUIBackend(Protocol): """Protocol for TUI backends that can render agent streaming output.""" @@ -33,6 +36,7 @@ class StreamingTUIBackend(Protocol): ask_user_prompt_fn: Callable[[dict], dict] | None = None, cancel_scope: str | None = None, gateway: GraphGateway, + runtime: AsyncRuntime | None = None, ) -> str: """Run streaming and return final response text.""" @@ -61,6 +65,7 @@ class RichStreamingBackend: ask_user_prompt_fn: Callable[[dict], dict] | None = None, cancel_scope: str | None = None, gateway: GraphGateway, + runtime: AsyncRuntime | None = None, ) -> str: return _run_streaming( agent=agent, @@ -78,4 +83,5 @@ class RichStreamingBackend: ask_user_prompt_fn=ask_user_prompt_fn, cancel_scope=cancel_scope, gateway=gateway, + runtime=runtime, ) diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index f784907..928b814 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -78,6 +78,9 @@ from .status_bar import ( make_usage_status_snapshot, ) +if TYPE_CHECKING: + from ..runtime import AsyncRuntime + _channel_logger = logging.getLogger(__name__) if TYPE_CHECKING: @@ -483,6 +486,7 @@ def run_textual_interactive( load_agent: Callable[..., Any], create_session_workspace: Callable[[str | None], str], config: Any | None = None, + async_runtime: AsyncRuntime | None = None, ) -> None: """Run full-screen Textual interactive chat loop.""" if config is None: @@ -1363,7 +1367,7 @@ def run_textual_interactive( Returns the ``ApprovalWidget.Decided`` message, or ``None`` on timeout / cancellation. """ - self._approval_future = asyncio.get_event_loop().create_future() + self._approval_future = asyncio.get_running_loop().create_future() try: return await asyncio.wait_for(self._approval_future, timeout=300) except (TimeoutError, asyncio.CancelledError): @@ -1409,7 +1413,7 @@ def run_textual_interactive( Returns the selected thread_id, or ``None`` on cancel/timeout. """ - self._picker_future = asyncio.get_event_loop().create_future() + self._picker_future = asyncio.get_running_loop().create_future() try: return await asyncio.wait_for(self._picker_future, timeout=120) except (TimeoutError, asyncio.CancelledError): @@ -1437,7 +1441,7 @@ def run_textual_interactive( Returns list of install sources, or None on cancel/timeout. """ - self._browser_future = asyncio.get_event_loop().create_future() + self._browser_future = asyncio.get_running_loop().create_future() try: return await asyncio.wait_for(self._browser_future, timeout=300) except (TimeoutError, asyncio.CancelledError): @@ -1464,7 +1468,7 @@ def run_textual_interactive( async def _wait_for_mcp_browse(self, browser_widget) -> list | None: """Wait for user to complete MCP server browsing.""" - self._mcp_browser_future = asyncio.get_event_loop().create_future() + self._mcp_browser_future = asyncio.get_running_loop().create_future() try: return await asyncio.wait_for(self._mcp_browser_future, timeout=300) except (TimeoutError, asyncio.CancelledError): @@ -1492,7 +1496,7 @@ def run_textual_interactive( Returns ``(name, provider)`` or ``None`` on cancel/timeout. """ - self._model_picker_future = asyncio.get_event_loop().create_future() + self._model_picker_future = asyncio.get_running_loop().create_future() try: return await asyncio.wait_for(self._model_picker_future, timeout=120) except (TimeoutError, asyncio.CancelledError): @@ -1554,6 +1558,7 @@ def run_textual_interactive( """ from ..stream.display import ( is_stream_cancel_requested, + iter_with_stream_cancel, ) container = self.query_one("#chat", VerticalScroll) @@ -1780,16 +1785,21 @@ def run_textual_interactive( summarization_w = None try: _anchor_engaged = False - async for event in graph_gateway.stream_events( - RunRequest( - message=_stream_input, - thread_id=thread_id_override or self._conversation_tid, - metadata=metadata, - target=GraphTarget( - local_graph=agent, - workspace_dir=self._workspace_dir, - ), - ) + async for event in iter_with_stream_cancel( + graph_gateway.stream_events( + RunRequest( + message=_stream_input, + thread_id=( + thread_id_override or self._conversation_tid + ), + metadata=metadata, + target=GraphTarget( + local_graph=agent, + workspace_dir=self._workspace_dir, + ), + ) + ), + cancel_scope, ): if is_stream_cancel_requested(cancel_scope): response = await _mark_cancelled_response() @@ -2380,6 +2390,11 @@ def run_textual_interactive( cancelled = False response = "" try: + # Foreground turns share the legacy default scope. Reset it at + # the turn boundary; scoped channel stop requests remain armed. + from ..stream.display import clear_stream_cancel + + clear_stream_cancel() self._busy = True self._turn_started_at = datetime.now() self._status_phase = ResearchPhase.THINKING @@ -2414,6 +2429,17 @@ def run_textual_interactive( ) except asyncio.CancelledError: cancelled = True + try: + from ..middleware.code_interpreter import ( + aclose_code_interpreters, + ) + + await aclose_code_interpreters() + except Exception: + _channel_logger.debug( + "code interpreter cleanup after cancellation failed", + exc_info=True, + ) self._append_system("\nInterrupted by user", style="dim italic #ffe082") finally: self._busy = False @@ -2547,6 +2573,7 @@ def run_textual_interactive( on_cmd_completed=self._on_channel_cmd_completed, channel_runtime=self._channel_runtime, graph_gateway=self._runtime_gateways.graph_gateway, + async_runtime=async_runtime, ) if _slash_handled: # A channel-issued /new or /resume rotates the thread in @@ -3091,6 +3118,7 @@ def run_textual_interactive( input_tokens_hint=self._status_last_input_tokens, channel_runtime=self._channel_runtime, graph_gateway=self._runtime_gateways.graph_gateway, + async_runtime=async_runtime, ) if await cmd_manager.execute(command, ctx): @@ -3237,6 +3265,9 @@ def run_textual_interactive( self._queued_messages.clear() self._render_queue_indicator() if self._run_task is not None and not self._run_task.done(): + from ..stream.display import request_stream_cancel + + request_stream_cancel() self._run_task.cancel() else: # Edge case: busy but no task — force reset @@ -3604,6 +3635,18 @@ def run_textual_interactive( finally: from .resume_hint import print_resume_hint + try: + from ..middleware.code_interpreter import ( + aclose_code_interpreters, + ) + + await aclose_code_interpreters() + except Exception: + _channel_logger.debug( + "code interpreter cleanup failed", + exc_info=True, + ) + # Best-effort resume hint — guarded so failures here (e.g. # DB teardown race during abnormal shutdown) cannot shadow # the original run_async traceback. @@ -3623,12 +3666,4 @@ def run_textual_interactive( except Exception: _channel_logger.debug("print_resume_hint failed", exc_info=True) - import nest_asyncio # type: ignore[import-untyped] - - nest_asyncio.apply() - try: - loop = asyncio.get_event_loop() - except RuntimeError: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - loop.run_until_complete(_amain()) + asyncio.run(_amain()) diff --git a/EvoScientist/cli/tui_runtime.py b/EvoScientist/cli/tui_runtime.py index f028506..b903d1c 100644 --- a/EvoScientist/cli/tui_runtime.py +++ b/EvoScientist/cli/tui_runtime.py @@ -2,14 +2,20 @@ from __future__ import annotations +import asyncio from collections.abc import Callable -from typing import Any +from typing import TYPE_CHECKING, Any from ..gateway import GraphGateway +from ..runtime import AsyncRuntimeError from ..stream.console import console from .tui_backends import RichStreamingBackend, StreamingTUIBackend +if TYPE_CHECKING: + from ..runtime import AsyncRuntime + DEFAULT_UI_BACKEND = "cli" +STREAM_CANCEL_SETTLE_TIMEOUT = 5.0 # "webui" launches the browser front-end instead of an in-terminal UI; it is # intercepted earlier (cli/commands.py:_main_callback) and never reaches the # streaming backends, but is listed here so normalize/resolve preserve it @@ -18,6 +24,41 @@ SUPPORTED_UI_BACKENDS = ("cli", "tui", "webui") _LEGACY_BACKEND_MAP = {"textual": "tui", "rich": "cli"} +class StreamCancellationTimeout(RuntimeError): + """A blocking renderer did not settle after its turn was cancelled.""" + + +def _consume_late_worker_result(worker: asyncio.Task[Any]) -> None: + """Retrieve a detached worker result so eventual failure is not unhandled.""" + try: + worker.exception() + except asyncio.CancelledError: + pass + + +async def settle_cancelled_worker( + worker: asyncio.Task[Any], + *, + on_cancel: Callable[[], Any], +) -> Any: + """Request cooperative cancellation and wait a bounded time for settlement.""" + on_cancel() + done, _ = await asyncio.wait( + {worker}, + timeout=STREAM_CANCEL_SETTLE_TIMEOUT, + ) + if not done: + worker.add_done_callback(_consume_late_worker_result) + raise StreamCancellationTimeout( + "The active turn did not stop within " + f"{STREAM_CANCEL_SETTLE_TIMEOUT:g} seconds after cancellation." + ) + try: + return worker.result() + except Exception: + return "" + + def normalize_ui_backend(value: str | None) -> str: """Normalize user-provided backend name with a safe default.""" if not value: @@ -81,6 +122,7 @@ def run_streaming( ask_user_prompt_fn: Callable[[dict], dict] | None = None, cancel_scope: str | None = None, gateway: GraphGateway, + runtime: AsyncRuntime | None = None, ) -> str: """Run streaming with the selected backend.""" backend = get_backend(ui_backend, warn_fallback=True) @@ -101,7 +143,10 @@ def run_streaming( ask_user_prompt_fn=ask_user_prompt_fn, cancel_scope=cancel_scope, gateway=gateway, + runtime=runtime, ) + except AsyncRuntimeError: + raise except RuntimeError: requested = normalize_ui_backend(ui_backend) if requested == "tui": @@ -124,5 +169,40 @@ def run_streaming( ask_user_prompt_fn=ask_user_prompt_fn, cancel_scope=cancel_scope, gateway=gateway, + runtime=runtime, ) raise + + +async def run_streaming_async( + *, + recover_on_cancel: bool = False, + **kwargs: Any, +) -> str: + """Run the synchronous Rich renderer without blocking a frontend loop. + + Cancellation requests the matching stream scope and gives the worker a + bounded interval to unwind. Foreground interactive turns may opt into + recovering the frontend task after cleanup so Ctrl+C returns to the prompt. + """ + from ..stream.display import request_stream_cancel + + worker = asyncio.create_task(asyncio.to_thread(run_streaming, **kwargs)) + try: + return await asyncio.shield(worker) + except asyncio.CancelledError: + try: + response = await settle_cancelled_worker( + worker, + on_cancel=lambda: request_stream_cancel(kwargs.get("cancel_scope")), + ) + finally: + from ..middleware.code_interpreter import aclose_code_interpreters + + await aclose_code_interpreters() + if recover_on_cancel: + current = asyncio.current_task() + if current is not None and current.uncancel() > 0: + raise + return response + raise diff --git a/EvoScientist/commands/base.py b/EvoScientist/commands/base.py index c125996..19d4a7e 100644 --- a/EvoScientist/commands/base.py +++ b/EvoScientist/commands/base.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Protocol, runtime_checkable if TYPE_CHECKING: from ..gateway import GraphGateway + from ..runtime import AsyncRuntime @dataclass @@ -91,6 +92,7 @@ class CommandContext: config: Any = None channel_runtime: ChannelRuntime | None = None graph_gateway: GraphGateway | None = None + async_runtime: AsyncRuntime | None = None command_error: str | None = None # Real LLM input token count from last usage_metadata (includes system # prompt + tool schemas). Used by /compact for accurate display. diff --git a/EvoScientist/commands/implementation/mcp_install.py b/EvoScientist/commands/implementation/mcp_install.py index 1217cba..1bfd04e 100644 --- a/EvoScientist/commands/implementation/mcp_install.py +++ b/EvoScientist/commands/implementation/mcp_install.py @@ -37,7 +37,7 @@ class InstallMCPCommand(Command): try: import asyncio - servers = await asyncio.get_event_loop().run_in_executor( + servers = await asyncio.get_running_loop().run_in_executor( None, fetch_marketplace_index ) except Exception as e: diff --git a/EvoScientist/commands/implementation/model.py b/EvoScientist/commands/implementation/model.py index 348d2b5..7f10663 100644 --- a/EvoScientist/commands/implementation/model.py +++ b/EvoScientist/commands/implementation/model.py @@ -130,6 +130,7 @@ class ModelCommand(Command): *, save: bool = False, ) -> None: + import asyncio import copy from ...cli.agent import _load_agent @@ -139,6 +140,7 @@ class ModelCommand(Command): set_active_config, set_chat_model_instance, ) + from ...runtime import AsyncRuntime cfg = _ensure_config() @@ -158,12 +160,19 @@ class ModelCommand(Command): try: new_chat_model = _build_chat_model(temp_cfg) - new_agent = _load_agent( - workspace_dir=ctx.workspace_dir, - checkpointer=ctx.checkpointer, - config=temp_cfg, - chat_model=new_chat_model, - events=events, + load_kwargs = { + "workspace_dir": ctx.workspace_dir, + "checkpointer": ctx.checkpointer, + "config": temp_cfg, + "chat_model": new_chat_model, + "events": events, + } + async_runtime = getattr(ctx, "async_runtime", None) + if isinstance(async_runtime, AsyncRuntime): + load_kwargs["runtime"] = async_runtime + new_agent = await asyncio.to_thread( + _load_agent, + **load_kwargs, ) except Exception as e: ctx.ui.append_system(f"Failed to switch model: {e}", style="red") diff --git a/EvoScientist/config/onboard/channels.py b/EvoScientist/config/onboard/channels.py index 9a50929..478f649 100644 --- a/EvoScientist/config/onboard/channels.py +++ b/EvoScientist/config/onboard/channels.py @@ -10,6 +10,7 @@ from __future__ import annotations import questionary from questionary import Choice +from ...runtime import AsyncRuntime from ..settings import EvoScientistConfig from .helpers import ( _setup_imessage, @@ -21,7 +22,11 @@ from .style import ( ) -def _step_channels(config: EvoScientistConfig) -> dict[str, object]: +def _step_channels( + config: EvoScientistConfig, + *, + runtime: AsyncRuntime | None = None, +) -> dict[str, object]: """Step: Select channels to enable on startup. Presents a multi-select list of supported channels. @@ -35,6 +40,12 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: Dict mapping config field names to their new values. Empty dict when the user skips or selects nothing. """ + # Direct/programmatic callers still get a single owned runtime for the + # whole step. CLI callers pass their application-scoped runtime instead. + if runtime is None: + with AsyncRuntime(thread_name="evosci-onboard-runtime") as owned_runtime: + return _step_channels(config, runtime=owned_runtime) + # Currently enabled channels _currently_enabled = { t.strip() @@ -592,11 +603,9 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: f" to {_accounts_path}.[/dim]" ) try: - import asyncio - from ...channels.wechat.personal import qr_login - creds = asyncio.run(qr_login()) + creds = runtime.run_sync(qr_login) except Exception as exc: console.print(f" [red]✗ Scan failed: {exc}[/red]") creds = None @@ -783,7 +792,7 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]: updates[senders_field] = senders.strip() # Probe validation - _probe_channel(ch_name, config, updates) + _probe_channel(ch_name, config, updates, runtime=runtime) enabled_channels.append(ch_name) @@ -820,12 +829,13 @@ def _probe_channel( ch_name: str, config: EvoScientistConfig, updates: dict[str, object], + *, + runtime: AsyncRuntime, ) -> None: """Run the probe for a channel type and print the result. Non-fatal: prints a warning on failure but does not prevent enabling. """ - import asyncio def _val(key: str, fallback: str = "") -> str: """Get a value from updates first, then config, then fallback.""" @@ -928,17 +938,7 @@ def _probe_channel( return True, "No probe available" try: - try: - loop = asyncio.get_event_loop() - if loop.is_running(): - import nest_asyncio # type: ignore[import-untyped] - - nest_asyncio.apply() - except RuntimeError: - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - - ok, detail = loop.run_until_complete(_run()) + ok, detail = runtime.run_sync(_run) if ok: console.print(f" [green]✓ {detail}[/green]") else: diff --git a/EvoScientist/config/onboard/wizard.py b/EvoScientist/config/onboard/wizard.py index 3994db3..8accfe7 100644 --- a/EvoScientist/config/onboard/wizard.py +++ b/EvoScientist/config/onboard/wizard.py @@ -9,6 +9,7 @@ import questionary from rich.panel import Panel from rich.text import Text +from ...runtime import AsyncRuntime from ..settings import ( EvoScientistConfig, get_config_path, @@ -475,6 +476,7 @@ def run_onboard( skip_validation: bool = False, prompter=None, only_sections: set[str] | frozenset[str] | None = None, + runtime: AsyncRuntime | None = None, ) -> bool: """Run the interactive onboarding wizard. @@ -487,6 +489,9 @@ def run_onboard( only_sections: If given, restrict the wizard to exactly these section ids — the Keep/Modify/Reset prompt is skipped. Used by ``EvoSci configure
`` to re-run a single phase. + runtime: Optional application-scoped async runtime used by channel + login and credential probes. Direct callers may omit it; the + channel step then owns a runtime for the duration of that step. Returns: True if configuration was saved, False if cancelled. @@ -883,7 +888,7 @@ def run_onboard( _step_tinytex() if "channels" in sections_to_run: - for key, value in _step_channels(config).items(): + for key, value in _step_channels(config, runtime=runtime).items(): setattr(config, key, value) _autosave(config) diff --git a/EvoScientist/mcp/client.py b/EvoScientist/mcp/client.py index bf931d7..b10dcf7 100644 --- a/EvoScientist/mcp/client.py +++ b/EvoScientist/mcp/client.py @@ -19,6 +19,8 @@ from typing import Any import yaml +from ..runtime import AsyncRuntime, AsyncRuntimeError + logger = logging.getLogger(__name__) @@ -838,6 +840,7 @@ def load_mcp_tools( config: dict[str, Any] | None = None, *, on_progress: ProgressCallback | None = None, + runtime: AsyncRuntime | None = None, ) -> dict[str, list]: """Load MCP tools and return them grouped by target agent. @@ -854,6 +857,10 @@ def load_mcp_tools( warnings when the caller has already loaded the config. on_progress: Optional callback invoked per server with ``(event, server_name, detail)``. See :data:`ProgressCallback`. + runtime: Runtime that owns MCP discovery work. When omitted, this + function creates one scoped to this call. The returned adapters + open a fresh MCP session for each tool call and do not retain the + discovery loop. Returns: Dict mapping agent name -> list of LangChain ``BaseTool`` objects. @@ -865,19 +872,28 @@ def load_mcp_tools( if not config: return {} - try: - loop = asyncio.get_running_loop() - except RuntimeError: - loop = None + if runtime is None: + with AsyncRuntime(thread_name="evosci-mcp-runtime") as owned_runtime: + return load_mcp_tools( + config, + on_progress=on_progress, + runtime=owned_runtime, + ) try: - if loop and loop.is_running(): - # Inside an already-running event loop (e.g. Jupyter) — - # nest_asyncio patches the loop so asyncio.run() works. - import nest_asyncio - - nest_asyncio.apply() - server_tools = asyncio.run(_load_tools(config, on_progress=on_progress)) + server_tools = runtime.run_sync( + lambda: _load_tools(config, on_progress=on_progress) + ) + except AsyncRuntimeError as exc: + # A bridge lifecycle/call-site error is not an MCP availability + # failure. In particular, hiding a running-loop violation here makes + # callers cache an empty tool set for the rest of the process. + if "cannot block a running event loop" in str(exc): + raise AsyncRuntimeError( + "load_mcp_tools() cannot run inside an async context; use " + "`await aload_mcp_tools(config, on_progress=...)` instead" + ) from exc + raise except Exception as exc: logger.warning("MCP tool loading failed: %s", exc) return {} diff --git a/EvoScientist/middleware/code_interpreter.py b/EvoScientist/middleware/code_interpreter.py index 49b1222..5e40eda 100644 --- a/EvoScientist/middleware/code_interpreter.py +++ b/EvoScientist/middleware/code_interpreter.py @@ -29,7 +29,9 @@ Usage:: from __future__ import annotations +import asyncio import contextlib +import logging import weakref from langchain.agents.middleware.types import ModelRequest @@ -40,6 +42,9 @@ from langchain_quickjs import CodeInterpreterMiddleware # values; tests / ad-hoc callers can omit and get sensible defaults. _DEFAULT_TIMEOUT_SECONDS: float = 60.0 _DEFAULT_MAX_RESULT_CHARS: int = 10000 +_CLOSE_TIMEOUT_SECONDS: float = 10.0 + +logger = logging.getLogger(__name__) _MEMORY_FIRST_INTERPRETER_PROMPT = ( "\n\nWhen memory tools (search_observations, read_memory) are available, use " @@ -83,10 +88,34 @@ class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware): _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() +async def aclose_code_interpreters( + *, + timeout: float = _CLOSE_TIMEOUT_SECONDS, +) -> None: + """Close live QuickJS middleware without blocking application shutdown.""" + middlewares = tuple(_live_interpreters) + if not middlewares: + return + + close_tasks = [middleware.aclose() for middleware in middlewares] + try: + results = await asyncio.wait_for( + asyncio.gather(*close_tasks, return_exceptions=True), + timeout=timeout, + ) + except TimeoutError: + logger.warning( + "code interpreter cleanup did not finish within %g seconds", + timeout, + ) + return + + for result in results: + if isinstance(result, BaseException): + logger.debug( + "code interpreter cleanup failed", + exc_info=(type(result), result, result.__traceback__), + ) # Read-only, batchable tools that benefit from being callable inside JS. diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py index 4deb2ac..6481713 100644 --- a/EvoScientist/middleware/model_fallback.py +++ b/EvoScientist/middleware/model_fallback.py @@ -287,6 +287,76 @@ async def _try_fallbacks( _raise_normalized(last_failing_request, last_exc) +def _try_fallbacks_sync( + request: ModelRequest, + invoke: Callable[[ModelRequest], ModelResponse], + primary_exc: Exception, + events: MiddlewareEventSink, +) -> ModelResponse: + """Synchronous counterpart to :func:`_try_fallbacks`. + + The synchronous middleware path calls a synchronous model handler. Keeping + that traversal synchronous avoids manufacturing an event loop solely to + share the async implementation. + """ + from ..llm.models import get_chat_model + + events.emit_fallback_notice( + f"Primary model failed: {type(primary_exc).__name__}: {primary_exc}", + "yellow", + ) + logger.warning( + "Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc + ) + + last_exc = primary_exc + last_failing_request = request + + for model_name, provider in get_fallback_chain(): + 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 = invoke(fb_request) + events.emit_fallback_notice( + f" Fallback to {model_name} ({provider}) succeeded", + "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: + events.emit_fallback_notice( + f" {model_name} hit non-fallbackable error ({reason}) " + f"-- aborting fallback chain", + "red", + ) + _raise_normalized(fb_request, fb_exc) + last_exc = fb_exc + last_failing_request = fb_request + events.emit_fallback_notice( + f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}", + "red", + ) + logger.warning( + "Fallback %s (provider=%s) failed: %s: %s", + model_name, + provider, + type(fb_exc).__name__, + fb_exc, + ) + + events.emit_fallback_notice( + " All fallbacks exhausted -- re-raising last error", "red" + ) + _raise_normalized(last_failing_request, last_exc) + + def _raise_normalized(request: ModelRequest, exc: Exception) -> None: """Wrap *exc* in a ``ProviderStreamError`` attributed to ``request.model`` and raise, so the outer chain sees the failure @@ -334,6 +404,23 @@ def _guard_and_fallback( return _try_fallbacks(request, invoke, primary_exc, events) +def _guard_and_fallback_sync( + primary_exc: Exception, + request: ModelRequest, + invoke: Callable[[ModelRequest], ModelResponse], + events: MiddlewareEventSink, +) -> ModelResponse: + """Validate and run the native synchronous fallback traversal.""" + reason = _is_non_fallbackable(primary_exc) + if reason is not None: + events.emit_fallback_notice( + f"Model error ({reason}) -- not eligible for fallback, re-raising", + "red", + ) + _raise_normalized(request, primary_exc) + return _try_fallbacks_sync(request, invoke, primary_exc, events) + + class ModelFallbackMiddleware(AgentMiddleware): """LangChain AgentMiddleware that retries failed model calls on fallbacks. @@ -363,15 +450,7 @@ class ModelFallbackMiddleware(AgentMiddleware): try: return handler(request) except Exception as exc: - - async def _sync_invoke(r: ModelRequest) -> ModelResponse: - return handler(r) - - import asyncio - - return asyncio.run( - _guard_and_fallback(exc, request, _sync_invoke, self._events) - ) + return _guard_and_fallback_sync(exc, request, handler, self._events) async def awrap_model_call( self, diff --git a/EvoScientist/middleware/tool_selector.py b/EvoScientist/middleware/tool_selector.py index a73d01a..47a6b47 100644 --- a/EvoScientist/middleware/tool_selector.py +++ b/EvoScientist/middleware/tool_selector.py @@ -18,6 +18,7 @@ Usage:: from __future__ import annotations import logging +import threading from collections.abc import Awaitable, Callable, Iterable from typing import Any @@ -146,6 +147,8 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware): # Agent tools are fixed after graph construction, so the filtered # always-include set is stable for this middleware instance. self._selector: AgentMiddleware | None = None + self._fallback_warning_emitted = False + self._fallback_warning_lock = threading.Lock() def _build_selector(self, request: ModelRequest) -> AgentMiddleware: if self._selector is None: @@ -157,6 +160,19 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware): def _selected_names(request: ModelRequest) -> list[str]: return [name for tool in request.tools if (name := _tool_name(tool))] + def _report_selector_failure(self, exc: Exception) -> None: + """Expose selector degradation once without flooding normal logs.""" + with self._fallback_warning_lock: + emit_warning = not self._fallback_warning_emitted + self._fallback_warning_emitted = True + if emit_warning: + logger.warning( + "tool_selector.fallback error_type=%s using_all_tools=true; " + "details and subsequent failures are logged at DEBUG", + type(exc).__name__, + ) + logger.debug("Tool selector failed, using all tools", exc_info=True) + def wrap_model_call( self, request: ModelRequest, @@ -197,19 +213,13 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware): except Exception as exc: if _handler_called: raise # Error from downstream model — don't retry - from ..llm.errors import ProviderStreamError - from .error_normalization import _is_provider_error - - if isinstance(exc, ProviderStreamError) or _is_provider_error(exc): - # Auth / quota / connection failures on the selector's - # own model. Falling back to "use all tools" would hit - # the same provider anyway (same client, likely same - # credentials). Surface it instead so the user sees - # the real cause. - raise - # Structured-output shape / config failure — gracefully - # degrade to using all tools. - logger.debug("Tool selector failed, using all tools", exc_info=True) + # The selector is an optimization, so every selector-only failure + # degrades to all tools. This includes provider failures: the + # downstream model-fallback middleware may replace the request's + # primary model, but it cannot replace this selector's fixed + # auxiliary model. Re-raising here would make a healthy fallback + # retry the same failed selector and never reach the model call. + self._report_selector_failure(exc) _end_selection() return handler(request) finally: @@ -252,14 +262,7 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware): except Exception as exc: if _handler_called: raise - from ..llm.errors import ProviderStreamError - from .error_normalization import _is_provider_error - - if isinstance(exc, ProviderStreamError) or _is_provider_error(exc): - # See sync path — surface provider errors, degrade only - # on shape / config failures. - raise - logger.debug("Tool selector failed, using all tools", exc_info=True) + self._report_selector_failure(exc) _end_selection() return await handler(request) finally: diff --git a/EvoScientist/runtime.py b/EvoScientist/runtime.py new file mode 100644 index 0000000..8dc01a6 --- /dev/null +++ b/EvoScientist/runtime.py @@ -0,0 +1,565 @@ +"""Application-scoped ownership for EvoScientist async work. + +``AsyncRuntime`` owns one continuously running event loop on a dedicated +thread. Callers share the runtime, never its raw loop: + +* synchronous code uses :meth:`AsyncRuntime.run_sync`; +* code already on another event loop uses :meth:`AsyncRuntime.run_async`; +* durable background work uses :meth:`AsyncRuntime.spawn`. + +The API accepts factories rather than pre-created coroutines so construction +happens on the owned loop. There is deliberately no module singleton: an +application bootstrap owns an instance, passes it to consumers, and closes it. +""" + +from __future__ import annotations + +import asyncio +import concurrent.futures +import contextvars +import logging +import threading +import time +from collections.abc import Awaitable, Callable +from typing import Any, Generic, TypeVar + +from EvoScientist._winloop import ensure_proactor_event_loop_policy + +logger = logging.getLogger(__name__) + +T = TypeVar("T") +AsyncFactory = Callable[[], Awaitable[T]] + + +class AsyncRuntimeError(RuntimeError): + """Base error raised by :class:`AsyncRuntime`.""" + + +class AsyncRuntimeClosedError(AsyncRuntimeError): + """Raised when work is submitted after shutdown begins.""" + + +class RuntimeHandle(concurrent.futures.Future[T], Generic[T]): + """Cross-thread result with a separate coroutine-settlement signal. + + Cancelling a concurrent future marks it done immediately, while the + asyncio task may still be running ``finally`` blocks. ``wait_settled`` + distinguishes those two moments for cancellation and shutdown paths. + """ + + def __init__(self, *, name: str | None = None) -> None: + super().__init__() + self._name = name + self._settled: concurrent.futures.Future[None] = concurrent.futures.Future() + self._settle_lock = threading.Lock() + + @property + def name(self) -> str | None: + return self._name + + @property + def settled(self) -> bool: + return self._settled.done() + + def wait_settled(self, timeout: float | None = None) -> bool: + """Block for task settlement; return ``False`` on timeout.""" + try: + self._settled.result(timeout) + except concurrent.futures.TimeoutError: + return False + return True + + async def wait_settled_async(self) -> None: + """Wait for task settlement without blocking the caller's loop.""" + # ``wrap_future`` propagates cancellation back to the concurrent + # future. Settlement is a shared, one-way runtime signal rather than + # work owned by any individual waiter, so a cancelled waiter must not + # cancel or falsely complete it for everyone else. + await asyncio.shield(asyncio.wrap_future(self._settled)) + + def _mark_settled(self) -> None: + # The lock makes the check-and-set atomic during forced shutdown. + with self._settle_lock: + if not self._settled.done(): + self._settled.set_result(None) + + +class AsyncRuntime: + """Own a persistent asyncio loop and its task lifecycle. + + Instances start lazily on first submission or eagerly through + :meth:`start`. A closed instance is permanently sealed; create a new + instance for a new application lifetime. + """ + + def __init__( + self, + *, + thread_name: str = "evosci-async-runtime", + start_timeout: float = 5.0, + cancellation_timeout: float = 2.0, + ) -> None: + if start_timeout <= 0: + raise ValueError("start_timeout must be greater than zero") + if cancellation_timeout < 0: + raise ValueError("cancellation_timeout must not be negative") + + self._thread_name = thread_name + self._start_timeout = start_timeout + self._cancellation_timeout = cancellation_timeout + + self._lock = threading.RLock() + self._ready = threading.Event() + self._stopped = threading.Event() + self._closed = False + self._failure: BaseException | None = None + self._loop: asyncio.AbstractEventLoop | None = None + self._thread: threading.Thread | None = None + + # Runtime-loop-only mapping. Strong references prevent pending tasks + # and their handles from being garbage-collected. + self._loop_tasks: dict[asyncio.Task[Any], RuntimeHandle[Any]] = {} + + def __enter__(self) -> AsyncRuntime: + self.start() + return self + + def __exit__(self, *exc_info: object) -> None: + self.close() + + @property + def is_running(self) -> bool: + with self._lock: + thread = self._thread + return ( + not self._closed + and thread is not None + and thread.is_alive() + and self._ready.is_set() + and not self._stopped.is_set() + ) + + def start(self) -> None: + """Start the loop thread idempotently and wait until it serves work.""" + with self._lock: + self._ensure_started_locked() + + def _ensure_started_locked(self) -> asyncio.AbstractEventLoop: + if self._closed: + raise AsyncRuntimeClosedError(f"{self._thread_name} is closed") + + thread = self._thread + if thread is not None: + if thread.is_alive() and self._ready.is_set() and self._loop is not None: + return self._loop + error = AsyncRuntimeError(f"{self._thread_name} stopped unexpectedly") + if self._failure is not None: + raise error from self._failure + raise error + + thread = threading.Thread( + target=self._thread_main, + name=self._thread_name, + daemon=True, + ) + self._thread = thread + try: + thread.start() + except BaseException: + self._thread = None + raise + + if not self._ready.wait(self._start_timeout): + # A late-created loop observes this seal in _thread_main and exits + # instead of becoming an orphan after its owner saw startup fail. + self._closed = True + raise TimeoutError( + f"{self._thread_name} did not start within {self._start_timeout:.1f}s" + ) + if self._failure is not None: + raise AsyncRuntimeError(f"{self._thread_name} failed to start") from ( + self._failure + ) + if self._loop is None or not thread.is_alive(): + raise AsyncRuntimeError(f"{self._thread_name} failed to start") + return self._loop + + def _thread_main(self) -> None: + loop: asyncio.AbstractEventLoop | None = None + try: + ensure_proactor_event_loop_policy() + loop = asyncio.new_event_loop() + self._loop = loop + asyncio.set_event_loop(loop) + + # A startup timeout may let close() seal the instance before loop + # creation finishes. Do not leave a late-starting daemon behind. + if self._closed: + self._ready.set() + return + + # The callback proves run_forever is serving work before start() + # returns; merely allocating a loop is not sufficient. + loop.call_soon(self._ready.set) + loop.run_forever() + except BaseException as exc: + self._failure = exc + self._ready.set() + logger.exception("%s loop failed", self._thread_name) + finally: + if loop is not None: + self._settle_abandoned_tasks() + loop.close() + asyncio.set_event_loop(None) + self._loop = None + self._stopped.set() + + def _settle_abandoned_tasks(self) -> None: + """Resolve handles when a stopped loop cannot unwind further.""" + for task, handle in list(self._loop_tasks.items()): + if not task.done(): + task.cancel() + if not handle.done(): + handle.cancel() + handle._mark_settled() + self._loop_tasks.clear() + + def submit(self, factory: AsyncFactory[T]) -> RuntimeHandle[T]: + """Schedule a factory on the owned loop and return its handle. + + Submission is atomic with :meth:`close`: work is either queued before + shutdown is sealed or rejected. + """ + return self._submit(factory, name=None) + + def _submit( + self, + factory: AsyncFactory[T], + *, + name: str | None, + ) -> RuntimeHandle[T]: + if not callable(factory): + raise TypeError("factory must be callable") + + handle: RuntimeHandle[T] = RuntimeHandle(name=name) + context = contextvars.copy_context() + + def create_task() -> None: + if handle.cancelled(): + handle._mark_settled() + return + + async def invoke_factory() -> T: + return await factory() + + try: + task = asyncio.create_task(invoke_factory(), name=name) + except BaseException as exc: + self._set_handle_exception(handle, exc) + handle._mark_settled() + return + + self._loop_tasks[task] = handle + + def cancel_task(done: concurrent.futures.Future[T]) -> None: + if not done.cancelled() or task.done(): + return + try: + task.get_loop().call_soon_threadsafe(task.cancel) + except RuntimeError: + handle._mark_settled() + + handle.add_done_callback(cancel_task) + task.add_done_callback(self._copy_task_result) + if handle.cancelled() and not task.done(): + task.cancel() + + try: + with self._lock: + loop = self._ensure_started_locked() + self._enqueue_locked(loop, create_task, context) + except BaseException: + handle._mark_settled() + raise + return handle + + def _enqueue_locked( + self, + loop: asyncio.AbstractEventLoop, + callback: Callable[[], None], + context: contextvars.Context, + ) -> None: + """Enqueue under the lifecycle lock; isolated for race testing.""" + try: + loop.call_soon_threadsafe(callback, context=context) + except RuntimeError as exc: + raise AsyncRuntimeError( + f"{self._thread_name} stopped while submitting work" + ) from exc + + def _copy_task_result(self, task: asyncio.Task[Any]) -> None: + handle = self._loop_tasks.pop(task) + try: + result = task.result() + except asyncio.CancelledError: + handle.cancel() + except BaseException as exc: + self._set_handle_exception(handle, exc) + else: + self._set_handle_result(handle, result) + finally: + handle._mark_settled() + + @staticmethod + def _set_handle_result(handle: RuntimeHandle[Any], result: Any) -> None: + try: + handle.set_result(result) + except concurrent.futures.InvalidStateError: + pass # External cancellation won; discard the completed result. + + @staticmethod + def _set_handle_exception(handle: RuntimeHandle[Any], exc: BaseException) -> None: + try: + handle.set_exception(exc) + except concurrent.futures.InvalidStateError: + pass + + def run_sync( + self, + factory: AsyncFactory[T], + *, + timeout: float | None = None, + on_submitted: Callable[[RuntimeHandle[T]], None] | None = None, + ) -> T: + """Run async work from sync code, blocking for its result. + + Any thread already running an event loop must use :meth:`run_async`; + blocking it would freeze that frontend even if it is not the owned loop. + ``on_submitted`` may retain the handle for cross-thread cancellation; + it runs after submission and before this method starts blocking. + """ + try: + asyncio.get_running_loop() + except RuntimeError: + pass + else: + raise AsyncRuntimeError( + "run_sync() cannot block a running event loop; use " + "`await runtime.run_async(...)` instead" + ) + + handle = self.submit(factory) + if on_submitted is not None: + try: + on_submitted(handle) + except BaseException: + handle.cancel() + self._wait_for_cancellation(handle) + raise + try: + return handle.result(timeout) + except BaseException: + if not handle.done(): + handle.cancel() + self._wait_for_cancellation(handle) + elif handle.cancelled(): + self._wait_for_cancellation(handle) + raise + + async def run_async(self, factory: AsyncFactory[T]) -> T: + """Await owned work without blocking the caller's event loop. + + Current UI adapters cross through ``to_thread`` and :meth:`run_sync`. + This public bridge is retained for embedders and future async surfaces + whose event loop must remain responsive while the owned loop does work. + """ + caller_loop = asyncio.get_running_loop() + with self._lock: + runtime_loop = self._loop + if caller_loop is runtime_loop: + raise AsyncRuntimeError( + "run_async() called from the owned loop; await directly instead" + ) + + handle = self.submit(factory) + try: + return await asyncio.wrap_future(handle) + except asyncio.CancelledError: + handle.cancel() + try: + await asyncio.wait_for( + asyncio.shield(handle.wait_settled_async()), + timeout=self._cancellation_timeout, + ) + except TimeoutError: + logger.warning( + "%s task did not settle within %.1fs after cancellation", + self._thread_name, + self._cancellation_timeout, + ) + raise + + def spawn( + self, + factory: AsyncFactory[Any], + *, + name: str, + ) -> RuntimeHandle[Any]: + """Start durable work, retaining it and logging unhandled failures. + + This public primitive is reserved for runtime-owned background + services; scoped request work should continue to use :meth:`submit`. + """ + handle = self._submit(factory, name=name) + handle.add_done_callback(self._on_background_done) + return handle + + @staticmethod + def _on_background_done(handle: concurrent.futures.Future[Any]) -> None: + if handle.cancelled(): + return + try: + exc = handle.exception() + except concurrent.futures.CancelledError: + return + if exc is not None: + name = getattr(handle, "name", None) + logger.error( + "unhandled exception in runtime task %r", + name, + exc_info=(type(exc), exc, exc.__traceback__), + ) + + def _wait_for_cancellation(self, handle: RuntimeHandle[Any]) -> None: + if not handle.wait_settled(self._cancellation_timeout): + logger.warning( + "%s task did not settle within %.1fs after cancellation", + self._thread_name, + self._cancellation_timeout, + ) + + @staticmethod + async def _drain() -> None: + current = asyncio.current_task() + pending = [task for task in asyncio.all_tasks() if task is not current] + for task in pending: + task.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + loop = asyncio.get_running_loop() + await loop.shutdown_asyncgens() + # Cancelling a task awaiting ``to_thread`` / ``run_in_executor`` does + # not stop its underlying callable. Do not report a clean runtime + # shutdown until the owned loop's default executor is actually idle. + await loop.shutdown_default_executor() + + def close(self, *, timeout: float = 5.0) -> None: + """Seal intake, settle pending tasks, stop the loop, and join its thread. + + Shutdown is bounded by ``timeout``. Executor work cannot be preempted by + asyncio cancellation; if it outlives the deadline this call raises + :class:`TimeoutError` and the sealed runtime finishes shutting down in + the background. A later ``close()`` waits for that shutdown and only + succeeds after the executor is idle and the loop thread has stopped. + """ + if timeout < 0: + raise ValueError("timeout must not be negative") + deadline = time.monotonic() + timeout + + with self._lock: + thread = self._thread + if thread is threading.current_thread(): + raise AsyncRuntimeError( + "close() cannot join the owned loop thread; close the " + "runtime from its application owner" + ) + + if self._closed: + wait_for_existing_close = thread is not None and thread.is_alive() + drain = None + loop = self._loop + else: + self._closed = True + wait_for_existing_close = False + loop = self._loop + if thread is None: + self._stopped.set() + return + if loop is None: + # start() timed out while the runtime thread was still + # creating its loop. _thread_main observes _closed and + # exits as soon as creation finishes. + drain = None + else: + # The lifecycle lock orders this after every accepted + # task-creation callback queued by submit(). + try: + drain = asyncio.run_coroutine_threadsafe(self._drain(), loop) + except RuntimeError: + drain = None + + if wait_for_existing_close: + if not self._stopped.wait(max(0.0, deadline - time.monotonic())): + raise TimeoutError( + f"timed out waiting for {self._thread_name} shutdown" + ) + return + if thread is None: + return + + if loop is None: + thread.join(max(0.0, deadline - time.monotonic())) + if thread.is_alive(): + raise TimeoutError( + f"{self._thread_name} did not stop within {timeout:.1f}s" + ) + return + + drain_timed_out = False + drain_error: BaseException | None = None + if drain is not None: + try: + drain.result(max(0.0, deadline - time.monotonic())) + except concurrent.futures.TimeoutError: + drain_timed_out = True + + def stop_after_drain( + _done: concurrent.futures.Future[None], + ) -> None: + try: + loop.call_soon_threadsafe(loop.stop) + except RuntimeError: + pass + + # Keep the loop alive while executor work finishes. This + # callback completes the already-sealed shutdown afterward. + drain.add_done_callback(stop_after_drain) + except BaseException as exc: + drain_error = exc + + if drain_timed_out: + raise TimeoutError( + f"{self._thread_name} tasks did not settle within {timeout:.1f}s" + ) + + try: + loop.call_soon_threadsafe(loop.stop) + except RuntimeError: + pass + thread.join(max(0.0, deadline - time.monotonic())) + + if thread.is_alive(): + raise TimeoutError( + f"{self._thread_name} did not stop within {timeout:.1f}s" + ) + if drain_error is not None: + raise drain_error + + +__all__ = [ + "AsyncFactory", + "AsyncRuntime", + "AsyncRuntimeClosedError", + "AsyncRuntimeError", + "RuntimeHandle", +] diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index dda0c10..4d817ac 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -6,13 +6,15 @@ Also provides the shared console and formatter globals. """ import asyncio +import concurrent.futures import inspect import logging import os import re import threading -from collections.abc import Callable -from typing import TYPE_CHECKING, Any +from collections.abc import AsyncIterator, Callable, Iterator +from contextlib import contextmanager +from typing import TYPE_CHECKING, Any, TypeVar from rich.console import Group # type: ignore[import-untyped] from rich.live import Live # type: ignore[import-untyped] @@ -21,8 +23,10 @@ from rich.panel import Panel # type: ignore[import-untyped] from rich.spinner import Spinner # type: ignore[import-untyped] from rich.text import Text # type: ignore[import-untyped] +from ..cancellation import bind_cancel_event from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest from ..paths import resolve_virtual_path +from ..runtime import AsyncRuntime, RuntimeHandle from .console import console from .diff_format import build_edit_diff from .formatter import ToolResultFormatter @@ -49,6 +53,24 @@ if TYPE_CHECKING: # Media file extensions that should trigger on_file_write callback _MEDIA_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".svg", ".pdf"} +_T = TypeVar("_T") +_QuestionRunner = Callable[[Any], Any] + + +class _StreamPromptCancelled(Exception): + """Internal control flow for an owned terminal prompt cancellation.""" + + +def _update_final_live_frame( + live: Any, + final_display: Any, + stream_handle: RuntimeHandle[Any] | None, +) -> None: + """Render the final frame unless this owned stream was cancelled.""" + if stream_handle is not None and stream_handle.cancelled(): + return + live.update(final_display) + live.refresh() def _graph_target_for_local_agent( @@ -119,10 +141,18 @@ formatter = ToolResultFormatter() # per-message scope so `/stop` only affects that message's run; scope-less # callers retain the legacy process-wide default event. _DEFAULT_STREAM_CANCEL_SCOPE = "__default__" -_stream_cancel_lock = threading.Lock() +# Serve signal handlers may request cancellation while the main thread is in a +# scoped cancellation lookup. Re-entrancy prevents the Python signal callback +# from deadlocking if it interrupts one of these short registry sections. +_stream_cancel_lock = threading.RLock() _stream_cancel_events: dict[str, threading.Event] = { _DEFAULT_STREAM_CANCEL_SCOPE: threading.Event() } +# The flag remains useful for cancellation requested before a stream starts and +# at synchronous HITL boundaries. Active Rich streams additionally register +# their owned-runtime handle so a request can interrupt a blocked ``__anext__`` +# immediately instead of waiting for the model to emit another event. +_stream_cancel_handles: dict[str, set[RuntimeHandle[Any]]] = {} # Backward-compat alias used by older tests and direct imports. _stream_cancel_event = _stream_cancel_events[_DEFAULT_STREAM_CANCEL_SCOPE] @@ -146,13 +176,83 @@ def _get_stream_cancel_event( def request_stream_cancel(cancel_scope: str | None = None) -> bool: - """Signal a specific in-flight stream to terminate.""" - event = _get_stream_cancel_event(cancel_scope, create=True) - already_requested = event.is_set() - event.set() + """Signal a stream and directly cancel any active owned coroutine.""" + scope_key = _stream_cancel_scope_key(cancel_scope) + with _stream_cancel_lock: + event = _stream_cancel_events.get(scope_key) + if event is None: + event = threading.Event() + _stream_cancel_events[scope_key] = event + already_requested = event.is_set() + event.set() + handles = tuple(_stream_cancel_handles.get(scope_key, ())) + + # Future.cancel() is thread-safe. Do it outside the registry lock because + # cancellation callbacks may settle quickly and unregister the handle. + for handle in handles: + handle.cancel() + # Sync tools may remain active in an executor after their awaiting graph + # task is cancelled, so stop their owned subprocesses explicitly. + from ..backends import cancel_active_shell_processes + + cancel_active_shell_processes(event) return not already_requested +def _register_stream_cancel_handle( + cancel_scope: str | None, + handle: RuntimeHandle[Any], +) -> None: + """Register an active owned task, honoring a pre-start stop request.""" + scope_key = _stream_cancel_scope_key(cancel_scope) + with _stream_cancel_lock: + handles = _stream_cancel_handles.setdefault(scope_key, set()) + handles.add(handle) + event = _stream_cancel_events.get(scope_key) + cancel_now = event is not None and event.is_set() + if cancel_now: + handle.cancel() + + +def _unregister_stream_cancel_handle( + cancel_scope: str | None, + handle: RuntimeHandle[Any], +) -> None: + scope_key = _stream_cancel_scope_key(cancel_scope) + with _stream_cancel_lock: + handles = _stream_cancel_handles.get(scope_key) + if handles is None: + return + handles.discard(handle) + if not handles: + _stream_cancel_handles.pop(scope_key, None) + + +def _run_owned_questionary_prompt( + question: Any, + *, + runtime: AsyncRuntime, + cancel_scope: str | None, +) -> Any: + """Run a terminal prompt as owned async work so stop can cancel it.""" + prompt_handle: RuntimeHandle[Any] | None = None + + def _register(handle: RuntimeHandle[Any]) -> None: + nonlocal prompt_handle + prompt_handle = handle + _register_stream_cancel_handle(cancel_scope, handle) + + try: + return runtime.run_sync(question.ask_async, on_submitted=_register) + except concurrent.futures.CancelledError as exc: + if is_stream_cancel_requested(cancel_scope): + raise _StreamPromptCancelled from exc + raise + finally: + if prompt_handle is not None: + _unregister_stream_cancel_handle(cancel_scope, prompt_handle) + + def is_stream_cancel_requested(cancel_scope: str | None = None) -> bool: event = _get_stream_cancel_event(cancel_scope) return event.is_set() if event is not None else False @@ -165,6 +265,40 @@ def clear_stream_cancel(cancel_scope: str | None = None) -> None: event.clear() +@contextmanager +def bind_stream_cancel(cancel_scope: str | None = None) -> Iterator[None]: + """Bind a stream's stop event for cancellation-aware blocking tools.""" + event = _get_stream_cancel_event(cancel_scope, create=True) + assert event is not None + with bind_cancel_event(event): + yield + + +async def iter_with_stream_cancel( + events: AsyncIterator[_T], + cancel_scope: str | None = None, +) -> AsyncIterator[_T]: + """Iterate graph events with the matching blocking-tool cancel context.""" + iterator = aiter(events) + try: + while True: + try: + # ContextVar tokens cannot safely straddle ``yield``: async + # generator finalization may run in a different task/context. + # The cancellation binding is only needed while requesting the + # next graph event, which includes any nested tool execution. + with bind_stream_cancel(cancel_scope): + event = await anext(iterator) + except StopAsyncIteration: + return + yield event + finally: + aclose = getattr(iterator, "aclose", None) + if aclose is not None: + with bind_stream_cancel(cancel_scope): + await aclose() + + def discard_stream_cancel(cancel_scope: str | None = None) -> None: """Drop a scope's stop signal after the owning request is fully done.""" scope_key = _stream_cancel_scope_key(cancel_scope) @@ -1026,6 +1160,8 @@ def _matches_shell_allow_list(command: str, allow_list: list[str]) -> bool: def _resolve_hitl_approval( interrupt_data: dict, prompt_fn: Callable[[list], list[dict] | None] | None = None, + *, + question_runner: _QuestionRunner | None = None, ) -> list[dict] | None: """Resolve HITL approval for an interrupt. @@ -1082,10 +1218,17 @@ def _resolve_hitl_approval( if prompt_fn is not None: return prompt_fn(action_requests) - return _prompt_hitl_approval(action_requests) + return _prompt_hitl_approval( + action_requests, + question_runner=question_runner, + ) -def _prompt_hitl_approval(action_requests: list) -> list[dict] | None: +def _prompt_hitl_approval( + action_requests: list, + *, + question_runner: _QuestionRunner | None = None, +) -> list[dict] | None: """Display approval prompt and get user decision. Returns list of decisions if approved, None if rejected. @@ -1130,11 +1273,14 @@ def _prompt_hitl_approval(action_requests: list) -> list[dict] | None: auto_label = "Approve all (session)" try: - selected = questionary.select( + question = questionary.select( "Approval required", choices=[approve_label, reject_label, auto_label], style=_PICKER_STYLE, - ).ask() + ) + selected = ( + question_runner(question) if question_runner is not None else question.ask() + ) except (EOFError, KeyboardInterrupt): console.print("[dim] Rejected.[/dim]") return None @@ -1152,37 +1298,11 @@ def _prompt_hitl_approval(action_requests: list) -> list[dict] | None: return None -# --------------------------------------------------------------------------- -# Async-to-sync bridge -# --------------------------------------------------------------------------- - - -def _create_event_loop() -> asyncio.AbstractEventLoop: - """Create and set the event loop for asyncio. - - Returns: - The created event loop. - """ - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - return loop - - -def _get_event_loop() -> asyncio.AbstractEventLoop: - """Get the event loop for asyncio. - - If no event loop is set, a new one is created. - - Returns: - The current event loop. - """ - loop = asyncio.get_event_loop() - if loop.is_closed(): - loop = _create_event_loop() - return loop - - -def _resolve_ask_user_prompt(ask_user_data: dict) -> dict: +def _resolve_ask_user_prompt( + ask_user_data: dict, + *, + question_runner: _QuestionRunner | None = None, +) -> dict: """Interactive console Q&A for ask_user events. Presents multiple-choice questions with arrow-key navigation via @@ -1235,11 +1355,16 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict: other_label = "Other (type your answer)" choice_labels.append(other_label) - selected = questionary.select( + question = questionary.select( prompt_text, choices=choice_labels, style=_PICKER_STYLE, - ).ask() + ) + selected = ( + question_runner(question) + if question_runner is not None + else question.ask() + ) if selected is None: # Ctrl+C raise KeyboardInterrupt @@ -1250,22 +1375,32 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict: continue if selected == other_label: - selected = questionary.text( + question = questionary.text( "Your answer:", validate=_make_validator(required), style=_PICKER_STYLE, - ).ask() + ) + selected = ( + question_runner(question) + if question_runner is not None + else question.ask() + ) if selected is None: raise KeyboardInterrupt answers.append(selected) else: - answer = questionary.text( + question = questionary.text( prompt_text, validate=_make_validator(required), style=_PICKER_STYLE, - ).ask() + ) + answer = ( + question_runner(question) + if question_runner is not None + else question.ask() + ) if answer is None: # Ctrl+C raise KeyboardInterrupt @@ -1298,6 +1433,7 @@ def _run_streaming( cancel_scope: str | None = None, *, gateway: GraphGateway, + runtime: AsyncRuntime | None = None, _state: StreamState | None = None, _hitl_depth: int = 0, _media_sent: set[str] | None = None, @@ -1329,6 +1465,31 @@ def _run_streaming( Returns: The final response text. """ + if runtime is None: + with AsyncRuntime(thread_name="evosci-stream-runtime") as owned_runtime: + return _run_streaming( + agent=agent, + message=message, + thread_id=thread_id, + show_thinking=show_thinking, + interactive=interactive, + on_thinking=on_thinking, + on_todo=on_todo, + on_file_write=on_file_write, + on_stream_event=on_stream_event, + status_footer_builder=status_footer_builder, + metadata=metadata, + hitl_prompt_fn=hitl_prompt_fn, + ask_user_prompt_fn=ask_user_prompt_fn, + cancel_scope=cancel_scope, + gateway=gateway, + runtime=owned_runtime, + _state=_state, + _hitl_depth=_hitl_depth, + _media_sent=_media_sent, + _sent_thinking_text=_sent_thinking_text, + ) + # Scope-less callers keep the legacy single-event semantics. Scoped # callers use unique per-request scopes, so pre-start `/stop` must # remain armed until this run consumes it. @@ -1348,108 +1509,114 @@ def _run_streaming( async def _consume() -> None: nonlocal _sent_thinking_text, _todo_sent - async for event in gateway.stream_events( + event_stream = gateway.stream_events( RunRequest( message=message, thread_id=thread_id, metadata=metadata, target=_graph_target_for_local_agent(agent, metadata), ) - ): - if is_stream_cancel_requested(cancel_scope): - _stopped_response() - return - event_type = state.handle_event(event) + ) + try: + async for event in event_stream: + if is_stream_cancel_requested(cancel_scope): + _stopped_response() + return + event_type = state.handle_event(event) - # Relay thinking to channel when transitioning away from - # thinking phase. Uses content comparison so that replayed - # thinking after resume is skipped, but genuinely new - # thinking is still delivered. - if ( - on_thinking - and event_type != "thinking" - and state.thinking_text - and len(state.thinking_text) >= _MIN_THINKING_LEN - ): - current = state.thinking_text.rstrip() - if current != _sent_thinking_text: - on_thinking(current) - _sent_thinking_text = current + # Relay thinking to channel when transitioning away from + # thinking phase. Uses content comparison so that replayed + # thinking after resume is skipped, but genuinely new + # thinking is still delivered. + if ( + on_thinking + and event_type != "thinking" + and state.thinking_text + and len(state.thinking_text) >= _MIN_THINKING_LEN + ): + current = state.thinking_text.rstrip() + if current != _sent_thinking_text: + on_thinking(current) + _sent_thinking_text = current - # Send todo list to channel on first write_todos tool_call - if ( - on_todo - and not _todo_sent - and event_type == "tool_call" - and event.get("name") == "write_todos" - and state.todo_items - ): - on_todo(state.todo_items) - _todo_sent = True + # Send todo list to channel on first write_todos tool_call + if ( + on_todo + and not _todo_sent + and event_type == "tool_call" + and event.get("name") == "write_todos" + and state.todo_items + ): + on_todo(state.todo_items) + _todo_sent = True - # Send media file to channel when write_file succeeds - if ( - on_file_write - and event_type == "tool_result" - and event.get("name") == "write_file" - and event.get("success") - ): - wf_path = "" - for tc in reversed(state.tool_calls): - if tc.get("name") == "write_file": - p = tc.get("args", {}).get("path", "") - if p and p not in _media_sent: - wf_path = p - break - if wf_path: - ext = os.path.splitext(wf_path)[1].lower() - if ext in _MEDIA_EXTENSIONS: - real_path = str(resolve_virtual_path(wf_path)) - if os.path.isfile(real_path): - _media_sent.add(wf_path) - on_file_write(real_path) + # Send media file to channel when write_file succeeds + if ( + on_file_write + and event_type == "tool_result" + and event.get("name") == "write_file" + and event.get("success") + ): + wf_path = "" + for tc in reversed(state.tool_calls): + if tc.get("name") == "write_file": + p = tc.get("args", {}).get("path", "") + if p and p not in _media_sent: + wf_path = p + break + if wf_path: + ext = os.path.splitext(wf_path)[1].lower() + if ext in _MEDIA_EXTENSIONS: + real_path = str(resolve_virtual_path(wf_path)) + if os.path.isfile(real_path): + _media_sent.add(wf_path) + on_file_write(real_path) - # Send media file to channel when read_file returns an image - if ( - on_file_write - and event_type == "tool_result" - and event.get("name") == "read_file" - and event.get("success") - ): - rf_path = "" - for tc in reversed(state.tool_calls): - if tc.get("name") == "read_file": - p = tc.get("args", {}).get("file_path", "") or tc.get( - "args", {} - ).get("path", "") - if p and p not in _media_sent: - rf_path = p - break - if rf_path: - ext = os.path.splitext(rf_path)[1].lower() - if ext in _MEDIA_EXTENSIONS: - real_path = rf_path - if not os.path.isfile(real_path): - real_path = str(resolve_virtual_path(rf_path)) - if os.path.isfile(real_path): - _media_sent.add(rf_path) - on_file_write(real_path) + # Send media file to channel when read_file returns an image + if ( + on_file_write + and event_type == "tool_result" + and event.get("name") == "read_file" + and event.get("success") + ): + rf_path = "" + for tc in reversed(state.tool_calls): + if tc.get("name") == "read_file": + p = tc.get("args", {}).get("file_path", "") or tc.get( + "args", {} + ).get("path", "") + if p and p not in _media_sent: + rf_path = p + break + if rf_path: + ext = os.path.splitext(rf_path)[1].lower() + if ext in _MEDIA_EXTENSIONS: + real_path = rf_path + if not os.path.isfile(real_path): + real_path = str(resolve_virtual_path(rf_path)) + if os.path.isfile(real_path): + _media_sent.add(rf_path) + on_file_write(real_path) - if on_stream_event is not None: - callback_result = on_stream_event(event_type, state) - if inspect.isawaitable(callback_result): - await callback_result + if on_stream_event is not None: + callback_result = on_stream_event(event_type, state) + if inspect.isawaitable(callback_result): + await callback_result - live.update( - create_streaming_display( - **state.get_display_args(), - show_thinking=show_thinking, - response_markdown=state.get_response_markdown(), - status_footer=( - status_footer_builder() if status_footer_builder else None - ), + live.update( + create_streaming_display( + **state.get_display_args(), + show_thinking=show_thinking, + response_markdown=state.get_response_markdown(), + status_footer=( + status_footer_builder() if status_footer_builder else None + ), + ) ) - ) + finally: + aclose = getattr(event_stream, "aclose", None) + if aclose is not None: + await aclose() try: if is_stream_cancel_requested(cancel_scope): @@ -1469,32 +1636,6 @@ def _run_streaming( ), ) ) - # Determine how to run the async streaming coroutine. - # - In TUI mode (Textual), there's already a running event loop; - # nest_asyncio is needed to allow run_until_complete inside it. - # - In serve/CLI mode, the main thread has no running loop; - # use a fresh event loop directly (no nest_asyncio needed or wanted, - # since nest_asyncio.apply() patches globally and breaks the bus - # thread's event loop Task-context detection). - try: - running_loop = asyncio.get_running_loop() - except RuntimeError: - running_loop = None - - if running_loop is not None: - # Already inside a running loop (TUI) — must use nest_asyncio. - # NOTE: nest_asyncio.apply() is global and irreversible within - # the process; avoid mixing TUI and serve modes in one process. - import nest_asyncio # type: ignore[import-untyped] - - nest_asyncio.apply() - loop = running_loop - else: - # No running loop (serve/CLI) — create a fresh one - try: - loop = _get_event_loop() - except RuntimeError: - loop = _create_event_loop() async def _run_with_refresh() -> None: async def _periodic_refresh() -> None: @@ -1507,7 +1648,8 @@ def _run_streaming( refresh_task = asyncio.ensure_future(_periodic_refresh()) try: - await _consume() + with bind_stream_cancel(cancel_scope): + await _consume() finally: refresh_task.cancel() try: @@ -1552,10 +1694,24 @@ def _run_streaming( interactive, status_footer_builder ), ) - live.update(final_display) - live.refresh() + _update_final_live_frame(live, final_display, stream_handle) - loop.run_until_complete(_run_with_refresh()) + stream_handle: RuntimeHandle[None] | None = None + + def _register(handle: RuntimeHandle[None]) -> None: + nonlocal stream_handle + stream_handle = handle + _register_stream_cancel_handle(cancel_scope, handle) + + try: + runtime.run_sync(_run_with_refresh, on_submitted=_register) + except concurrent.futures.CancelledError: + if not is_stream_cancel_requested(cancel_scope): + raise + _stopped_response() + finally: + if stream_handle is not None: + _unregister_stream_cancel_handle(cancel_scope, stream_handle) # Flush any remaining thinking that wasn't sent during streaming. if on_thinking and state.thinking_text: @@ -1571,7 +1727,17 @@ def _run_streaming( if ask_user_prompt_fn is not None: result = ask_user_prompt_fn(state.pending_ask_user) else: - result = _resolve_ask_user_prompt(state.pending_ask_user) + try: + result = _resolve_ask_user_prompt( + state.pending_ask_user, + question_runner=lambda question: _run_owned_questionary_prompt( + question, + runtime=runtime, + cancel_scope=cancel_scope, + ), + ) + except _StreamPromptCancelled: + return _stopped_response() from langgraph.types import Command # type: ignore[import-untyped] state.pending_ask_user = None @@ -1594,6 +1760,7 @@ def _run_streaming( ask_user_prompt_fn=ask_user_prompt_fn, cancel_scope=cancel_scope, gateway=gateway, + runtime=runtime, _state=state, _hitl_depth=_hitl_depth + 1, _media_sent=_media_sent, @@ -1604,10 +1771,22 @@ def _run_streaming( if state.pending_interrupt is not None and _hitl_depth < _MAX_HITL_ITERATIONS: if is_stream_cancel_requested(cancel_scope): return _stopped_response() - decisions = _resolve_hitl_approval( - state.pending_interrupt, - prompt_fn=hitl_prompt_fn, - ) + try: + decisions = _resolve_hitl_approval( + state.pending_interrupt, + prompt_fn=hitl_prompt_fn, + question_runner=( + None + if hitl_prompt_fn is not None + else lambda question: _run_owned_questionary_prompt( + question, + runtime=runtime, + cancel_scope=cancel_scope, + ) + ), + ) + except _StreamPromptCancelled: + return _stopped_response() if is_stream_cancel_requested(cancel_scope): return _stopped_response() if decisions is not None: @@ -1633,6 +1812,7 @@ def _run_streaming( ask_user_prompt_fn=ask_user_prompt_fn, cancel_scope=cancel_scope, gateway=gateway, + runtime=runtime, _state=state, _hitl_depth=_hitl_depth + 1, _media_sent=_media_sent, diff --git a/pyproject.toml b/pyproject.toml index 0e06e2d..67e22f1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -42,7 +42,6 @@ dependencies = [ "filelock>=3.16", "lazy-loader>=0.5", "markdownify>=1.2", - "nest-asyncio>=1.6", "tzlocal>=5.0", "langchain-mcp-adapters>=0.2", # <8.2.7: 8.2.7 Kitty "report-all-keys" breaks CJK input on iTerm2 diff --git a/tests/test_backends.py b/tests/test_backends.py index 8fbdf03..6200c1e 100644 --- a/tests/test_backends.py +++ b/tests/test_backends.py @@ -2,7 +2,9 @@ import re import shlex +import subprocess import sys +import time from pathlib import Path import pytest @@ -1125,6 +1127,20 @@ class TestSandboxId: # === execute() literal cwd sanitization === +class TestExecuteValidation: + @pytest.mark.parametrize("command", ["", None, 123]) + def test_execute_rejects_empty_or_non_string_commands(self, command, tmp_workspace): + backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) + + response = backend.execute(command) + + assert response == backends.ExecuteResponse( + output="Error: Command must be a non-empty string.", + exit_code=1, + truncated=False, + ) + + class TestExecuteCwdSanitization: def test_literal_workspace_path_replaced(self, tmp_workspace, monkeypatch): """``prepare_sandbox_command`` must rewrite a literal workspace-root @@ -1140,7 +1156,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) command = f"mkdir -p {tmp_workspace}/test-sanitized && echo ok" @@ -1161,7 +1179,9 @@ class TestExecuteCwdSanitization: captured["timeout"] = timeout return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) command = ( "ssh -p 2222 -i key host " @@ -1183,7 +1203,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) workspace = tmp_path / "ws" workspace.mkdir() backend = CustomSandboxBackend(root_dir=str(workspace), virtual_mode=True) @@ -1213,7 +1235,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) resp = backend.execute("ssh -N host", timeout=30) @@ -1256,7 +1280,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) command = "ssh host 'echo $(cat /etc/passwd)'" @@ -1274,7 +1300,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) resp = backend.execute( @@ -1297,7 +1325,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) resp = backend.execute("ssh host 'pwd' > /tmp/out", timeout=30) @@ -1314,7 +1344,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) command = "echo __EVOSCI_SSH_REMOTE_0__ && ssh host 'ls /home'" @@ -1356,7 +1388,9 @@ class TestExecuteCwdSanitization: captured["timeout"] = timeout return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) command = "ssh host 'ls /home/username/project'" @@ -1390,7 +1424,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) resp = backend.execute(f"{ssh_path} host ls /home/username/project", timeout=30) @@ -1407,7 +1443,9 @@ class TestExecuteCwdSanitization: captured["command"] = command return backends.ExecuteResponse(output="ok", exit_code=0, truncated=False) - monkeypatch.setattr(backends.LocalShellBackend, "execute", fake_execute) + monkeypatch.setattr( + CustomSandboxBackend, "_execute_prepared_command", fake_execute + ) backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) command = "cat /data/file.txt && ssh host 'ls /home/username/project'" @@ -1514,6 +1552,84 @@ class TestExecuteTimeout: execute_accepts_timeout.cache_clear() assert execute_accepts_timeout(CustomSandboxBackend) is True + @pytest.mark.skipif( + sys.platform == "win32", + reason="POSIX process-group regression", + ) + def test_timeout_kills_descendants_after_shell_leader_exits(self, tmp_workspace): + """A dead shell leader must not hide descendants retaining its pipes.""" + backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True) + + started = time.monotonic() + response = backend.execute("sleep 2 &", timeout=0.1) + elapsed = time.monotonic() - started + + assert response.exit_code == 124 + assert elapsed < 1 + + @pytest.mark.skipif( + sys.platform == "win32", + reason="POSIX detached-process regression", + ) + def test_timeout_bounds_drain_when_detached_descendant_holds_pipes( + self, + tmp_workspace, + monkeypatch, + ): + """An escaped descendant cannot hold execute() open through inherited pipes.""" + monkeypatch.setattr(backends, "_PROCESS_DRAIN_GRACE_SECONDS", 0.05) + backend = CustomSandboxBackend( + root_dir=tmp_workspace, + virtual_mode=True, + env={"EVOSCI_TEST_PYTHON": sys.executable}, + ) + code = ( + "import os,time; " + "pid=os.fork(); " + "os._exit(0) if pid else (os.setsid(), time.sleep(1), os._exit(0))" + ) + # Pass the absolute executable through the environment so virtual-path + # normalization does not reinterpret it as a workspace path. + command = f'"$EVOSCI_TEST_PYTHON" -c {shlex.quote(code)}' + + started = time.monotonic() + response = backend.execute(command, timeout=0.05) + elapsed = time.monotonic() - started + + assert response.exit_code == 124 + assert response.truncated is True + assert elapsed < 0.4 + + +def test_active_shell_registry_lock_allows_signal_handler_reentry(): + """A signal handler can re-enter registry code on the interrupted thread.""" + lock = backends._active_shell_processes_lock + assert lock.acquire(timeout=0.1) + try: + assert lock.acquire(timeout=0.1) + lock.release() + finally: + lock.release() + + +def test_terminate_process_tree_does_not_target_reaped_pid(monkeypatch): + """A completed Popen PID must not be reused as a process-group target.""" + process = subprocess.Popen([sys.executable, "-c", "pass"]) + process.wait(timeout=5) + termination_attempted = False + + def fail_termination(*args, **kwargs): + nonlocal termination_attempted + termination_attempted = True + + monkeypatch.setattr(backends.os, "killpg", fail_termination, raising=False) + monkeypatch.setattr(backends.subprocess, "run", fail_termination) + monkeypatch.setattr(process, "kill", fail_termination) + + backends._terminate_process_tree(process) + + assert termination_attempted is False + # === '..' traversal false-positive fix === diff --git a/tests/test_channel_comprehensive.py b/tests/test_channel_comprehensive.py index d556653..f4e7bfc 100644 --- a/tests/test_channel_comprehensive.py +++ b/tests/test_channel_comprehensive.py @@ -15,6 +15,7 @@ Test groups: from __future__ import annotations import asyncio +import threading from datetime import datetime from unittest.mock import AsyncMock, MagicMock @@ -39,6 +40,7 @@ from EvoScientist.channels.consumer import InboundConsumer from EvoScientist.channels.formatter import convert_markdown from EvoScientist.channels.middleware import DedupCache, MentionGatingMiddleware from EvoScientist.channels.retry import RetryConfig, RetryInfo, retry_async +from EvoScientist.runtime import AsyncRuntime, AsyncRuntimeError # ═══════════════════════════════════════════════════════════════════ # Helpers @@ -74,6 +76,41 @@ async def _wait_for_async(predicate) -> None: await asyncio.sleep(0) +class TestInboundSyncAdapter: + def test_reuses_explicit_runtime(self): + channel = StubChannel() + channel._inbound_middlewares = [] + raw = RawIncoming(sender_id="user", chat_id="chat", text="hello") + + with AsyncRuntime(thread_name="test-channel-adapter") as runtime: + message = channel._build_inbound(raw, runtime=runtime) + + assert message is not None + assert message.content == "hello" + + def test_direct_sync_call_scopes_and_closes_runtime(self): + channel = StubChannel() + channel._inbound_middlewares = [] + raw = RawIncoming(sender_id="user", chat_id="chat", text="hello") + + message = channel._build_inbound(raw) + + assert message is not None + assert not any( + thread.name == "evosci-channel-adapter-runtime" and thread.is_alive() + for thread in threading.enumerate() + ) + + async def test_async_caller_must_use_async_api(self): + channel = StubChannel() + channel._inbound_middlewares = [] + raw = RawIncoming(sender_id="user", chat_id="chat", text="hello") + + with AsyncRuntime(thread_name="test-channel-adapter") as runtime: + with pytest.raises(AsyncRuntimeError, match="running event loop"): + channel._build_inbound(raw, runtime=runtime) + + # ═══════════════════════════════════════════════════════════════════ # 1. DedupCache # ═══════════════════════════════════════════════════════════════════ diff --git a/tests/test_channel_debug.py b/tests/test_channel_debug.py index f81af68..0cc66a6 100644 --- a/tests/test_channel_debug.py +++ b/tests/test_channel_debug.py @@ -2,6 +2,7 @@ import asyncio import logging +import threading from unittest.mock import AsyncMock, MagicMock, patch from EvoScientist.channels.debug import ( @@ -359,3 +360,25 @@ def test_emit_debug_event_warns_on_level_mismatch(caplog): # Reset for other tests dbg._warned_debug_level_mismatch = False + + +async def test_standalone_agent_construction_runs_off_channel_loop(monkeypatch): + import EvoScientist.EvoScientist as agent_module + from EvoScientist.channels.standalone import _create_standalone_agent + + channel_thread = threading.current_thread() + sentinel = object() + + def fake_create_cli_agent(): + assert threading.current_thread() is not channel_thread + try: + asyncio.get_running_loop() + except RuntimeError: + pass + else: # pragma: no cover - assertion branch + raise AssertionError("agent construction inherited the channel loop") + return sentinel + + monkeypatch.setattr(agent_module, "create_cli_agent", fake_create_cli_agent) + + assert await _create_standalone_agent() is sentinel diff --git a/tests/test_channel_sends.py b/tests/test_channel_sends.py new file mode 100644 index 0000000..ff86f77 --- /dev/null +++ b/tests/test_channel_sends.py @@ -0,0 +1,75 @@ +"""Behavioral tests for channel sends crossing frontend event loops.""" + +from __future__ import annotations + +import asyncio +import logging +import threading + +import pytest + +from EvoScientist.cli.channel_sends import PendingChannelSends +from EvoScientist.runtime import AsyncRuntime + + +@pytest.mark.asyncio +async def test_pending_send_does_not_stall_owned_runtime() -> None: + """A blocked channel transport must not block unrelated runtime work.""" + bus_loop = asyncio.get_running_loop() + send_started = asyncio.Event() + release_send = asyncio.Event() + runtime_progressed = threading.Event() + sends = PendingChannelSends(bus_loop, logging.getLogger(__name__)) + + async def _blocked_send() -> None: + send_started.set() + await release_send.wait() + + async def _stream_callback_and_probe() -> None: + sends.submit(_blocked_send(), "Thinking") + await asyncio.sleep(0) + runtime_progressed.set() + + with AsyncRuntime(thread_name="test-channel-send-runtime") as runtime: + callback = runtime.submit(_stream_callback_and_probe) + await asyncio.wait_for(send_started.wait(), timeout=1) + assert runtime_progressed.wait(timeout=1) + callback.result(timeout=1) + + settle = asyncio.create_task(sends.settle_async()) + await asyncio.sleep(0) + assert not settle.done() + + release_send.set() + await asyncio.wait_for(settle, timeout=1) + + +@pytest.mark.asyncio +async def test_async_settlement_waits_for_every_scheduled_send() -> None: + """The channel response can wait for all callback delivery off-loop.""" + first_started = asyncio.Event() + first_release = asyncio.Event() + events: list[str] = [] + sends = PendingChannelSends(asyncio.get_running_loop(), logging.getLogger(__name__)) + + async def _first() -> None: + events.append("first-started") + first_started.set() + await first_release.wait() + events.append("first-finished") + + async def _second() -> None: + events.append("second-finished") + + sends.submit(_first(), "First") + sends.submit(_second(), "Second") + + settle = asyncio.create_task(sends.settle_async()) + await asyncio.wait_for(first_started.wait(), timeout=1) + assert not settle.done() + assert events == ["first-started"] + + first_release.set() + await asyncio.wait_for(settle, timeout=1) + + assert events == ["first-started", "first-finished", "second-finished"] diff --git a/tests/test_cli_async_runtime.py b/tests/test_cli_async_runtime.py new file mode 100644 index 0000000..2d53fdb --- /dev/null +++ b/tests/test_cli_async_runtime.py @@ -0,0 +1,131 @@ +"""The first bounded CLI adoption of the owned async runtime.""" + +import asyncio +import threading + +import pytest +from typer.testing import CliRunner + +import EvoScientist.cli.commands # noqa: F401 - registers commands on app +from EvoScientist.cli import commands +from EvoScientist.cli._app import app + + +@pytest.mark.parametrize("args", [["sessions"], ["sessions", "stats"]]) +def test_sessions_stats_uses_and_closes_cli_owned_runtime(monkeypatch, args): + execution: dict[str, object] = {} + + async def fake_db_stats(): + execution["thread"] = threading.current_thread().name + execution["loop"] = asyncio.get_running_loop() + return { + "db_path": "/tmp/sessions.db", + "size_bytes": 0, + "thread_count": 0, + "checkpoint_count": 0, + "write_count": 0, + "top_threads": [], + } + + monkeypatch.setattr("EvoScientist.sessions.db_stats", fake_db_stats) + + result = CliRunner().invoke(app, args) + + assert result.exit_code == 0, result.exception + assert execution["thread"] == "evosci-async-runtime" + assert isinstance(execution["loop"], asyncio.AbstractEventLoop) + assert not any( + thread.name == "evosci-async-runtime" and thread.is_alive() + for thread in threading.enumerate() + ) + + +@pytest.mark.parametrize( + ("args", "patch_target"), + [ + (["onboard"], "EvoScientist.config.onboard.run_onboard"), + (["configure", "channels"], "EvoScientist.cli.commands._run_onboard_cli"), + ], +) +def test_onboarding_commands_share_and_close_cli_runtime( + monkeypatch, args, patch_target +): + execution: dict[str, object] = {} + + async def record_execution(): + execution["thread"] = threading.current_thread().name + execution["loop"] = asyncio.get_running_loop() + + def fake_onboard(**kwargs): + runtime = kwargs["runtime"] + execution["runtime"] = runtime + runtime.run_sync(record_execution) + return True + + monkeypatch.setattr(patch_target, fake_onboard) + + result = CliRunner().invoke(app, args) + + assert result.exit_code == 0, result.exception + assert execution["thread"] == "evosci-async-runtime" + assert isinstance(execution["loop"], asyncio.AbstractEventLoop) + assert not any( + thread.name == "evosci-async-runtime" and thread.is_alive() + for thread in threading.enumerate() + ) + + +def test_channel_setup_shares_and_closes_cli_runtime(monkeypatch): + execution: dict[str, object] = {} + + async def record_execution(): + execution["thread"] = threading.current_thread().name + execution["loop"] = asyncio.get_running_loop() + + def fake_step_channels(_config, *, runtime): + runtime.run_sync(record_execution) + return {} + + monkeypatch.setattr("EvoScientist.config.load_config", object) + monkeypatch.setattr( + "EvoScientist.config.onboard.channels._step_channels", fake_step_channels + ) + + result = CliRunner().invoke(app, ["channel", "setup"]) + + assert result.exit_code == 0, result.exception + assert execution["thread"] == "evosci-async-runtime" + assert isinstance(execution["loop"], asyncio.AbstractEventLoop) + assert not any( + thread.name == "evosci-async-runtime" and thread.is_alive() + for thread in threading.enumerate() + ) + + +def test_cli_reports_runtime_close_timeout_without_raw_exception(monkeypatch): + class _TimeoutRuntime: + def run_sync(self, factory): + return asyncio.run(factory()) + + def close(self): + raise TimeoutError("executor work still active") + + async def fake_db_stats(): + return { + "db_path": "/tmp/sessions.db", + "size_bytes": 0, + "thread_count": 0, + "checkpoint_count": 0, + "write_count": 0, + "top_threads": [], + } + + monkeypatch.setattr(commands, "AsyncRuntime", _TimeoutRuntime) + monkeypatch.setattr("EvoScientist.sessions.db_stats", fake_db_stats) + + result = CliRunner().invoke(app, ["sessions", "stats"]) + + assert result.exit_code == 1 + assert "Async runtime shutdown did not complete" in result.output + assert "executor work still active" in result.output + assert not isinstance(result.exception, TimeoutError) diff --git a/tests/test_cli_interrupts.py b/tests/test_cli_interrupts.py new file mode 100644 index 0000000..494c7c2 --- /dev/null +++ b/tests/test_cli_interrupts.py @@ -0,0 +1,85 @@ +"""Signal-level regression tests for Rich CLI turn cancellation.""" + +import asyncio +import signal +import threading + +import pytest + +from EvoScientist.cli import interactive + + +@pytest.mark.asyncio +async def test_session_turns_are_serialized() -> None: + """A channel turn cannot start while a foreground turn owns the session.""" + turn_lock = asyncio.Lock() + first_started = asyncio.Event() + release_first = asyncio.Event() + second_started = asyncio.Event() + order: list[str] = [] + + async def first_turn() -> None: + order.append("first-started") + first_started.set() + await release_first.wait() + order.append("first-finished") + + async def second_turn() -> None: + order.append("second-started") + second_started.set() + + first = asyncio.create_task(interactive._run_serialized_turn(turn_lock, first_turn)) + await first_started.wait() + second = asyncio.create_task( + interactive._run_serialized_turn(turn_lock, second_turn) + ) + await asyncio.sleep(0) + + assert not second_started.is_set() + release_first.set() + await asyncio.gather(first, second) + assert order == ["first-started", "first-finished", "second-started"] + + +@pytest.mark.skipif( + threading.current_thread() is not threading.main_thread(), + reason="process signal handlers require the main thread", +) +def test_ctrl_c_can_cancel_two_separate_rich_cli_turns(monkeypatch): + """A recovered turn must not consume asyncio.run's force-quit budget.""" + started = asyncio.Event() + calls = 0 + + async def fake_run_streaming_async(**kwargs): + nonlocal calls + assert kwargs["recover_on_cancel"] is True + calls += 1 + started.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + current = asyncio.current_task() + assert current is not None + current.uncancel() + return "[Stopped.]" + + monkeypatch.setattr(interactive, "run_streaming_async", fake_run_streaming_async) + original_sigint = signal.getsignal(signal.SIGINT) + + async def cancel_started_turn() -> None: + await started.wait() + signal.raise_signal(signal.SIGINT) + + async def scenario() -> None: + runner_sigint = signal.getsignal(signal.SIGINT) + for _ in range(2): + started.clear() + sender = asyncio.create_task(cancel_started_turn()) + assert await interactive._run_rich_cli_streaming_turn() == "[Stopped.]" + await sender + assert signal.getsignal(signal.SIGINT) is runner_sigint + + asyncio.run(scenario()) + + assert calls == 2 + assert signal.getsignal(signal.SIGINT) is original_sigint diff --git a/tests/test_cli_serve.py b/tests/test_cli_serve.py index d19234b..21dd646 100644 --- a/tests/test_cli_serve.py +++ b/tests/test_cli_serve.py @@ -3,10 +3,12 @@ from __future__ import annotations import os +import signal from types import SimpleNamespace from EvoScientist.cli import commands from EvoScientist.config import MemoryObservationWriter +from EvoScientist.runtime import AsyncRuntime def _make_config( @@ -37,6 +39,7 @@ def _make_config( memory_observation_writer=MemoryObservationWriter.ALL, memory_workers_enabled=False, memory_skill_synthesis_enabled=False, + model="test-model", provider="anthropic", anthropic_auth_mode="api_key", openai_auth_mode="api_key", @@ -55,6 +58,8 @@ def _run_serve_once( auto_mode: bool = False, ask_user: bool = False, dangerous: bool = False, + message_queue=None, + process_message=None, ): import EvoScientist.config as config_mod @@ -67,8 +72,11 @@ def _run_serve_once( def _fake_ensure_dirs(): order.append(("ensure_dirs", None)) - def _fake_load_agent(workspace_dir=None, checkpointer=None, config=None): + def _fake_load_agent( + workspace_dir=None, checkpointer=None, config=None, *, runtime=None + ): captured["workspace_dir"] = workspace_dir + captured["async_runtime"] = runtime return object() def _fake_start_channels_bus_mode(cfg, agent, thread_id, *, send_thinking=None): @@ -92,7 +100,13 @@ def _run_serve_once( commands, "_start_channels_bus_mode", _fake_start_channels_bus_mode ) monkeypatch.setattr(commands, "_channels_stop", _fake_channels_stop) - monkeypatch.setattr(commands, "_message_queue", _InterruptQueue()) + monkeypatch.setattr( + commands, + "_message_queue", + message_queue if message_queue is not None else _InterruptQueue(), + ) + if process_message is not None: + monkeypatch.setattr(commands, "_serve_process_message", process_message) def _fake_get_effective_config(cli_overrides=None): captured["cli_overrides"] = dict(cli_overrides or {}) @@ -106,15 +120,18 @@ def _run_serve_once( if cwd is not None: monkeypatch.setattr(commands.os, "getcwd", lambda: cwd) - commands.serve( - no_thinking=no_thinking, - workdir=workdir, - debug=debug, - auto_approve=auto_approve, - auto_mode=auto_mode, - ask_user=ask_user, - dangerous=dangerous, - ) + with AsyncRuntime(thread_name="test-serve-runtime") as runtime: + monkeypatch.setattr(commands, "_get_cli_async_runtime", lambda _ctx: runtime) + commands.serve( + object(), + no_thinking=no_thinking, + workdir=workdir, + debug=debug, + auto_approve=auto_approve, + auto_mode=auto_mode, + ask_user=ask_user, + dangerous=dangerous, + ) return order, captured @@ -266,3 +283,87 @@ def test_serve_dangerous_sets_dangerous_mode(monkeypatch, tmp_path): ) assert captured["cli_overrides"] == {"dangerous_mode": True} + + +def test_serve_sigterm_cancels_active_message_scope_before_shutdown( + monkeypatch, tmp_path +): + handlers = {} + message = commands.ChannelMessage( + msg_id="message-1", + content="run a long command", + sender="user", + channel_type="telegram", + chat_id="chat-1", + ) + + class _OneMessageQueue: + def get(self, timeout=None): + return message + + def _fake_signal(signum, handler): + previous = handlers.get(signum, signal.SIG_DFL) + handlers[signum] = handler + return previous + + cancelled_scopes = [] + + def _fake_process_message(*args, **kwargs): + handlers[signal.SIGTERM](signal.SIGTERM, None) + + monkeypatch.setattr(signal, "signal", _fake_signal) + monkeypatch.setattr( + "EvoScientist.stream.display.request_stream_cancel", + cancelled_scopes.append, + ) + + _run_serve_once( + monkeypatch, + _make_config(default_workdir=str(tmp_path)), + message_queue=_OneMessageQueue(), + process_message=_fake_process_message, + ) + + assert cancelled_scopes == ["channel:telegram:chat-1:message-1"] + + +def test_serve_sigint_cancels_active_message_scope_before_interrupt( + monkeypatch, tmp_path +): + handlers = {} + message = commands.ChannelMessage( + msg_id="message-2", + content="run a long command", + sender="user", + channel_type="telegram", + chat_id="chat-2", + ) + + class _OneMessageQueue: + def get(self, timeout=None): + return message + + def _fake_signal(signum, handler): + previous = handlers.get(signum, signal.SIG_DFL) + handlers[signum] = handler + return previous + + cancelled_scopes = [] + + def _fake_process_message(*args, **kwargs): + handlers[signal.SIGINT](signal.SIGINT, None) + + monkeypatch.setattr(signal, "signal", _fake_signal) + monkeypatch.setattr( + "EvoScientist.stream.display.request_stream_cancel", + cancelled_scopes.append, + ) + + _run_serve_once( + monkeypatch, + _make_config(default_workdir=str(tmp_path)), + message_queue=_OneMessageQueue(), + process_message=_fake_process_message, + ) + + assert cancelled_scopes == ["channel:telegram:chat-2:message-2"] diff --git a/tests/test_code_interpreter_middleware.py b/tests/test_code_interpreter_middleware.py index 68fddf5..9ed6c00 100644 --- a/tests/test_code_interpreter_middleware.py +++ b/tests/test_code_interpreter_middleware.py @@ -9,6 +9,7 @@ allowlist (``task()`` stays reachable as the REPL global, with responseSchema). from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -63,6 +64,28 @@ async def test_aclose_code_interpreters_closes_registered_instances(monkeypatch) close.assert_awaited_once_with() +@pytest.mark.asyncio +async def test_aclose_code_interpreters_bounds_stalled_cleanup(monkeypatch, caplog): + middleware = create_code_interpreter_middleware() + cancelled = asyncio.Event() + + async def stalled_close(): + try: + await asyncio.Event().wait() + finally: + cancelled.set() + + monkeypatch.setattr(middleware, "aclose", stalled_close) + + await asyncio.wait_for( + aclose_code_interpreters(timeout=0.01), + timeout=0.2, + ) + + assert cancelled.is_set() + assert "cleanup did not finish within 0.01 seconds" in caplog.text + + 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_event_loop.py b/tests/test_event_loop.py index 186e66f..b1047b7 100644 --- a/tests/test_event_loop.py +++ b/tests/test_event_loop.py @@ -1,325 +1,149 @@ -"""Tests for event loop management in streaming display.""" +"""Owned-runtime tests for the synchronous Rich streaming adapter.""" import asyncio +import threading from unittest.mock import Mock, patch import pytest -from EvoScientist.stream.display import _create_event_loop, _get_event_loop +from EvoScientist.runtime import AsyncRuntime, AsyncRuntimeError +from EvoScientist.stream.display import _run_streaming from tests.fakes import FakeGraphGateway -class _TrackingEventLoopPolicy(asyncio.DefaultEventLoopPolicy): - """Event loop policy that records loops created by one test.""" +def _text_stream(loops, response="test response"): + async def _stream(_request): + loops.append(asyncio.get_running_loop()) + yield {"type": "text", "content": response} + yield {"type": "done", "response": response} - def __init__(self): - super().__init__() - self.created_loops: list[asyncio.AbstractEventLoop] = [] - - def new_event_loop(self) -> asyncio.AbstractEventLoop: - loop = super().new_event_loop() - self.created_loops.append(loop) - return loop + return _stream -@pytest.fixture(autouse=True) -def isolated_event_loop_policy(): - previous_policy = asyncio.get_event_loop_policy() - test_policy = _TrackingEventLoopPolicy() - asyncio.set_event_loop_policy(test_policy) - try: - yield - finally: - try: - for loop in test_policy.created_loops: - if not loop.is_closed(): - loop.close() - finally: - asyncio.set_event_loop_policy(previous_policy) +def test_sequential_streams_reuse_application_runtime_loop(): + loops: list[asyncio.AbstractEventLoop] = [] + gateway = FakeGraphGateway(stream=_text_stream(loops)) - -class TestCreateEventLoop: - """Tests for _create_event_loop helper.""" - - def test_creates_new_loop(self): - """Should create a new event loop and set it as current.""" - # Get initial loop (if any) - try: - initial_loop = asyncio.get_event_loop() - initial_loop.close() - except RuntimeError: - pass - - # Create new loop - loop = _create_event_loop() - - assert loop is not None - assert not loop.is_closed() - assert asyncio.get_event_loop() is loop - - # Cleanup - loop.close() - - def test_replaces_closed_loop(self): - """Should replace a closed loop.""" - old_loop = asyncio.new_event_loop() - asyncio.set_event_loop(old_loop) - old_loop.close() - - new_loop = _create_event_loop() - - assert new_loop is not old_loop - assert not new_loop.is_closed() - assert asyncio.get_event_loop() is new_loop - - # Cleanup - new_loop.close() - - -class TestGetEventLoop: - """Tests for _get_event_loop helper.""" - - def test_returns_existing_open_loop(self): - """Should return existing event loop if it's open.""" - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - - result = _get_event_loop() - - assert result is loop - assert not result.is_closed() - - # Cleanup - loop.close() - - def test_creates_new_loop_when_closed(self): - """Should create new event loop if current one is closed.""" - old_loop = asyncio.new_event_loop() - asyncio.set_event_loop(old_loop) - old_loop.close() - - result = _get_event_loop() - - assert result is not old_loop - assert not result.is_closed() - - # Cleanup - result.close() - - def test_handles_no_event_loop(self): - """Should handle RuntimeError when no event loop exists (edge case).""" - # This test simulates what happens in a worker thread - # In practice, get_event_loop() returns a closed loop, not RuntimeError - # But we handle the RuntimeError case defensively - loop = _get_event_loop() - assert loop is not None - assert not loop.is_closed() - - # Cleanup - loop.close() - - -class TestMultipleStreamingCalls: - """Tests for the main bug fix: multiple _run_streaming calls.""" - - def test_sequential_streaming_calls(self): - """Multiple sequential calls should work without 'Event loop is closed' error.""" - from EvoScientist.stream.display import _run_streaming - - # Mock agent that returns simple events - mock_agent = Mock() - - async def mock_stream(_request): - """Mock event stream.""" - yield {"type": "text", "content": "test response"} - yield {"type": "done", "response": "test response"} - - # Clean up any existing event loop to start fresh - try: - existing_loop = asyncio.get_event_loop() - if not existing_loop.is_closed(): - existing_loop.close() - except RuntimeError: - pass - - gateway = FakeGraphGateway(stream=mock_stream) - - # Patch Live to avoid terminal output during tests - with patch("EvoScientist.stream.display.Live"): - # First call - _run_streaming( - agent=mock_agent, - message="test message 1", - thread_id="thread1", - show_thinking=False, - interactive=True, - gateway=gateway, - ) - - # Second call - this would fail with "Event loop is closed" before the fix - _run_streaming( - agent=mock_agent, - message="test message 2", - thread_id="thread1", - show_thinking=False, - interactive=True, - gateway=gateway, - ) - - # Third call for good measure - _run_streaming( - agent=mock_agent, - message="test message 3", - thread_id="thread1", - show_thinking=False, - interactive=True, - gateway=gateway, - ) - - def test_loop_reused_across_calls(self): - """Event loop should be reused across multiple calls.""" - # Create a fresh loop - loop = _create_event_loop() - - # Simulate multiple calls - for _ in range(3): - current_loop = _get_event_loop() - assert not current_loop.is_closed() - - # Run a simple coroutine - async def dummy(): - return "ok" - - result = current_loop.run_until_complete(dummy()) - assert result == "ok" - - # Loop should still be open - assert not loop.is_closed() - - # Cleanup - loop.close() - - def test_closed_loop_recovery(self): - """If loop gets closed, next call should create a new one.""" - # Create and close a loop - loop1 = _create_event_loop() - loop1.close() - - # Next call should detect closed loop and create new one - loop2 = _get_event_loop() - - assert loop2 is not loop1 - assert not loop2.is_closed() - - # Should be able to use the new loop - async def dummy(): - return "success" - - result = loop2.run_until_complete(dummy()) - assert result == "success" - - # Cleanup - loop2.close() - - def test_recursive_streaming_does_not_resend_same_thinking(self): - """Resumed runs should not replay the original thinking to channels.""" - from EvoScientist.stream.display import _run_streaming - - mock_agent = Mock() - thinking = "Initial plan. " * 20 - stream_calls = 0 - - async def mock_stream(_request): - nonlocal stream_calls - stream_calls += 1 - if stream_calls == 1: - yield {"type": "thinking", "content": thinking} - yield { - "type": "ask_user", - "interrupt_id": "ask-1", - "tool_call_id": "tc-1", - "questions": [{"question": "Continue?"}], - } - return - - yield {"type": "text", "content": "final answer"} - yield {"type": "done", "response": "final answer"} - - sent_thinking: list[str] = [] - - with patch("EvoScientist.stream.display.Live"): + with ( + AsyncRuntime(thread_name="test-stream-runtime") as runtime, + patch("EvoScientist.stream.display.Live"), + ): + for index in range(3): result = _run_streaming( - agent=mock_agent, - message="test message", + agent=Mock(), + message=f"message {index}", thread_id="thread1", show_thinking=False, interactive=True, - on_thinking=sent_thinking.append, - ask_user_prompt_fn=lambda _data: { - "answers": ["yes"], - "status": "answered", - }, - gateway=FakeGraphGateway(stream=mock_stream), + gateway=gateway, + runtime=runtime, ) + assert result == "test response" - assert result == "final answer" - assert sent_thinking == [thinking.rstrip()] - - def test_recursive_streaming_sends_new_thinking_after_resume(self): - """Genuinely new thinking in resumed rounds should be relayed.""" - from EvoScientist.stream.display import _run_streaming - - mock_agent = Mock() - thinking_r1 = "Initial plan. " * 20 - thinking_r2 = "Revised plan. " * 20 - stream_calls = 0 - - async def mock_stream(_request): - nonlocal stream_calls - stream_calls += 1 - if stream_calls == 1: - yield {"type": "thinking", "content": thinking_r1} - yield { - "type": "ask_user", - "interrupt_id": "ask-1", - "tool_call_id": "tc-1", - "questions": [{"question": "Continue?"}], - } - return - - yield {"type": "thinking", "content": thinking_r2} - yield {"type": "text", "content": "final answer"} - yield {"type": "done", "response": "final answer"} - - sent_thinking: list[str] = [] - - with patch("EvoScientist.stream.display.Live"): - result = _run_streaming( - agent=mock_agent, - message="test message", - thread_id="thread1", - show_thinking=False, - interactive=True, - on_thinking=sent_thinking.append, - ask_user_prompt_fn=lambda _data: { - "answers": ["yes"], - "status": "answered", - }, - gateway=FakeGraphGateway(stream=mock_stream), - ) - - assert result == "final answer" - assert sent_thinking == [thinking_r1.rstrip(), thinking_r2.rstrip()] + assert len(loops) == 3 + assert loops[0] is loops[1] is loops[2] -class TestEventLoopThreadSafety: - """Tests for thread safety edge cases.""" +def test_direct_streaming_call_scopes_and_closes_runtime(): + execution: dict[str, object] = {} - def test_main_thread_normal_case(self): - """Normal case in main thread should work.""" - loop = _get_event_loop() - assert loop is not None - assert not loop.is_closed() + async def stream(_request): + execution["thread"] = threading.current_thread().name + execution["loop"] = asyncio.get_running_loop() + yield {"type": "done", "response": "ok"} - # Cleanup - loop.close() + with patch("EvoScientist.stream.display.Live"): + result = _run_streaming( + agent=Mock(), + message="message", + thread_id="thread1", + show_thinking=False, + interactive=True, + gateway=FakeGraphGateway(stream=stream), + ) + + assert result == "ok" + assert execution["thread"] == "evosci-stream-runtime" + assert isinstance(execution["loop"], asyncio.AbstractEventLoop) + assert not any( + thread.name == "evosci-stream-runtime" and thread.is_alive() + for thread in threading.enumerate() + ) + + +async def test_async_caller_must_offload_synchronous_renderer(): + gateway = FakeGraphGateway(stream=_text_stream([])) + + with ( + AsyncRuntime(thread_name="test-stream-runtime") as runtime, + patch("EvoScientist.stream.display.Live"), + pytest.raises(AsyncRuntimeError, match="running event loop"), + ): + _run_streaming( + agent=Mock(), + message="message", + thread_id="thread1", + show_thinking=False, + interactive=True, + gateway=gateway, + runtime=runtime, + ) + + +@pytest.mark.parametrize( + ("second_thinking", "expected_count"), + [(None, 1), ("Revised plan. " * 20, 2)], +) +def test_recursive_streaming_reuses_runtime_and_deduplicates_thinking( + second_thinking, expected_count +): + initial_thinking = "Initial plan. " * 20 + stream_calls = 0 + loops: list[asyncio.AbstractEventLoop] = [] + + async def stream(_request): + nonlocal stream_calls + loops.append(asyncio.get_running_loop()) + stream_calls += 1 + if stream_calls == 1: + yield {"type": "thinking", "content": initial_thinking} + yield { + "type": "ask_user", + "interrupt_id": "ask-1", + "tool_call_id": "tc-1", + "questions": [{"question": "Continue?"}], + } + return + if second_thinking is not None: + yield {"type": "thinking", "content": second_thinking} + else: + yield {"type": "thinking", "content": initial_thinking} + yield {"type": "text", "content": "final answer"} + yield {"type": "done", "response": "final answer"} + + sent_thinking: list[str] = [] + with ( + AsyncRuntime(thread_name="test-stream-runtime") as runtime, + patch("EvoScientist.stream.display.Live"), + ): + result = _run_streaming( + agent=Mock(), + message="test message", + thread_id="thread1", + show_thinking=False, + interactive=True, + on_thinking=sent_thinking.append, + ask_user_prompt_fn=lambda _data: { + "answers": ["yes"], + "status": "answered", + }, + gateway=FakeGraphGateway(stream=stream), + runtime=runtime, + ) + + assert result == "final answer" + assert len(sent_thinking) == expected_count + assert sent_thinking[0] == initial_thinking.rstrip() + if second_thinking is not None: + assert sent_thinking[1] == second_thinking.rstrip() + assert loops[0] is loops[1] diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 5e1bb62..499c505 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -1,6 +1,8 @@ """Tests for EvoScientist.mcp module.""" +import asyncio import textwrap +import threading from pathlib import Path from types import SimpleNamespace @@ -16,10 +18,12 @@ from EvoScientist.mcp.client import ( add_mcp_server, edit_mcp_server, load_mcp_config, + load_mcp_tools, parse_mcp_add_args, parse_mcp_edit_args, remove_mcp_server, ) +from EvoScientist.runtime import AsyncRuntime, AsyncRuntimeError # ---- _interpolate_env ---- @@ -283,6 +287,58 @@ def _make_tool(name: str): return SimpleNamespace(name=name) +class TestOwnedRuntimeLoading: + def test_reuses_caller_runtime_for_discovery(self, monkeypatch): + executions: list[tuple[str, asyncio.AbstractEventLoop]] = [] + tool = _make_tool("search") + + async def fake_load(_config, *, on_progress=None): + executions.append( + (threading.current_thread().name, asyncio.get_running_loop()) + ) + return {"server": [tool]} + + monkeypatch.setattr("EvoScientist.mcp.client._load_tools", fake_load) + config = {"server": {"transport": "http", "url": "http://example.test"}} + + with AsyncRuntime(thread_name="test-mcp-runtime") as runtime: + first = load_mcp_tools(config, runtime=runtime) + second = load_mcp_tools(config, runtime=runtime) + + assert first == {"main": [tool]} + assert second == {"main": [tool]} + assert [thread for thread, _loop in executions] == [ + "test-mcp-runtime", + "test-mcp-runtime", + ] + assert executions[0][1] is executions[1][1] + + def test_direct_caller_gets_a_scoped_runtime(self, monkeypatch): + execution: dict[str, object] = {} + + async def fake_load(_config, *, on_progress=None): + execution["thread"] = threading.current_thread().name + execution["loop"] = asyncio.get_running_loop() + return {"server": []} + + monkeypatch.setattr("EvoScientist.mcp.client._load_tools", fake_load) + config = {"server": {"transport": "http", "url": "http://example.test"}} + + assert load_mcp_tools(config) == {"main": []} + assert execution["thread"] == "evosci-mcp-runtime" + assert isinstance(execution["loop"], asyncio.AbstractEventLoop) + assert not any( + thread.name == "evosci-mcp-runtime" and thread.is_alive() + for thread in threading.enumerate() + ) + + async def test_direct_caller_does_not_hide_running_loop_violation(self): + config = {"server": {"transport": "http", "url": "http://example.test"}} + + with pytest.raises(AsyncRuntimeError, match=r"await aload_mcp_tools\(config"): + load_mcp_tools(config) + + class TestFilterTools: def test_none_allowlist_passes_all(self): tools = [_make_tool("a"), _make_tool("b"), _make_tool("c")] diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index 66a1f8b..08f0ddb 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -16,8 +16,10 @@ from langchain_core.messages import AIMessage, HumanMessage from EvoScientist.middleware.events import NoOpSink from EvoScientist.middleware.model_fallback import ( _guard_and_fallback, + _guard_and_fallback_sync, _is_non_fallbackable, _try_fallbacks, + _try_fallbacks_sync, add_fallback, clear_fallbacks, ) @@ -410,6 +412,49 @@ class TestGuardAndFallback: invoke.assert_awaited_once() +class TestSynchronousFallback: + """The sync middleware path must not create or nest an event loop.""" + + def test_first_fallback_succeeds_without_async_bridge(self): + add_fallback("fb-model", "fb-provider") + req = _fake_request() + invoke = MagicMock(return_value=AI_RESPONSE) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + result = _try_fallbacks_sync(req, invoke, Exception("503 boom"), _SINK) + + assert result is AI_RESPONSE + invoke.assert_called_once() + + def test_guard_rejects_non_fallbackable_error_before_handler(self): + add_fallback("fb", "prov") + req = _fake_request() + invoke = MagicMock() + + with pytest.raises(ContextOverflowError): + _guard_and_fallback_sync( + ContextOverflowError("overflow"), req, invoke, _SINK + ) + + invoke.assert_not_called() + + def test_middleware_sync_entrypoint_uses_native_traversal(self): + from EvoScientist.middleware.model_fallback import ModelFallbackMiddleware + + add_fallback("fb", "prov") + req = _fake_request() + response = AI_RESPONSE + handler = MagicMock(side_effect=[Exception("503 primary"), response]) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + result = ModelFallbackMiddleware().wrap_model_call(req, handler) + + assert result is response + assert handler.call_count == 2 + + # ═════════════════════════════════════════════════════════════════ # 4. UI emit callback # ═════════════════════════════════════════════════════════════════ diff --git a/tests/test_onboard_async_runtime.py b/tests/test_onboard_async_runtime.py new file mode 100644 index 0000000..d52e924 --- /dev/null +++ b/tests/test_onboard_async_runtime.py @@ -0,0 +1,47 @@ +"""Owned-runtime coverage for bounded onboarding async work.""" + +import asyncio +import threading + +from EvoScientist.config import EvoScientistConfig +from EvoScientist.config.onboard.channels import _probe_channel +from EvoScientist.runtime import AsyncRuntime + + +def test_channel_probes_reuse_the_provided_runtime(monkeypatch): + executions: list[tuple[str, asyncio.AbstractEventLoop]] = [] + + async def validate_telegram(_token, _proxy): + executions.append((threading.current_thread().name, asyncio.get_running_loop())) + return True, "telegram ok" + + async def validate_discord(_token, _proxy): + executions.append((threading.current_thread().name, asyncio.get_running_loop())) + return True, "discord ok" + + monkeypatch.setattr( + "EvoScientist.channels.telegram.probe.validate_telegram_token", + validate_telegram, + ) + monkeypatch.setattr( + "EvoScientist.channels.discord.probe.validate_discord_token", + validate_discord, + ) + monkeypatch.setattr( + "EvoScientist.config.onboard.channels.console.print", lambda *_a, **_k: None + ) + + config = EvoScientistConfig() + updates = { + "telegram_bot_token": "telegram-token", + "discord_bot_token": "discord-token", + } + with AsyncRuntime(thread_name="test-onboard-runtime") as runtime: + _probe_channel("telegram", config, updates, runtime=runtime) + _probe_channel("discord", config, updates, runtime=runtime) + + assert [thread for thread, _loop in executions] == [ + "test-onboard-runtime", + "test-onboard-runtime", + ] + assert executions[0][1] is executions[1][1] diff --git a/tests/test_runtime.py b/tests/test_runtime.py new file mode 100644 index 0000000..4d5ed4d --- /dev/null +++ b/tests/test_runtime.py @@ -0,0 +1,622 @@ +"""Focused tests for the application-scoped owned async runtime.""" + +from __future__ import annotations + +import asyncio +import concurrent.futures +import contextvars +import logging +import threading +import time +from typing import Any + +import pytest + +from EvoScientist.runtime import ( + AsyncRuntime, + AsyncRuntimeClosedError, + AsyncRuntimeError, + RuntimeHandle, +) + + +def _wait_until(predicate, *, timeout: float = 5.0) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if predicate(): + return + time.sleep(0.005) + assert predicate(), "condition was not met before timeout" + + +@pytest.fixture +def runtime(): + instance = AsyncRuntime(cancellation_timeout=1.0) + try: + yield instance + finally: + instance.close(timeout=5.0) + + +def test_constructor_validates_timeouts(): + with pytest.raises(ValueError, match="start_timeout"): + AsyncRuntime(start_timeout=0) + with pytest.raises(ValueError, match="cancellation_timeout"): + AsyncRuntime(cancellation_timeout=-1) + + +def test_start_is_idempotent_and_waits_until_loop_runs(runtime): + runtime.start() + first_thread = runtime._thread + runtime.start() + + assert runtime._thread is first_thread + assert first_thread is not None + assert first_thread.name == "evosci-async-runtime" + assert first_thread.daemon + assert first_thread.is_alive() + assert runtime.is_running + + +def test_context_manager_owns_start_and_close(): + with AsyncRuntime(thread_name="context-runtime") as runtime: + thread = runtime._thread + assert runtime.is_running + assert runtime.run_sync(lambda: asyncio.sleep(0, result=3)) == 3 + + assert thread is not None + assert not thread.is_alive() + assert not runtime.is_running + + +def test_start_applies_windows_policy_before_creating_loop(monkeypatch): + import EvoScientist.runtime as runtime_module + + calls: list[str] = [] + real_new_event_loop = asyncio.new_event_loop + + def policy_spy() -> bool: + calls.append("policy") + return False + + def loop_spy() -> asyncio.AbstractEventLoop: + calls.append("loop") + return real_new_event_loop() + + monkeypatch.setattr(runtime_module, "ensure_proactor_event_loop_policy", policy_spy) + monkeypatch.setattr(runtime_module.asyncio, "new_event_loop", loop_spy) + + runtime = AsyncRuntime() + try: + runtime.start() + assert calls == ["policy", "loop"] + finally: + runtime.close() + + +def test_close_after_startup_timeout_does_not_leave_runtime_thread(monkeypatch): + import EvoScientist.runtime as runtime_module + + loop_creation_started = threading.Event() + release_loop_creation = threading.Event() + real_new_event_loop = asyncio.new_event_loop + + def delayed_new_event_loop() -> asyncio.AbstractEventLoop: + loop_creation_started.set() + assert release_loop_creation.wait(5) + return real_new_event_loop() + + monkeypatch.setattr( + runtime_module.asyncio, "new_event_loop", delayed_new_event_loop + ) + runtime = AsyncRuntime(start_timeout=0.01) + + with pytest.raises(TimeoutError, match="did not start"): + runtime.start() + assert loop_creation_started.is_set() + + release_loop_creation.set() + runtime.close(timeout=5) + + assert runtime._thread is not None + assert not runtime._thread.is_alive() + + +def test_runtime_is_instance_scoped_not_a_module_singleton(): + import EvoScientist.runtime as runtime_module + + assert not hasattr(runtime_module, "runtime") + assert AsyncRuntime() is not AsyncRuntime() + + +def test_submit_invokes_factory_on_owned_thread(runtime): + factory_thread: list[str] = [] + + async def identify() -> str: + return threading.current_thread().name + + def factory(): + factory_thread.append(threading.current_thread().name) + return identify() + + handle = runtime.submit(factory) + + assert isinstance(handle, RuntimeHandle) + assert handle.result(5) == "evosci-async-runtime" + assert handle.wait_settled(5) + assert factory_thread == ["evosci-async-runtime"] + + +def test_submissions_share_one_owned_loop(runtime): + async def current_loop() -> asyncio.AbstractEventLoop: + return asyncio.get_running_loop() + + first = runtime.run_sync(current_loop) + second = runtime.run_sync(current_loop) + + assert first is second + + +def test_submit_propagates_context_variables(runtime): + request_id: contextvars.ContextVar[str] = contextvars.ContextVar("request_id") + token = request_id.set("request-42") + try: + handle = runtime.submit( + lambda: asyncio.sleep(0, result=request_id.get("missing")) + ) + request_id.set("changed-after-submit") + assert handle.result(5) == "request-42" + finally: + request_id.reset(token) + + +def test_submit_is_safe_from_worker_threads(runtime): + results: list[int] = [] + + def worker(value: int) -> None: + result = runtime.submit(lambda: asyncio.sleep(0, result=value * 2)).result(5) + results.append(result) + + threads = [threading.Thread(target=worker, args=(value,)) for value in range(4)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(5) + + assert all(not thread.is_alive() for thread in threads) + assert sorted(results) == [0, 2, 4, 6] + + +def test_handle_distinguishes_public_cancellation_from_task_settlement(runtime): + started = threading.Event() + release_cleanup = threading.Event() + + async def blocked() -> None: + started.set() + try: + await asyncio.Event().wait() + finally: + while not release_cleanup.is_set(): + await asyncio.sleep(0.005) + + handle = runtime.submit(blocked) + assert started.wait(5) + + assert handle.cancel() + assert handle.done() + assert handle.cancelled() + assert not handle.settled + + release_cleanup.set() + assert handle.wait_settled(5) + + +async def test_cancelling_async_waiter_does_not_cancel_settlement_signal(runtime): + started = threading.Event() + release = threading.Event() + + async def blocked() -> None: + started.set() + while not release.is_set(): + await asyncio.sleep(0.005) + + handle = runtime.submit(blocked) + assert await asyncio.to_thread(started.wait, 5) + + waiter = asyncio.create_task(handle.wait_settled_async()) + await asyncio.sleep(0) + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await waiter + + assert not handle.settled + assert handle.wait_settled(0) is False + + release.set() + assert handle.result(5) is None + assert handle.wait_settled(5) + + +def test_cancelling_before_task_creation_never_invokes_factory(runtime): + loop_blocked = threading.Event() + release_loop = threading.Event() + factory_called = False + + def block_loop() -> None: + loop_blocked.set() + assert release_loop.wait(5) + + runtime.start() + assert runtime._loop is not None + runtime._loop.call_soon_threadsafe(block_loop) + assert loop_blocked.wait(5) + + async def operation() -> None: + nonlocal factory_called + factory_called = True + + handle = runtime.submit(operation) + assert handle.cancel() + release_loop.set() + + assert handle.wait_settled(5) + assert not factory_called + + +def test_run_sync_returns_result_and_propagates_exception(runtime): + assert runtime.run_sync(lambda: asyncio.sleep(0, result=42)) == 42 + + async def fail() -> None: + raise ValueError("broken") + + with pytest.raises(ValueError, match="broken"): + runtime.run_sync(fail) + + +async def test_run_sync_rejects_every_running_event_loop(runtime): + called = False + + async def operation() -> None: + nonlocal called + called = True + + with pytest.raises(AsyncRuntimeError, match="cannot block a running event loop"): + runtime.run_sync(operation) + + assert not called + + +def test_run_sync_timeout_cancels_and_waits_for_cleanup(runtime): + cleanup_finished = threading.Event() + + async def blocked() -> None: + try: + await asyncio.Event().wait() + finally: + await asyncio.sleep(0.02) + cleanup_finished.set() + + with pytest.raises(concurrent.futures.TimeoutError): + runtime.run_sync(blocked, timeout=0.01) + + assert cleanup_finished.is_set() + + +def test_run_sync_interrupt_cancels_and_waits_for_cleanup(runtime, monkeypatch): + started = threading.Event() + cleanup_finished = threading.Event() + + async def blocked() -> None: + started.set() + try: + await asyncio.Event().wait() + finally: + await asyncio.sleep(0.02) + cleanup_finished.set() + + real_submit = runtime.submit + + class InterruptingResult: + def __init__(self, handle: RuntimeHandle[Any]) -> None: + self._handle = handle + + def result(self, timeout: float | None = None) -> Any: + assert started.wait(5) + raise KeyboardInterrupt + + def done(self) -> bool: + return self._handle.done() + + def cancelled(self) -> bool: + return self._handle.cancelled() + + def cancel(self) -> bool: + return self._handle.cancel() + + def wait_settled(self, timeout: float | None = None) -> bool: + return self._handle.wait_settled(timeout) + + monkeypatch.setattr( + runtime, + "submit", + lambda factory: InterruptingResult(real_submit(factory)), + ) + + with pytest.raises(KeyboardInterrupt): + runtime.run_sync(blocked) + + assert cleanup_finished.is_set() + + +async def test_run_async_bridges_without_blocking_callers_loop(runtime): + caller_loop = asyncio.get_running_loop() + runtime_loop, thread_name = await runtime.run_async(lambda: _loop_and_thread()) + + assert runtime_loop is not caller_loop + assert thread_name == "evosci-async-runtime" + + +async def _loop_and_thread() -> tuple[asyncio.AbstractEventLoop, str]: + return asyncio.get_running_loop(), threading.current_thread().name + + +async def test_run_async_propagates_caller_context(runtime): + request_id: contextvars.ContextVar[str] = contextvars.ContextVar("async_request_id") + token = request_id.set("from-ui-loop") + try: + assert ( + await runtime.run_async( + lambda: asyncio.sleep(0, result=request_id.get("missing")) + ) + == "from-ui-loop" + ) + finally: + request_id.reset(token) + + +async def test_run_async_cancellation_waits_for_runtime_cleanup(runtime): + started = threading.Event() + cleanup_finished = threading.Event() + + async def blocked() -> None: + started.set() + try: + await asyncio.Event().wait() + finally: + await asyncio.sleep(0.02) + cleanup_finished.set() + + caller = asyncio.create_task(runtime.run_async(blocked)) + assert await asyncio.to_thread(started.wait, 5) + caller.cancel() + + with pytest.raises(asyncio.CancelledError): + await caller + + assert cleanup_finished.is_set() + + +def test_run_async_rejects_calls_from_owned_loop(runtime): + factory_called = False + + async def operation() -> None: + nonlocal factory_called + factory_called = True + + async def invoke_from_runtime() -> None: + with pytest.raises(AsyncRuntimeError, match="owned loop"): + await runtime.run_async(operation) + + runtime.run_sync(invoke_from_runtime) + assert not factory_called + + +def test_spawn_runs_durable_work_and_returns_named_handle(runtime): + release = threading.Event() + finished = threading.Event() + + async def background() -> str: + while not release.is_set(): + await asyncio.sleep(0.005) + finished.set() + return "complete" + + handle = runtime.spawn(background, name="durable-work") + assert handle.name == "durable-work" + assert not handle.done() + + release.set() + assert handle.result(5) == "complete" + assert handle.wait_settled(5) + assert finished.is_set() + + +def test_spawn_logs_unhandled_failures(runtime, caplog): + async def fail() -> None: + raise RuntimeError("background exploded") + + with caplog.at_level(logging.ERROR, logger="EvoScientist.runtime"): + handle = runtime.spawn(fail, name="failing-background") + with pytest.raises(RuntimeError, match="background exploded"): + handle.result(5) + assert handle.wait_settled(5) + _wait_until( + lambda: any( + "failing-background" in record.getMessage() for record in caplog.records + ) + ) + + record = next( + record + for record in caplog.records + if "failing-background" in record.getMessage() + ) + assert isinstance(record.exc_info[1], RuntimeError) + + +def test_spawn_cancellation_is_not_logged(runtime, caplog): + started = threading.Event() + + async def blocked() -> None: + started.set() + await asyncio.Event().wait() + + with caplog.at_level(logging.ERROR, logger="EvoScientist.runtime"): + handle = runtime.spawn(blocked, name="cancelled-background") + assert started.wait(5) + handle.cancel() + assert handle.wait_settled(5) + + assert not caplog.records + + +def test_close_cancels_and_settles_pending_work(runtime): + started = threading.Event() + cleanup_finished = threading.Event() + + async def blocked() -> None: + started.set() + try: + await asyncio.Queue().get() + finally: + cleanup_finished.set() + + handle = runtime.spawn(blocked, name="pending") + assert started.wait(5) + thread = runtime._thread + + runtime.close(timeout=5) + + assert handle.cancelled() + assert handle.wait_settled(5) + assert cleanup_finished.is_set() + assert thread is not None + assert not thread.is_alive() + assert not runtime.is_running + + +def test_close_waits_for_default_executor_work(): + runtime = AsyncRuntime() + started = threading.Event() + finished = threading.Event() + + def blocking_job() -> None: + started.set() + time.sleep(0.1) + finished.set() + + runtime.spawn( + lambda: asyncio.to_thread(blocking_job), + name="executor-job", + ) + assert started.wait(2) + + runtime.close(timeout=2) + + assert finished.is_set() + assert runtime._thread is not None + assert not runtime._thread.is_alive() + + +def test_close_timeout_never_reports_success_while_executor_is_active(): + runtime = AsyncRuntime() + started = threading.Event() + release = threading.Event() + finished = threading.Event() + + def blocking_job() -> None: + started.set() + release.wait() + finished.set() + + runtime.spawn( + lambda: asyncio.to_thread(blocking_job), + name="blocked-executor-job", + ) + assert started.wait(2) + + with pytest.raises(TimeoutError, match="did not settle"): + runtime.close(timeout=0.05) + + assert not finished.is_set() + assert runtime._thread is not None + assert runtime._thread.is_alive() + + release.set() + runtime.close(timeout=2) + + assert finished.is_set() + assert not runtime._thread.is_alive() + + +def test_close_is_idempotent_and_close_before_start_seals_runtime(): + runtime = AsyncRuntime() + runtime.close() + runtime.close() + + with pytest.raises(AsyncRuntimeClosedError, match="closed"): + runtime.start() + with pytest.raises(AsyncRuntimeClosedError, match="closed"): + runtime.submit(lambda: asyncio.sleep(0)) + + +def test_close_rejects_calls_from_runtime_thread(runtime): + async def close_from_runtime() -> str: + with pytest.raises(AsyncRuntimeError, match="application owner"): + runtime.close() + return threading.current_thread().name + + assert runtime.run_sync(close_from_runtime) == "evosci-async-runtime" + assert runtime.run_sync(lambda: asyncio.sleep(0, result="still alive")) == ( + "still alive" + ) + + +def test_submit_enqueue_is_atomic_with_close(): + submit_holds_lock = threading.Event() + release_submit = threading.Event() + + class PausedRuntime(AsyncRuntime): + def _enqueue_locked(self, loop, callback, context): + submit_holds_lock.set() + assert release_submit.wait(5) + super()._enqueue_locked(loop, callback, context) + + runtime = PausedRuntime() + submitted: dict[str, RuntimeHandle[str]] = {} + + def submit() -> None: + submitted["handle"] = runtime.submit( + lambda: asyncio.sleep(0, result="accepted") + ) + + submit_thread = threading.Thread(target=submit) + close_thread = threading.Thread(target=lambda: runtime.close(timeout=5)) + submit_thread.start() + assert submit_holds_lock.wait(5) + close_thread.start() + + release_submit.set() + submit_thread.join(5) + close_thread.join(5) + + assert not submit_thread.is_alive() + assert not close_thread.is_alive() + handle = submitted["handle"] + assert handle.done() + assert handle.wait_settled(5) + + +def test_submission_after_started_runtime_is_closed_never_invokes_factory(runtime): + runtime.start() + runtime.close() + called = False + + async def operation() -> None: + nonlocal called + called = True + + with pytest.raises(AsyncRuntimeClosedError, match="closed"): + runtime.submit(operation) + + assert not called diff --git a/tests/test_serve_agent_holder.py b/tests/test_serve_agent_holder.py index 37d71c7..b880f11 100644 --- a/tests/test_serve_agent_holder.py +++ b/tests/test_serve_agent_holder.py @@ -8,6 +8,8 @@ captured at startup. from __future__ import annotations +import asyncio +import time from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -27,6 +29,7 @@ from EvoScientist.cli.commands import ( from EvoScientist.commands.base import ChannelRuntime from EvoScientist.config import EvoScientistConfig from EvoScientist.gateway import RuntimeGateways, ThreadStore +from EvoScientist.runtime import AsyncRuntime from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -59,6 +62,7 @@ def _runtime_state( config: EvoScientistConfig | None = None, thread_store: ThreadStore | None = None, runtime_gateways: RuntimeGateways | None = None, + async_runtime: AsyncRuntime | None = None, ) -> ServeRuntimeState: store = thread_store or _thread_store() return ServeRuntimeState( @@ -67,9 +71,22 @@ def _runtime_state( workspace_dir=workspace_dir, config=config, runtime_gateways=runtime_gateways or _runtime_gateways(store), + async_runtime=async_runtime or MagicMock(spec=AsyncRuntime), ) +def test_serve_runtime_state_requires_owned_runtime(): + """Message processing cannot be constructed without its runtime owner.""" + with pytest.raises(TypeError, match="async_runtime"): + ServeRuntimeState( + agent=_agent(), + thread_id="tid", + workspace_dir=None, + config=None, + runtime_gateways=_runtime_gateways(), + ) + + async def test_hook_updates_runtime_state_on_agent_swap(): """``/model`` mutates ``ctx.agent`` to a new handle — the hook must push that handle into the shared runtime state so the outer poll loop sees @@ -203,7 +220,11 @@ async def test_hook_updates_workspace_dir_on_resume(): await hook(ctx, old_agent, cmd) sync_server.assert_awaited_once_with(cfg, workspace_dir="/restored-ws") - load_agent.assert_called_once_with(workspace_dir="/restored-ws", config=cfg) + load_agent.assert_called_once_with( + workspace_dir="/restored-ws", + config=cfg, + runtime=state.async_runtime, + ) assert state.workspace_dir == "/restored-ws" assert state.agent is reloaded_agent @@ -366,7 +387,11 @@ async def test_serve_resume_callback_syncs_reloads_and_adopts_workspace(): await cb("new-tid", "/new-ws") sync_server.assert_awaited_once_with(cfg, workspace_dir="/new-ws") - load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg) + load_agent.assert_called_once_with( + workspace_dir="/new-ws", + config=cfg, + runtime=state.async_runtime, + ) assert call_order == ["load", "sync"] assert state.thread_id == "new-tid" assert state.workspace_dir == "/new-ws" @@ -443,7 +468,11 @@ async def test_serve_resume_callback_preserves_state_when_sync_fails(): ): await cb("new-tid", "/new-ws") - load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg) + load_agent.assert_called_once_with( + workspace_dir="/new-ws", + config=cfg, + runtime=state.async_runtime, + ) set_active.assert_called_once_with("/old-ws") assert state.agent is old_agent assert state.resume_warning_thread_id is None @@ -480,7 +509,11 @@ async def test_serve_resume_callback_load_failure_does_not_sync_or_adopt(): ): await cb("new-tid", "/new-ws") - load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg) + load_agent.assert_called_once_with( + workspace_dir="/new-ws", + config=cfg, + runtime=state.async_runtime, + ) set_active.assert_called_once_with("/old-ws") sync_server.assert_not_awaited() assert state.resume_warning_thread_id is None @@ -536,6 +569,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback(): ) with ( + AsyncRuntime(thread_name="test-serve-runtime") as async_runtime, patch( "EvoScientist.cli.commands.dispatch_channel_slash_command", new=AsyncMock(side_effect=RuntimeError("slash broke")), @@ -543,6 +577,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback(): patch("EvoScientist.cli.commands._set_channel_response") as mock_set_resp, patch("EvoScientist.cli.tui_runtime.run_streaming") as mock_run_streaming, ): + state.async_runtime = async_runtime _register_channel_request(msg) _serve_process_message( msg, @@ -588,6 +623,7 @@ def test_serve_process_message_uses_runtime_workspace_from_state(): return {} with ( + AsyncRuntime(thread_name="test-serve-runtime") as async_runtime, patch( "EvoScientist.cli.commands.dispatch_channel_slash_command", new=AsyncMock(side_effect=_fake_dispatch), @@ -598,6 +634,7 @@ def test_serve_process_message_uses_runtime_workspace_from_state(): ), patch("EvoScientist.cli.tui_runtime.run_streaming", return_value="ok"), ): + state.async_runtime = async_runtime _register_channel_request(msg) _serve_process_message( msg, @@ -609,3 +646,78 @@ def test_serve_process_message_uses_runtime_workspace_from_state(): assert captured["slash_workspace"] == "/restored-workspace" assert captured["meta_workspace"] == "/restored-workspace" + + +def test_serve_channel_send_does_not_block_owned_runtime_loop(): + """Channel I/O is scheduled on the bus loop and settled before the reply.""" + from EvoScientist.cli import channel as channel_mod + + events: list[str] = [] + callback_elapsed: list[float] = [] + + class _ChannelRef: + send_thinking = True + + async def send_thinking_message(self, **_kwargs): + events.append("send-started") + await asyncio.sleep(0.05) + events.append("send-finished") + + msg = ChannelMessage( + msg_id="msg-nonblocking-send", + content="hello", + sender="channel-user", + channel_type="telegram", + metadata={}, + channel_ref=_ChannelRef(), + bus_ref=None, + chat_id="channel-user", + message_id="ts-send", + ) + state = _runtime_state(agent=_agent(), thread_id="tid") + + def _fake_run_streaming(**kwargs): + async def _invoke_callback() -> None: + started = time.monotonic() + kwargs["on_thinking"]("x" * 250) + callback_elapsed.append(time.monotonic() - started) + events.append("callback-returned") + + kwargs["runtime"].run_sync(_invoke_callback) + return "ok" + + def _capture_response(_msg_id: str, _response: str) -> None: + events.append("response-set") + + with AsyncRuntime(thread_name="test-serve-send-runtime") as runtime: + runtime.submit(lambda: asyncio.sleep(0)).result(timeout=1) + state.async_runtime = runtime + assert runtime._loop is not None + + with ( + patch.object(channel_mod, "_bus_loop", runtime._loop), + patch( + "EvoScientist.cli.commands.dispatch_channel_slash_command", + new=AsyncMock(return_value=False), + ), + patch( + "EvoScientist.cli.tui_runtime.run_streaming", + side_effect=_fake_run_streaming, + ), + patch( + "EvoScientist.cli.commands._set_channel_response", + side_effect=_capture_response, + ), + ): + _register_channel_request(msg) + _serve_process_message( + msg, + runtime_state=state, + model="model", + workspace_dir="/tmp", + show_thinking=True, + ) + + assert callback_elapsed[0] < 0.5 + assert events.index("callback-returned") < events.index("send-started") + assert events.index("send-finished") < events.index("response-set") diff --git a/tests/test_stream_cancel.py b/tests/test_stream_cancel.py index 16e20c1..eda3ba2 100644 --- a/tests/test_stream_cancel.py +++ b/tests/test_stream_cancel.py @@ -2,10 +2,18 @@ from __future__ import annotations +import asyncio +import sys +import threading +import time +from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from EvoScientist.backends import CustomSandboxBackend +from EvoScientist.cancellation import current_cancel_event +from EvoScientist.runtime import AsyncRuntime from EvoScientist.stream import display as display_mod from tests.fakes import FakeGraphGateway @@ -19,6 +27,7 @@ def _clean_cancel_event(): display_mod._stream_cancel_events[display_mod._DEFAULT_STREAM_CANCEL_SCOPE] = ( display_mod._stream_cancel_event ) + display_mod._stream_cancel_handles.clear() yield with display_mod._stream_cancel_lock: display_mod._stream_cancel_event.clear() @@ -26,6 +35,7 @@ def _clean_cancel_event(): display_mod._stream_cancel_events[display_mod._DEFAULT_STREAM_CANCEL_SCOPE] = ( display_mod._stream_cancel_event ) + display_mod._stream_cancel_handles.clear() # --------------------------------------------------------------------------- @@ -64,6 +74,227 @@ def test_consume_breaks_on_cancel_event(): assert "[Stopped.]" in result +def test_cancel_interrupts_stalled_stream_and_closes_it_in_consumer_task(): + """Cancellation must not wait for a stalled gateway to yield again.""" + cancel_scope = "scope:stalled" + stream_started = threading.Event() + stream_closed = threading.Event() + tasks: dict[str, asyncio.Task[object] | None] = {} + result: dict[str, str] = {} + + async def _stalled_stream(_request): + tasks["consumer"] = asyncio.current_task() + stream_started.set() + try: + await asyncio.Event().wait() + if False: + yield {} + finally: + tasks["closer"] = asyncio.current_task() + stream_closed.set() + + def _run() -> None: + result["response"] = display_mod._run_streaming( + agent=MagicMock(), + message="hello", + thread_id="t1", + show_thinking=False, + interactive=True, + cancel_scope=cancel_scope, + gateway=FakeGraphGateway(stream=_stalled_stream), + ) + + worker = threading.Thread(target=_run) + worker.start() + assert stream_started.wait(2) + + display_mod.request_stream_cancel(cancel_scope) + worker.join(2) + + assert not worker.is_alive() + assert stream_closed.is_set() + assert tasks["closer"] is tasks["consumer"] + assert result["response"] == "[Stopped.]" + + +def test_cancel_interrupts_owned_questionary_prompt(): + """A terminal prompt must settle instead of outliving the frontend turn.""" + cancel_scope = "scope:questionary" + started = threading.Event() + closed = threading.Event() + result: dict[str, object] = {} + + class _BlockingQuestion: + async def ask_async(self): + started.set() + try: + await asyncio.Event().wait() + finally: + closed.set() + + def _run(runtime: AsyncRuntime) -> None: + try: + display_mod._run_owned_questionary_prompt( + _BlockingQuestion(), + runtime=runtime, + cancel_scope=cancel_scope, + ) + except display_mod._StreamPromptCancelled: + result["cancelled"] = True + + with AsyncRuntime(thread_name="test-questionary-runtime") as runtime: + worker = threading.Thread(target=_run, args=(runtime,)) + worker.start() + assert started.wait(2) + + display_mod.request_stream_cancel(cancel_scope) + worker.join(2) + + assert not worker.is_alive() + assert closed.is_set() + assert result == {"cancelled": True} + + +def test_cancelled_stream_does_not_repaint_final_live_frame(): + """Late stream cleanup must not overwrite a newer frontend frame.""" + live = MagicMock() + handle = display_mod.RuntimeHandle() + handle.cancel() + + display_mod._update_final_live_frame(live, object(), handle) + + live.update.assert_not_called() + live.refresh.assert_not_called() + + +def test_cancel_unwinds_hitl_prompt_and_renderer(monkeypatch): + """The real Rich HITL branch must release its prompt before returning.""" + cancel_scope = "scope:hitl-questionary" + started = threading.Event() + closed = threading.Event() + result: dict[str, str] = {} + + class _BlockingQuestion: + async def ask_async(self): + started.set() + try: + await asyncio.Event().wait() + finally: + closed.set() + + def ask(self): # pragma: no cover - the owned path must use ask_async + raise AssertionError("blocking questionary.ask() was used") + + monkeypatch.setitem( + sys.modules, + "questionary", + SimpleNamespace(select=lambda *_args, **_kwargs: _BlockingQuestion()), + ) + monkeypatch.setattr( + "EvoScientist.config.settings.load_config", + lambda: SimpleNamespace(auto_approve=False, shell_allow_list=""), + ) + + async def _empty_stream(_request): + if False: + yield {} + + state = display_mod.StreamState() + state.response_text = "Partial answer" + state.pending_interrupt = { + "action_requests": [{"name": "execute", "args": {"command": "echo hi"}}] + } + + def _run(runtime: AsyncRuntime) -> None: + result["response"] = display_mod._run_streaming( + agent=MagicMock(), + message="hello", + thread_id="t1", + show_thinking=False, + interactive=True, + cancel_scope=cancel_scope, + _state=state, + gateway=FakeGraphGateway(stream=_empty_stream), + runtime=runtime, + ) + + with AsyncRuntime(thread_name="test-hitl-runtime") as runtime: + worker = threading.Thread(target=_run, args=(runtime,)) + worker.start() + assert started.wait(2) + + display_mod.request_stream_cancel(cancel_scope) + worker.join(2) + + assert not worker.is_alive() + assert closed.is_set() + assert result["response"] == "Partial answer\n[Stopped.]" + + +def test_cancel_terminates_active_shell_process_tree(tmp_path): + """A cancelled turn must not leave delayed shell side effects running.""" + cancel_scope = "scope:shell" + backend = CustomSandboxBackend(root_dir=str(tmp_path), virtual_mode=True) + started = tmp_path / "started.txt" + forbidden = tmp_path / "forbidden.txt" + result: dict[str, object] = {} + + if sys.platform == "win32": + command = ( + "echo started> started.txt & " + "ping -n 11 127.0.0.1 > nul & " + "echo late> forbidden.txt" + ) + else: + command = "printf started > started.txt; sleep 10; printf late > forbidden.txt" + + async def _events(): + result["response"] = await asyncio.to_thread(backend.execute, command) + if False: + yield {} + + async def _consume() -> None: + async for _ in display_mod.iter_with_stream_cancel(_events(), cancel_scope): + pass + + worker = threading.Thread(target=lambda: asyncio.run(_consume())) + worker.start() + deadline = time.monotonic() + 3 + while not started.exists() and time.monotonic() < deadline: + time.sleep(0.02) + assert started.read_text().strip() == "started" + + display_mod.request_stream_cancel(cancel_scope) + worker.join(3) + + assert not worker.is_alive() + assert result["response"].exit_code == 130 + assert not forbidden.exists() + + +async def test_stream_cancel_binding_can_close_in_different_task_context(): + """No ContextVar token may survive across an async-generator yield.""" + closed = False + + async def _events(): + nonlocal closed + try: + yield {"type": "text", "content": "one"} + await asyncio.Event().wait() + finally: + closed = True + + wrapped = display_mod.iter_with_stream_cancel(_events(), "scope:cross-context") + async for _ in wrapped: + break + + assert current_cancel_event() is None + await asyncio.create_task(wrapped.aclose()) + + assert closed is True + assert current_cancel_event() is None + + # --------------------------------------------------------------------------- # 2. fresh _run_streaming clears stale set event # --------------------------------------------------------------------------- diff --git a/tests/test_tool_selector_middleware.py b/tests/test_tool_selector_middleware.py index e769f71..74ad16d 100644 --- a/tests/test_tool_selector_middleware.py +++ b/tests/test_tool_selector_middleware.py @@ -1,7 +1,7 @@ """Tests for LLMToolSelectorMiddleware integration and the event-sink handoff.""" from typing import Any -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from langchain.agents.middleware.types import ModelRequest @@ -230,6 +230,60 @@ def test_selector_failure_reports_ended_without_selection(): assert sink.calls[-1] == ("ended",) +def test_selector_failure_warns_once_per_middleware_instance(caplog): + """Repeated degradation stays visible without warning on every request.""" + mock_selector = MagicMock() + mock_selector.wrap_model_call.side_effect = RuntimeError("revoked credentials") + cond = _ConditionalToolSelectorMiddleware( + selector_factory=MagicMock(return_value=mock_selector), + threshold=5, + ) + request = _request([_tool(f"t{i}") for i in range(10)]) + + caplog.set_level("WARNING", logger="EvoScientist.middleware.tool_selector") + cond.wrap_model_call(request, MagicMock()) + cond.wrap_model_call(request, MagicMock()) + + warnings = [ + record + for record in caplog.records + if "tool_selector.fallback" in record.getMessage() + ] + assert len(warnings) == 1 + assert "RuntimeError" in warnings[0].getMessage() + + +@pytest.mark.asyncio +async def test_selector_provider_failure_allows_downstream_model_fallback(caplog): + """A failed fixed selector model must not block a healthy request fallback.""" + from EvoScientist.llm.errors import ProviderStreamError + + mock_selector = MagicMock() + mock_selector.awrap_model_call = AsyncMock( + side_effect=ProviderStreamError( + provider="openrouter", + class_qualname="openrouter.ProviderError", + message="primary unavailable", + ) + ) + cond = _ConditionalToolSelectorMiddleware( + selector_factory=MagicMock(return_value=mock_selector), + threshold=5, + ) + request = _request([_tool(f"t{i}") for i in range(10)]) + response = MagicMock() + handler = AsyncMock(return_value=response) + + caplog.set_level("WARNING", logger="EvoScientist.middleware.tool_selector") + result = await cond.awrap_model_call(request, handler) + + assert result is response + handler.assert_awaited_once_with(request) + assert any( + "tool_selector.fallback" in record.getMessage() for record in caplog.records + ) + + def test_selector_failure_ends_before_sync_fallback_handler(): """All-tools fallback must not run while selector suppression is active.""" mock_selector = MagicMock() diff --git a/tests/test_tui_banner_position.py b/tests/test_tui_banner_position.py index 969476c..82771dc 100644 --- a/tests/test_tui_banner_position.py +++ b/tests/test_tui_banner_position.py @@ -16,6 +16,7 @@ pilot so they exercise the actual production code path. from __future__ import annotations +import asyncio from contextlib import asynccontextmanager from pathlib import Path from unittest.mock import AsyncMock @@ -31,13 +32,12 @@ pytest.importorskip("textual") # ``EvoTextualInteractiveApp`` is defined inside ``run_textual_interactive``, # so it is not reachable as ``tui_interactive.EvoTextualInteractiveApp``. We # grab it by invoking the factory once with a patched ``App.run_async`` that -# captures the freshly-built instance. The factory must be invoked from a -# fresh top-level event loop (it pulls in nest_asyncio and the global loop), -# so each test boots it via :func:`_capture_app`. +# captures the freshly-built instance. The synchronous factory owns its +# top-level event loop, so each test boots it via :func:`_capture_app`. # --------------------------------------------------------------------------- -def _capture_app(monkeypatch) -> object: +async def _capture_app(monkeypatch) -> object: """Build an ``EvoTextualInteractiveApp`` without entering its main loop.""" from textual.app import App @@ -83,10 +83,11 @@ def _capture_app(monkeypatch) -> object: # and the module-level ``create_session_workspace`` / ``load_agent`` # symbols never get a chance to run. - # The factory is synchronous at the outer level — it drives its own loop - # via ``nest_asyncio`` + ``loop.run_until_complete`` internally. + # The factory is synchronous at the outer level and owns its top-level + # loop via ``asyncio.run``. try: - tui_mod.run_textual_interactive( + await asyncio.to_thread( + tui_mod.run_textual_interactive, show_thinking=False, channel_send_thinking=False, workspace_dir=None, @@ -131,7 +132,7 @@ async def test_clear_chat_resets_scroll_after_long_anchored_conversation( from textual.containers import VerticalScroll from textual.widgets import Static - app = _capture_app(monkeypatch) + app = await _capture_app(monkeypatch) async with app.run_test(size=(80, 24)) as pilot: await pilot.pause() chat = app.query_one("#chat", VerticalScroll) @@ -162,7 +163,7 @@ async def test_clear_chat_with_anchor_released_also_resets(monkeypatch): from textual.containers import VerticalScroll from textual.widgets import Static - app = _capture_app(monkeypatch) + app = await _capture_app(monkeypatch) async with app.run_test(size=(80, 24)) as pilot: await pilot.pause() chat = app.query_one("#chat", VerticalScroll) @@ -190,7 +191,7 @@ async def test_clear_chat_short_conversation_anchored(monkeypatch): from textual.containers import VerticalScroll from textual.widgets import Static - app = _capture_app(monkeypatch) + app = await _capture_app(monkeypatch) async with app.run_test(size=(80, 24)) as pilot: await pilot.pause() chat = app.query_one("#chat", VerticalScroll) @@ -227,7 +228,7 @@ async def test_clear_chat_then_full_user_turn_keeps_banner_at_top(monkeypatch): from textual.containers import VerticalScroll from textual.widgets import Static - app = _capture_app(monkeypatch) + app = await _capture_app(monkeypatch) # Tall-ish terminal: welcome + a few messages must fit in the # viewport, mirroring the user's manual-test setup. async with app.run_test(size=(80, 40)) as pilot: @@ -301,7 +302,7 @@ async def test_short_turn_keeps_banner_at_top_after_layout_refresh(monkeypatch): from EvoScientist.cli.widgets.assistant_message import AssistantMessage from EvoScientist.cli.widgets.user_message import UserMessage - app = _capture_app(monkeypatch) + app = await _capture_app(monkeypatch) # Tall terminal: welcome + a short exchange fits with room to spare, # which is exactly the bug condition (content < viewport). async with app.run_test(size=(80, 40)) as pilot: @@ -339,7 +340,7 @@ async def test_long_turn_keeps_viewport_pinned_to_bottom(monkeypatch): from textual.containers import VerticalScroll from textual.widgets import Static - app = _capture_app(monkeypatch) + app = await _capture_app(monkeypatch) async with app.run_test(size=(80, 24)) as pilot: await pilot.pause() chat = app.query_one("#chat", VerticalScroll) diff --git a/tests/test_ui_runtime.py b/tests/test_ui_runtime.py index 19b84ae..9d5ba54 100644 --- a/tests/test_ui_runtime.py +++ b/tests/test_ui_runtime.py @@ -1,12 +1,20 @@ """Tests for UI backend runtime selection.""" +import asyncio +import threading +import time from dataclasses import dataclass +import pytest + from EvoScientist.cli.tui_runtime import ( + StreamCancellationTimeout, normalize_ui_backend, resolve_ui_backend, run_streaming, + run_streaming_async, ) +from EvoScientist.runtime import AsyncRuntimeError from tests.fakes import FakeGraphGateway @@ -77,3 +85,151 @@ def test_run_streaming_falls_back_to_cli_on_runtime_error(monkeypatch): gateway=FakeGraphGateway(), ) assert result == "fallback-ok" + + +def test_run_streaming_does_not_retry_on_owned_runtime_error(monkeypatch): + attempts = 0 + + class _RuntimeFailureBackend: + def run_streaming(self, **kwargs): + nonlocal attempts + attempts += 1 + raise AsyncRuntimeError("owned runtime failed") + + monkeypatch.setattr( + "EvoScientist.cli.tui_runtime.get_backend", + lambda *a, **k: _RuntimeFailureBackend(), + ) + monkeypatch.setattr( + "EvoScientist.cli.tui_runtime.RichStreamingBackend", + lambda: _RuntimeFailureBackend(), + ) + + with pytest.raises(AsyncRuntimeError, match="owned runtime failed"): + run_streaming( + ui_backend="tui", + agent=object(), + message="hello", + thread_id="t1", + show_thinking=False, + interactive=True, + gateway=FakeGraphGateway(), + ) + + assert attempts == 1 + + +async def test_async_streaming_cancellation_stops_and_joins_worker(monkeypatch): + from EvoScientist.stream.display import ( + discard_stream_cancel, + is_stream_cancel_requested, + ) + + scope = "test:async-renderer-cancel" + started = threading.Event() + finished = threading.Event() + + def fake_run_streaming(**kwargs): + assert kwargs["cancel_scope"] == scope + started.set() + while not is_stream_cancel_requested(scope): + time.sleep(0.001) + finished.set() + return "stopped" + + monkeypatch.setattr( + "EvoScientist.cli.tui_runtime.run_streaming", fake_run_streaming + ) + + task = asyncio.create_task(run_streaming_async(cancel_scope=scope)) + assert await asyncio.to_thread(started.wait, 1) + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task + + assert finished.is_set() + discard_stream_cancel(scope) + + +async def test_async_streaming_can_recover_foreground_task_after_cancel(monkeypatch): + from EvoScientist.stream.display import ( + discard_stream_cancel, + is_stream_cancel_requested, + ) + + scope = "test:async-renderer-recover" + started = threading.Event() + + def fake_run_streaming(**kwargs): + started.set() + while not is_stream_cancel_requested(scope): + time.sleep(0.001) + return "[Stopped.]" + + cleanup_called = False + + async def fake_cleanup() -> None: + nonlocal cleanup_called + cleanup_called = True + + monkeypatch.setattr( + "EvoScientist.cli.tui_runtime.run_streaming", fake_run_streaming + ) + monkeypatch.setattr( + "EvoScientist.middleware.code_interpreter.aclose_code_interpreters", + fake_cleanup, + ) + + task = asyncio.create_task( + run_streaming_async(cancel_scope=scope, recover_on_cancel=True) + ) + assert await asyncio.to_thread(started.wait, 1) + task.cancel() + + assert await task == "[Stopped.]" + assert cleanup_called + assert not task.cancelled() + discard_stream_cancel(scope) + + +async def test_noncooperative_worker_reports_settlement_timeout(monkeypatch): + """Cancellation timeout is an ordinary lifecycle error, not BaseException.""" + from EvoScientist.cli import tui_runtime + from EvoScientist.stream.display import discard_stream_cancel + + scope = "test:async-renderer-timeout" + started = threading.Event() + release = threading.Event() + finished = threading.Event() + + def fake_run_streaming(**_kwargs): + started.set() + release.wait() + finished.set() + return "late" + + async def fake_cleanup() -> None: + return None + + monkeypatch.setattr(tui_runtime, "run_streaming", fake_run_streaming) + monkeypatch.setattr(tui_runtime, "STREAM_CANCEL_SETTLE_TIMEOUT", 0.05) + monkeypatch.setattr( + "EvoScientist.middleware.code_interpreter.aclose_code_interpreters", + fake_cleanup, + ) + + task = asyncio.create_task( + run_streaming_async(cancel_scope=scope, recover_on_cancel=True) + ) + assert await asyncio.to_thread(started.wait, 1) + task.cancel() + + with pytest.raises(StreamCancellationTimeout, match="did not stop"): + await task + + assert not finished.is_set() + release.set() + assert await asyncio.to_thread(finished.wait, 1) + await asyncio.sleep(0) + discard_stream_cancel(scope) diff --git a/uv.lock b/uv.lock index c6e855b..ea0a46b 100644 --- a/uv.lock +++ b/uv.lock @@ -972,7 +972,6 @@ dependencies = [ { name = "langgraph-sdk" }, { name = "lazy-loader" }, { name = "markdownify" }, - { name = "nest-asyncio" }, { name = "openrouter" }, { name = "prompt-toolkit" }, { name = "psutil" }, @@ -1083,7 +1082,6 @@ requires-dist = [ { name = "lark-oapi", marker = "extra == 'feishu'", specifier = ">=1.4.0" }, { name = "lazy-loader", specifier = ">=0.5" }, { name = "markdownify", specifier = ">=1.2" }, - { name = "nest-asyncio", specifier = ">=1.6" }, { name = "openrouter", specifier = ">=0.10.8,<0.11.0" }, { name = "pre-commit", marker = "extra == 'dev'", specifier = ">=3.5.0" }, { name = "prompt-toolkit", specifier = ">=3.0" }, @@ -2694,15 +2692,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/81/08/7036c080d7117f28a4af526d794aab6a84463126db031b007717c1a6676e/multidict-6.7.1-py3-none-any.whl", hash = "sha256:55d97cc6dae627efa6a6e548885712d4864b81110ac76fa4e534c03819fa4a56", size = 12319, upload-time = "2026-01-26T02:46:44.004Z" }, ] -[[package]] -name = "nest-asyncio" -version = "1.6.0" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/83/f8/51569ac65d696c8ecbee95938f89d4abf00f47d58d48f6fbabfe8f0baefe/nest_asyncio-1.6.0.tar.gz", hash = "sha256:6f172d5449aca15afd6c646851f4e31e02c598d553a667e38cafa997cfec55fe", size = 7418, upload-time = "2024-01-21T14:25:19.227Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/a0/c4/c2971a3ba4c6103a3d10c4b0f24f461ddc027f0f09763220cf35ca1401b3/nest_asyncio-1.6.0-py3-none-any.whl", hash = "sha256:87af6efd6b5e897c81050477ef65c62e2b2f35d51703cae01aff2905b1852e1c", size = 5195, upload-time = "2024-01-21T14:25:17.223Z" }, -] - [[package]] name = "nodeenv" version = "1.10.0"