refactor(runtime): centralize async bridges under an owned runtime (#376)
* feat(runtime): add application-scoped async runtime * refactor(cli): use owned runtime for session stats * refactor(onboard): use the owned async runtime * docs(runtime): record async bridge ownership * refactor(middleware): keep sync fallback synchronous * refactor(mcp): load tools on an owned runtime * refactor(cli): share owned runtime across entry points * refactor(channels): make inbound sync bridge explicit * refactor(stream): run Rich streaming on owned runtime * chore(runtime): remove nest-asyncio dependency * refactor(asyncio): require active loops in async code * docs(runtime): document final event loop ownership * fix(stream): cancel stalled owned streams * fix(cli): recover cleanly from stream cancellation * fix(runtime): drain executor work before shutdown * fix(runtime): terminate cancelled shell process trees * fix(models): let fallback bypass selector failures * fix(cli): reset interrupt handling between turns * docs: rm implementation spec * fix(serve): cancel active turns during shutdown * fix(runtime): protect settlement from waiter cancellation * fix(backends): reject empty shell commands * fix(runtime): terminate descendants after shell exit * fix(mcp): keep standalone discovery off channel loop * fix(cli): own and settle interactive prompt cancellation * fix(serve): keep channel sends off runtime loop * fix(stream): scope cancel context to iterator steps * refactor(serve): require the owned async runtime * fix(channels): keep interactive sends off runtime loop * fix(selector): surface fallback without log spam * test(runtime): normalize Windows shell marker * fix(cli): serialize interactive session turns * fix(shell): bound output drain after termination * fix(ui): do not retry owned runtime failures * fix(shell): allow signal-safe registry reentry * fix(shell): avoid terminating reused process ids * fix(channels): preserve streaming send order * fix(cli): report runtime shutdown timeouts cleanly * fix(mcp): guide async callers to async loader * docs(runtime): clarify reserved async bridge APIs * fix(runtime): bound code interpreter cleanup * test(shell): use active Python for drain regression --------- Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
This commit is contained in:
@@ -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(
|
||||
|
||||
+241
-3
@@ -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 "<no output>"
|
||||
|
||||
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:
|
||||
|
||||
@@ -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()
|
||||
@@ -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).
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
+140
-87
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
+133
-33
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 <section>`` 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)
|
||||
|
||||
|
||||
+27
-11
@@ -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 {}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
+352
-172
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
+127
-11
@@ -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 ===
|
||||
|
||||
|
||||
@@ -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
|
||||
# ═══════════════════════════════════════════════════════════════════
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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
|
||||
+112
-11
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
+126
-302
@@ -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]
|
||||
|
||||
@@ -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")]
|
||||
|
||||
@@ -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
|
||||
# ═════════════════════════════════════════════════════════════════
|
||||
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user