mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 16:21:50 +03:00
fix(dingtalk): drain inbound background tasks
This commit is contained in:
@@ -196,7 +196,7 @@ class NanobotDingTalkHandler(_CallbackHandlerBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
self.channel._background_tasks.add(task)
|
self.channel._background_tasks.add(task)
|
||||||
task.add_done_callback(self.channel._background_tasks.discard)
|
task.add_done_callback(self.channel._on_background_task_done)
|
||||||
|
|
||||||
return AckMessage.STATUS_OK, "OK"
|
return AckMessage.STATUS_OK, "OK"
|
||||||
|
|
||||||
@@ -257,6 +257,16 @@ class DingTalkChannel(BaseChannel):
|
|||||||
# Hold references to background tasks to prevent GC
|
# Hold references to background tasks to prevent GC
|
||||||
self._background_tasks: set[asyncio.Task[None]] = set()
|
self._background_tasks: set[asyncio.Task[None]] = set()
|
||||||
|
|
||||||
|
def _on_background_task_done(self, task: asyncio.Task[None]) -> None:
|
||||||
|
self._background_tasks.discard(task)
|
||||||
|
if task.cancelled():
|
||||||
|
return
|
||||||
|
exception = task.exception()
|
||||||
|
if exception is not None:
|
||||||
|
self.logger.opt(exception=exception).error(
|
||||||
|
"DingTalk inbound message task failed"
|
||||||
|
)
|
||||||
|
|
||||||
async def start(self) -> None:
|
async def start(self) -> None:
|
||||||
"""Start the DingTalk bot with Stream Mode."""
|
"""Start the DingTalk bot with Stream Mode."""
|
||||||
current_task = asyncio.current_task()
|
current_task = asyncio.current_task()
|
||||||
@@ -326,8 +336,11 @@ class DingTalkChannel(BaseChannel):
|
|||||||
await self._http.aclose()
|
await self._http.aclose()
|
||||||
self._http = None
|
self._http = None
|
||||||
# Cancel outstanding background tasks
|
# Cancel outstanding background tasks
|
||||||
for task in self._background_tasks:
|
background_tasks = tuple(self._background_tasks)
|
||||||
|
for task in background_tasks:
|
||||||
task.cancel()
|
task.cancel()
|
||||||
|
if background_tasks:
|
||||||
|
await asyncio.gather(*background_tasks, return_exceptions=True)
|
||||||
self._background_tasks.clear()
|
self._background_tasks.clear()
|
||||||
|
|
||||||
async def _close_stream_client(self) -> None:
|
async def _close_stream_client(self) -> None:
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import json
|
|||||||
import zipfile
|
import zipfile
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
@@ -402,6 +402,61 @@ async def test_handler_uses_voice_recognition_text_when_text_is_empty(monkeypatc
|
|||||||
assert msg.chat_id == "group:conv123"
|
assert msg.chat_id == "group:conv123"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_handler_retrieves_background_message_failure(monkeypatch) -> None:
|
||||||
|
bus = MessageBus()
|
||||||
|
channel = DingTalkChannel(
|
||||||
|
DingTalkConfig(client_id="app", client_secret="secret", allow_from=["user1"]),
|
||||||
|
bus,
|
||||||
|
)
|
||||||
|
handler = NanobotDingTalkHandler(channel)
|
||||||
|
failure = RuntimeError("inbound dispatch failed")
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
channel.logger = mock_logger
|
||||||
|
|
||||||
|
class _FakeChatbotMessage:
|
||||||
|
text = SimpleNamespace(content="hello")
|
||||||
|
extensions = {}
|
||||||
|
sender_staff_id = "user1"
|
||||||
|
sender_id = "fallback-user"
|
||||||
|
sender_nick = "Alice"
|
||||||
|
message_type = "text"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def from_dict(_data):
|
||||||
|
return _FakeChatbotMessage()
|
||||||
|
|
||||||
|
async def fail(*_args) -> None:
|
||||||
|
raise failure
|
||||||
|
|
||||||
|
monkeypatch.setattr(dingtalk_module, "ChatbotMessage", _FakeChatbotMessage)
|
||||||
|
monkeypatch.setattr(dingtalk_module, "AckMessage", SimpleNamespace(STATUS_OK="OK"))
|
||||||
|
monkeypatch.setattr(channel, "_on_message", fail)
|
||||||
|
event_loop = asyncio.get_running_loop()
|
||||||
|
previous_handler = event_loop.get_exception_handler()
|
||||||
|
loop_errors: list[dict[str, object]] = []
|
||||||
|
event_loop.set_exception_handler(lambda _loop, context: loop_errors.append(context))
|
||||||
|
|
||||||
|
try:
|
||||||
|
status, body = await handler.process(
|
||||||
|
SimpleNamespace(data={"conversationType": "1", "text": {"content": "hello"}})
|
||||||
|
)
|
||||||
|
for _ in range(10):
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
if not channel._background_tasks:
|
||||||
|
break
|
||||||
|
finally:
|
||||||
|
event_loop.set_exception_handler(previous_handler)
|
||||||
|
|
||||||
|
assert (status, body) == ("OK", "OK")
|
||||||
|
assert not channel._background_tasks
|
||||||
|
assert not loop_errors
|
||||||
|
mock_logger.opt.assert_called_once_with(exception=failure)
|
||||||
|
mock_logger.opt.return_value.error.assert_called_once_with(
|
||||||
|
"DingTalk inbound message task failed"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_handler_processes_file_message(monkeypatch) -> None:
|
async def test_handler_processes_file_message(monkeypatch) -> None:
|
||||||
"""Test that file messages are handled and forwarded with downloaded path."""
|
"""Test that file messages are handled and forwarded with downloaded path."""
|
||||||
@@ -650,6 +705,41 @@ async def test_stop_cancels_stream_client_after_sdk_swallows_first_cancel(monkey
|
|||||||
assert start_task.cancelled()
|
assert start_task.cancelled()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stop_waits_for_background_message_tasks() -> None:
|
||||||
|
channel = DingTalkChannel(
|
||||||
|
DingTalkConfig(client_id="app", client_secret="secret", allow_from=["*"]),
|
||||||
|
MessageBus(),
|
||||||
|
)
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
channel.logger = mock_logger
|
||||||
|
started = asyncio.Event()
|
||||||
|
cancelled = asyncio.Event()
|
||||||
|
|
||||||
|
async def wait_forever() -> None:
|
||||||
|
started.set()
|
||||||
|
try:
|
||||||
|
await asyncio.Future()
|
||||||
|
finally:
|
||||||
|
cancelled.set()
|
||||||
|
|
||||||
|
task = asyncio.create_task(wait_forever())
|
||||||
|
channel._background_tasks.add(task)
|
||||||
|
task.add_done_callback(channel._on_background_task_done)
|
||||||
|
await started.wait()
|
||||||
|
|
||||||
|
try:
|
||||||
|
await channel.stop()
|
||||||
|
assert task.done()
|
||||||
|
assert cancelled.is_set()
|
||||||
|
assert not channel._background_tasks
|
||||||
|
mock_logger.opt.assert_not_called()
|
||||||
|
finally:
|
||||||
|
if not task.done():
|
||||||
|
task.cancel()
|
||||||
|
await asyncio.gather(task, return_exceptions=True)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_download_dingtalk_file(tmp_path, monkeypatch) -> None:
|
async def test_download_dingtalk_file(tmp_path, monkeypatch) -> None:
|
||||||
"""Test the two-step file download flow (get URL then download content)."""
|
"""Test the two-step file download flow (get URL then download content)."""
|
||||||
|
|||||||
Reference in New Issue
Block a user