mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 00:03:01 +03:00
fix(memory): preserve replay prefix on session reset
This commit is contained in:
@@ -306,16 +306,15 @@ async def cmd_new(ctx: CommandContext) -> OutboundMessage:
|
||||
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)
|
||||
snapshot = session.messages[session.last_consolidated:]
|
||||
snapshot = list(session.messages)
|
||||
archive_snapshot = None
|
||||
runtime = None
|
||||
if snapshot:
|
||||
if session.last_consolidated < len(snapshot):
|
||||
runtime = ctx.runtime or loop.runtime_for_session(session)
|
||||
archive_snapshot = replace(
|
||||
session,
|
||||
messages=snapshot,
|
||||
metadata=dict(session.metadata),
|
||||
last_consolidated=0,
|
||||
provider_state=None,
|
||||
)
|
||||
session.clear()
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
"""Test session management with cache-friendly message handling."""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Coroutine
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
@@ -488,12 +490,13 @@ class TestNewCommandArchival:
|
||||
def _make_loop(tmp_path: Path):
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
from nanobot.providers.base import GenerationSettings, LLMResponse
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
provider.estimate_prompt_tokens.return_value = (10_000, "test")
|
||||
provider.generation = GenerationSettings(max_tokens=100)
|
||||
loop = AgentLoop(
|
||||
bus=bus,
|
||||
provider=provider,
|
||||
@@ -542,29 +545,35 @@ class TestNewCommandArchival:
|
||||
assert call_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_archives_only_unconsolidated_messages(self, tmp_path: Path) -> None:
|
||||
async def test_new_reuses_replay_prefix_and_archives_only_unconsolidated_messages(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
from nanobot.bus.events import InboundMessage
|
||||
|
||||
loop = self._make_loop(tmp_path)
|
||||
loop.set_runtime_context_window(128_000)
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
for i in range(15):
|
||||
for i in range(5):
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
session.last_consolidated = len(session.messages) - 3
|
||||
session.last_consolidated = len(session.messages) - 2
|
||||
ordinary_history = session.get_history()
|
||||
assert [message["content"] for message in ordinary_history] == [
|
||||
"msg1",
|
||||
"resp1",
|
||||
"msg2",
|
||||
"resp2",
|
||||
"msg3",
|
||||
"resp3",
|
||||
"msg4",
|
||||
"resp4",
|
||||
]
|
||||
loop.sessions.save(session)
|
||||
|
||||
archived_count = -1
|
||||
archived_session_key = None
|
||||
expected_runtime = loop.llm_runtime()
|
||||
|
||||
async def _fake_summarize(session, *, archive_end, runtime) -> str:
|
||||
nonlocal archived_count, archived_session_key
|
||||
assert runtime is expected_runtime
|
||||
archived_count = len(session.messages[:archive_end])
|
||||
archived_session_key = session.key
|
||||
return "Summary."
|
||||
|
||||
loop.consolidator.archive_session = _fake_summarize # type: ignore[method-assign]
|
||||
scheduled: list[Coroutine[Any, Any, object]] = []
|
||||
loop.schedule_background = scheduled.append # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||
@@ -572,9 +581,12 @@ class TestNewCommandArchival:
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
|
||||
assert len(scheduled) == 1
|
||||
await scheduled[0]
|
||||
await loop.aclose()
|
||||
assert archived_count == 3
|
||||
assert archived_session_key == "cli:test"
|
||||
sent = loop.provider.chat_with_retry.call_args.kwargs["messages"]
|
||||
assert sent[1:-1] == ordinary_history
|
||||
assert "final 2 conversation messages" in sent[-1]["content"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_clears_session_and_responds(self, tmp_path: Path) -> None:
|
||||
|
||||
Reference in New Issue
Block a user