Introduce a shared ConversationService (steward/bot/core.py) that owns the LLM call, history, knowledge-base search, and thread-memory keying behind a normalized ThreadKey, so both Telegram and Matrix drive the same pipeline. - Add steward/bot/matrix.py: a mautrix-python appservice bot that receives Synapse transactions and replies via the client-server API. - Refactor telegram.py handlers into thin wrappers over ConversationService. - Generalize ThreadMemoryStore/ThreadSummary to platform-scoped keys with legacy chat_id:thread_id migration. - Add a matrix config section (homeserver, tokens, room/user allowlists). - Rewrite main.py as async, starting Telegram and/or Matrix on one event loop. - Add mautrix>=0.21.0 dependency. Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
582 lines
20 KiB
Python
582 lines
20 KiB
Python
"""Tests for thread-aware memory: ThreadMemoryStore and /flush, /recall handlers."""
|
||
|
||
from pathlib import Path
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
import pytest
|
||
from telegram import Chat, Message, Update, User
|
||
from telegram.ext import CallbackContext
|
||
|
||
from steward.bot.core import ConversationService
|
||
from steward.bot.telegram import (
|
||
flush_handler,
|
||
message_handler,
|
||
recall_handler,
|
||
)
|
||
from steward.bot.thread_key import ThreadKey
|
||
from steward.config import Settings
|
||
from steward.llm.client import LLMClient
|
||
from steward.memory.thread_store import ThreadMemoryStore, ThreadSummary
|
||
from tests.conftest import make_settings
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Helpers
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _make_settings(**kwargs) -> Settings:
|
||
"""Wrapper for test compatibility."""
|
||
defaults = dict(telegram_bot_token="tok", openai_api_key="key")
|
||
defaults.update(kwargs)
|
||
return make_settings(**defaults)
|
||
|
||
|
||
def _make_update(
|
||
user_id: int = 1,
|
||
text: str = "hi",
|
||
chat_id: int = 100,
|
||
thread_id: int | None = None,
|
||
) -> Update:
|
||
user = MagicMock(spec=User)
|
||
user.id = user_id
|
||
|
||
message = MagicMock(spec=Message)
|
||
message.text = text
|
||
message.reply_text = AsyncMock()
|
||
message.message_thread_id = thread_id
|
||
|
||
chat = MagicMock(spec=Chat)
|
||
chat.id = chat_id
|
||
|
||
update = MagicMock(spec=Update)
|
||
update.effective_user = user
|
||
update.effective_chat = chat
|
||
update.message = message
|
||
return update
|
||
|
||
|
||
def _make_context(
|
||
settings: Settings,
|
||
llm: LLMClient | None = None,
|
||
store: ThreadMemoryStore | None = None,
|
||
) -> CallbackContext: # type: ignore[type-arg]
|
||
ctx = MagicMock(spec=CallbackContext)
|
||
if store is None:
|
||
mock_store = MagicMock(spec=ThreadMemoryStore)
|
||
mock_store.search.return_value = []
|
||
store = mock_store
|
||
if llm is None:
|
||
llm = MagicMock(spec=LLMClient)
|
||
service = ConversationService(settings, llm, store)
|
||
ctx.bot_data = {
|
||
"settings": settings,
|
||
"service": service,
|
||
"llm": llm,
|
||
}
|
||
ctx.args = []
|
||
return ctx
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# ThreadMemoryStore unit tests
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class TestThreadMemoryStore:
|
||
def _store(self, tmp_path: Path) -> ThreadMemoryStore:
|
||
return ThreadMemoryStore(tmp_path / "mem.json")
|
||
|
||
def _key(self, scope: str = "1", thread: str | None = "2") -> ThreadKey:
|
||
return ThreadKey(platform="telegram", scope=scope, thread=thread)
|
||
|
||
def _summary(self, scope: str = "1", thread: str | None = "2", **kwargs) -> ThreadSummary:
|
||
defaults = dict(summary="A recap.", message_count=5)
|
||
defaults.update(kwargs)
|
||
return ThreadSummary(platform="telegram", scope=scope, thread=thread, **defaults)
|
||
|
||
def test_get_missing_returns_none(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
assert store.get(self._key()) is None
|
||
|
||
def test_save_and_get_roundtrip(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
summary = self._summary()
|
||
store.save(summary)
|
||
|
||
retrieved = store.get(self._key())
|
||
assert retrieved is not None
|
||
assert retrieved.summary == "A recap."
|
||
assert retrieved.message_count == 5
|
||
|
||
def test_save_overwrites_previous_entry(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(self._summary(summary="old", message_count=2))
|
||
store.save(self._summary(summary="new", message_count=4))
|
||
assert store.get(self._key()).summary == "new" # type: ignore[union-attr]
|
||
|
||
def test_persisted_to_disk(self, tmp_path: Path):
|
||
path = tmp_path / "mem.json"
|
||
store = ThreadMemoryStore(path)
|
||
store.save(self._summary(scope="5", thread="9", summary="saved", message_count=1))
|
||
|
||
# Load a fresh store from the same file
|
||
store2 = ThreadMemoryStore(path)
|
||
assert store2.get(self._key("5", "9")) is not None
|
||
|
||
def test_all_returns_newest_first(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(
|
||
self._summary(
|
||
thread="1",
|
||
summary="first",
|
||
message_count=1,
|
||
flushed_at="2026-01-01T00:00:00+00:00",
|
||
)
|
||
)
|
||
store.save(
|
||
self._summary(
|
||
thread="2",
|
||
summary="second",
|
||
message_count=1,
|
||
flushed_at="2026-06-01T00:00:00+00:00",
|
||
)
|
||
)
|
||
results = store.all()
|
||
assert results[0].summary == "second"
|
||
assert results[1].summary == "first"
|
||
|
||
def test_corrupt_file_starts_fresh(self, tmp_path: Path):
|
||
path = tmp_path / "mem.json"
|
||
path.write_text("not-json", encoding="utf-8")
|
||
store = ThreadMemoryStore(path)
|
||
assert store.all() == []
|
||
|
||
def test_format_for_telegram_contains_thread_id(self, tmp_path: Path):
|
||
s = self._summary(thread="77", summary="recap", message_count=3)
|
||
text = s.format_for_telegram()
|
||
assert "77" in text
|
||
assert "recap" in text
|
||
|
||
def test_format_for_telegram_shows_tags(self, tmp_path: Path):
|
||
s = self._summary(thread="77", summary="recap", message_count=3, tags=["api", "auth"])
|
||
text = s.format_for_telegram()
|
||
assert "api" in text
|
||
assert "auth" in text
|
||
|
||
def test_from_dict_tolerates_missing_tags(self, tmp_path: Path):
|
||
"""from_dict must handle legacy entries that predate the tags field."""
|
||
raw = {
|
||
"chat_id": 1,
|
||
"thread_id": 2,
|
||
"summary": "old",
|
||
"message_count": 3,
|
||
"flushed_at": "2026-01-01T00:00:00+00:00",
|
||
}
|
||
s = ThreadSummary.from_dict(raw)
|
||
assert s.tags == []
|
||
|
||
def test_search_returns_matching_summaries(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(
|
||
self._summary(
|
||
thread="1", summary="API work", message_count=2, tags=["api", "design", "auth"]
|
||
)
|
||
)
|
||
store.save(
|
||
self._summary(
|
||
thread="2", summary="Database work", message_count=2, tags=["database", "schema"]
|
||
)
|
||
)
|
||
results = store.search("api")
|
||
assert len(results) == 1
|
||
assert results[0].thread_id == "1"
|
||
|
||
def test_search_no_match_returns_empty(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(self._summary(thread="1", summary="recap", message_count=1, tags=["database"]))
|
||
assert store.search("deployment") == []
|
||
|
||
def test_search_empty_query_returns_empty(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(self._summary(thread="1", summary="recap", message_count=1, tags=["database"]))
|
||
assert store.search("") == []
|
||
|
||
def test_search_case_insensitive(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(self._summary(thread="1", summary="recap", message_count=1, tags=["API"]))
|
||
assert len(store.search("api")) == 1
|
||
|
||
def test_search_skips_untagged_summaries(self, tmp_path: Path):
|
||
store = self._store(tmp_path)
|
||
store.save(self._summary(thread="1", summary="recap", message_count=1, tags=[]))
|
||
assert store.search("api") == []
|
||
|
||
def test_legacy_store_roundtrip(self, tmp_path: Path):
|
||
"""Summaries without tags survive a save/load cycle via from_dict."""
|
||
path = tmp_path / "mem.json"
|
||
import json as _json
|
||
|
||
path.write_text(
|
||
_json.dumps(
|
||
{
|
||
"1:1": {
|
||
"chat_id": 1,
|
||
"thread_id": 1,
|
||
"summary": "old",
|
||
"message_count": 1,
|
||
"flushed_at": "2026-01-01T00:00:00+00:00",
|
||
}
|
||
}
|
||
),
|
||
encoding="utf-8",
|
||
)
|
||
store = ThreadMemoryStore(path)
|
||
s = store.get(self._key("1", "1"))
|
||
assert s is not None
|
||
assert s.tags == []
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Thread-aware message_handler tests
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_thread_message_stored_in_thread_history():
|
||
"""Messages in a thread go to the thread's history, not the user's."""
|
||
settings = _make_settings()
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
mock_llm.chat = AsyncMock(return_value="thread reply")
|
||
|
||
ctx = _make_context(settings, llm=mock_llm)
|
||
service: ConversationService = ctx.bot_data["service"]
|
||
thread_key = ThreadKey(platform="telegram", scope="100", thread="55")
|
||
user_key = ThreadKey(platform="telegram", scope="1")
|
||
|
||
update = _make_update(user_id=1, chat_id=100, thread_id=55, text="thread message")
|
||
await message_handler(update, ctx)
|
||
|
||
assert len(service._histories[thread_key]) == 2
|
||
assert service._histories[thread_key][0]["role"] == "user"
|
||
# Regular user history untouched
|
||
assert not service.has_history(user_key)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_thread_history_is_unbounded():
|
||
"""Thread history never gets trimmed regardless of how many turns there are."""
|
||
from steward.bot.core import _MAX_HISTORY
|
||
|
||
settings = _make_settings()
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
mock_llm.chat = AsyncMock(return_value="reply")
|
||
|
||
ctx = _make_context(settings, llm=mock_llm)
|
||
service: ConversationService = ctx.bot_data["service"]
|
||
thread_key = ThreadKey(platform="telegram", scope="200", thread="66")
|
||
# Pre-fill well beyond the cap used for non-thread history
|
||
service._histories[thread_key] = [
|
||
{"role": "user" if i % 2 == 0 else "assistant", "content": f"m{i}"}
|
||
for i in range(_MAX_HISTORY * 4) # 4× the normal cap
|
||
]
|
||
prior_len = len(service._histories[thread_key])
|
||
|
||
update = _make_update(user_id=2, chat_id=200, thread_id=66, text="more")
|
||
await message_handler(update, ctx)
|
||
|
||
# Should have grown by exactly 2 (user + assistant), never trimmed
|
||
assert len(service._histories[thread_key]) == prior_len + 2
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# /flush handler tests
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_flush_outside_thread_warns():
|
||
"""flush_handler outside a thread should warn the user."""
|
||
settings = _make_settings()
|
||
store = MagicMock(spec=ThreadMemoryStore)
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
mock_llm.chat = AsyncMock(return_value="summary")
|
||
|
||
update = _make_update(thread_id=None) # no thread
|
||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||
await flush_handler(update, ctx)
|
||
|
||
store.save.assert_not_called()
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "flush" in text.lower() or "thread" in text.lower()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_flush_empty_thread_warns():
|
||
"""flush_handler with no history should tell the user there is nothing to flush."""
|
||
settings = _make_settings()
|
||
store = MagicMock(spec=ThreadMemoryStore)
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
mock_llm.chat = AsyncMock(return_value="summary")
|
||
|
||
update = _make_update(chat_id=300, thread_id=88)
|
||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||
await flush_handler(update, ctx)
|
||
|
||
store.save.assert_not_called()
|
||
mock_llm.chat.assert_not_awaited()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_flush_summarises_stores_and_compresses():
|
||
"""flush_handler should summarise, persist, and compress in-memory history."""
|
||
settings = _make_settings()
|
||
store = MagicMock(spec=ThreadMemoryStore)
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
# First call → summary, second call → tags
|
||
mock_llm.chat = AsyncMock(side_effect=["Great summary of the thread.", "api, design, testing"])
|
||
|
||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||
service: ConversationService = ctx.bot_data["service"]
|
||
key = ThreadKey(platform="telegram", scope="400", thread="99")
|
||
service._histories[key] = [
|
||
{"role": "user", "content": "question one"},
|
||
{"role": "assistant", "content": "answer one"},
|
||
{"role": "user", "content": "question two"},
|
||
{"role": "assistant", "content": "answer two"},
|
||
]
|
||
|
||
update = _make_update(chat_id=400, thread_id=99)
|
||
await flush_handler(update, ctx)
|
||
|
||
# LLM called twice: once for summary, once for tags
|
||
assert mock_llm.chat.await_count == 2
|
||
# Summary persisted
|
||
store.save.assert_called_once()
|
||
saved: ThreadSummary = store.save.call_args.args[0]
|
||
assert saved.platform == "telegram"
|
||
assert saved.scope == "400"
|
||
assert saved.thread == "99"
|
||
assert saved.summary == "Great summary of the thread."
|
||
assert saved.message_count == 2 # 2 user turns
|
||
assert saved.tags == ["api", "design", "testing"]
|
||
|
||
# In-memory history replaced with compressed context
|
||
assert len(service._histories[key]) == 1
|
||
assert service._histories[key][0]["role"] == "system"
|
||
assert "Great summary" in service._histories[key][0]["content"]
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# /recall handler tests
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_recall_in_thread_returns_stored_summary(tmp_path):
|
||
"""recall_handler inside a thread returns the stored summary for that thread."""
|
||
settings = _make_settings()
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
store.save(
|
||
ThreadSummary(
|
||
platform="telegram", scope="500", thread="11", summary="recap text", message_count=3
|
||
)
|
||
)
|
||
|
||
update = _make_update(chat_id=500, thread_id=11)
|
||
ctx = _make_context(settings, store=store)
|
||
await recall_handler(update, ctx)
|
||
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "recap text" in text
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_recall_in_thread_no_summary_guides_user(tmp_path):
|
||
"""recall_handler inside a thread with no summary tells user to /flush first."""
|
||
settings = _make_settings()
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
|
||
update = _make_update(chat_id=600, thread_id=22)
|
||
ctx = _make_context(settings, store=store)
|
||
await recall_handler(update, ctx)
|
||
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "flush" in text.lower()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_recall_outside_thread_lists_all(tmp_path):
|
||
"""recall_handler outside a thread lists all stored summaries."""
|
||
settings = _make_settings()
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
store.save(
|
||
ThreadSummary(platform="telegram", scope="1", thread="1", summary="alpha", message_count=1)
|
||
)
|
||
store.save(
|
||
ThreadSummary(platform="telegram", scope="1", thread="2", summary="beta", message_count=2)
|
||
)
|
||
|
||
update = _make_update(thread_id=None) # no thread
|
||
ctx = _make_context(settings, store=store)
|
||
await recall_handler(update, ctx)
|
||
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "alpha" in text or "beta" in text
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_recall_outside_thread_empty_store(tmp_path):
|
||
"""recall_handler outside a thread with no summaries guides user."""
|
||
settings = _make_settings()
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
|
||
update = _make_update(thread_id=None)
|
||
ctx = _make_context(settings, store=store)
|
||
await recall_handler(update, ctx)
|
||
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "flush" in text.lower()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_recall_with_query_returns_matching_summaries(tmp_path):
|
||
"""recall_handler with a query argument searches the knowledge base by keyword."""
|
||
settings = _make_settings()
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
store.save(
|
||
ThreadSummary(
|
||
platform="telegram",
|
||
scope="1",
|
||
thread="1",
|
||
summary="API authentication discussion",
|
||
message_count=2,
|
||
tags=["api", "auth"],
|
||
)
|
||
)
|
||
store.save(
|
||
ThreadSummary(
|
||
platform="telegram",
|
||
scope="1",
|
||
thread="2",
|
||
summary="Database schema planning",
|
||
message_count=3,
|
||
tags=["database", "schema"],
|
||
)
|
||
)
|
||
|
||
update = _make_update(thread_id=None)
|
||
ctx = _make_context(settings, store=store)
|
||
ctx.args = ["api"]
|
||
await recall_handler(update, ctx)
|
||
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "API authentication" in text
|
||
assert "Database" not in text
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_recall_with_query_no_match(tmp_path):
|
||
"""recall_handler with a query that matches nothing tells the user."""
|
||
settings = _make_settings()
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
store.save(
|
||
ThreadSummary(
|
||
platform="telegram",
|
||
scope="1",
|
||
thread="1",
|
||
summary="recap",
|
||
message_count=1,
|
||
tags=["database"],
|
||
)
|
||
)
|
||
|
||
update = _make_update(thread_id=None)
|
||
ctx = _make_context(settings, store=store)
|
||
ctx.args = ["deployment"]
|
||
await recall_handler(update, ctx)
|
||
|
||
update.message.reply_text.assert_awaited_once()
|
||
text = update.message.reply_text.call_args.args[0]
|
||
assert "deployment" in text.lower() or "no memories" in text.lower()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_message_handler_injects_kb_context_transiently(tmp_path):
|
||
"""message_handler injects relevant KB summaries as transient context without storing them."""
|
||
settings = _make_settings()
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
mock_llm.chat = AsyncMock(return_value="reply with context")
|
||
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
store.save(
|
||
ThreadSummary(
|
||
platform="telegram",
|
||
scope="1",
|
||
thread="1",
|
||
summary="Previous API discussion",
|
||
message_count=2,
|
||
tags=["api"],
|
||
)
|
||
)
|
||
|
||
user_id = 999
|
||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||
service: ConversationService = ctx.bot_data["service"]
|
||
user_key = ThreadKey(platform="telegram", scope=str(user_id))
|
||
|
||
update = _make_update(user_id=user_id, text="tell me about the api work")
|
||
await message_handler(update, ctx)
|
||
|
||
# LLM was called
|
||
mock_llm.chat.assert_awaited_once()
|
||
call_kwargs = mock_llm.chat.call_args
|
||
# The history passed to the LLM should contain the KB context message
|
||
history_arg = call_kwargs.kwargs.get("history") or (
|
||
call_kwargs.args[1] if len(call_kwargs.args) > 1 else None
|
||
)
|
||
assert history_arg is not None
|
||
assert any("knowledge base" in m.get("content", "").lower() for m in history_arg)
|
||
|
||
# But the KB message must NOT be stored in history
|
||
assert all(
|
||
"knowledge base" not in m.get("content", "").lower() for m in service._histories[user_key]
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_message_handler_no_kb_injection_when_no_match(tmp_path):
|
||
"""message_handler does not inject KB context when no summaries match."""
|
||
settings = _make_settings()
|
||
mock_llm = MagicMock(spec=LLMClient)
|
||
mock_llm.chat = AsyncMock(return_value="plain reply")
|
||
|
||
store = ThreadMemoryStore(tmp_path / "mem.json")
|
||
# Store a summary with unrelated tags
|
||
store.save(
|
||
ThreadSummary(
|
||
platform="telegram",
|
||
scope="1",
|
||
thread="1",
|
||
summary="Database recap",
|
||
message_count=1,
|
||
tags=["database"],
|
||
)
|
||
)
|
||
|
||
user_id = 888
|
||
ctx = _make_context(settings, llm=mock_llm, store=store)
|
||
|
||
update = _make_update(user_id=user_id, text="what is the weather like?")
|
||
await message_handler(update, ctx)
|
||
|
||
call_kwargs = mock_llm.chat.call_args
|
||
history_arg = call_kwargs.kwargs.get("history") or []
|
||
# No KB context injected
|
||
assert not any("knowledge base" in m.get("content", "").lower() for m in history_arg)
|