"""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 from tests.conftest import make_settings # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _make_settings(**kwargs) -> Settings: """Wrapper for test compatibility.""" defaults = dict(telegram_bot_token="tok", openai_api_key="key") defaults.update(kwargs) return make_settings(**defaults) 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)