diff --git a/.env.example b/.env.example index ff0520c..4e4571f 100644 --- a/.env.example +++ b/.env.example @@ -21,3 +21,6 @@ OPENAI_MODEL=gpt-4o # ANALYSIS_TARGET_API_KEY= # ANALYSIS_CRON_HOUR=8 # ANALYSIS_CRON_MINUTE=0 + +# Thread memory store path (JSON file for persisted thread summaries) +# THREAD_MEMORY_PATH=thread_memory.json diff --git a/steward/bot/telegram.py b/steward/bot/telegram.py index 32a3b3b..034d94d 100644 --- a/steward/bot/telegram.py +++ b/steward/bot/telegram.py @@ -15,13 +15,41 @@ from telegram.ext import ( from steward.config import Settings from steward.llm.client import LLMClient +from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary from steward.proposals.generator import Proposal, ProposalGenerator logger = logging.getLogger(__name__) -# Per-user conversation history (in-memory for the MVP) +# Per-user conversation history for non-threaded messages (capped at _MAX_HISTORY turns). _history: dict[int, list[dict[str, str]]] = defaultdict(list) -_MAX_HISTORY = 20 # keep the last N turns per user +_MAX_HISTORY = 20 + +# Per-thread conversation history: (chat_id, thread_id) → full history (unbounded). +# Messages belonging to a Telegram message thread are kept in their entirety here +# until explicitly flushed by the /flush command. +_thread_history: dict[tuple[int, int], list[dict[str, str]]] = defaultdict(list) + +_FLUSH_SYSTEM_PROMPT = ( + "You are Steward. The following is a complete Telegram message thread conversation. " + "Produce a concise but comprehensive summary that captures:\n" + "- The main topics discussed\n" + "- Key decisions or conclusions reached\n" + "- Any outstanding actions or open questions\n" + "- Important context that would help recall this conversation later\n\n" + "Be precise. Omit pleasantries." +) + + +def _thread_key(update: Update) -> tuple[int, int] | None: + """Return the (chat_id, thread_id) key if the message is part of a thread, else None.""" + msg = update.message + chat = update.effective_chat + if msg is None or chat is None: + return None + thread_id = msg.message_thread_id + if thread_id is None: + return None + return (chat.id, thread_id) def _is_allowed(user_id: int, settings: Settings) -> bool: @@ -32,7 +60,7 @@ def _is_allowed(user_id: int, settings: Settings) -> bool: async def _send_long(update: Update, text: str) -> None: - """Send text, splitting if it exceeds Telegram's 4096-char limit.""" + """Send text, splitting across messages if it exceeds Telegram's 4096-char limit.""" limit = 4096 for i in range(0, len(text), limit): await update.message.reply_text( # type: ignore[union-attr] @@ -53,7 +81,9 @@ async def start_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> N "I'm your AI-assisted personal operations platform.\n" "Talk to me naturally, or use:\n" "/help – show available commands\n" - "/clear – reset our conversation history\n" + "/clear – reset conversation history\n" + "/flush – summarise and archive this thread's memory\n" + "/recall – retrieve archived thread summaries\n" "/analyse – run a manual API analysis right now", parse_mode=ParseMode.MARKDOWN, ) @@ -70,21 +100,150 @@ async def help_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> No "*Steward commands*\n\n" "/start – greeting\n" "/help – this message\n" - "/clear – reset conversation history\n" + "/clear – reset conversation history for this context\n" + "/flush – summarise the current thread, store the summary, and compress memory\n" + " _(only available inside a message thread)_\n" + "/recall – show the stored summary for this thread, or list all summaries\n" "/analyse – trigger an immediate API analysis and proposal", parse_mode=ParseMode.MARKDOWN, ) async def clear_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle /clear – wipe conversation history for this user.""" + """Handle /clear – wipe conversation history for this context. + + Inside a message thread: clears the thread's unbounded history. + Outside a thread: clears the per-user capped history. + """ settings: Settings = context.bot_data["settings"] user = update.effective_user if user is None or not _is_allowed(user.id, settings): return - _history[user.id].clear() - await update.message.reply_text("Conversation history cleared.") # type: ignore[union-attr] + key = _thread_key(update) + if key is not None: + _thread_history[key].clear() + await update.message.reply_text( # type: ignore[union-attr] + "Thread conversation history cleared." + ) + else: + _history[user.id].clear() + await update.message.reply_text( # type: ignore[union-attr] + "Conversation history cleared." + ) + + +async def flush_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: + """Handle /flush – summarise thread memory, persist it, and compress in-memory history. + + Steps: + 1. Verify the command is issued inside a message thread. + 2. Summarise the full thread history via the LLM. + 3. Persist the summary in the ThreadMemoryStore (keyed by chat_id + thread_id). + 4. Replace the in-memory thread history with a single compressed context message + so conversation can continue with the summary as background. + """ + settings: Settings = context.bot_data["settings"] + llm: LLMClient = context.bot_data["llm"] + store: ThreadMemoryStore = context.bot_data["thread_store"] + user = update.effective_user + + if user is None or not _is_allowed(user.id, settings): + return + + key = _thread_key(update) + if key is None: + await update.message.reply_text( # type: ignore[union-attr] + "\u26a0\ufe0f /flush can only be used inside a message thread." + ) + return + + chat_id, thread_id = key + history = _thread_history[key] + + if not history: + await update.message.reply_text( # type: ignore[union-attr] + "This thread has no conversation history to flush." + ) + return + + await update.message.reply_text( # type: ignore[union-attr] + "\U0001f4be Summarising thread memory\u2026" + ) + + # Build a readable transcript for the LLM to summarise + transcript_lines = [] + for msg in history: + role_label = "User" if msg["role"] == "user" else "Steward" + transcript_lines.append(f"{role_label}: {msg['content']}") + transcript = "\n".join(transcript_lines) + + summary_text = await llm.chat( + f"Thread transcript:\n\n{transcript}", + system_prompt=_FLUSH_SYSTEM_PROMPT, + ) + + message_count = sum(1 for m in history if m["role"] == "user") + thread_summary = ThreadSummary( + chat_id=chat_id, + thread_id=thread_id, + summary=summary_text, + message_count=message_count, + ) + store.save(thread_summary) + + # Compress: replace history with a single system-context entry so the thread + # can continue with the summary as background knowledge. + _thread_history[key] = [ + {"role": "system", "content": f"Summary of earlier conversation:\n{summary_text}"} + ] + + await update.message.reply_text( # type: ignore[union-attr] + f"\u2705 Thread memory flushed and stored " + f"(thread `{thread_id}`, {message_count} messages summarised).", + parse_mode=ParseMode.MARKDOWN, + ) + + +async def recall_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: + """Handle /recall – retrieve stored thread summaries. + + Inside a thread: shows the stored summary for this thread (if any). + Outside a thread: lists all stored summaries (newest first). + """ + settings: Settings = context.bot_data["settings"] + store: ThreadMemoryStore = context.bot_data["thread_store"] + user = update.effective_user + + if user is None or not _is_allowed(user.id, settings): + return + + key = _thread_key(update) + if key is not None: + chat_id, thread_id = key + stored = store.get(chat_id, thread_id) + if stored is None: + await update.message.reply_text( # type: ignore[union-attr] + "No stored summary for this thread yet. Use /flush to create one." + ) + else: + await _send_long(update, stored.format_for_telegram()) + return + + # Outside a thread: list all stored summaries + all_summaries = store.all() + if not all_summaries: + await update.message.reply_text( # type: ignore[union-attr] + "No thread summaries stored yet. Use /flush inside a message thread." + ) + return + + lines = ["\U0001f4da *Stored thread summaries*\n"] + for s in all_summaries: + date = s.flushed_at[:10] + first_line = s.summary.split("\n")[0][:80] + lines.append(f"\u2022 Thread `{s.thread_id}` ({date}): {first_line}\u2026") + await _send_long(update, "\n".join(lines)) async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: @@ -109,7 +268,11 @@ async def analyse_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: - """Handle plain text messages – forward to LLM and reply.""" + """Handle plain text messages – forward to LLM and reply. + + Thread messages: history is stored unbounded under the (chat_id, thread_id) key. + Non-thread messages: history is capped at _MAX_HISTORY turns per user. + """ settings: Settings = context.bot_data["settings"] llm: LLMClient = context.bot_data["llm"] user = update.effective_user @@ -120,32 +283,39 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> if not text: return - history = _history[user.id] - reply = await llm.chat(text, history=history) - - # Update history - history.append({"role": "user", "content": text}) - history.append({"role": "assistant", "content": reply}) - # Trim to keep only the most recent turns (2 messages per turn) - if len(history) > _MAX_HISTORY * 2: - _history[user.id] = history[-( _MAX_HISTORY * 2):] + key = _thread_key(update) + if key is not None: + # Thread message: unbounded history + history = _thread_history[key] + reply = await llm.chat(text, history=history) + history.append({"role": "user", "content": text}) + history.append({"role": "assistant", "content": reply}) + else: + # Non-thread message: capped history per user + history = _history[user.id] + reply = await llm.chat(text, history=history) + history.append({"role": "user", "content": text}) + history.append({"role": "assistant", "content": reply}) + if len(history) > _MAX_HISTORY * 2: + _history[user.id] = history[-(_MAX_HISTORY * 2) :] await _send_long(update, reply) -def build_application(settings: Settings, llm: LLMClient) -> Application: # type: ignore[type-arg] +def build_application( + settings: Settings, llm: LLMClient, thread_store: ThreadMemoryStore | None = None +) -> Application: # type: ignore[type-arg] """Build and return the Telegram Application.""" - app = ( - Application.builder() - .token(settings.telegram_bot_token) - .build() - ) + app = Application.builder().token(settings.telegram_bot_token).build() app.bot_data["settings"] = settings app.bot_data["llm"] = llm + app.bot_data["thread_store"] = thread_store or ThreadMemoryStore(settings.thread_memory_path) app.add_handler(CommandHandler("start", start_handler)) app.add_handler(CommandHandler("help", help_handler)) app.add_handler(CommandHandler("clear", clear_handler)) + app.add_handler(CommandHandler("flush", flush_handler)) + app.add_handler(CommandHandler("recall", recall_handler)) app.add_handler(CommandHandler("analyse", analyse_handler)) app.add_handler(MessageHandler(filters.TEXT & ~filters.COMMAND, message_handler)) diff --git a/steward/config.py b/steward/config.py index 5c36303..a1a70d5 100644 --- a/steward/config.py +++ b/steward/config.py @@ -29,6 +29,9 @@ class Settings(BaseSettings): analysis_cron_hour: int = 8 analysis_cron_minute: int = 0 + # Thread memory + thread_memory_path: str = "thread_memory.json" + def get_settings() -> Settings: """Return application settings singleton.""" diff --git a/steward/main.py b/steward/main.py index 0d03b2d..8d0cb9c 100644 --- a/steward/main.py +++ b/steward/main.py @@ -8,6 +8,7 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler from steward.bot.telegram import build_application, send_proposal from steward.config import get_settings from steward.llm.client import LLMClient +from steward.memory.thread_store import ThreadMemoryStore from steward.proposals.generator import ProposalGenerator logging.basicConfig( @@ -44,7 +45,8 @@ def main() -> None: sys.exit(1) llm = LLMClient(settings) - app = build_application(settings, llm) + thread_store = ThreadMemoryStore(settings.thread_memory_path) + app = build_application(settings, llm, thread_store) generator = ProposalGenerator(settings, llm) scheduler = AsyncIOScheduler() diff --git a/steward/memory/__init__.py b/steward/memory/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/steward/memory/thread_store.py b/steward/memory/thread_store.py new file mode 100644 index 0000000..2bf2a6e --- /dev/null +++ b/steward/memory/thread_store.py @@ -0,0 +1,84 @@ +"""Persistent storage for flushed thread summaries.""" + +from __future__ import annotations + +import json +import logging +from dataclasses import asdict, dataclass, field +from datetime import UTC, datetime +from pathlib import Path + +logger = logging.getLogger(__name__) + + +@dataclass +class ThreadSummary: + """A persisted summary of a flushed Telegram message thread.""" + + chat_id: int + thread_id: int + summary: str + message_count: int + flushed_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat()) + + @property + def key(self) -> str: + return f"{self.chat_id}:{self.thread_id}" + + def format_for_telegram(self) -> str: + """Return a concise Telegram-formatted recall card.""" + flushed = self.flushed_at[:19].replace("T", " ") + header = ( + f"\U0001f9e0 *Thread summary* (thread `{self.thread_id}`)\n" + f"_Flushed: {flushed} UTC \u2014 {self.message_count} messages_" + ) + return f"{header}\n\n{self.summary}" + + +class ThreadMemoryStore: + """JSON-backed store for flushed thread summaries. + + The store is intentionally lightweight for the MVP. Each flush overwrites + any previous summary for the same (chat_id, thread_id) pair. + """ + + def __init__(self, path: str | Path = "thread_memory.json") -> None: + self._path = Path(path) + self._data: dict[str, dict[str, int | str]] = self._load() + + def _load(self) -> dict[str, dict[str, int | str]]: + if self._path.exists(): + try: + raw = json.loads(self._path.read_text(encoding="utf-8")) + if isinstance(raw, dict): + return raw # type: ignore[return-value] + except (json.JSONDecodeError, OSError): + logger.warning( + "Could not read thread memory store at %s; starting fresh", self._path + ) + return {} + + def _save(self) -> None: + try: + self._path.write_text( + json.dumps(self._data, indent=2, ensure_ascii=False), encoding="utf-8" + ) + except OSError: + logger.exception("Failed to write thread memory store to %s", self._path) + + def save(self, summary: ThreadSummary) -> None: + """Persist a thread summary, replacing any previous entry for this thread.""" + self._data[summary.key] = asdict(summary) + self._save() + + def get(self, chat_id: int, thread_id: int) -> ThreadSummary | None: + """Return the stored summary for a thread, or None if not found.""" + raw = self._data.get(f"{chat_id}:{thread_id}") + if raw is None: + return None + return ThreadSummary(**raw) # type: ignore[arg-type] + + def all(self) -> list[ThreadSummary]: + """Return all stored summaries, newest first.""" + entries = [ThreadSummary(**v) for v in self._data.values()] # type: ignore[arg-type] + return sorted(entries, key=lambda s: s.flushed_at, reverse=True) diff --git a/tests/test_bot.py b/tests/test_bot.py index 53a2a5a..9ecda8e 100644 --- a/tests/test_bot.py +++ b/tests/test_bot.py @@ -3,18 +3,20 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from telegram import Message, Update, User +from telegram import Chat, Message, Update, User from telegram.ext import CallbackContext from steward.bot.telegram import ( _history, _is_allowed, + _thread_history, clear_handler, message_handler, start_handler, ) from steward.config import Settings from steward.llm.client import LLMClient +from steward.memory.thread_store import ThreadMemoryStore def _make_settings(**kwargs) -> Settings: @@ -26,23 +28,41 @@ def _make_settings(**kwargs) -> Settings: return Settings(**defaults) -def _make_update(user_id: int = 12345, text: str = "hello") -> Update: +def _make_update( + user_id: int = 12345, + text: str = "hello", + chat_id: int | None = None, + 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 if chat_id is not None else user_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) -> CallbackContext: # type: ignore[type-arg] +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} + ctx.bot_data = { + "settings": settings, + "llm": llm, + "thread_store": store or MagicMock(spec=ThreadMemoryStore), + } return ctx @@ -85,12 +105,13 @@ async def test_start_handler_ignores_disallowed_user(): @pytest.mark.asyncio -async def test_clear_handler_clears_history(): +async def test_clear_handler_clears_user_history(): + """clear_handler without a thread clears per-user history.""" settings = _make_settings() user_id = 42 _history[user_id] = [{"role": "user", "content": "old msg"}] - update = _make_update(user_id=user_id) + update = _make_update(user_id=user_id) # no thread_id ctx = _make_context(settings) await clear_handler(update, ctx) @@ -99,6 +120,22 @@ async def test_clear_handler_clears_history(): update.message.reply_text.assert_awaited_once() +@pytest.mark.asyncio +async def test_clear_handler_clears_thread_history(): + """clear_handler inside a thread clears that thread's history.""" + settings = _make_settings() + key = (100, 7) + _thread_history[key] = [{"role": "user", "content": "thread msg"}] + + update = _make_update(user_id=1, chat_id=100, thread_id=7) + ctx = _make_context(settings) + + await clear_handler(update, ctx) + + assert _thread_history[key] == [] + update.message.reply_text.assert_awaited_once() + + @pytest.mark.asyncio async def test_message_handler_calls_llm_and_replies(): settings = _make_settings() @@ -113,7 +150,30 @@ async def test_message_handler_calls_llm_and_replies(): mock_llm.chat.assert_awaited_once() update.message.reply_text.assert_awaited_once() - # History should now contain the user/assistant turn assert len(_history[77]) == 2 assert _history[77][0]["role"] == "user" assert _history[77][1]["role"] == "assistant" + + +@pytest.mark.asyncio +async def test_message_handler_non_thread_trims_history(): + """Non-thread history is trimmed to _MAX_HISTORY turns.""" + from steward.bot.telegram import _MAX_HISTORY + + settings = _make_settings() + mock_llm = MagicMock(spec=LLMClient) + mock_llm.chat = AsyncMock(return_value="reply") + + user_id = 200 + # Pre-fill exactly at the limit + _history[user_id] = [ + {"role": "user" if i % 2 == 0 else "assistant", "content": f"msg{i}"} + for i in range(_MAX_HISTORY * 2) + ] + + update = _make_update(user_id=user_id, text="new question") + ctx = _make_context(settings, llm=mock_llm) + await message_handler(update, ctx) + + # After appending 2 new messages, should trim back to _MAX_HISTORY * 2 + assert len(_history[user_id]) == _MAX_HISTORY * 2 diff --git a/tests/test_thread_memory.py b/tests/test_thread_memory.py new file mode 100644 index 0000000..6df8de0 --- /dev/null +++ b/tests/test_thread_memory.py @@ -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()