steward_mirror/tests/test_thread_memory.py

501 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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)
if store is None:
mock_store = MagicMock(spec=ThreadMemoryStore)
mock_store.search.return_value = []
store = mock_store
ctx.bot_data = {
"settings": settings,
"llm": llm or MagicMock(spec=LLMClient),
"thread_store": store,
}
ctx.args = []
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
def test_format_for_telegram_shows_tags(self, tmp_path: Path):
s = ThreadSummary(
chat_id=1, thread_id=77, summary="recap", message_count=3,
tags=["api", "auth"],
)
text = s.format_for_telegram()
assert "api" in text
assert "auth" in text
def test_from_dict_tolerates_missing_tags(self, tmp_path: Path):
"""from_dict must handle legacy entries that predate the tags field."""
raw = {"chat_id": 1, "thread_id": 2, "summary": "old", "message_count": 3,
"flushed_at": "2026-01-01T00:00:00+00:00"}
s = ThreadSummary.from_dict(raw)
assert s.tags == []
def test_search_returns_matching_summaries(self, tmp_path: Path):
store = self._store(tmp_path)
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="API work", message_count=2,
tags=["api", "design", "auth"]))
store.save(ThreadSummary(chat_id=1, thread_id=2, summary="Database work", message_count=2,
tags=["database", "schema"]))
results = store.search("api")
assert len(results) == 1
assert results[0].thread_id == 1
def test_search_no_match_returns_empty(self, tmp_path: Path):
store = self._store(tmp_path)
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
tags=["database"]))
assert store.search("deployment") == []
def test_search_empty_query_returns_empty(self, tmp_path: Path):
store = self._store(tmp_path)
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
tags=["database"]))
assert store.search("") == []
def test_search_case_insensitive(self, tmp_path: Path):
store = self._store(tmp_path)
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
tags=["API"]))
assert len(store.search("api")) == 1
def test_search_skips_untagged_summaries(self, tmp_path: Path):
store = self._store(tmp_path)
store.save(ThreadSummary(chat_id=1, thread_id=1, summary="recap", message_count=1,
tags=[]))
assert store.search("api") == []
def test_legacy_store_roundtrip(self, tmp_path: Path):
"""Summaries without tags survive a save/load cycle via from_dict."""
path = tmp_path / "mem.json"
import json as _json
path.write_text(
_json.dumps({
"1:1": {"chat_id": 1, "thread_id": 1, "summary": "old",
"message_count": 1, "flushed_at": "2026-01-01T00:00:00+00:00"}
}),
encoding="utf-8",
)
store = ThreadMemoryStore(path)
s = store.get(1, 1)
assert s is not None
assert s.tags == []
# ---------------------------------------------------------------------------
# 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)
# First call → summary, second call → tags
mock_llm.chat = AsyncMock(side_effect=["Great summary of the thread.", "api, design, testing"])
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 twice: once for summary, once for tags
assert mock_llm.chat.await_count == 2
# 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
assert saved.tags == ["api", "design", "testing"]
# 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()
@pytest.mark.asyncio
async def test_recall_with_query_returns_matching_summaries(tmp_path):
"""recall_handler with a query argument searches the knowledge base by keyword."""
settings = _make_settings()
store = ThreadMemoryStore(tmp_path / "mem.json")
store.save(ThreadSummary(
chat_id=1, thread_id=1, summary="API authentication discussion",
message_count=2, tags=["api", "auth"],
))
store.save(ThreadSummary(
chat_id=1, thread_id=2, summary="Database schema planning",
message_count=3, tags=["database", "schema"],
))
update = _make_update(thread_id=None)
ctx = _make_context(settings, store=store)
ctx.args = ["api"]
await recall_handler(update, ctx)
update.message.reply_text.assert_awaited_once()
text = update.message.reply_text.call_args.args[0]
assert "API authentication" in text
assert "Database" not in text
@pytest.mark.asyncio
async def test_recall_with_query_no_match(tmp_path):
"""recall_handler with a query that matches nothing tells the user."""
settings = _make_settings()
store = ThreadMemoryStore(tmp_path / "mem.json")
store.save(ThreadSummary(
chat_id=1, thread_id=1, summary="recap", message_count=1, tags=["database"],
))
update = _make_update(thread_id=None)
ctx = _make_context(settings, store=store)
ctx.args = ["deployment"]
await recall_handler(update, ctx)
update.message.reply_text.assert_awaited_once()
text = update.message.reply_text.call_args.args[0]
assert "deployment" in text.lower() or "no memories" in text.lower()
@pytest.mark.asyncio
async def test_message_handler_injects_kb_context_transiently(tmp_path):
"""message_handler injects relevant KB summaries as transient context without storing them."""
settings = _make_settings()
mock_llm = MagicMock(spec=LLMClient)
mock_llm.chat = AsyncMock(return_value="reply with context")
store = ThreadMemoryStore(tmp_path / "mem.json")
store.save(ThreadSummary(
chat_id=1, thread_id=1, summary="Previous API discussion",
message_count=2, tags=["api"],
))
user_id = 999
from steward.bot.telegram import _history
_history[user_id].clear()
update = _make_update(user_id=user_id, text="tell me about the api work")
ctx = _make_context(settings, llm=mock_llm, store=store)
await message_handler(update, ctx)
# LLM was called
mock_llm.chat.assert_awaited_once()
call_kwargs = mock_llm.chat.call_args
# The history passed to the LLM should contain the KB context message
history_arg = call_kwargs.kwargs.get("history") or (
call_kwargs.args[1] if len(call_kwargs.args) > 1 else None
)
assert history_arg is not None
assert any("knowledge base" in m.get("content", "").lower() for m in history_arg)
# But the KB message must NOT be stored in _history
assert all(
"knowledge base" not in m.get("content", "").lower()
for m in _history[user_id]
)
@pytest.mark.asyncio
async def test_message_handler_no_kb_injection_when_no_match(tmp_path):
"""message_handler does not inject KB context when no summaries match."""
settings = _make_settings()
mock_llm = MagicMock(spec=LLMClient)
mock_llm.chat = AsyncMock(return_value="plain reply")
store = ThreadMemoryStore(tmp_path / "mem.json")
# Store a summary with unrelated tags
store.save(ThreadSummary(
chat_id=1, thread_id=1, summary="Database recap", message_count=1, tags=["database"],
))
user_id = 888
from steward.bot.telegram import _history
_history[user_id].clear()
update = _make_update(user_id=user_id, text="what is the weather like?")
ctx = _make_context(settings, llm=mock_llm, store=store)
await message_handler(update, ctx)
call_kwargs = mock_llm.chat.call_args
history_arg = call_kwargs.kwargs.get("history") or []
# No KB context injected
assert not any("knowledge base" in m.get("content", "").lower() for m in history_arg)