diff --git a/tests/tui_gateway/test_compute_host_phase1.py b/tests/tui_gateway/test_compute_host_phase1.py index 07e1efd564..5b9ccb68c5 100644 --- a/tests/tui_gateway/test_compute_host_phase1.py +++ b/tests/tui_gateway/test_compute_host_phase1.py @@ -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 diff --git a/tui_gateway/compute_host.py b/tui_gateway/compute_host.py index ed951fa97a..706be07b08 100644 --- a/tui_gateway/compute_host.py +++ b/tui_gateway/compute_host.py @@ -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 "")