Fix all ty check errors without using type: ignore

- messaging/telegram.py: Remove unused type: ignore, fix retry_after typing
  with isinstance(timedelta), use local app variable for None narrowing
- messaging/tree_data.py: Replace _queue._queue access with drain-and-restore
  approach for get_queue_snapshot (avoids private API)
- tests/test_api.py: Use APIError instead of RuntimeError for status_code test
- tests/test_config.py: Use cast(Any, ...) for invalid validation tests
- tests/test_dependencies.py: Add isinstance check for NvidiaNimProvider
- tests/test_handler_markdown_and_status_edges.py: Use patch.object for
  tree_queue method mocks
- tests/test_response_models.py: Add isinstance narrowing for content blocks,
  use Literal list for stop_reason parametrization
- tests/test_restart_reply_restore.py: Use patch.object for enqueue mock
- tests/test_server_module.py: Use patch.object for uvicorn.run and
  get_settings
- tests/test_telegram_edge_cases.py: Use patch.object for method mocks
- tests/test_tree_concurrency.py: Add None assertions for get_node/get_tree

Co-authored-by: Ali Khokhar <alishahryar2@gmail.com>
This commit is contained in:
Cursor Agent
2026-02-15 02:28:33 +00:00
parent d0ea0a450f
commit 1cface7f8c
11 changed files with 135 additions and 93 deletions
+19 -13
View File
@@ -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)
+10 -3
View File
@@ -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."""
+2 -2
View File
@@ -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",
+6 -2
View File
@@ -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:
+2
View File
@@ -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"))
+23 -21
View File
@@ -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"})
+27 -9
View File
@@ -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",
+4 -5
View File
@@ -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")
+12 -12
View File
@@ -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
+25 -26
View File
@@ -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
+5
View File
@@ -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