fix(tui_gateway): retain live-turn sessions unfinalized when the drain deadline expires

The drain reserves a slice of the shutdown budget so flush_all_sessions
still runs when in-flight turns outlast the window. But that flush was
unconditional: a session whose turn was still running got its one-shot
_finalize_session spent mid-turn, and the executor.shutdown(wait=False,
cancel_futures=True) immediately after does not join the turn. The
session was then permanently un-finalizable and its active-session lease
had been released out from under live work — the same persistence and
lifecycle race the drain exists to close, just relocated past the
deadline instead of removed.

Give _turn_futures a session association (Future -> sid, the same key
space as server._sessions) at both submit sites, and on deadline expiry
exclude the sids whose futures are still running from the flush. Those
sessions are retained unfinalized and therefore recoverable; sessions
with no live turn finalize exactly as before. The done-callback now pops
under the lock, since a bare dict.pop is not the drop-in set.discard was.

wait semantics, the reserve math and the bounded per-tick sleep are
unchanged, so this adds no shutdown latency. All three shutdown callers
(orphan, sigterm, and the tight stdin_closed wait=2.0 path) funnel
through this one function and are covered.
This commit is contained in:
briandevans
2026-08-02 16:34:08 -07:00
committed by kshitij
parent e239184521
commit eb4f514b2d
2 changed files with 128 additions and 33 deletions

View File

