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:
@@ -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
|
||||
|
||||
@@ -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 "")
|
||||
|
||||
Reference in New Issue
Block a user