fix(agent): bound per-session file state

This commit is contained in:
yu-xin-c
2026-08-15 23:49:20 +08:00
committed by Xubin Ren
parent ecef2b055d
commit 42afebb0cb
8 changed files with 73 additions and 7 deletions
+7 -1
View File
@@ -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:
+15 -5
View File
@@ -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()
+1
View File
@@ -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
+2
View File
@@ -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:
+2
View File
@@ -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()
+11
View File
@@ -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),
) )
+5
View File
@@ -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
+29
View File
@@ -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)