mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 00:03:01 +03:00
fix(agent): bound per-session file state
This commit is contained in:
@@ -76,6 +76,7 @@ from nanobot.session.goal_state import (
|
|||||||
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
|
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
|
||||||
from nanobot.session.keys import UNIFIED_SESSION_KEY, remember_last_channel
|
from nanobot.session.keys import UNIFIED_SESSION_KEY, remember_last_channel
|
||||||
from nanobot.session.manager import (
|
from nanobot.session.manager import (
|
||||||
|
SESSION_CACHE_MAX_SIZE,
|
||||||
Session,
|
Session,
|
||||||
SessionManager,
|
SessionManager,
|
||||||
replay_max_messages_for_context,
|
replay_max_messages_for_context,
|
||||||
@@ -380,7 +381,7 @@ class AgentLoop:
|
|||||||
self.tools = tool_registry if tool_registry is not None else ToolRegistry()
|
self.tools = tool_registry if tool_registry is not None else ToolRegistry()
|
||||||
# One file-read/write tracker per logical session. The tool registry is
|
# One file-read/write tracker per logical session. The tool registry is
|
||||||
# shared by this loop, so tools resolve the active state via contextvars.
|
# shared by this loop, so tools resolve the active state via contextvars.
|
||||||
self._file_state_store = FileStateStore()
|
self._file_state_store = FileStateStore(max_sessions=SESSION_CACHE_MAX_SIZE)
|
||||||
self._exec_session_manager = ExecSessionManager()
|
self._exec_session_manager = ExecSessionManager()
|
||||||
self.runner = AgentRunner()
|
self.runner = AgentRunner()
|
||||||
self.subagents = SubagentManager(
|
self.subagents = SubagentManager(
|
||||||
@@ -818,8 +819,13 @@ class AgentLoop:
|
|||||||
self.sessions.invalidate(key)
|
self.sessions.invalidate(key)
|
||||||
await self._cancel_active_tasks(key)
|
await self._cancel_active_tasks(key)
|
||||||
finally:
|
finally:
|
||||||
|
self.discard_session_file_state(key)
|
||||||
self._discarding_sessions.discard(key)
|
self._discarding_sessions.discard(key)
|
||||||
|
|
||||||
|
def discard_session_file_state(self, key: str) -> None:
|
||||||
|
"""Forget ephemeral file-read state for a reset or removed session."""
|
||||||
|
self._file_state_store.discard(key)
|
||||||
|
|
||||||
def _effective_session_key(self, msg: InboundMessage) -> str:
|
def _effective_session_key(self, msg: InboundMessage) -> str:
|
||||||
"""Return the session key used for task routing and mid-turn injections."""
|
"""Return the session key used for task routing and mid-turn injections."""
|
||||||
if self._unified_session and not msg.session_key_override:
|
if self._unified_session and not msg.session_key_override:
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import hashlib
|
import hashlib
|
||||||
import os
|
import os
|
||||||
|
from collections import OrderedDict
|
||||||
from contextvars import ContextVar, Token
|
from contextvars import ContextVar, Token
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -135,21 +136,30 @@ class FileStates:
|
|||||||
|
|
||||||
|
|
||||||
class FileStateStore:
|
class FileStateStore:
|
||||||
"""Lookup table for per-session file read/write state."""
|
"""Bounded lookup table for per-session file read/write state."""
|
||||||
|
|
||||||
__slots__ = ("_states_by_key",)
|
__slots__ = ("_max_sessions", "_states_by_key")
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self, *, max_sessions: int = 128) -> None:
|
||||||
self._states_by_key: dict[str, FileStates] = {}
|
if max_sessions <= 0:
|
||||||
|
raise ValueError("max_sessions must be positive")
|
||||||
|
self._max_sessions = max_sessions
|
||||||
|
self._states_by_key: OrderedDict[str, FileStates] = OrderedDict()
|
||||||
|
|
||||||
def for_session(self, session_key: str | None) -> FileStates:
|
def for_session(self, session_key: str | None) -> FileStates:
|
||||||
key = session_key or "__default__"
|
key = session_key or "__default__"
|
||||||
states = self._states_by_key.get(key)
|
states = self._states_by_key.pop(key, None)
|
||||||
if states is None:
|
if states is None:
|
||||||
states = FileStates()
|
states = FileStates()
|
||||||
self._states_by_key[key] = states
|
self._states_by_key[key] = states
|
||||||
|
while len(self._states_by_key) > self._max_sessions:
|
||||||
|
self._states_by_key.popitem(last=False)
|
||||||
return states
|
return states
|
||||||
|
|
||||||
|
def discard(self, session_key: str | None) -> None:
|
||||||
|
"""Forget file state when a session is reset or removed."""
|
||||||
|
self._states_by_key.pop(session_key or "__default__", None)
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
self._states_by_key.clear()
|
self._states_by_key.clear()
|
||||||
|
|
||||||
|
|||||||
@@ -302,6 +302,7 @@ async def cmd_new(ctx: CommandContext) -> OutboundMessage:
|
|||||||
"""Stop active task and start a fresh session."""
|
"""Stop active task and start a fresh session."""
|
||||||
loop = ctx.loop
|
loop = ctx.loop
|
||||||
await loop._cancel_active_tasks(ctx.key) # pyright: ignore[reportPrivateUsage]
|
await loop._cancel_active_tasks(ctx.key) # pyright: ignore[reportPrivateUsage]
|
||||||
|
loop.discard_session_file_state(ctx.key)
|
||||||
session = ctx.session or loop.sessions.get_or_create(ctx.key)
|
session = ctx.session or loop.sessions.get_or_create(ctx.key)
|
||||||
snapshot = session.messages[session.last_consolidated:]
|
snapshot = session.messages[session.last_consolidated:]
|
||||||
runtime = None
|
runtime = None
|
||||||
|
|||||||
@@ -138,6 +138,7 @@ class SessionClient:
|
|||||||
|
|
||||||
def clear(self, session_key: str) -> SessionSnapshot:
|
def clear(self, session_key: str) -> SessionSnapshot:
|
||||||
"""Clear one session and persist the empty session."""
|
"""Clear one session and persist the empty session."""
|
||||||
|
self._loop.discard_session_file_state(session_key)
|
||||||
session = self._loop.sessions.get_or_create(session_key)
|
session = self._loop.sessions.get_or_create(session_key)
|
||||||
session.clear()
|
session.clear()
|
||||||
self._loop.sessions.save(session)
|
self._loop.sessions.save(session)
|
||||||
@@ -145,6 +146,7 @@ class SessionClient:
|
|||||||
|
|
||||||
def delete(self, session_key: str) -> bool:
|
def delete(self, session_key: str) -> bool:
|
||||||
"""Delete one session from disk and cache."""
|
"""Delete one session from disk and cache."""
|
||||||
|
self._loop.discard_session_file_state(session_key)
|
||||||
return self._loop.sessions.delete_session(session_key)
|
return self._loop.sessions.delete_session(session_key)
|
||||||
|
|
||||||
def flush(self) -> int:
|
def flush(self) -> int:
|
||||||
|
|||||||
@@ -131,6 +131,7 @@ async def test_session_discard_control_cancels_active_turn(tmp_path, monkeypatch
|
|||||||
terminate_exec_sessions,
|
terminate_exec_sessions,
|
||||||
)
|
)
|
||||||
key = "websocket:transient-cancelled"
|
key = "websocket:transient-cancelled"
|
||||||
|
previous_file_state = loop._file_state_store.for_session(key)
|
||||||
loop.sessions.get_or_create_transient(
|
loop.sessions.get_or_create_transient(
|
||||||
key,
|
key,
|
||||||
disabled_tools={"create_goal", "update_goal", "spawn", "cron"},
|
disabled_tools={"create_goal", "update_goal", "spawn", "cron"},
|
||||||
@@ -157,6 +158,7 @@ async def test_session_discard_control_cancels_active_turn(tmp_path, monkeypatch
|
|||||||
await asyncio.wait_for(active_task, timeout=2)
|
await asyncio.wait_for(active_task, timeout=2)
|
||||||
await asyncio.wait_for(wait_for_discard(key), timeout=2)
|
await asyncio.wait_for(wait_for_discard(key), timeout=2)
|
||||||
assert loop.sessions.get_cached(key) is None
|
assert loop.sessions.get_cached(key) is None
|
||||||
|
assert loop._file_state_store.for_session(key) is not previous_file_state
|
||||||
terminate_exec_sessions.assert_awaited_once_with(key)
|
terminate_exec_sessions.assert_awaited_once_with(key)
|
||||||
|
|
||||||
loop.stop()
|
loop.stop()
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.agent.tools.file_state import FileStateStore
|
||||||
from nanobot.bus.events import InboundMessage
|
from nanobot.bus.events import InboundMessage
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.command.builtin import cmd_new, register_builtin_commands
|
from nanobot.command.builtin import cmd_new, register_builtin_commands
|
||||||
@@ -250,10 +251,16 @@ class TestCmdNewUnifiedSession:
|
|||||||
# asyncio.create_task(). Mirror that exactly so the coroutine is consumed
|
# asyncio.create_task(). Mirror that exactly so the coroutine is consumed
|
||||||
# and no RuntimeWarning is emitted.
|
# and no RuntimeWarning is emitted.
|
||||||
admitted_runtime = MagicMock(name="admitted_runtime")
|
admitted_runtime = MagicMock(name="admitted_runtime")
|
||||||
|
file_state_store = FileStateStore()
|
||||||
|
previous_file_state = file_state_store.for_session("unified:default")
|
||||||
|
tracked_file = tmp_path / "tracked.txt"
|
||||||
|
tracked_file.write_text("tracked", encoding="utf-8")
|
||||||
|
previous_file_state.record_read(tracked_file)
|
||||||
loop = SimpleNamespace(
|
loop = SimpleNamespace(
|
||||||
sessions=sessions,
|
sessions=sessions,
|
||||||
consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)),
|
consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)),
|
||||||
_cancel_active_tasks=AsyncMock(return_value=0),
|
_cancel_active_tasks=AsyncMock(return_value=0),
|
||||||
|
discard_session_file_state=file_state_store.discard,
|
||||||
llm_runtime=MagicMock(return_value=MagicMock()),
|
llm_runtime=MagicMock(return_value=MagicMock()),
|
||||||
schedule_background=lambda coro: asyncio.ensure_future(coro),
|
schedule_background=lambda coro: asyncio.ensure_future(coro),
|
||||||
)
|
)
|
||||||
@@ -278,6 +285,9 @@ class TestCmdNewUnifiedSession:
|
|||||||
sessions.invalidate("unified:default")
|
sessions.invalidate("unified:default")
|
||||||
reloaded = sessions.get_or_create("unified:default")
|
reloaded = sessions.get_or_create("unified:default")
|
||||||
assert reloaded.messages == []
|
assert reloaded.messages == []
|
||||||
|
reset_file_state = file_state_store.for_session("unified:default")
|
||||||
|
assert reset_file_state is not previous_file_state
|
||||||
|
assert reset_file_state.is_unchanged(tracked_file) is False
|
||||||
loop.consolidator.archive.assert_called_once_with(
|
loop.consolidator.archive.assert_called_once_with(
|
||||||
expected_snapshot,
|
expected_snapshot,
|
||||||
runtime=admitted_runtime,
|
runtime=admitted_runtime,
|
||||||
@@ -302,6 +312,7 @@ class TestCmdNewUnifiedSession:
|
|||||||
sessions=sessions,
|
sessions=sessions,
|
||||||
consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)),
|
consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)),
|
||||||
_cancel_active_tasks=AsyncMock(return_value=0),
|
_cancel_active_tasks=AsyncMock(return_value=0),
|
||||||
|
discard_session_file_state=MagicMock(),
|
||||||
runtime_for_session=MagicMock(return_value=MagicMock()),
|
runtime_for_session=MagicMock(return_value=MagicMock()),
|
||||||
schedule_background=lambda coro: asyncio.ensure_future(coro),
|
schedule_background=lambda coro: asyncio.ensure_future(coro),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1500,10 +1500,15 @@ async def test_session_helpers_get_list_export_clear_delete_flush(tmp_path):
|
|||||||
exported.messages[0]["content"] = "mutated copy"
|
exported.messages[0]["content"] = "mutated copy"
|
||||||
assert bot.sessions.get("sdk:first").messages[0]["content"] == "hello"
|
assert bot.sessions.get("sdk:first").messages[0]["content"] == "hello"
|
||||||
|
|
||||||
|
state_before_clear = bot._loop._file_state_store.for_session("sdk:first")
|
||||||
cleared = bot.sessions.clear("sdk:first")
|
cleared = bot.sessions.clear("sdk:first")
|
||||||
assert cleared.messages == []
|
assert cleared.messages == []
|
||||||
|
state_after_clear = bot._loop._file_state_store.for_session("sdk:first")
|
||||||
|
assert state_after_clear is not state_before_clear
|
||||||
assert bot.sessions.flush() >= 1
|
assert bot.sessions.flush() >= 1
|
||||||
|
state_before_delete = state_after_clear
|
||||||
assert bot.sessions.delete("sdk:first") is True
|
assert bot.sessions.delete("sdk:first") is True
|
||||||
|
assert bot._loop._file_state_store.for_session("sdk:first") is not state_before_delete
|
||||||
assert bot.sessions.get("sdk:first") is None
|
assert bot.sessions.get("sdk:first") is None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.agent.tools.file_state import FileStateStore
|
||||||
|
|
||||||
|
|
||||||
|
def test_file_state_store_evicts_least_recently_used_session() -> None:
|
||||||
|
store = FileStateStore(max_sessions=2)
|
||||||
|
first = store.for_session("first")
|
||||||
|
second = store.for_session("second")
|
||||||
|
|
||||||
|
assert store.for_session("first") is first
|
||||||
|
store.for_session("third")
|
||||||
|
|
||||||
|
assert store.for_session("first") is first
|
||||||
|
assert store.for_session("second") is not second
|
||||||
|
|
||||||
|
|
||||||
|
def test_file_state_store_discards_reset_session() -> None:
|
||||||
|
store = FileStateStore()
|
||||||
|
previous = store.for_session("websocket:chat")
|
||||||
|
|
||||||
|
store.discard("websocket:chat")
|
||||||
|
|
||||||
|
assert store.for_session("websocket:chat") is not previous
|
||||||
|
|
||||||
|
|
||||||
|
def test_file_state_store_requires_positive_capacity() -> None:
|
||||||
|
with pytest.raises(ValueError, match="max_sessions must be positive"):
|
||||||
|
FileStateStore(max_sessions=0)
|
||||||
Reference in New Issue
Block a user