fix(tui): settle session close against active turns
This commit is contained in:
@@ -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
|
||||
|
||||
+39
-5
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user