steward_mirror/tests/test_thread_memory.py
Andrew Ridgway 76777c98eb
feat: add Matrix appservice bot and platform-agnostic conversation core
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>
2026-08-18 21:38:01 +10:00

582 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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