steward_mirror/tests/test_bot.py
Daniel 293591e48a
feat: improve Telegram conversation heuristics (#14)
Co-authored-by: openhands <openhands@all-hands.dev>
2026-07-26 21:39:59 +10:00

267 lines
8.3 KiB
Python

"""Tests for steward.bot.telegram (handler logic)."""
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 (
_history,
_is_allowed,
_is_chat_enabled,
_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
from tests.conftest import make_settings
def _make_settings(**kwargs) -> Settings:
"""Wrapper for test compatibility."""
defaults = dict(
telegram_bot_token="test-token",
openai_api_key="test-key",
)
defaults.update(kwargs)
return make_settings(**defaults)
def _make_update(
user_id: int = 12345,
text: str = "hello",
chat_id: int | None = None,
thread_id: int | None = None,
chat_type: str = "private",
) -> Update:
user = MagicMock(spec=User)
user.id = user_id
message = MagicMock(spec=Message)
message.text = text
message.reply_text = AsyncMock()
message.reply_poll = AsyncMock()
message.message_thread_id = thread_id
chat = MagicMock(spec=Chat)
chat.id = chat_id if chat_id is not None else user_id
chat.type = chat_type
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
ctx.bot_data = {
"settings": settings,
"llm": llm,
"thread_store": store,
}
ctx.args = []
return ctx
class TestIsAllowed:
def test_no_allowlist_allows_everyone(self):
settings = _make_settings(telegram_allowed_user_ids=[])
assert _is_allowed(999, settings) is True
def test_allowlist_accepts_known_user(self):
settings = _make_settings(telegram_allowed_user_ids=[1, 2, 3])
assert _is_allowed(2, settings) is True
def test_allowlist_rejects_unknown_user(self):
settings = _make_settings(telegram_allowed_user_ids=[1, 2, 3])
assert _is_allowed(999, settings) is False
class TestIsChatEnabled:
def test_private_chat_ignores_group_allowlist(self):
settings = _make_settings(telegram_group_ids=[-100123])
update = _make_update(user_id=123, chat_type="private")
assert _is_chat_enabled(update.effective_chat, settings) is True
def test_group_chat_accepts_configured_group(self):
settings = _make_settings(telegram_group_ids=[-5308306472])
update = _make_update(chat_id=-5308306472, chat_type="group")
assert _is_chat_enabled(update.effective_chat, settings) is True
def test_group_chat_rejects_unconfigured_group(self):
settings = _make_settings(telegram_group_ids=[-1005308306472])
update = _make_update(chat_id=-5308306472, chat_type="group")
assert _is_chat_enabled(update.effective_chat, settings) is False
@pytest.mark.asyncio
async def test_start_handler_replies(monkeypatch):
settings = _make_settings()
update = _make_update()
ctx = _make_context(settings)
await start_handler(update, ctx)
update.message.reply_text.assert_awaited_once()
text = update.message.reply_text.call_args.args[0]
assert "Steward" in text
@pytest.mark.asyncio
async def test_start_handler_ignores_disallowed_user():
settings = _make_settings(telegram_allowed_user_ids=[9999])
update = _make_update(user_id=1111)
ctx = _make_context(settings)
await start_handler(update, ctx)
update.message.reply_text.assert_not_awaited()
@pytest.mark.asyncio
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) # no thread_id
ctx = _make_context(settings)
await clear_handler(update, ctx)
assert _history[user_id] == []
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()
mock_llm = MagicMock(spec=LLMClient)
mock_llm.chat = AsyncMock(return_value="LLM response")
update = _make_update(user_id=77, text="What is the weather?")
ctx = _make_context(settings, llm=mock_llm)
_history[77].clear()
await message_handler(update, ctx)
mock_llm.chat.assert_awaited_once()
update.message.reply_text.assert_awaited_once()
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_sends_multiple_reply_messages():
settings = _make_settings()
mock_llm = MagicMock(spec=LLMClient)
mock_llm.chat = AsyncMock(
return_value="[MESSAGE]\nFirst thought.\n[MESSAGE]\nWhat outcome do you want?"
)
update = _make_update(user_id=78, text="Can you help with this vague thing?")
ctx = _make_context(settings, llm=mock_llm)
_history[78].clear()
await message_handler(update, ctx)
assert update.message.reply_text.await_count == 2
assert update.message.reply_text.await_args_list[0].args[0] == "First thought."
assert update.message.reply_text.await_args_list[1].args[0] == "What outcome do you want?"
update.message.reply_poll.assert_not_awaited()
assert _history[78][1]["content"] == "First thought.\n\nWhat outcome do you want?"
@pytest.mark.asyncio
async def test_message_handler_sends_native_poll_from_directive():
settings = _make_settings()
mock_llm = MagicMock(spec=LLMClient)
mock_llm.chat = AsyncMock(
return_value=(
"[MESSAGE]\nI can turn that into a vote.\n"
"[POLL]\n"
"question: Which implementation should we choose?\n"
"- Minimal change\n"
"- Full redesign\n"
"[/POLL]\n"
"[MESSAGE]\nI'll use the winning option."
)
)
update = _make_update(user_id=79, text="Should we do a minimal change or redesign?")
ctx = _make_context(settings, llm=mock_llm)
_history[79].clear()
await message_handler(update, ctx)
assert update.message.reply_text.await_count == 2
update.message.reply_poll.assert_awaited_once_with(
question="Which implementation should we choose?",
options=["Minimal change", "Full redesign"],
is_anonymous=False,
)
assert _history[79][1]["content"] == (
"I can turn that into a vote.\n\n"
"Poll: Which implementation should we choose? (Minimal change; Full redesign)\n\n"
"I'll use the winning option."
)
@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