From 8856ae68f6a553a35b3f3b225a838833464ab995 Mon Sep 17 00:00:00 2001 From: mac Date: Wed, 26 Aug 2026 23:13:27 +0300 Subject: [PATCH] server: guarantee abort delivery on stream cancellation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit stream_with_cancellation reacted to a client disconnect with an unowned asyncio.create_task(abort_user(uid)): request teardown could outrun delivery and a failure inside the task degraded to a never-retrieved-exception warning, so the scheduler kept decoding for a client that was gone. Await the abort inline behind asyncio.shield (a second cancellation cannot kill the delivery task), and make abort_user claim the uid first — exactly one AbortMsg even if cancellation runs twice, and none at all when the stream already finished normally. Found via the freetoken-mlx downstream audit (docs/AUDIT.md, defect 2). --- python/freetoken/server/api_server.py | 13 ++- tests/server/test_stream_cancellation.py | 120 +++++++++++++++++++++++ 2 files changed, 128 insertions(+), 5 deletions(-) create mode 100644 tests/server/test_stream_cancellation.py diff --git a/python/freetoken/server/api_server.py b/python/freetoken/server/api_server.py index 3e2acc854..add3c6605 100644 --- a/python/freetoken/server/api_server.py +++ b/python/freetoken/server/api_server.py @@ -364,15 +364,18 @@ async def stream_with_cancellation(self, generator, request: Request, uid: int): raise asyncio.CancelledError yield chunk except asyncio.CancelledError: - asyncio.create_task(self.abort_user(uid)) + try: + await asyncio.shield(self.abort_user(uid)) + except Exception: # noqa: BLE001 + logger.exception("Failed to deliver abort for user %s", uid) raise async def abort_user(self, uid: int): + claimed = self.event_map.pop(uid, None) is not None + self.ack_map.pop(uid, None) + if not claimed: + return await asyncio.sleep(0.1) - if uid in self.ack_map: - del self.ack_map[uid] - if uid in self.event_map: - del self.event_map[uid] self.stats.on_abort(uid) logger.warning("Aborting request for user %s", uid) await self.send_one(AbortMsg(uid=uid)) diff --git a/tests/server/test_stream_cancellation.py b/tests/server/test_stream_cancellation.py new file mode 100644 index 000000000..120c50456 --- /dev/null +++ b/tests/server/test_stream_cancellation.py @@ -0,0 +1,120 @@ +"""Cancellation-path tests for FrontendManager.stream_with_cancellation / abort_user.""" + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace + +import pytest + +from freetoken.message import AbortMsg +from freetoken.server.api_server import FrontendManager + + +class _Stats: + def __init__(self): + self.aborts = [] + + def on_abort(self, uid): + self.aborts.append(uid) + + +def _state(send_impl=None): + st = SimpleNamespace( + ack_map={7: [object()]}, + event_map={7: asyncio.Event()}, + stats=_Stats(), + ) + sent = [] + + async def default_send(msg): + sent.append(msg) + + st.send_one = send_impl or default_send + st.sent = sent + return st + + +class _Request: + def __init__(self, disconnected=False): + self._disconnected = disconnected + + async def is_disconnected(self): + return self._disconnected + + +async def _consume(state, request, uid=7): + state.abort_user = lambda request_uid: FrontendManager.abort_user(state, request_uid) + async for _ in FrontendManager.stream_with_cancellation(state, _never(), request, uid): + pass + + +async def _never(): + await asyncio.sleep(3600) + yield b"" # pragma: no cover + + +def test_cancellation_sends_one_abort_and_cleans_maps_inline(): + async def run(): + state = _state() + task = asyncio.create_task(_consume(state, _Request())) + await asyncio.sleep(0.01) + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task + + assert asyncio.all_tasks() - {asyncio.current_task()} == set() + assert len(state.sent) == 1 + assert isinstance(state.sent[0], AbortMsg) + assert state.sent[0].uid == 7 + assert state.ack_map == {} + assert state.event_map == {} + assert state.stats.aborts == [7] + + asyncio.run(run()) + + +def test_abort_delivery_failure_preserves_cancellation(): + async def boom(msg): + raise RuntimeError("zmq down") + + async def run(): + state = _state(send_impl=boom) + task = asyncio.create_task(_consume(state, _Request())) + await asyncio.sleep(0.01) + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task + + asyncio.run(run()) + + +def test_abort_user_is_idempotent(): + async def run(): + state = _state() + + await FrontendManager.abort_user(state, 7) + assert len(state.sent) == 1 + + await FrontendManager.abort_user(state, 7) + assert len(state.sent) == 1 + assert state.stats.aborts == [7] + + asyncio.run(run()) + + +def test_normal_completion_sends_no_abort(): + async def run(): + state = _state() + + async def one(): + yield b"data: x\n\n" + + async for _ in FrontendManager.stream_with_cancellation(state, one(), _Request(), 7): + pass + assert state.sent == [] + assert 7 in state.ack_map # wait_for_ack owns normal-path cleanup, not the stream + + asyncio.run(run())