fix(tui): avoid saving empty sessions

This commit is contained in:
Xubin Ren
2026-08-24 10:40:16 +08:00
parent 8344066696
commit d50a2fab32
4 changed files with 108 additions and 7 deletions
+2 -2
View File
@@ -906,7 +906,7 @@ class WebSocketChannel(BaseChannel):
) )
if scope is None: if scope is None:
return return
self._workspaces.persist_scope(new_id, scope) self._workspaces.stage_scope(new_id, scope)
self._attach(connection, new_id) self._attach(connection, new_id)
await self._send_event( await self._send_event(
connection, connection,
@@ -1023,7 +1023,7 @@ class WebSocketChannel(BaseChannel):
) )
if scope is None: if scope is None:
return return
self._workspaces.persist_scope(cid, scope) self._workspaces.stage_scope(cid, scope)
# Other clients on the same gateway only need an invalidation; they # Other clients on the same gateway only need an invalidation; they
# can reload the authoritative session row without receiving a # can reload the authoritative session row without receiving a
# local project path that belongs to another connection. # local project path that belongs to another connection.
@@ -1511,6 +1511,7 @@ async def test_webui_message_scope_inherits_persisted_session_scope(
}, },
}, },
) )
assert sessions.list_sessions() == []
await channel._dispatch_envelope( await channel._dispatch_envelope(
conn, conn,
"webui-client", "webui-client",
@@ -1524,6 +1525,44 @@ async def test_webui_message_scope_inherits_persisted_session_scope(
} }
@pytest.mark.asyncio
async def test_new_chat_without_message_does_not_create_session(
bus: MagicMock,
tmp_path,
) -> None:
sessions = SessionManager(tmp_path / "sessions")
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"], "host": "127.0.0.1"},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
conn = AsyncMock()
conn.remote_address = ("127.0.0.1", 50123)
await channel._dispatch_envelope(
conn,
"tui-client",
{
"type": "new_chat",
"workspace_scope": {
"project_path": str(tmp_path),
"access_mode": "full",
},
},
)
attached = json.loads(conn.send.await_args_list[0].args[0])
assert attached["event"] == "attached"
assert sessions.list_sessions() == []
assert channel._workspaces.scope_for_session_key(
f"websocket:{attached['chat_id']}"
).access_mode == "full"
await channel._cleanup_connection(conn)
assert sessions.list_sessions() == []
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_workspace_scope_change_invalidates_other_attached_clients( async def test_workspace_scope_change_invalidates_other_attached_clients(
bus: MagicMock, bus: MagicMock,
@@ -1730,6 +1769,10 @@ async def test_webui_set_workspace_scope_rejects_running_chat(bus: MagicMock, tm
}, },
}, },
) )
channel._workspaces.persist_scope(
"chat-running",
channel._workspaces.scope_for_session_key("websocket:chat-running"),
)
conn.send.reset_mock() conn.send.reset_mock()
wth._WEBSOCKET_TURN_WALL_STARTED_AT["chat-running"] = 123.0 wth._WEBSOCKET_TURN_WALL_STARTED_AT["chat-running"] = 123.0
@@ -1796,6 +1839,13 @@ async def test_remote_webui_scope_allows_access_reduction(
payload = json.loads(conn.send.await_args.args[0]) payload = json.loads(conn.send.await_args.args[0])
assert payload["event"] == "session_updated" assert payload["event"] == "session_updated"
assert payload["workspace_scope"]["access_mode"] == "restricted" assert payload["workspace_scope"]["access_mode"] == "restricted"
assert sessions.list_sessions() == []
await channel._dispatch_envelope(
conn,
"webui-client",
{"type": "message", "chat_id": "chat-remote", "content": "hello", "webui": True},
)
saved = sessions.read_session_file("websocket:chat-remote") saved = sessions.read_session_file("websocket:chat-remote")
assert saved["metadata"]["workspace_scope"] == { assert saved["metadata"]["workspace_scope"] == {
"project_path": str(default_workspace.resolve()), "project_path": str(default_workspace.resolve()),
@@ -1865,8 +1915,10 @@ async def test_remote_access_reduction_rejects_stale_in_flight_message_scope(
release_hydrate.set() release_hydrate.set()
await message_task await message_task
saved = sessions.read_session_file(f"websocket:{chat_id}") assert sessions.read_session_file(f"websocket:{chat_id}") is None
assert saved["metadata"]["workspace_scope"]["access_mode"] == "restricted" assert channel._workspaces.scope_for_session_key(
f"websocket:{chat_id}"
).access_mode == "restricted"
payload = json.loads(message_conn.send.await_args.args[0]) payload = json.loads(message_conn.send.await_args.args[0])
assert payload["event"] == "error" assert payload["event"] == "error"
assert payload["detail"] == "workspace_scope_rejected" assert payload["detail"] == "workspace_scope_rejected"
@@ -1954,8 +2006,10 @@ async def test_native_webui_scope_allows_custom_scope_without_loopback(
assert payload["workspace_scope"]["restrict_to_workspace"] is False assert payload["workspace_scope"]["restrict_to_workspace"] is False
assert payload["workspace_scope"]["sandbox_status"]["restrict_to_workspace"] is False assert payload["workspace_scope"]["sandbox_status"]["restrict_to_workspace"] is False
assert payload["workspace_scope"]["sandbox_status"]["workspace_root"] == str(project.resolve()) assert payload["workspace_scope"]["sandbox_status"]["workspace_root"] == str(project.resolve())
saved = sessions.read_session_file("websocket:chat-native") assert sessions.read_session_file("websocket:chat-native") is None
assert saved["metadata"]["workspace_scope"] == { assert channel._workspaces.scope_for_session_key(
"websocket:chat-native"
).metadata() == {
"project_path": str(project.resolve()), "project_path": str(project.resolve()),
"access_mode": "full", "access_mode": "full",
} }
+24 -1
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import json import json
import os import os
import time import time
from collections import OrderedDict
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
@@ -28,6 +29,7 @@ _MAX_STATE_FILE_BYTES = 128 * 1024
_DEFAULT_ACCESS_MODES = {"default", "full"} _DEFAULT_ACCESS_MODES = {"default", "full"}
_LEGACY_RESTRICTED_DEFAULT_ACCESS_MODE = "restricted" _LEGACY_RESTRICTED_DEFAULT_ACCESS_MODE = "restricted"
_WEBUI_SCOPE_CHANNEL = "websocket" _WEBUI_SCOPE_CHANNEL = "websocket"
_MAX_DRAFT_SCOPES = 128
def _scope_change_is_non_escalating(current: WorkspaceScope, requested: WorkspaceScope) -> bool: def _scope_change_is_non_escalating(current: WorkspaceScope, requested: WorkspaceScope) -> bool:
@@ -186,6 +188,7 @@ class WebUIWorkspaceController:
self._sessions = session_manager self._sessions = session_manager
self._default_workspace = default_workspace self._default_workspace = default_workspace
self._default_restrict_to_workspace = default_restrict_to_workspace self._default_restrict_to_workspace = default_restrict_to_workspace
self._draft_scopes: OrderedDict[str, WorkspaceScope] = OrderedDict()
def default_scope(self) -> WorkspaceScope: def default_scope(self) -> WorkspaceScope:
return default_scope_for_webui( return default_scope_for_webui(
@@ -230,6 +233,10 @@ class WebUIWorkspaceController:
return self._scope_from_metadata_value(raw_scope, default_scope=default_scope) return self._scope_from_metadata_value(raw_scope, default_scope=default_scope)
def scope_for_session_key(self, session_key: str) -> WorkspaceScope: def scope_for_session_key(self, session_key: str) -> WorkspaceScope:
draft = self._draft_scopes.get(session_key)
if draft is not None:
self._draft_scopes.move_to_end(session_key)
return draft
if self._sessions is None: if self._sessions is None:
return self.default_scope() return self.default_scope()
data = self._sessions.read_session_metadata(session_key) data = self._sessions.read_session_metadata(session_key)
@@ -328,8 +335,24 @@ class WebUIWorkspaceController:
return scope return scope
def persist_scope(self, chat_id: str, scope: WorkspaceScope) -> None: def persist_scope(self, chat_id: str, scope: WorkspaceScope) -> None:
session_key = f"websocket:{chat_id}"
self._draft_scopes.pop(session_key, None)
if self._sessions is not None: if self._sessions is not None:
session = self._sessions.get_or_create(f"websocket:{chat_id}") session = self._sessions.get_or_create(session_key)
session.metadata["webui"] = True session.metadata["webui"] = True
session.metadata[WORKSPACE_SCOPE_METADATA_KEY] = scope.metadata() session.metadata[WORKSPACE_SCOPE_METADATA_KEY] = scope.metadata()
self._sessions.save(session) self._sessions.save(session)
def stage_scope(self, chat_id: str, scope: WorkspaceScope) -> None:
"""Keep a new chat's scope transient until its first accepted message."""
session_key = f"websocket:{chat_id}"
if (
self._sessions is not None
and self._sessions.read_session_metadata(session_key) is not None
):
self.persist_scope(chat_id, scope)
return
self._draft_scopes[session_key] = scope
self._draft_scopes.move_to_end(session_key)
while len(self._draft_scopes) > _MAX_DRAFT_SCOPES:
self._draft_scopes.popitem(last=False)
+24
View File
@@ -234,6 +234,30 @@ def test_scope_for_session_key_reads_metadata_without_full_history(
assert scope.access_mode == "full" assert scope.access_mode == "full"
def test_new_chat_scope_is_persisted_only_after_first_message(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui")
default = tmp_path / "default"
project = tmp_path / "project"
default.mkdir()
project.mkdir()
sessions = SessionManager(tmp_path / "sessions")
controller = WebUIWorkspaceController(
session_manager=sessions,
default_workspace=default,
default_restrict_to_workspace=True,
)
scope = default_workspace_scope(project, restrict_to_workspace=False)
controller.stage_scope("draft-chat", scope)
assert sessions.list_sessions() == []
assert controller.scope_for_session_key("websocket:draft-chat") == scope
controller.persist_scope("draft-chat", scope)
assert [item["key"] for item in sessions.list_sessions()] == ["websocket:draft-chat"]
def test_scope_for_session_key_always_reads_the_active_store(tmp_path, monkeypatch) -> None: def test_scope_for_session_key_always_reads_the_active_store(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui") monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui")
default = tmp_path / "default" default = tmp_path / "default"