fix(dingtalk): stop stream task on shutdown

This commit is contained in:
chengyongru 2026-07-02 14:51:03 +08:00 committed by Xubin Ren
parent 8b645135bc
commit c579551bb1
2 changed files with 92 additions and 0 deletions

View File

@ -6,6 +6,8 @@ import mimetypes
import os import os
import time import time
import zipfile import zipfile
from contextlib import suppress
from inspect import isawaitable
from io import BytesIO from io import BytesIO
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@ -204,6 +206,7 @@ class DingTalkChannel(BaseChannel):
self.config: DingTalkConfig = config self.config: DingTalkConfig = config
self._client: Any = None self._client: Any = None
self._http: httpx.AsyncClient | None = None self._http: httpx.AsyncClient | None = None
self._start_task: asyncio.Task | None = None
# Access Token management for sending messages # Access Token management for sending messages
self._access_token: str | None = None self._access_token: str | None = None
@ -214,6 +217,8 @@ class DingTalkChannel(BaseChannel):
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()
self._start_task = current_task
try: try:
if not DINGTALK_AVAILABLE: if not DINGTALK_AVAILABLE:
self.logger.error( self.logger.error(
@ -255,10 +260,25 @@ class DingTalkChannel(BaseChannel):
except Exception: except Exception:
self.logger.exception("Failed to start channel") self.logger.exception("Failed to start channel")
finally:
self._running = False
if self._start_task is current_task:
self._start_task = None
async def stop(self) -> None: async def stop(self) -> None:
"""Stop the DingTalk bot.""" """Stop the DingTalk bot."""
self._running = False self._running = False
await self._close_stream_client()
start_task = self._start_task
if start_task and start_task is not asyncio.current_task() and not start_task.done():
start_task.cancel()
await asyncio.sleep(0)
if not start_task.done():
start_task.cancel()
with suppress(asyncio.CancelledError):
await start_task
self._client = None
# Close the shared HTTP client # Close the shared HTTP client
if self._http: if self._http:
await self._http.aclose() await self._http.aclose()
@ -268,6 +288,23 @@ class DingTalkChannel(BaseChannel):
task.cancel() task.cancel()
self._background_tasks.clear() self._background_tasks.clear()
async def _close_stream_client(self) -> None:
client = self._client
if client is None:
return
close = getattr(client, "close", None)
if close is None:
websocket = getattr(client, "websocket", None)
close = getattr(websocket, "close", None)
if close is None:
return
try:
result = close()
if isawaitable(result):
await result
except Exception:
self.logger.debug("DingTalk stream client close failed", exc_info=True)
async def _get_access_token(self) -> str | None: async def _get_access_token(self) -> str | None:
"""Get or refresh Access Token.""" """Get or refresh Access Token."""
if self._access_token and time.time() < self._token_expiry: if self._access_token and time.time() < self._token_expiry:

View File

@ -403,6 +403,61 @@ async def test_start_configures_http_timeout(monkeypatch) -> None:
await channel.stop() await channel.stop()
@pytest.mark.asyncio
async def test_stop_cancels_stream_client_after_sdk_swallows_first_cancel(monkeypatch) -> None:
channel = DingTalkChannel(
DingTalkConfig(client_id="app", client_secret="secret", allow_from=["*"]),
MessageBus(),
)
created: dict[str, object] = {}
class _FakeWebsocket:
def __init__(self) -> None:
self.closed = False
async def close(self) -> None:
self.closed = True
class _CancelSwallowingStreamClient:
def __init__(self, _credential):
self.websocket = _FakeWebsocket()
self.started = asyncio.Event()
self.cancelled_once = asyncio.Event()
created["client"] = self
def register_callback_handler(self, _topic, _handler):
pass
async def start(self):
self.started.set()
while True:
try:
await asyncio.Future()
except asyncio.CancelledError:
self.cancelled_once.set()
await asyncio.sleep(3600)
monkeypatch.setattr(dingtalk_module, "DINGTALK_AVAILABLE", True)
monkeypatch.setattr(dingtalk_module, "Credential", lambda *a, **k: object())
monkeypatch.setattr(dingtalk_module, "DingTalkStreamClient", _CancelSwallowingStreamClient)
monkeypatch.setattr(dingtalk_module, "ChatbotMessage", SimpleNamespace(TOPIC="topic"))
start_task = asyncio.create_task(channel.start())
while "client" not in created:
await asyncio.sleep(0)
client = created["client"]
await asyncio.wait_for(client.started.wait(), timeout=0.5)
start_task.cancel()
await asyncio.wait_for(client.cancelled_once.wait(), timeout=0.5)
assert not start_task.done()
await asyncio.wait_for(channel.stop(), timeout=0.5)
assert client.websocket.closed is True
assert start_task.cancelled()
@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)."""