hermes-agent/tests/gateway/test_sse_agent_cancel.py

192 lines
6.5 KiB
Python

"""Tests for SSE client disconnect → agent task cancellation.
When a streaming /v1/chat/completions client disconnects mid-stream
(network drop, browser tab close), the agent is interrupted via
agent.interrupt() so it stops making LLM API calls, and the asyncio
task wrapper is cancelled.
"""
import asyncio
import queue
from unittest.mock import AsyncMock, MagicMock, patch
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_adapter():
"""Build a minimal APIServerAdapter with mocked internals."""
from gateway.platforms.api_server import APIServerAdapter
from gateway.config import PlatformConfig
config = PlatformConfig(enabled=True, token="test-key")
adapter = APIServerAdapter(config)
return adapter
def _make_request():
"""Build a mock aiohttp request."""
req = MagicMock()
req.headers = {}
return req
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestSSEAgentCancelOnDisconnect:
"""gateway/platforms/api_server.py — _write_sse_chat_completion()"""
def test_agent_task_cancelled_on_client_disconnect(self):
"""When response.write raises ConnectionResetError (client dropped),
the agent task must be cancelled."""
adapter = _make_adapter()
stream_q = queue.Queue()
stream_q.put("hello ") # Some data already queued
# Agent task that runs forever (simulates a long LLM call)
agent_done = asyncio.Event()
async def fake_agent():
await agent_done.wait()
return {"final_response": "done"}, {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}
async def run():
from aiohttp import web
agent_task = asyncio.ensure_future(fake_agent())
# Mock response that raises ConnectionResetError on second write
mock_response = AsyncMock(spec=web.StreamResponse)
call_count = 0
async def write_side_effect(data):
nonlocal call_count
call_count += 1
if call_count >= 2:
raise ConnectionResetError("client disconnected")
mock_response.write = AsyncMock(side_effect=write_side_effect)
mock_response.prepare = AsyncMock()
with patch.object(type(adapter), '_write_sse_chat_completion',
adapter._write_sse_chat_completion):
# Patch StreamResponse creation
with patch("gateway.platforms.api_server.web.StreamResponse",
return_value=mock_response):
await adapter._write_sse_chat_completion(
_make_request(), "cmpl-123", "gpt-4", 1234567890,
stream_q, agent_task,
)
# The critical assertion: agent_task must be cancelled
assert agent_task.cancelled() or agent_task.done()
# Clean up
agent_done.set()
asyncio.run(run())
def test_broken_pipe_also_cancels_agent(self):
"""BrokenPipeError (another disconnect variant) also cancels the task."""
adapter = _make_adapter()
stream_q = queue.Queue()
async def fake_agent():
await asyncio.sleep(0.2) # Never completes
return {}, {}
async def run():
from aiohttp import web
agent_task = asyncio.ensure_future(fake_agent())
mock_response = AsyncMock(spec=web.StreamResponse)
mock_response.write = AsyncMock(side_effect=BrokenPipeError("pipe broken"))
mock_response.prepare = AsyncMock()
with patch("gateway.platforms.api_server.web.StreamResponse",
return_value=mock_response):
await adapter._write_sse_chat_completion(
_make_request(), "cmpl-789", "gpt-4", 1234567890,
stream_q, agent_task,
)
assert agent_task.cancelled() or agent_task.done()
asyncio.run(run())
def _capturing_response():
"""Mock StreamResponse that records all written SSE bytes as text."""
from aiohttp import web
chunks: list = []
resp = AsyncMock(spec=web.StreamResponse)
resp.prepare = AsyncMock()
async def _write(data):
chunks.append(data.decode() if isinstance(data, (bytes, bytearray)) else data)
resp.write = AsyncMock(side_effect=_write)
return resp, chunks
def _finish_reason(chunks: list):
"""Extract the terminal finish_reason and its chunk from captured SSE."""
import json
sse = "".join(chunks)
finish = None
for line in sse.splitlines():
if line.startswith("data: ") and '"finish_reason"' in line:
obj = json.loads(line[6:])
if obj["choices"][0].get("finish_reason") is not None:
finish = obj
return (finish["choices"][0]["finish_reason"] if finish else None), finish, sse
class TestSSEAgentFailureFinishReason:
"""gateway/platforms/api_server.py — _write_sse_chat_completion()
A clean stream-queue termination (sentinel received) followed by an agent
failure must NOT report finish_reason: "stop". Both failure modes — an
``agent_task`` that raises and a ``result`` dict flagged failed — surface
as finish_reason: "error", mirroring the non-streaming path. Issue #12422.
"""
def _run(self, fake_agent, queue_items=("partial",)):
adapter = _make_adapter()
stream_q = queue.Queue()
for item in queue_items:
stream_q.put(item)
stream_q.put(None) # clean end-of-stream sentinel
async def run():
agent_task = asyncio.ensure_future(fake_agent())
resp, chunks = _capturing_response()
with patch("gateway.platforms.api_server.web.StreamResponse",
return_value=resp):
await adapter._write_sse_chat_completion(
_make_request(), "cmpl-fail", "gpt-4", 1234567890,
stream_q, agent_task,
)
return _finish_reason(chunks)
return asyncio.run(run())
def test_agent_task_raises_reports_error_not_stop(self):
async def crash():
raise RuntimeError("boom from agent")
reason, finish, sse = self._run(crash)
assert reason == "error"
assert "error" in finish
assert "data: [DONE]" in sse