From 609af9fe6701e8d391040b44d23282b27d866485 Mon Sep 17 00:00:00 2001 From: h Date: Fri, 28 Aug 2026 19:33:26 +0200 Subject: [PATCH] fix(core): a turn cut by shutdown keeps running_turn for recover --- src/beaver_gateway/core/conversations.py | 15 ++++++++++++--- tests/test_conversations.py | 16 ++++++++++++++++ 2 files changed, 28 insertions(+), 3 deletions(-) diff --git a/src/beaver_gateway/core/conversations.py b/src/beaver_gateway/core/conversations.py index 052f6f6..c3e2e38 100644 --- a/src/beaver_gateway/core/conversations.py +++ b/src/beaver_gateway/core/conversations.py @@ -825,6 +825,7 @@ class Conversations: item_origin=item_origin, ) stop = "error" + cut = False try: events = backend.complete( agent=self._claude_agent(conv.agent_name), @@ -841,9 +842,12 @@ class Conversations: async for event in events: yield event stop = "interrupted" if capture.interrupted else "end_turn" + except asyncio.CancelledError: + cut = True + raise finally: runner.turn_id = None - await self._mark_done(conv, capture) + await self._mark_done(conv, capture, cut=cut) self._bus.publish( "turn.end", conversation_id=conv.external_id, @@ -1173,9 +1177,14 @@ class Conversations: await self._update(conv, apply) - async def _mark_done(self, conv: Conversation, capture: TurnCapture) -> None: + async def _mark_done( + self, conv: Conversation, capture: TurnCapture, *, cut: bool = False + ) -> None: + """Close the turn; a cancelled one keeps ``running_turn`` for ``recover``.""" + async def apply(row: Conversation) -> None: - row.running_turn = None + if not cut: + row.running_turn = None row.last_activity_at = datetime.now(UTC) if capture.session_id is not None: row.session_id = capture.session_id diff --git a/tests/test_conversations.py b/tests/test_conversations.py index 95b195a..01331b8 100644 --- a/tests/test_conversations.py +++ b/tests/test_conversations.py @@ -496,6 +496,22 @@ async def test_recover_closes_open_tool_use_and_injects_interrupted( assert ScriptedClient.instances == [] +async def test_graceful_stop_keeps_running_turn_for_recover(world: World) -> None: + conv = await world.conversations.create(kind="master", agent="a", origin="test") + ScriptedClient.hold = asyncio.Event() + await world.conversations.post(conv, "long") + await asyncio.sleep(0.2) + assert (await world.conversations.get(conv.external_id)).running_turn is not None + await world.conversations.stop() + row = await world.conversations.get(conv.external_id) + assert row.running_turn is not None + cut = await world.conversations.recover() + assert [c.external_id for c in cut] == [conv.external_id] + assert (await world.conversations.get(conv.external_id)).running_turn is None + assert await world.statuses(conv) == [("user", "interrupted"), ("normal", "queued")] + ScriptedClient.hold = None + + async def test_read_and_bindings(world: World) -> None: sid = str(uuid.uuid4()) history = [