"""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.core import ConversationService from steward.bot.telegram import ( flush_handler, message_handler, recall_handler, ) from steward.bot.thread_key import ThreadKey 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 if llm is None: llm = MagicMock(spec=LLMClient) service = ConversationService(settings, llm, store) ctx.bot_data = { "settings": settings, "service": service, "llm": llm, } ctx.args = [] return ctx # --------------------------------------------------------------------------- # ThreadMemoryStore unit tests # --------------------------------------------------------------------------- class TestThreadMemoryStore: def _store(self, tmp_path: Path) -> ThreadMemoryStore: return ThreadMemoryStore(tmp_path / "mem.json") def _key(self, scope: str = "1", thread: str | None = "2") -> ThreadKey: return ThreadKey(platform="telegram", scope=scope, thread=thread) def _summary(self, scope: str = "1", thread: str | None = "2", **kwargs) -> ThreadSummary: defaults = dict(summary="A recap.", message_count=5) defaults.update(kwargs) return ThreadSummary(platform="telegram", scope=scope, thread=thread, **defaults) def test_get_missing_returns_none(self, tmp_path: Path): store = self._store(tmp_path) assert store.get(self._key()) is None def test_save_and_get_roundtrip(self, tmp_path: Path): store = self._store(tmp_path) summary = self._summary() store.save(summary) retrieved = store.get(self._key()) 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(self._summary(summary="old", message_count=2)) store.save(self._summary(summary="new", message_count=4)) assert store.get(self._key()).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(self._summary(scope="5", thread="9", summary="saved", message_count=1)) # Load a fresh store from the same file store2 = ThreadMemoryStore(path) assert store2.get(self._key("5", "9")) is not None def test_all_returns_newest_first(self, tmp_path: Path): store = self._store(tmp_path) store.save( self._summary( thread="1", summary="first", message_count=1, flushed_at="2026-01-01T00:00:00+00:00", ) ) store.save( self._summary( thread="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 = self._summary(thread="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 = self._summary(thread="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( self._summary( thread="1", summary="API work", message_count=2, tags=["api", "design", "auth"] ) ) store.save( self._summary( thread="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(self._summary(thread="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(self._summary(thread="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(self._summary(thread="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(self._summary(thread="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(self._key("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 the thread's history, not the user's.""" settings = _make_settings() mock_llm = MagicMock(spec=LLMClient) mock_llm.chat = AsyncMock(return_value="thread reply") ctx = _make_context(settings, llm=mock_llm) service: ConversationService = ctx.bot_data["service"] thread_key = ThreadKey(platform="telegram", scope="100", thread="55") user_key = ThreadKey(platform="telegram", scope="1") update = _make_update(user_id=1, chat_id=100, thread_id=55, text="thread message") await message_handler(update, ctx) assert len(service._histories[thread_key]) == 2 assert service._histories[thread_key][0]["role"] == "user" # Regular user history untouched assert not service.has_history(user_key) @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.core import _MAX_HISTORY settings = _make_settings() mock_llm = MagicMock(spec=LLMClient) mock_llm.chat = AsyncMock(return_value="reply") ctx = _make_context(settings, llm=mock_llm) service: ConversationService = ctx.bot_data["service"] thread_key = ThreadKey(platform="telegram", scope="200", thread="66") # Pre-fill well beyond the cap used for non-thread history service._histories[thread_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(service._histories[thread_key]) update = _make_update(user_id=2, chat_id=200, thread_id=66, text="more") await message_handler(update, ctx) # Should have grown by exactly 2 (user + assistant), never trimmed assert len(service._histories[thread_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") 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"]) ctx = _make_context(settings, llm=mock_llm, store=store) service: ConversationService = ctx.bot_data["service"] key = ThreadKey(platform="telegram", scope="400", thread="99") service._histories[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) 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.platform == "telegram" assert saved.scope == "400" assert saved.thread == "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(service._histories[key]) == 1 assert service._histories[key][0]["role"] == "system" assert "Great summary" in service._histories[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( platform="telegram", scope="500", thread="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(platform="telegram", scope="1", thread="1", summary="alpha", message_count=1) ) store.save( ThreadSummary(platform="telegram", scope="1", thread="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( platform="telegram", scope="1", thread="1", summary="API authentication discussion", message_count=2, tags=["api", "auth"], ) ) store.save( ThreadSummary( platform="telegram", scope="1", thread="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( platform="telegram", scope="1", thread="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( platform="telegram", scope="1", thread="1", summary="Previous API discussion", message_count=2, tags=["api"], ) ) user_id = 999 ctx = _make_context(settings, llm=mock_llm, store=store) service: ConversationService = ctx.bot_data["service"] user_key = ThreadKey(platform="telegram", scope=str(user_id)) update = _make_update(user_id=user_id, text="tell me about the api work") 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 service._histories[user_key] ) @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( platform="telegram", scope="1", thread="1", summary="Database recap", message_count=1, tags=["database"], ) ) user_id = 888 ctx = _make_context(settings, llm=mock_llm, store=store) update = _make_update(user_id=user_id, text="what is the weather like?") 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)