diff --git a/csp_bot/commands/agent.py b/csp_bot/commands/agent.py index a640a87..37b9f79 100644 --- a/csp_bot/commands/agent.py +++ b/csp_bot/commands/agent.py @@ -18,7 +18,7 @@ import os import threading from abc import abstractmethod -from collections.abc import Sequence +from collections.abc import Callable, Sequence from concurrent.futures import Future, ThreadPoolExecutor from dataclasses import dataclass, field from datetime import UTC, datetime, timedelta @@ -50,6 +50,33 @@ def _utc_now() -> datetime: return datetime.now(UTC) +class _AgentRunControl: + """Thread-safe cancellation bridge to an active asyncio agent run.""" + + def __init__(self) -> None: + self._lock = threading.Lock() + self._cancel_requested = False + self._cancel: Callable[[], Any] | None = None + + def bind(self, cancel: Callable[[], Any]) -> None: + with self._lock: + self._cancel = cancel + cancel_requested = self._cancel_requested + if cancel_requested: + cancel() + + def cancel(self) -> None: + with self._lock: + self._cancel_requested = True + cancel = self._cancel + if cancel is not None: + cancel() + + def clear(self) -> None: + with self._lock: + self._cancel = None + + @dataclass class AgentSession: """Tracks a multi-turn conversation between a user and an agent command.""" @@ -225,6 +252,7 @@ def _run_agent( prompt: str | Sequence[Any], loop: asyncio.AbstractEventLoop | None = None, message_history: Sequence[Any] | None = None, + control: _AgentRunControl | None = None, ) -> Any: """Run an agent on an event loop (for use in thread pool). @@ -239,14 +267,25 @@ def _run_agent( coro = agent.run(prompt, message_history=message_history) if loop is not None and loop.is_running(): future = asyncio.run_coroutine_threadsafe(coro, loop) - return future.result() + if control is not None: + control.bind(future.cancel) + try: + return future.result() + finally: + if control is not None: + control.clear() owns_loop = loop is None if owns_loop: loop = asyncio.new_event_loop() + task = loop.create_task(coro) + if control is not None: + control.bind(lambda: loop.call_soon_threadsafe(task.cancel)) try: - return loop.run_until_complete(coro) + return loop.run_until_complete(task) finally: + if control is not None: + control.clear() if owns_loop: loop.close() @@ -287,6 +326,7 @@ def build_prompt(self, command): _backends: ClassVar[dict[str, BackendBase]] = {} _backend_loops: ClassVar[dict[str, asyncio.AbstractEventLoop]] = {} _futures: ClassVar[dict[str, Future]] = {} + _run_controls: ClassVar[dict[str, _AgentRunControl]] = {} _sessions: ClassVar[SessionStore] = SessionStore(ttl_seconds=900.0) model_name: str = "claude-sonnet-4-6" @@ -655,8 +695,10 @@ def preexecute(self, command: BotCommand) -> BotCommand: # Use the backend's event loop so aiohttp sessions stay valid backend_loop = self._backend_loops.get(command.backend) - future = _executor.submit(_run_agent, agent, prompt, backend_loop, history) + control = _AgentRunControl() + future = _executor.submit(_run_agent, agent, prompt, backend_loop, history, control) self._futures[key] = future + self._run_controls[key] = control log.info( "AgentCommand[%s] submitted for user %s (session history: %d msgs)", self.command(), @@ -693,6 +735,9 @@ def execute(self, command: BotCommand) -> Message | list[Message | BaseCommand] elapsed = command.times_run * self.poll_interval if elapsed >= self.timeout: self._futures.pop(key, None) + control = self._run_controls.pop(key, None) + if control is not None: + control.cancel() future.cancel() return Message( content="Sorry, the AI request timed out. Please try again.", @@ -732,6 +777,7 @@ def execute(self, command: BotCommand) -> Message | list[Message | BaseCommand] # Future is done — get result self._futures.pop(key, None) + self._run_controls.pop(key, None) try: result = future.result() output = str(result.output) if hasattr(result, "output") else str(result) diff --git a/csp_bot/tests/test_agent_command.py b/csp_bot/tests/test_agent_command.py index 6ed5ef5..c7d61d9 100644 --- a/csp_bot/tests/test_agent_command.py +++ b/csp_bot/tests/test_agent_command.py @@ -3,7 +3,7 @@ import asyncio import socket import threading -from concurrent.futures import Future +from concurrent.futures import Future, ThreadPoolExecutor from datetime import UTC, datetime, timedelta from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -12,7 +12,7 @@ from chatom import Message, User from chatom.backend import BackendBase -from csp_bot.commands.agent import AgentCommand, AgentCommandModel, AgentSession, SessionStore, _run_agent +from csp_bot.commands.agent import AgentCommand, AgentCommandModel, AgentSession, SessionStore, _AgentRunControl, _run_agent from csp_bot.structs import BotCommand, CommandVariant @@ -41,7 +41,9 @@ def build_prompt(self, command): def cmd(): """Fresh AgentCommand instance with clean state.""" AgentCommand._futures = {} + AgentCommand._run_controls = {} AgentCommand._backends = {} + AgentCommand._backend_loops = {} AgentCommand._sessions = SessionStore(ttl_seconds=900.0) return ConcreteAgentCommand() @@ -402,6 +404,59 @@ def test_timeout_cancels_future(self, cmd, bot_command): mock_future.cancel.assert_called_once() assert key not in AgentCommand._futures + @pytest.mark.parametrize("use_backend_loop", [False, True]) + def test_timeout_stops_active_agent_and_releases_worker(self, cmd, bot_command, use_backend_loop): + started = threading.Event() + cancelled = threading.Event() + release = threading.Event() + + class BlockingAgent: + async def run(self, prompt, message_history=None): + started.set() + try: + while not release.is_set(): + await asyncio.sleep(0.01) + except asyncio.CancelledError: + cancelled.set() + raise + + backend_loop = None + loop_thread = None + if use_backend_loop: + backend_loop = asyncio.new_event_loop() + + def run_loop(): + asyncio.set_event_loop(backend_loop) + backend_loop.run_forever() + + loop_thread = threading.Thread(target=run_loop, daemon=True) + loop_thread.start() + AgentCommand._backend_loops = {bot_command.backend: backend_loop} + + executor = ThreadPoolExecutor(max_workers=1) + try: + with patch("csp_bot.commands.agent._executor", executor), patch.object(cmd, "build_agent", return_value=BlockingAgent()): + cmd.preexecute(bot_command) + assert started.wait(timeout=1) + + bot_command.times_run = cmd.timeout // cmd.poll_interval + 1 + result = cmd.execute(bot_command) + + assert isinstance(result, Message) + assert "timed out" in result.content.lower() + assert cancelled.wait(timeout=1) + assert cmd._command_key(bot_command) not in AgentCommand._run_controls + assert executor.submit(lambda: "released").result(timeout=1) == "released" + finally: + release.set() + executor.shutdown(wait=True) + if backend_loop is not None: + backend_loop.call_soon_threadsafe(backend_loop.stop) + if loop_thread is not None: + loop_thread.join(timeout=1) + if backend_loop is not None: + backend_loop.close() + def test_no_future_returns_error(self, cmd, bot_command): bot_command.args = () # Clear any error args result = cmd.execute(bot_command) @@ -410,6 +465,27 @@ def test_no_future_returns_error(self, cmd, bot_command): class TestRunAgent: + def test_honors_cancellation_requested_before_task_binding(self): + release = threading.Event() + + class BlockingAgent: + async def run(self, prompt, message_history=None): + while not release.is_set(): + await asyncio.sleep(0.01) + + control = _AgentRunControl() + control.cancel() + executor = ThreadPoolExecutor(max_workers=1) + try: + future = executor.submit(_run_agent, BlockingAgent(), "hello", None, None, control) + + with pytest.raises(asyncio.CancelledError): + future.result(timeout=1) + assert executor.submit(lambda: "released").result(timeout=1) == "released" + finally: + release.set() + executor.shutdown(wait=True) + def test_uses_threadsafe_submission_for_running_backend_loop(self): class FakeAgent: async def run(self, prompt, message_history=None):