diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index d193f66c47..75a2163adf 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -4455,6 +4455,59 @@ def test_session_close_releases_resume_lock_before_slow_teardown(monkeypatch): assert response["result"] == {"closed": True} +def test_session_close_settles_active_turn_before_teardown(monkeypatch): + """Close must not tear down agent resources while their turn is unwinding.""" + turn_started = threading.Event() + release_turn = threading.Event() + teardown_started = threading.Event() + response = {} + + def _turn(): + turn_started.set() + assert release_turn.wait(timeout=2.0) + + def _teardown(_session, *, end_reason="tui_close"): + if end_reason == "tui_close": + teardown_started.set() + + session = _session() + run_thread = threading.Thread(target=_turn) + session["_run_thread"] = run_thread + server._sessions["settle-close"] = session + monkeypatch.setattr(server, "_teardown_session", _teardown) + monkeypatch.setattr( + server, "_TURN_SETTLE_BEFORE_CLOSE_SECONDS", 1.0, raising=False + ) + + close_thread = threading.Thread( + target=lambda: response.update( + server.handle_request( + { + "id": "close", + "method": "session.close", + "params": {"session_id": "settle-close"}, + } + ) + ) + ) + run_thread.start() + close_thread.start() + try: + assert turn_started.wait(timeout=1.0) + assert not teardown_started.wait(timeout=0.1) + release_turn.set() + close_thread.join(timeout=2.0) + finally: + release_turn.set() + run_thread.join(timeout=2.0) + close_thread.join(timeout=2.0) + server._sessions.pop("settle-close", None) + + assert not close_thread.is_alive() + assert teardown_started.is_set() + assert response["result"] == {"closed": True} + + def test_ws_orphan_reap_closes_worker_when_session_stays_detached(monkeypatch): """A detached WS session past its grace window has its slash_worker closed. @@ -6188,6 +6241,56 @@ class _RecordingAgent: return {"final_response": "", "messages": []} +def test_run_prompt_submit_rejects_worker_when_close_wins_publication( + monkeypatch, tmp_path +): + """A close claimed during message.start must prevent the worker from running.""" + _configure_immediate_prompt_run(monkeypatch, tmp_path, immediate_threads=False) + emit_entered = threading.Event() + release_emit = threading.Event() + dispatch_results = [] + turns = [] + popped = [] + sid = "close-wins-publication" + session = _session( + session_key="close-wins-publication-key", + agent=_RecordingAgent(turns), + running=True, + ) + + def _blocking_emit(event, *_args, **_kwargs): + if event == "message.start": + emit_entered.set() + assert release_emit.wait(timeout=2.0) + + monkeypatch.setattr(server, "_emit", _blocking_emit) + server._sessions[sid] = session + dispatch_thread = threading.Thread( + target=lambda: dispatch_results.append( + server._run_prompt_submit("rid", sid, session, "turn") + ) + ) + + try: + dispatch_thread.start() + assert emit_entered.wait(timeout=1.0) + popped.append(server._pop_session_by_id(sid)) + assert popped == [session] + release_emit.set() + dispatch_thread.join(timeout=2.0) + finally: + release_emit.set() + dispatch_thread.join(timeout=2.0) + run_thread = session.get("_run_thread") + if run_thread is not None and run_thread.is_alive(): + run_thread.join(timeout=2.0) + server._sessions.pop(sid, None) + + assert dispatch_results == [False] + assert session["running"] is False + assert turns == [] + + @pytest.mark.parametrize("exit_code", [0, 7]) def test_run_prompt_submit_requeues_foreign_completion( monkeypatch, tmp_path, exit_code diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 3116bd2338..b176628c90 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -178,6 +178,7 @@ try: except (ValueError, TypeError): _ws_orphan_reap_grace = 20.0 _WS_ORPHAN_REAP_GRACE_S = max(0.0, _ws_orphan_reap_grace) +_TURN_SETTLE_BEFORE_CLOSE_SECONDS = 5.0 _DETAIL_SECTION_NAMES = ("thinking", "tools", "subagents", "activity") _DETAIL_MODES = frozenset({"hidden", "collapsed", "expanded"}) @@ -964,7 +965,10 @@ def _pop_session_by_id(sid: str) -> dict | None: the global ``_session_resume_lock``. """ with _sessions_lock: - session = _sessions.pop(sid, None) + session = _sessions.get(sid) + if session is not None: + session["_closing"] = True + _sessions.pop(sid, None) if session is None: return None # The session is already out of _sessions here, so downstream teardown @@ -980,6 +984,22 @@ def _teardown_popped_session( """Finish a close after the caller has atomically detached the session.""" if session is None: return False + run_thread = session.get("_run_thread") + if ( + end_reason != "tui_shutdown" + and run_thread is not None + and run_thread is not threading.current_thread() + ): + try: + if run_thread.is_alive(): + run_thread.join(timeout=_TURN_SETTLE_BEFORE_CLOSE_SECONDS) + if run_thread.is_alive(): + logger.warning( + "session turn thread still alive after %.1fs teardown grace", + _TURN_SETTLE_BEFORE_CLOSE_SECONDS, + ) + except Exception: + logger.debug("failed waiting for session turn thread", exc_info=True) _teardown_session(session, end_reason=end_reason) return True @@ -10384,14 +10404,17 @@ def _run_prompt_submit( display_metadata: dict | None = None, image_paths: list[str] | None = None, queued_prompt_generation: int | None = None, -) -> None: +) -> bool: with session["history_lock"]: + if session.get("_closing"): + session["running"] = False + return False if ( queued_prompt_generation is not None and int(session.get("_queued_prompt_generation", 0)) != queued_prompt_generation ): session["running"] = False - return + return False if image_paths is None: images = list(session.get("attached_images", [])) session["attached_images"] = [] @@ -11317,8 +11340,19 @@ def _run_prompt_submit( ) run_thread = threading.Thread(target=run, daemon=True) - session["_run_thread"] = run_thread - run_thread.start() + with _sessions_lock: + registered = _sessions.get(sid) + can_start = ( + not session.get("_closing") + and (registered is None or registered is session) + ) + if can_start: + session["_run_thread"] = run_thread + run_thread.start() + if not can_start: + with session["history_lock"]: + session["running"] = False + return can_start # Byte-upload attach caps. 25 MB matches Anthropic's per-image limit; 50 MB / 25