@@ -133,23 +133,28 @@ def _make_compress_host_session(events: list) -> dict:
}
def _record_finalize(monkeypatch, events: list[str]) -> None:
"""Give ``flush_all_sessions`` one session and record when it finalizes."""
monkeypatch.setattr(server, "_sessions", {"s1": {"session_key": "s1"}}, raising=False)
def _record_finalize(monkeypatch, events: list[str], *sids: str) -> None:
"""Give ``flush_all_sessions`` sessions and record which ones finalize."""
keys = sids or ("s1",)
monkeypatch.setattr(
server,
"_sessions",
{sid: {"session_key": sid} for sid in keys},
raising=False,
)
monkeypatch.setattr(
server,
"_finalize_session",
lambda _session, end_reason="tui_close": events.append(f"finalize:{end_reason}"),
lambda _session, end_reason="tui_close": events.append(
f"finalize:{_session['session_key']}:{end_reason}"
),
raising=False,
)
def _register_turn(host: ComputeHost, fn) -> None:
def _register_turn(host: ComputeHost, fn, sid: str = "s1") -> None:
"""Submit a turn exactly the way ``_handle_turn_start`` does."""
future = host._executor.submit(fn)
with host._turn_futures_lock:
host._turn_futures.add(future)
future.add_done_callback(host._turn_futures.discard)
host._track_turn_future(host._executor.submit(fn), sid)
def test_shutdown_drains_in_flight_turn_before_finalizing_sessions(monkeypatch):
@@ -164,20 +169,29 @@ def test_shutdown_drains_in_flight_turn_before_finalizing_sessions(monkeypatch):
time.sleep(0.3)
events.append("turn_end")
_register_turn(host, _turn)
_register_turn(host, _turn, sid="s1")
assert running.wait(timeout=5.0)
host.shutdown(reason="sigterm", wait=3.0)
# ``_finalize_session`` latches on ``session["_finalized"]``, so its single
# run has to observe the finished turn or the tail is unpersistable.
assert events == ["turn_end", "finalize:compute_host_sigterm"]
# run has to observe the finished turn or the tail is unpersistable. A turn
# that *did* drain must still finalize — the live-turn skip must not
# over-reach into sessions whose work is done.
assert events == ["turn_end", "finalize:s1:compute_host_sigterm"]
# The done-callback still has to remove the entry now that the container is
# a dict: ``set.discard`` was a valid bare callback, ``dict.pop`` is not.
deadline = time.monotonic() + 2.0
while host._turn_futures and time.monotonic() < deadline:
time.sleep(0.01)
assert host._turn_futures == {}, "in-flight turns must not accumulate"
def test_shutdown_still_finalizes_when_the_drain_deadline_expires(monkeypatch):
def test_shutdown_retains_a_live_turns_session_when_the_drain_deadline_expires(monkeypatch):
wait = 1.0
events: list[str] = []
_record_finalize(monkeypatch, events)
_record_finalize(monkeypatch, events, "live", "idle")
host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0)
release = threading.Event()
@@ -187,7 +201,7 @@ def test_shutdown_still_finalizes_when_the_drain_deadline_expires(monkeypatch):
running.set()
release.wait(timeout=30.0)
_register_turn(host, _stuck_turn)
_register_turn(host, _stuck_turn, sid="live")
assert running.wait(timeout=5.0)
try:
@@ -197,10 +211,53 @@ def test_shutdown_still_finalizes_when_the_drain_deadline_expires(monkeypatch):
finally:
release.set()
# A turn that outlives the window must not cost the flush entirely: the
# supervisor's SIGKILL lands on the same deadline this budget comes from,
# so the drain has to stop short and leave the finalize room to run.
assert events == ["finalize:compute_host_sigterm"]
# ``_finalize_session`` is one-shot, and the ``shutdown(wait=False)`` that
# follows does not join the turn. Spending "live"'s single latch mid-turn
# would leave it permanently un-finalizable and release its active-session
# lease out from under running work — the same lifecycle race the drain
# exists to close, just moved past the deadline. It is retained unfinalized
# for recovery instead. A turn outliving the window must not cost the flush
# for anyone else, so "idle" still finalizes in the same pass.
assert events == ["finalize:idle:compute_host_sigterm"]
assert elapsed < wait
def test_shutdown_retains_live_sessions_within_the_stdin_closed_budget(monkeypatch):
"""The tightest real budget any caller uses is ``wait=2.0``.
``run_host`` finalizes through ``host.shutdown(reason="stdin_closed",
wait=2.0)``, which is where the reserve — ``wait`` minus
``min(_FLUSH_RESERVE_SECS, wait / 2)`` — has the least room to work with.
The retain-live-sessions rule must hold there without costing the flush for
idle sessions and without pushing the call past the budget the supervisor's
kill escalation is timed against.
"""
wait = 2.0
drain_budget = wait - min(compute_host._FLUSH_RESERVE_SECS, wait / 2.0)
events: list[str] = []
_record_finalize(monkeypatch, events, "live", "idle")
host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0)
release = threading.Event()
running = threading.Event()
def _stuck_turn() -> None:
running.set()
release.wait(timeout=30.0)
_register_turn(host, _stuck_turn, sid="live")
assert running.wait(timeout=5.0)
try:
started = time.monotonic()
host.shutdown(reason="stdin_closed", wait=wait)
elapsed = time.monotonic() - started
finally:
release.set()
assert events == ["finalize:idle:compute_host_stdin_closed"]
assert elapsed >= drain_budget - 1e-6, "the drain must use its full window"
assert elapsed < wait
@@ -218,7 +275,7 @@ def test_shutdown_drain_sleep_never_overshoots_the_reserve(monkeypatch):
drain_budget = wait - min(compute_host._FLUSH_RESERVE_SECS, wait / 2.0)
events: list[str] = []
_record_finalize(monkeypatch, events)
_record_finalize(monkeypatch, events, "idle")
slept: list[float] = []
real_sleep = time.sleep
@@ -237,7 +294,7 @@ def test_shutdown_drain_sleep_never_overshoots_the_reserve(monkeypatch):
running.set()
release.wait(timeout=30.0)
_register_turn(host, _stuck_turn)
_register_turn(host, _stuck_turn, sid="live")
assert running.wait(timeout=5.0)
try:
@@ -245,6 +302,6 @@ def test_shutdown_drain_sleep_never_overshoots_the_reserve(monkeypatch):
finally:
release.set()
assert events == ["finalize:compute_host_sigterm"]
assert events == ["finalize:idle:compute_host_sigterm"]
assert slept, "the drain loop should have ticked at least once"
assert sum(slept) <= drain_budget + 1e-6

View File

