"""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.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_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