feat: thread-aware memory with /flush and /recall commands

This commit is contained in:
copilot-swe-agent[bot] 2026-07-25 12:09:38 +00:00 committed by GitHub
parent 8759bea963
commit c74c8b0d8b
8 changed files with 672 additions and 32 deletions

View File

@ -21,3 +21,6 @@ OPENAI_MODEL=gpt-4o
# ANALYSIS_TARGET_API_KEY= # ANALYSIS_TARGET_API_KEY=
# ANALYSIS_CRON_HOUR=8 # ANALYSIS_CRON_HOUR=8
# ANALYSIS_CRON_MINUTE=0 # ANALYSIS_CRON_MINUTE=0
# Thread memory store path (JSON file for persisted thread summaries)
# THREAD_MEMORY_PATH=thread_memory.json

View File

@ -15,13 +15,41 @@ from telegram.ext import (
from steward.config import Settings from steward.config import Settings
from steward.llm.client import LLMClient from steward.llm.client import LLMClient
from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary
from steward.proposals.generator import Proposal, ProposalGenerator from steward.proposals.generator import Proposal, ProposalGenerator
logger = logging.getLogger(__name__) 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) _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: 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: 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 limit = 4096
for i in range(0, len(text), limit): for i in range(0, len(text), limit):
await update.message.reply_text( # type: ignore[union-attr] 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" "I'm your AI-assisted personal operations platform.\n"
"Talk to me naturally, or use:\n" "Talk to me naturally, or use:\n"
"/help – show available commands\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", "/analyse – run a manual API analysis right now",
parse_mode=ParseMode.MARKDOWN, parse_mode=ParseMode.MARKDOWN,
) )
@ -70,21 +100,150 @@ async def help_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> No
"*Steward commands*\n\n" "*Steward commands*\n\n"
"/start – greeting\n" "/start – greeting\n"
"/help – this message\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", "/analyse – trigger an immediate API analysis and proposal",
parse_mode=ParseMode.MARKDOWN, parse_mode=ParseMode.MARKDOWN,
) )
async def clear_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: 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"] settings: Settings = context.bot_data["settings"]
user = update.effective_user user = update.effective_user
if user is None or not _is_allowed(user.id, settings): if user is None or not _is_allowed(user.id, settings):
return return
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() _history[user.id].clear()
await update.message.reply_text("Conversation history cleared.") # type: ignore[union-attr] 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: 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: 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"] settings: Settings = context.bot_data["settings"]
llm: LLMClient = context.bot_data["llm"] llm: LLMClient = context.bot_data["llm"]
user = update.effective_user user = update.effective_user
@ -120,32 +283,39 @@ async def message_handler(update: Update, context: ContextTypes.DEFAULT_TYPE) ->
if not text: if not text:
return return
history = _history[user.id] 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) reply = await llm.chat(text, history=history)
# Update history
history.append({"role": "user", "content": text}) history.append({"role": "user", "content": text})
history.append({"role": "assistant", "content": reply}) history.append({"role": "assistant", "content": reply})
# Trim to keep only the most recent turns (2 messages per turn)
if len(history) > _MAX_HISTORY * 2: if len(history) > _MAX_HISTORY * 2:
_history[user.id] = history[-(_MAX_HISTORY * 2) :] _history[user.id] = history[-(_MAX_HISTORY * 2) :]
await _send_long(update, reply) 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.""" """Build and return the Telegram Application."""
app = ( app = Application.builder().token(settings.telegram_bot_token).build()
Application.builder()
.token(settings.telegram_bot_token)
.build()
)
app.bot_data["settings"] = settings app.bot_data["settings"] = settings
app.bot_data["llm"] = llm 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("start", start_handler))
app.add_handler(CommandHandler("help", help_handler)) app.add_handler(CommandHandler("help", help_handler))
app.add_handler(CommandHandler("clear", clear_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(CommandHandler("analyse", analyse_handler))
app.add_handler(MessageHandler(filters.TEXT & ~filters.COMMAND, message_handler)) app.add_handler(MessageHandler(filters.TEXT & ~filters.COMMAND, message_handler))

View File

@ -29,6 +29,9 @@ class Settings(BaseSettings):
analysis_cron_hour: int = 8 analysis_cron_hour: int = 8
analysis_cron_minute: int = 0 analysis_cron_minute: int = 0
# Thread memory
thread_memory_path: str = "thread_memory.json"
def get_settings() -> Settings: def get_settings() -> Settings:
"""Return application settings singleton.""" """Return application settings singleton."""

View File

@ -8,6 +8,7 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler
from steward.bot.telegram import build_application, send_proposal from steward.bot.telegram import build_application, send_proposal
from steward.config import get_settings from steward.config import get_settings
from steward.llm.client import LLMClient from steward.llm.client import LLMClient
from steward.memory.thread_store import ThreadMemoryStore
from steward.proposals.generator import ProposalGenerator from steward.proposals.generator import ProposalGenerator
logging.basicConfig( logging.basicConfig(
@ -44,7 +45,8 @@ def main() -> None:
sys.exit(1) sys.exit(1)
llm = LLMClient(settings) 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) generator = ProposalGenerator(settings, llm)
scheduler = AsyncIOScheduler() scheduler = AsyncIOScheduler()

View File

View File

@ -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)

View File

@ -3,18 +3,20 @@
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from telegram import Message, Update, User from telegram import Chat, Message, Update, User
from telegram.ext import CallbackContext from telegram.ext import CallbackContext
from steward.bot.telegram import ( from steward.bot.telegram import (
_history, _history,
_is_allowed, _is_allowed,
_thread_history,
clear_handler, clear_handler,
message_handler, message_handler,
start_handler, start_handler,
) )
from steward.config import Settings from steward.config import Settings
from steward.llm.client import LLMClient from steward.llm.client import LLMClient
from steward.memory.thread_store import ThreadMemoryStore
def _make_settings(**kwargs) -> Settings: def _make_settings(**kwargs) -> Settings:
@ -26,23 +28,41 @@ def _make_settings(**kwargs) -> Settings:
return Settings(**defaults) 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 = MagicMock(spec=User)
user.id = user_id user.id = user_id
message = MagicMock(spec=Message) message = MagicMock(spec=Message)
message.text = text message.text = text
message.reply_text = AsyncMock() 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 = MagicMock(spec=Update)
update.effective_user = user update.effective_user = user
update.effective_chat = chat
update.message = message update.message = message
return update 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 = 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 return ctx
@ -85,12 +105,13 @@ async def test_start_handler_ignores_disallowed_user():
@pytest.mark.asyncio @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() settings = _make_settings()
user_id = 42 user_id = 42
_history[user_id] = [{"role": "user", "content": "old msg"}] _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) ctx = _make_context(settings)
await clear_handler(update, ctx) await clear_handler(update, ctx)
@ -99,6 +120,22 @@ async def test_clear_handler_clears_history():
update.message.reply_text.assert_awaited_once() 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 @pytest.mark.asyncio
async def test_message_handler_calls_llm_and_replies(): async def test_message_handler_calls_llm_and_replies():
settings = _make_settings() settings = _make_settings()
@ -113,7 +150,30 @@ async def test_message_handler_calls_llm_and_replies():
mock_llm.chat.assert_awaited_once() mock_llm.chat.assert_awaited_once()
update.message.reply_text.assert_awaited_once() update.message.reply_text.assert_awaited_once()
# History should now contain the user/assistant turn
assert len(_history[77]) == 2 assert len(_history[77]) == 2
assert _history[77][0]["role"] == "user" assert _history[77][0]["role"] == "user"
assert _history[77][1]["role"] == "assistant" 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

318
tests/test_thread_memory.py Normal file
View File

@ -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()