@@ -19,7 +19,7 @@ import time
import uuid
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Callable
from typing import Any, Callable, Collection
from agent.interrupt_compat import request_hard_interrupt
@@ -152,7 +152,10 @@ class ComputeHost:
self._boot_id = uuid.uuid4().hex
self._progress_counter = 0
self._progress_lock = threading.Lock()
self._turn_futures: set[concurrent.futures.Future] = set()
# Future -> the ``sid`` whose turn it is running. ``shutdown`` needs to
# know *whose* turn is still live, not merely that something is, so that
# it can leave those sessions unfinalized; a bare set cannot answer that.
self._turn_futures: dict[concurrent.futures.Future, str] = {}
self._turn_futures_lock = threading.Lock()
self._transport = _HostTransport(self.emit)
self._heartbeat_secs = (
@@ -192,6 +195,15 @@ class ComputeHost:
the drain, so the flush still runs when in-flight turns outlast the
window. ``wait`` itself is unchanged, so this adds no shutdown latency
and no new exposure to the supervisor's kill escalation.
Sessions whose turn is *still running* when the drain deadline expires
are excluded from that flush. Finalizing one would spend its single
latch mid-turn — ``shutdown(wait=False, cancel_futures=True)`` below
does not join the turn — leaving the session permanently
un-finalizable and its active-session lease released out from under
live work: exactly the race the drain exists to close, just moved later.
Leaving them unfinalized keeps them recoverable instead. Sessions with
no live turn finalize here as they always have.
"""
self._closed.set()
budget = max(0.0, wait)
@@ -208,15 +220,30 @@ class ComputeHost:
# deadline and eat into the reserve it is there to protect, which
# for a small ``wait`` can be the whole of it.
time.sleep(min(0.05, remaining))
self.flush_all_sessions(reason=reason)
with self._turn_futures_lock:
live_sids = {sid for future, sid in self._turn_futures.items() if sid and not future.done()}
self.flush_all_sessions(reason=reason, skip_sids=live_sids)
self._executor.shutdown(wait=False, cancel_futures=True)
def flush_all_sessions(self, *, reason: str = "shutdown") -> None:
def flush_all_sessions(
self,
*,
reason: str = "shutdown",
skip_sids: Collection[str] | None = None,
) -> None:
"""Finalize every server session except the ones named in ``skip_sids``.
``skip_sids`` carries the sessions whose turn is still live, which must
not spend their one-shot ``_finalize_session`` while running.
"""
try:
from tui_gateway import server
except Exception:
return
for session in list(getattr(server, "_sessions", {}).values()):
skip = set(skip_sids or ())
for sid, session in list(getattr(server, "_sessions", {}).items()):
if sid in skip:
continue
try:
server._finalize_session(session, end_reason=f"compute_host_{reason}")
except Exception:
@@ -262,15 +289,28 @@ class ComputeHost:
self._sessions[sid] = HostSession(sid=sid, agent=SpikeAgent(sid, list(history)))
self.emit({"type": "session.seeded", "sid": sid, "request_id": frame.get("request_id")})
def _track_turn_future(self, future: concurrent.futures.Future, sid: str) -> None:
"""Register an in-flight turn against the session running it.
The callback has to remove the entry under the lock — a bare
``dict.pop`` bound method is not the drop-in ``set.discard`` was — or
the mapping grows for the life of the host.
"""
with self._turn_futures_lock:
self._turn_futures[future] = sid
future.add_done_callback(self._untrack_turn_future)
def _untrack_turn_future(self, future: concurrent.futures.Future) -> None:
with self._turn_futures_lock:
self._turn_futures.pop(future, None)
def _handle_turn_start(self, frame: dict[str, Any]) -> None:
sid = str(frame.get("sid") or "")
if sid in self._sessions:
self._handle_spike_turn_start(frame)
return
future = self._executor.submit(self._run_real_turn, dict(frame))
with self._turn_futures_lock:
self._turn_futures.add(future)
future.add_done_callback(self._turn_futures.discard)
self._track_turn_future(future, sid)
def _handle_spike_turn_start(self, frame: dict[str, Any]) -> None:
sid = str(frame.get("sid") or "")
@@ -284,9 +324,7 @@ class ComputeHost:
return
session.running = True
future = self._executor.submit(self._run_spike_turn, session, dict(frame))
with self._turn_futures_lock:
self._turn_futures.add(future)
future.add_done_callback(self._turn_futures.discard)
self._track_turn_future(future, sid)
def _handle_interrupt(self, frame: dict[str, Any]) -> None:
sid = str(frame.get("sid") or "")