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
+67 -7
View File
@@ -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