feat: thread-aware memory with /flush and /recall commands
This commit is contained in:
committed by
GitHub
parent
8759bea963
commit
c74c8b0d8b
@@ -0,0 +1,318 @@
|
||||
"""Tests for thread-aware memory: ThreadMemoryStore and /flush, /recall handlers."""
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from telegram import Chat, Message, Update, User
|
||||
from telegram.ext import CallbackContext
|
||||
|
||||
from steward.bot.telegram import (
|
||||
_thread_history,
|
||||
flush_handler,
|
||||
message_handler,
|
||||
recall_handler,
|
||||
)
|
||||
from steward.config import Settings
|
||||
from steward.llm.client import LLMClient
|
||||
from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_settings(**kwargs) -> Settings:
|
||||
return Settings(telegram_bot_token="tok", openai_api_key="key", **kwargs)
|
||||
|
||||
|
||||
def _make_update(
|
||||
user_id: int = 1,
|
||||
text: str = "hi",
|
||||
chat_id: int = 100,
|
||||
thread_id: int | None = None,
|
||||
) -> Update:
|
||||
user = MagicMock(spec=User)
|
||||
user.id = user_id
|
||||
|
||||
message = MagicMock(spec=Message)
|
||||
message.text = text
|
||||
message.reply_text = AsyncMock()
|
||||
message.message_thread_id = thread_id
|
||||
|
||||
chat = MagicMock(spec=Chat)
|
||||
chat.id = chat_id
|
||||
|
||||
update = MagicMock(spec=Update)
|
||||
update.effective_user = user
|
||||
update.effective_chat = chat
|
||||
update.message = message
|
||||
return update
|
||||
|
||||
|
||||
def _make_context(
|
||||
settings: Settings,
|
||||
llm: LLMClient | None = None,
|
||||
store: ThreadMemoryStore | None = None,
|
||||
) -> CallbackContext: # type: ignore[type-arg]
|
||||
ctx = MagicMock(spec=CallbackContext)
|
||||
ctx.bot_data = {
|
||||
"settings": settings,
|
||||
"llm": llm or MagicMock(spec=LLMClient),
|
||||
"thread_store": store or MagicMock(spec=ThreadMemoryStore),
|
||||
}
|
||||
return ctx
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ThreadMemoryStore unit tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestThreadMemoryStore:
|
||||
def _store(self, tmp_path: Path) -> ThreadMemoryStore:
|
||||
return ThreadMemoryStore(tmp_path / "mem.json")
|
||||
|
||||
def test_get_missing_returns_none(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
assert store.get(1, 2) is None
|
||||
|
||||
def test_save_and_get_roundtrip(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
summary = ThreadSummary(chat_id=1, thread_id=42, summary="A recap.", message_count=5)
|
||||
store.save(summary)
|
||||
|
||||
retrieved = store.get(1, 42)
|
||||
assert retrieved is not None
|
||||
assert retrieved.summary == "A recap."
|
||||
assert retrieved.message_count == 5
|
||||
|
||||
def test_save_overwrites_previous_entry(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="old", message_count=2))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="new", message_count=4))
|
||||
assert store.get(1, 1).summary == "new" # type: ignore[union-attr]
|
||||
|
||||
def test_persisted_to_disk(self, tmp_path: Path):
|
||||
path = tmp_path / "mem.json"
|
||||
store = ThreadMemoryStore(path)
|
||||
store.save(ThreadSummary(chat_id=5, thread_id=9, summary="saved", message_count=1))
|
||||
|
||||
# Load a fresh store from the same file
|
||||
store2 = ThreadMemoryStore(path)
|
||||
assert store2.get(5, 9) is not None
|
||||
|
||||
def test_all_returns_newest_first(self, tmp_path: Path):
|
||||
store = self._store(tmp_path)
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="first", message_count=1,
|
||||
flushed_at="2026-01-01T00:00:00+00:00"))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=2, summary="second", message_count=1,
|
||||
flushed_at="2026-06-01T00:00:00+00:00"))
|
||||
results = store.all()
|
||||
assert results[0].summary == "second"
|
||||
assert results[1].summary == "first"
|
||||
|
||||
def test_corrupt_file_starts_fresh(self, tmp_path: Path):
|
||||
path = tmp_path / "mem.json"
|
||||
path.write_text("not-json", encoding="utf-8")
|
||||
store = ThreadMemoryStore(path)
|
||||
assert store.all() == []
|
||||
|
||||
def test_format_for_telegram_contains_thread_id(self, tmp_path: Path):
|
||||
s = ThreadSummary(chat_id=1, thread_id=77, summary="recap", message_count=3)
|
||||
text = s.format_for_telegram()
|
||||
assert "77" in text
|
||||
assert "recap" in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Thread-aware message_handler tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_thread_message_stored_in_thread_history():
|
||||
"""Messages in a thread go to _thread_history, not _history."""
|
||||
from steward.bot.telegram import _history
|
||||
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="thread reply")
|
||||
|
||||
key = (100, 55)
|
||||
_thread_history[key].clear()
|
||||
|
||||
update = _make_update(user_id=1, chat_id=100, thread_id=55, text="thread message")
|
||||
ctx = _make_context(settings, llm=mock_llm)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
assert len(_thread_history[key]) == 2
|
||||
assert _thread_history[key][0]["role"] == "user"
|
||||
# Regular user history untouched
|
||||
assert len(_history[1]) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_thread_history_is_unbounded():
|
||||
"""Thread history never gets trimmed regardless of how many turns there are."""
|
||||
from steward.bot.telegram import _MAX_HISTORY
|
||||
|
||||
settings = _make_settings()
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="reply")
|
||||
|
||||
key = (200, 66)
|
||||
# Pre-fill well beyond the cap used for non-thread history
|
||||
_thread_history[key] = [
|
||||
{"role": "user" if i % 2 == 0 else "assistant", "content": f"m{i}"}
|
||||
for i in range(_MAX_HISTORY * 4) # 4× the normal cap
|
||||
]
|
||||
prior_len = len(_thread_history[key])
|
||||
|
||||
update = _make_update(user_id=2, chat_id=200, thread_id=66, text="more")
|
||||
ctx = _make_context(settings, llm=mock_llm)
|
||||
await message_handler(update, ctx)
|
||||
|
||||
# Should have grown by exactly 2 (user + assistant), never trimmed
|
||||
assert len(_thread_history[key]) == prior_len + 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /flush handler tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_outside_thread_warns():
|
||||
"""flush_handler outside a thread should warn the user."""
|
||||
settings = _make_settings()
|
||||
store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="summary")
|
||||
|
||||
update = _make_update(thread_id=None) # no thread
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await flush_handler(update, ctx)
|
||||
|
||||
store.save.assert_not_called()
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "flush" in text.lower() or "thread" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_empty_thread_warns():
|
||||
"""flush_handler with no history should tell the user there is nothing to flush."""
|
||||
settings = _make_settings()
|
||||
store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="summary")
|
||||
|
||||
key = (300, 88)
|
||||
_thread_history[key].clear()
|
||||
|
||||
update = _make_update(chat_id=300, thread_id=88)
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await flush_handler(update, ctx)
|
||||
|
||||
store.save.assert_not_called()
|
||||
mock_llm.chat.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_summarises_stores_and_compresses():
|
||||
"""flush_handler should summarise, persist, and compress in-memory history."""
|
||||
settings = _make_settings()
|
||||
store = MagicMock(spec=ThreadMemoryStore)
|
||||
mock_llm = MagicMock(spec=LLMClient)
|
||||
mock_llm.chat = AsyncMock(return_value="Great summary of the thread.")
|
||||
|
||||
key = (400, 99)
|
||||
_thread_history[key] = [
|
||||
{"role": "user", "content": "question one"},
|
||||
{"role": "assistant", "content": "answer one"},
|
||||
{"role": "user", "content": "question two"},
|
||||
{"role": "assistant", "content": "answer two"},
|
||||
]
|
||||
|
||||
update = _make_update(chat_id=400, thread_id=99)
|
||||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||||
await flush_handler(update, ctx)
|
||||
|
||||
# LLM called to summarise
|
||||
mock_llm.chat.assert_awaited_once()
|
||||
# Summary persisted
|
||||
store.save.assert_called_once()
|
||||
saved: ThreadSummary = store.save.call_args.args[0]
|
||||
assert saved.chat_id == 400
|
||||
assert saved.thread_id == 99
|
||||
assert saved.summary == "Great summary of the thread."
|
||||
assert saved.message_count == 2 # 2 user turns
|
||||
|
||||
# In-memory history replaced with compressed context
|
||||
assert len(_thread_history[key]) == 1
|
||||
assert _thread_history[key][0]["role"] == "system"
|
||||
assert "Great summary" in _thread_history[key][0]["content"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /recall handler tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_in_thread_returns_stored_summary(tmp_path):
|
||||
"""recall_handler inside a thread returns the stored summary for that thread."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(chat_id=500, thread_id=11, summary="recap text", message_count=3))
|
||||
|
||||
update = _make_update(chat_id=500, thread_id=11)
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "recap text" in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_in_thread_no_summary_guides_user(tmp_path):
|
||||
"""recall_handler inside a thread with no summary tells user to /flush first."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
|
||||
update = _make_update(chat_id=600, thread_id=22)
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "flush" in text.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_outside_thread_lists_all(tmp_path):
|
||||
"""recall_handler outside a thread lists all stored summaries."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="alpha", message_count=1))
|
||||
store.save(ThreadSummary(chat_id=1, thread_id=2, summary="beta", message_count=2))
|
||||
|
||||
update = _make_update(thread_id=None) # no thread
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "alpha" in text or "beta" in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_outside_thread_empty_store(tmp_path):
|
||||
"""recall_handler outside a thread with no summaries guides user."""
|
||||
settings = _make_settings()
|
||||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||||
|
||||
update = _make_update(thread_id=None)
|
||||
ctx = _make_context(settings, store=store)
|
||||
await recall_handler(update, ctx)
|
||||
|
||||
update.message.reply_text.assert_awaited_once()
|
||||
text = update.message.reply_text.call_args.args[0]
|
||||
assert "flush" in text.lower()
|
||||
Reference in New Issue
Block a user