feat: thread-aware memory with /flush and /recall commands
This commit is contained in:
parent
8759bea963
commit
c74c8b0d8b
@ -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
|
||||||
|
|||||||
@ -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))
|
||||||
|
|
||||||
|
|||||||
@ -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."""
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
0
steward/memory/__init__.py
Normal file
0
steward/memory/__init__.py
Normal file
84
steward/memory/thread_store.py
Normal file
84
steward/memory/thread_store.py
Normal 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)
|
||||||
@ -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
318
tests/test_thread_memory.py
Normal 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()
|
||||||
Loading…
x
Reference in New Issue
Block a user