diff --git a/messaging/telegram.py b/messaging/telegram.py index 6196840f..1c4e458c 100644 --- a/messaging/telegram.py +++ b/messaging/telegram.py @@ -181,10 +181,12 @@ class TelegramPlatform(MessagingPlatform): ) raise except RetryAfter as e: - # Telegram explicitly tells us to wait + # Telegram explicitly tells us to wait (PTB_TIMEDELTA: retry_after is timedelta) + from datetime import timedelta + retry_after = e.retry_after - if hasattr(retry_after, "total_seconds"): - wait_secs = float(retry_after.total_seconds()) # type: ignore + if isinstance(retry_after, timedelta): + wait_secs = retry_after.total_seconds() else: wait_secs = float(retry_after) @@ -223,11 +225,12 @@ class TelegramPlatform(MessagingPlatform): parse_mode: Optional[str] = "MarkdownV2", ) -> str: """Send a message to a chat.""" - if not self._application or not self._application.bot: + app = self._application + if not app or not app.bot: raise RuntimeError("Telegram application or bot not initialized") async def _do_send(parse_mode=parse_mode): - bot = self._application.bot # type: ignore + bot = app.bot msg = await bot.send_message( chat_id=chat_id, text=text, @@ -246,11 +249,12 @@ class TelegramPlatform(MessagingPlatform): parse_mode: Optional[str] = "MarkdownV2", ) -> None: """Edit an existing message.""" - if not self._application or not self._application.bot: + app = self._application + if not app or not app.bot: raise RuntimeError("Telegram application or bot not initialized") async def _do_edit(parse_mode=parse_mode): - bot = self._application.bot # type: ignore + bot = app.bot await bot.edit_message_text( chat_id=chat_id, message_id=int(message_id), @@ -266,11 +270,12 @@ class TelegramPlatform(MessagingPlatform): message_id: str, ) -> None: """Delete a message from a chat.""" - if not self._application or not self._application.bot: + app = self._application + if not app or not app.bot: raise RuntimeError("Telegram application or bot not initialized") async def _do_delete(): - bot = self._application.bot # type: ignore + bot = app.bot await bot.delete_message(chat_id=chat_id, message_id=int(message_id)) await self._with_retry(_do_delete) @@ -279,11 +284,12 @@ class TelegramPlatform(MessagingPlatform): """Delete multiple messages (best-effort).""" if not message_ids: return - if not self._application or not self._application.bot: + app = self._application + if not app or not app.bot: raise RuntimeError("Telegram application or bot not initialized") # PTB supports bulk deletion via delete_messages; fall back to per-message. - bot = self._application.bot # type: ignore + bot = app.bot if hasattr(bot, "delete_messages"): async def _do_bulk(): @@ -296,7 +302,7 @@ class TelegramPlatform(MessagingPlatform): if not mids: return None # delete_messages accepts a sequence of ints (up to 100). - await bot.delete_messages(chat_id=chat_id, message_ids=mids) # type: ignore[attr-defined] + await bot.delete_messages(chat_id=chat_id, message_ids=mids) await self._with_retry(_do_bulk) return @@ -392,7 +398,7 @@ class TelegramPlatform(MessagingPlatform): def fire_and_forget(self, task: Awaitable[Any]) -> None: """Execute a coroutine without awaiting it.""" if asyncio.iscoroutine(task): - asyncio.create_task(task) # type: ignore + asyncio.create_task(task) else: asyncio.ensure_future(task) diff --git a/messaging/tree_data.py b/messaging/tree_data.py index 210bd6eb..079f813a 100644 --- a/messaging/tree_data.py +++ b/messaging/tree_data.py @@ -261,9 +261,16 @@ class MessageTree: List of node IDs in FIFO order. """ async with self._lock: - # asyncio.Queue stores its items in a deque at _queue. - # We copy it here for safe, consistent reads. - return list(self._queue._queue) + # Drain queue, copy items, then put them back to preserve order. + items: List[str] = [] + while True: + try: + items.append(self._queue.get_nowait()) + except asyncio.QueueEmpty: + break + for item in items: + self._queue.put_nowait(item) + return items def get_queue_size(self) -> int: """Get number of messages waiting in queue.""" diff --git a/tests/test_api.py b/tests/test_api.py index 09f8363b..7ddbc7de 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -3,6 +3,7 @@ from api.app import app from api.dependencies import get_provider from unittest.mock import AsyncMock, MagicMock from providers.nvidia_nim import NvidiaNimProvider +from providers.exceptions import APIError # Mock provider mock_provider = MagicMock(spec=NvidiaNimProvider) @@ -121,8 +122,7 @@ def test_generic_exception_returns_500(): def test_generic_exception_with_status_code(): """Exception with status_code attribute uses that status.""" - exc = RuntimeError("bad gateway") - exc.status_code = 502 + exc = APIError("bad gateway", status_code=502) mock_provider.complete.side_effect = exc response = client.post( "/v1/messages", diff --git a/tests/test_config.py b/tests/test_config.py index efa123db..db1f5d48 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -151,8 +151,10 @@ class TestNimSettingsInvalidBounds: NimSettings(min_tokens=-1) def test_reasoning_effort_invalid(self): + from typing import Any, cast + with pytest.raises(ValidationError): - NimSettings(reasoning_effort="invalid") + NimSettings(reasoning_effort=cast(Any, "invalid")) class TestNimSettingsValidators: @@ -187,8 +189,10 @@ class TestNimSettingsValidators: def test_extra_forbid_rejects_unknown_field(self): """NimSettings with extra='forbid' rejects unknown fields.""" + from typing import Any, cast + with pytest.raises(ValidationError): - NimSettings(unknown_field="value") + NimSettings(**cast(Any, {"unknown_field": "value"})) class TestSettingsOptionalStr: diff --git a/tests/test_dependencies.py b/tests/test_dependencies.py index ed4fbc84..cdc0160d 100644 --- a/tests/test_dependencies.py +++ b/tests/test_dependencies.py @@ -53,6 +53,7 @@ async def test_cleanup_provider(): mock_settings.return_value = _make_mock_settings() provider = get_provider() + assert isinstance(provider, NvidiaNimProvider) provider._client = AsyncMock() await cleanup_provider() @@ -90,6 +91,7 @@ async def test_cleanup_provider_aclose_raises(): mock_settings.return_value = _make_mock_settings() provider = get_provider() + assert isinstance(provider, NvidiaNimProvider) provider._client = AsyncMock() provider._client.aclose = AsyncMock(side_effect=RuntimeError("cleanup failed")) diff --git a/tests/test_handler_markdown_and_status_edges.py b/tests/test_handler_markdown_and_status_edges.py index ee7f7cc1..86f18ed7 100644 --- a/tests/test_handler_markdown_and_status_edges.py +++ b/tests/test_handler_markdown_and_status_edges.py @@ -1,5 +1,5 @@ import pytest -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch from messaging.handler import ClaudeMessageHandler from messaging.telegram_markdown import render_markdown_to_mdv2 @@ -79,14 +79,16 @@ def test_get_initial_status_branches(): session_store = MagicMock() handler = ClaudeMessageHandler(platform, cli_manager, session_store) - handler.tree_queue.is_node_tree_busy = MagicMock(return_value=True) - handler.tree_queue.get_queue_size = MagicMock(return_value=2) - s1 = handler._get_initial_status(tree=object(), parent_node_id="p") + with ( + patch.object(handler.tree_queue, "is_node_tree_busy", MagicMock(return_value=True)), + patch.object(handler.tree_queue, "get_queue_size", MagicMock(return_value=2)), + ): + s1 = handler._get_initial_status(tree=object(), parent_node_id="p") assert "Queued" in s1 assert "position 3" in s1 or "position 3" in s1.replace("\\", "") - handler.tree_queue.is_node_tree_busy = MagicMock(return_value=False) - s2 = handler._get_initial_status(tree=object(), parent_node_id="p") + with patch.object(handler.tree_queue, "is_node_tree_busy", MagicMock(return_value=False)): + s2 = handler._get_initial_status(tree=object(), parent_node_id="p") assert "Continuing" in s2 cli_manager.get_stats.return_value = {"active_sessions": 10, "max_sessions": 10} @@ -149,18 +151,17 @@ async def test_process_node_session_limit_marks_error_and_updates_ui(): fake_tree = MagicMock() fake_tree.update_state = AsyncMock() - handler.tree_queue.get_tree_for_node = MagicMock(return_value=fake_tree) + with patch.object(handler.tree_queue, "get_tree_for_node", MagicMock(return_value=fake_tree)): + incoming = IncomingMessage( + text="hi", + chat_id="c", + user_id="u", + message_id="n1", + platform="telegram", + ) + node = MessageNode(node_id="n1", incoming=incoming, status_message_id="s1") - incoming = IncomingMessage( - text="hi", - chat_id="c", - user_id="u", - message_id="n1", - platform="telegram", - ) - node = MessageNode(node_id="n1", incoming=incoming, status_message_id="s1") - - await handler._process_node("n1", node) + await handler._process_node("n1", node) assert platform.queue_edit_message.await_count >= 1 fake_tree.update_state.assert_awaited() @@ -189,13 +190,14 @@ async def test_stop_all_tasks_saves_tree_for_cancelled_nodes(): ) node = MessageNode(node_id="n1", incoming=incoming, status_message_id="s1") - handler.tree_queue.cancel_all = AsyncMock(return_value=[node]) tree = MagicMock() tree.root_id = "root" tree.to_dict = MagicMock(return_value={"root": "ok"}) - handler.tree_queue.get_tree_for_node = MagicMock(return_value=tree) - - count = await handler.stop_all_tasks() + with ( + patch.object(handler.tree_queue, "cancel_all", AsyncMock(return_value=[node])), + patch.object(handler.tree_queue, "get_tree_for_node", MagicMock(return_value=tree)), + ): + count = await handler.stop_all_tasks() assert count == 1 cli_manager.stop_all.assert_awaited_once() session_store.save_tree.assert_called_once_with("root", {"root": "ok"}) diff --git a/tests/test_response_models.py b/tests/test_response_models.py index dea653bd..0cab9b41 100644 --- a/tests/test_response_models.py +++ b/tests/test_response_models.py @@ -82,8 +82,10 @@ class TestMessagesResponse: usage=Usage(input_tokens=1, output_tokens=1), ) assert len(resp.content) == 1 - assert resp.content[0].type == "text" - assert resp.content[0].text == "response" + block = resp.content[0] + assert isinstance(block, ContentBlockText) + assert block.type == "text" + assert block.text == "response" def test_with_tool_use_content(self): resp = MessagesResponse( @@ -100,8 +102,10 @@ class TestMessagesResponse: usage=Usage(input_tokens=1, output_tokens=1), stop_reason="tool_use", ) - assert resp.content[0].type == "tool_use" - assert resp.content[0].name == "Read" + block = resp.content[0] + assert isinstance(block, ContentBlockToolUse) + assert block.type == "tool_use" + assert block.name == "Read" assert resp.stop_reason == "tool_use" def test_with_thinking_content(self): @@ -115,9 +119,13 @@ class TestMessagesResponse: usage=Usage(input_tokens=5, output_tokens=10), ) assert len(resp.content) == 2 - assert resp.content[0].type == "thinking" - assert resp.content[0].thinking == "Let me reason..." - assert resp.content[1].type == "text" + block0 = resp.content[0] + assert isinstance(block0, ContentBlockThinking) + assert block0.type == "thinking" + assert block0.thinking == "Let me reason..." + block1 = resp.content[1] + assert isinstance(block1, ContentBlockText) + assert block1.type == "text" def test_with_all_content_types(self): resp = MessagesResponse( @@ -143,11 +151,21 @@ class TestMessagesResponse: content=[{"type": "custom", "data": "value"}], usage=Usage(input_tokens=1, output_tokens=1), ) - assert resp.content[0]["type"] == "custom" + block = resp.content[0] + assert isinstance(block, dict) + assert block["type"] == "custom" def test_stop_reason_values(self): """All valid stop_reason values should be accepted.""" - for reason in ["end_turn", "max_tokens", "stop_sequence", "tool_use"]: + from typing import Literal + + reasons: list[Literal["end_turn", "max_tokens", "stop_sequence", "tool_use"]] = [ + "end_turn", + "max_tokens", + "stop_sequence", + "tool_use", + ] + for reason in reasons: resp = MessagesResponse( id="msg", model="model", diff --git a/tests/test_restart_reply_restore.py b/tests/test_restart_reply_restore.py index 6a76e944..db52c354 100644 --- a/tests/test_restart_reply_restore.py +++ b/tests/test_restart_reply_restore.py @@ -40,8 +40,6 @@ async def test_reply_to_old_status_message_after_restore_routes_to_parent( ) # Prevent background task scheduling; we only want to validate routing/tree mutation. - handler2.tree_queue.enqueue = AsyncMock(return_value=False) - mock_platform.queue_send_message = AsyncMock(return_value="status_reply") reply = IncomingMessage( @@ -53,7 +51,8 @@ async def test_reply_to_old_status_message_after_restore_routes_to_parent( reply_to_message_id="status_A", ) - await handler2.handle_message(reply) + with patch.object(handler2.tree_queue, "enqueue", AsyncMock(return_value=False)): + await handler2.handle_message(reply) restored_tree = handler2.tree_queue.get_tree_for_node("A") assert restored_tree is not None @@ -90,7 +89,6 @@ async def test_reply_to_old_status_message_without_mapping_creates_new_conversat queue_update_callback=handler2._update_queue_positions, node_started_callback=handler2._mark_node_processing, ) - handler2.tree_queue.enqueue = AsyncMock(return_value=False) mock_platform.queue_send_message = AsyncMock(return_value="status_reply") reply = IncomingMessage( @@ -102,7 +100,8 @@ async def test_reply_to_old_status_message_without_mapping_creates_new_conversat reply_to_message_id="status_A", ) - await handler2.handle_message(reply) + with patch.object(handler2.tree_queue, "enqueue", AsyncMock(return_value=False)): + await handler2.handle_message(reply) # Since the mapping is missing, this should be treated as a new conversation. new_tree = handler2.tree_queue.get_tree_for_node("R1") diff --git a/tests/test_server_module.py b/tests/test_server_module.py index 8f1f602d..386ab32b 100644 --- a/tests/test_server_module.py +++ b/tests/test_server_module.py @@ -8,25 +8,25 @@ def test_server_module_exports_app_and_create_app(): def test_server_main_invokes_uvicorn_run(monkeypatch): import runpy from types import SimpleNamespace - from unittest.mock import MagicMock + from unittest.mock import MagicMock, patch import config.settings as settings_mod import uvicorn as uvicorn_mod # Patch settings used by server.__main__ block. old_get_settings = settings_mod.get_settings - settings_mod.get_settings = lambda: SimpleNamespace(host="127.0.0.1", port=9999) - - old_run = uvicorn_mod.run - uvicorn_mod.run = MagicMock() + mock_settings = SimpleNamespace(host="127.0.0.1", port=9999) try: - runpy.run_module("server", run_name="__main__") - uvicorn_mod.run.assert_called_once() - _, kwargs = uvicorn_mod.run.call_args - assert kwargs["host"] == "127.0.0.1" - assert kwargs["port"] == 9999 - assert kwargs["log_level"] == "debug" + with ( + patch.object(settings_mod, "get_settings", lambda: mock_settings), + patch.object(uvicorn_mod, "run") as mock_run, + ): + runpy.run_module("server", run_name="__main__") + mock_run.assert_called_once() + call_kwargs = mock_run.call_args[1] + assert call_kwargs["host"] == "127.0.0.1" + assert call_kwargs["port"] == 9999 + assert call_kwargs["log_level"] == "debug" finally: - uvicorn_mod.run = old_run settings_mod.get_settings = old_get_settings diff --git a/tests/test_telegram_edge_cases.py b/tests/test_telegram_edge_cases.py index 35348cb7..26846fd0 100644 --- a/tests/test_telegram_edge_cases.py +++ b/tests/test_telegram_edge_cases.py @@ -121,9 +121,10 @@ async def test_queue_send_message_without_limiter_calls_send_message(): platform = TelegramPlatform(bot_token="t") platform._limiter = None - platform.send_message = AsyncMock(return_value="1") - assert await platform.queue_send_message("c", "t") == "1" - platform.send_message.assert_awaited_once() + with patch.object(platform, "send_message", new_callable=AsyncMock) as mock_send: + mock_send.return_value = "1" + assert await platform.queue_send_message("c", "t") == "1" + mock_send.assert_awaited_once() @pytest.mark.asyncio @@ -133,9 +134,9 @@ async def test_queue_edit_message_without_limiter_calls_edit_message(): platform = TelegramPlatform(bot_token="t") platform._limiter = None - platform.edit_message = AsyncMock() - await platform.queue_edit_message("c", "1", "t") - platform.edit_message.assert_awaited_once() + with patch.object(platform, "edit_message", new_callable=AsyncMock) as mock_edit: + await platform.queue_edit_message("c", "1", "t") + mock_edit.assert_awaited_once() def test_fire_and_forget_non_coroutine_uses_ensure_future(monkeypatch): @@ -157,14 +158,13 @@ async def test_on_start_command_replies_and_forwards(): from messaging.telegram import TelegramPlatform platform = TelegramPlatform(bot_token="t") - platform._on_telegram_message = AsyncMock() + with patch.object(platform, "_on_telegram_message", new_callable=AsyncMock) as mock_msg: + update = MagicMock() + update.message.reply_text = AsyncMock() - update = MagicMock() - update.message.reply_text = AsyncMock() - - await platform._on_start_command(update, MagicMock()) - update.message.reply_text.assert_awaited_once() - platform._on_telegram_message.assert_awaited_once() + await platform._on_start_command(update, MagicMock()) + update.message.reply_text.assert_awaited_once() + mock_msg.assert_awaited_once() @pytest.mark.asyncio @@ -173,22 +173,21 @@ async def test_on_telegram_message_handler_error_sends_error_message(): from messaging.telegram import TelegramPlatform platform = TelegramPlatform(bot_token="t", allowed_user_id="123") - platform.send_message = AsyncMock() + with patch.object(platform, "send_message", new_callable=AsyncMock) as mock_send: + async def _boom(_incoming): + raise RuntimeError("bad") - async def _boom(_incoming): - raise RuntimeError("bad") + platform.on_message(_boom) - platform.on_message(_boom) + update = MagicMock() + update.message.text = "hello" + update.message.message_id = 7 + update.message.reply_to_message = None + update.effective_user.id = 123 + update.effective_chat.id = 456 - update = MagicMock() - update.message.text = "hello" - update.message.message_id = 7 - update.message.reply_to_message = None - update.effective_user.id = 123 - update.effective_chat.id = 456 - - await platform._on_telegram_message(update, MagicMock()) - platform.send_message.assert_awaited_once() + await platform._on_telegram_message(update, MagicMock()) + mock_send.assert_awaited_once() @pytest.mark.asyncio diff --git a/tests/test_tree_concurrency.py b/tests/test_tree_concurrency.py index 84bf0689..fb1e76df 100644 --- a/tests/test_tree_concurrency.py +++ b/tests/test_tree_concurrency.py @@ -112,6 +112,7 @@ class TestMessageTreeConcurrency: for i in range(5): node = tree.get_node(f"n{i}") + assert node is not None assert node.state == MessageState.IN_PROGRESS @pytest.mark.asyncio @@ -321,6 +322,7 @@ class TestTreeQueueManagerConcurrency: results = await asyncio.gather(*[add_reply(i) for i in range(5)]) assert len(results) == 5 tree = mgr.get_tree("root") + assert tree is not None assert len(tree.all_nodes()) == 6 # root + 5 replies @pytest.mark.asyncio @@ -444,7 +446,9 @@ class TestTreeQueueManagerConcurrency: assert count == 2 root = tree.get_node("root") + assert root is not None assert root.state == MessageState.ERROR + assert root.error_message is not None assert "restart" in root.error_message @pytest.mark.asyncio @@ -459,6 +463,7 @@ class TestTreeQueueManagerConcurrency: # root + c1 + c2 should all be marked assert len(affected) >= 1 root = tree.get_node("root") + assert root is not None assert root.state == MessageState.ERROR @pytest.mark.asyncio