Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 80 additions & 0 deletions tests/test_tui_gateway_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,22 @@ def test_dispatch_rejects_non_object_params():
}


def test_session_create_marks_close_on_disconnect():
resp = server.handle_request(
{
"id": "1",
"method": "session.create",
"params": {"close_on_disconnect": "true"},
}
)

sid = resp["result"]["session_id"]
try:
assert server._sessions[sid]["close_on_disconnect"] is True
finally:
server._sessions.pop(sid, None)


def test_voice_toggle_returns_configured_record_key(monkeypatch):
monkeypatch.setattr(
server,
Expand Down Expand Up @@ -755,6 +771,70 @@ def test_session_close_commits_memory_and_fires_finalize_hook(monkeypatch):
server._sessions.pop("sid", None)


def test_close_sessions_for_transport_closes_sidecar_worker(monkeypatch):
closed_workers: list[str] = []
finalized: list[str] = []
closed_agents: list[str] = []
unregistered: list[str] = []
transport = object()

class _FakeWorker:
def close(self):
closed_workers.append("worker")

class _FakeAgent:
def close(self):
closed_agents.append("agent")

monkeypatch.setattr(
server,
"_finalize_session",
lambda session, end_reason="tui_close": finalized.append(end_reason),
)
import tools.approval as _approval

monkeypatch.setattr(
_approval,
"unregister_gateway_notify",
lambda key: unregistered.append(key),
)

server._sessions["sidecar"] = _session(
agent=_FakeAgent(),
session_key="sidecar-key",
slash_worker=_FakeWorker(),
transport=transport,
close_on_disconnect=True,
)

try:
server._close_sessions_for_transport(transport, end_reason="ws_disconnect")

assert "sidecar" not in server._sessions
assert finalized == ["ws_disconnect"]
assert closed_agents == ["agent"]
assert closed_workers == ["worker"]
assert unregistered == ["sidecar-key"]
finally:
server._sessions.pop("sidecar", None)


def test_close_sessions_for_transport_preserves_normal_session():
transport = object()
server._sessions["normal"] = _session(
transport=transport,
close_on_disconnect=False,
)

try:
server._close_sessions_for_transport(transport, end_reason="ws_disconnect")

assert "normal" in server._sessions
assert server._sessions["normal"]["transport"] is server._stdio_transport
finally:
server._sessions.pop("normal", None)


def test_init_session_fires_reset_hook(monkeypatch):
hooks = []

Expand Down
87 changes: 56 additions & 31 deletions tui_gateway/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -318,15 +318,57 @@ def _finalize_session(session: dict | None, end_reason: str = "tui_close") -> No
pass


def _coerce_close_on_disconnect(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
if value is None:
return False
return bool(value)


def _close_session_by_id(sid: str, *, end_reason: str = "tui_close") -> bool:
session = _sessions.pop(sid, None)
if not session:
return False
_finalize_session(session, end_reason=end_reason)
try:
from tools.approval import unregister_gateway_notify

unregister_gateway_notify(session["session_key"])
except Exception:
pass
try:
agent = session.get("agent")
if agent and hasattr(agent, "close"):
agent.close()
except Exception:
pass
try:
worker = session.get("slash_worker")
if worker:
worker.close()
except Exception:
pass
return True


def _close_sessions_for_transport(
transport: Any, *, end_reason: str = "ws_disconnect"
) -> None:
for sid, session in list(_sessions.items()):
if session.get("transport") is not transport:
continue
if session.get("close_on_disconnect"):
_close_session_by_id(sid, end_reason=end_reason)
else:
session["transport"] = _stdio_transport


def _shutdown_sessions() -> None:
for session in list(_sessions.values()):
_finalize_session(session, end_reason="tui_shutdown")
try:
worker = session.get("slash_worker")
if worker:
worker.close()
except Exception:
pass
for sid in list(_sessions.keys()):
_close_session_by_id(sid, end_reason="tui_shutdown")


atexit.register(_shutdown_sessions)
Expand Down Expand Up @@ -1896,6 +1938,7 @@ def _init_session(sid: str, key: str, agent, history: list, cols: int = 80):
"tool_progress_mode": _load_tool_progress_mode(),
"edit_snapshots": {},
"tool_started_at": {},
"close_on_disconnect": False,
# Pin async event emissions to whichever transport created the
# session (stdio for Ink, JSON-RPC WS for the dashboard sidebar).
"transport": current_transport() or _stdio_transport,
Expand Down Expand Up @@ -2065,6 +2108,9 @@ def _(rid, params: dict) -> dict:
sid = uuid.uuid4().hex[:8]
key = _new_session_key()
cols = int(params.get("cols", 80))
close_on_disconnect = _coerce_close_on_disconnect(
params.get("close_on_disconnect")
)
_enable_gateway_prompts()

ready = threading.Event()
Expand All @@ -2087,6 +2133,7 @@ def _(rid, params: dict) -> dict:
"slash_worker": None,
"tool_progress_mode": _load_tool_progress_mode(),
"tool_started_at": {},
"close_on_disconnect": close_on_disconnect,
"transport": current_transport() or _stdio_transport,
}

Expand Down Expand Up @@ -2605,29 +2652,7 @@ def _(rid, params: dict) -> dict:
@method("session.close")
def _(rid, params: dict) -> dict:
sid = params.get("session_id", "")
session = _sessions.pop(sid, None)
if not session:
return _ok(rid, {"closed": False})
_finalize_session(session)
try:
from tools.approval import unregister_gateway_notify

unregister_gateway_notify(session["session_key"])
except Exception:
pass
try:
agent = session.get("agent")
if agent and hasattr(agent, "close"):
agent.close()
except Exception:
pass
try:
worker = session.get("slash_worker")
if worker:
worker.close()
except Exception:
pass
return _ok(rid, {"closed": True})
return _ok(rid, {"closed": _close_session_by_id(sid)})


@method("session.branch")
Expand Down
9 changes: 4 additions & 5 deletions tui_gateway/ws.py
Original file line number Diff line number Diff line change
Expand Up @@ -162,11 +162,10 @@ async def handle_ws(ws: Any) -> None:
finally:
transport.close()

# Detach the transport from any sessions it owned so later emits
# fall back to stdio instead of crashing into a closed socket.
for _, sess in list(server._sessions.items()):
if sess.get("transport") is transport:
sess["transport"] = server._stdio_transport
# Preserve the historical "session survives reconnect" behavior for
# normal TUI sessions, but eagerly close explicit sidecar sessions
# whose slash worker should not outlive this websocket.
server._close_sessions_for_transport(transport, end_reason="ws_disconnect")

try:
await ws.close()
Expand Down
4 changes: 3 additions & 1 deletion web/src/components/ChatSidebar.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -119,7 +119,9 @@ export function ChatSidebar({ channel, className }: ChatSidebarProps) {
if (cancelled) {
return;
}
return gw.request<{ session_id: string }>("session.create", {});
return gw.request<{ session_id: string }>("session.create", {
close_on_disconnect: true,
});
})
.then((created) => {
if (cancelled || !created?.session_id) {
Expand Down