mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-04 02:01:48 +03:00
refactor: simplify cross-session messaging
This commit is contained in:
@@ -5,7 +5,6 @@ from __future__ import annotations
|
||||
import json
|
||||
from contextlib import AbstractContextManager
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -14,9 +13,8 @@ from nanobot.agent.tools.loader import ToolLoader
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.agent.tools.sessions import ReadSessionTool, SearchSessionsTool
|
||||
from nanobot.runtime_context import RuntimeContextBlock, append_runtime_context
|
||||
from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.session.session_handles import SessionHandleDirectory
|
||||
from nanobot.session.session_handles import session_handle_for_key
|
||||
from nanobot.webui.transcript import append_transcript_object
|
||||
|
||||
|
||||
@@ -43,14 +41,11 @@ def _decode(value: str) -> dict[str, object]:
|
||||
|
||||
def _webui_request(
|
||||
session_key: str = "websocket:current",
|
||||
*,
|
||||
workspace: Path | None = None,
|
||||
) -> AbstractContextManager[RequestContext]:
|
||||
return request_context(RequestContext(
|
||||
channel="websocket",
|
||||
chat_id=session_key.removeprefix("websocket:"),
|
||||
session_key=session_key,
|
||||
workspace=workspace,
|
||||
))
|
||||
|
||||
|
||||
@@ -141,7 +136,10 @@ async def test_search_sessions_has_no_hidden_content_scan_cutoff(tmp_path, monke
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_sessions_ranks_titles_before_message_matches(tmp_path):
|
||||
async def test_search_sessions_ranks_titles_before_message_matches(tmp_path, monkeypatch):
|
||||
webui_dir = tmp_path / "webui"
|
||||
monkeypatch.setattr("nanobot.webui.transcript.get_webui_dir", lambda: webui_dir)
|
||||
monkeypatch.setattr("nanobot.webui.session_list_index.get_webui_dir", lambda: webui_dir)
|
||||
manager = SessionManager(tmp_path)
|
||||
_save_session(
|
||||
manager,
|
||||
@@ -237,62 +235,6 @@ async def test_read_session_filters_by_query_and_returns_recent_matches(tmp_path
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_session_accepts_a_cross_workspace_session_handle(tmp_path: Path) -> None:
|
||||
project = tmp_path / "project"
|
||||
other = tmp_path / "other"
|
||||
project.mkdir()
|
||||
other.mkdir()
|
||||
manager = SessionManager(tmp_path / "state")
|
||||
for key, workspace, content in (
|
||||
("websocket:current", project, "current"),
|
||||
("websocket:handle", project, "handle answer"),
|
||||
("websocket:other", other, "other answer"),
|
||||
):
|
||||
_save_session(
|
||||
manager,
|
||||
key,
|
||||
title=key,
|
||||
messages=[{"role": "assistant", "content": content}],
|
||||
)
|
||||
session = manager.get_or_create(key)
|
||||
session.metadata.update({
|
||||
"webui": True,
|
||||
WORKSPACE_SCOPE_METADATA_KEY: {
|
||||
"project_path": str(workspace),
|
||||
"access_mode": "restricted",
|
||||
},
|
||||
})
|
||||
manager.save(session)
|
||||
directory = SessionHandleDirectory(manager)
|
||||
handles = directory.ensure_many([
|
||||
"websocket:current",
|
||||
"websocket:handle",
|
||||
"websocket:other",
|
||||
])
|
||||
current = handles["websocket:current"]
|
||||
handle = handles["websocket:handle"]
|
||||
outside = handles["websocket:other"]
|
||||
tool = ReadSessionTool(manager)
|
||||
|
||||
with _webui_request(workspace=project):
|
||||
result = _decode(await tool.execute(session_key=f"@{handle.name}"))
|
||||
outside_result = _decode(await tool.execute(session_key=f"@{outside.name}"))
|
||||
self_read = await tool.execute(session_key=f"@{current.name}")
|
||||
|
||||
assert result["handle"] == f"@{handle.name}"
|
||||
assert "session_key" not in result
|
||||
assert "session_ref" not in result
|
||||
assert "title" not in result
|
||||
assert "websocket:" not in json.dumps(result)
|
||||
assert result["messages"][0]["content"] == "handle answer"
|
||||
assert outside_result["handle"] == f"@{outside.name}"
|
||||
assert outside_result["messages"][0]["content"] == "other answer"
|
||||
assert "websocket:" not in json.dumps(outside_result)
|
||||
assert self_read.is_error and f"@{current.name}" in str(self_read)
|
||||
assert "websocket:" not in str(self_read)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_session_reports_invalid_requests(tmp_path):
|
||||
with _webui_request():
|
||||
@@ -330,22 +272,15 @@ async def test_session_tools_read_persisted_sessions_from_any_channel(tmp_path):
|
||||
messages=[{"role": "user", "content": "needle"}],
|
||||
)
|
||||
tools = SearchSessionsTool(manager), ReadSessionTool(manager)
|
||||
slack_handle = SessionHandleDirectory(manager).ensure_many(["slack:history"])[
|
||||
"slack:history"
|
||||
]
|
||||
|
||||
with request_context(RequestContext(
|
||||
channel="telegram",
|
||||
chat_id="external",
|
||||
session_key="telegram:external",
|
||||
workspace=tmp_path,
|
||||
)):
|
||||
search = _decode(await tools[0].execute(query="needle"))
|
||||
websocket_read = _decode(await tools[1].execute(session_key="websocket:visible"))
|
||||
slack_read = _decode(await tools[1].execute(session_key="slack:history"))
|
||||
slack_handle_read = _decode(
|
||||
await tools[1].execute(session_key=f"@{slack_handle.name}")
|
||||
)
|
||||
current_read = await tools[1].execute(session_key="telegram:external")
|
||||
|
||||
assert {row["session_key"] for row in search["results"]} == {
|
||||
@@ -354,13 +289,35 @@ async def test_session_tools_read_persisted_sessions_from_any_channel(tmp_path):
|
||||
}
|
||||
assert websocket_read["session_key"] == "websocket:visible"
|
||||
assert slack_read["session_key"] == "slack:history"
|
||||
assert slack_handle_read["handle"] == f"@{slack_handle.name}"
|
||||
assert slack_handle_read["messages"][0]["content"] == "needle"
|
||||
assert current_read.is_error and "session not found" in str(current_read)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_tools_work_without_request_context(tmp_path):
|
||||
async def test_read_session_accepts_a_persisted_session_handle(tmp_path):
|
||||
manager = SessionManager(tmp_path)
|
||||
_save_session(
|
||||
manager,
|
||||
"slack:history",
|
||||
title="Slack history",
|
||||
messages=[{"role": "user", "content": "needle"}],
|
||||
)
|
||||
handle = session_handle_for_key("slack:history")
|
||||
|
||||
with _webui_request():
|
||||
result = _decode(await ReadSessionTool(manager).execute(
|
||||
session_key=f"@{handle.name}",
|
||||
))
|
||||
|
||||
assert result["handle"] == f"@{handle.name}"
|
||||
assert [message["content"] for message in result["messages"]] == ["needle"]
|
||||
assert "session_key" not in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_tools_work_without_request_context(tmp_path, monkeypatch):
|
||||
webui_dir = tmp_path / "webui"
|
||||
monkeypatch.setattr("nanobot.webui.transcript.get_webui_dir", lambda: webui_dir)
|
||||
monkeypatch.setattr("nanobot.webui.session_list_index.get_webui_dir", lambda: webui_dir)
|
||||
manager = SessionManager(tmp_path)
|
||||
_save_session(
|
||||
manager,
|
||||
|
||||
Reference in New Issue
Block a user