test(api_server): drive the real _run_agent call site through handler cancellation

The salvaged regression exercises the worker counter and the shutdown gate
directly; this one drives the production entry point (APIServerAdapter._run_agent),
cancels the handler task while the executor thread is parked in the turn, and
asserts the worker-scoped count survives the cancellation and is released only when
the worker exits. Reverting just the api_server.py call-site hunk turns it red.

Refs #116535.
This commit is contained in:
teknium1
2026-09-19 23:34:40 -07:00
committed by Teknium
parent 589d4d05d4
commit a4716a8cd2

View File

@@ -369,6 +369,40 @@ class TestRunAgentRegistersForShutdownInterrupt:
assert list(observed["during"].values()) == [agent]
assert adapter._shutdown_interruptible_agents == {}
@pytest.mark.asyncio
async def test_cancelled_handler_keeps_the_worker_counted_until_the_turn_exits(self):
"""Cancelling the ``_run_agent`` handler task must not drop the worker count (#116535).
The handler-side ``_inflight_agent_runs`` legitimately drops in the handler's
``finally``; the shutdown SessionDB-close gate reads the worker-scoped count instead,
which the real ``_run_agent`` call site must hold until the executor thread exits.
"""
adapter = APIServerAdapter(PlatformConfig(enabled=True))
loop = asyncio.get_running_loop()
started, release = asyncio.Event(), threading.Event()
agent = _parked_agent(loop, started, release)
agent.interrupt.side_effect = None
baseline = _api_runs.api_worker_live_count()
with patch.object(adapter, "_create_agent", return_value=agent):
task = asyncio.ensure_future(
adapter._run_agent(user_message="hello", conversation_history=[], session_id="s1"))
await asyncio.wait_for(started.wait(), _TURN_UNBLOCK_TIMEOUT)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert adapter._inflight_agent_runs == 0
assert _api_runs.api_worker_live_count() == baseline + 1, (
"cancelled handler dropped the worker-scoped count while the turn was still running")
release.set()
for _ in range(200):
if _api_runs.api_worker_live_count() == baseline:
break
await asyncio.sleep(0.01)
assert _api_runs.api_worker_live_count() == baseline, "worker exit did not release the count"
@pytest.mark.asyncio
async def test_agent_is_unregistered_when_the_turn_raises(self):
adapter = APIServerAdapter(PlatformConfig(enabled=True))