Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 50 additions & 4 deletions csp_bot/commands/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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).

Expand All @@ -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()

Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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.",
Expand Down Expand Up @@ -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)
Expand Down
80 changes: 78 additions & 2 deletions csp_bot/tests/test_agent_command.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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


Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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)
Expand All @@ -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):
Expand Down
Loading