From 38ff44778a0d11d71597c2531bfd600df855307c Mon Sep 17 00:00:00 2001 From: AoHanBei Date: Tue, 11 Aug 2026 22:27:11 +0800 Subject: [PATCH] fix(wecom): serialize websocket shutdown (#4762) * fix(wecom): await connection task shutdown * fix(wecom): serialize websocket shutdown --- backend/AGENTS.md | 1 + backend/app/channels/wecom.py | 139 ++++++++++++------- backend/tests/test_channels.py | 237 +++++++++++++++++++++++++++++++++ 3 files changed, 332 insertions(+), 45 deletions(-) diff --git a/backend/AGENTS.md b/backend/AGENTS.md index ce3b28a84..01b33dbc1 100644 --- a/backend/AGENTS.md +++ b/backend/AGENTS.md @@ -1022,6 +1022,7 @@ The cached value is reused for both the blocking (`runs.wait`) and streaming (`_ - No public IP, OAuth callback URL, or provider webhook route is required by the current implementation. - Telegram uses a deep-link `/start ` flow over the existing long-polling worker. Slack, Discord, Feishu/Lark, DingTalk, WeChat, and WeCom use `/connect ` over their existing outbound channel workers. - WeChat timing settings (`polling_timeout`, `polling_retry_delay`, `qrcode_poll_interval`, `qrcode_poll_timeout`) accept only positive finite seconds; invalid values fall back to their defaults so polling cannot enter a hot loop or sleep forever. +- WeCom serializes `start()` and `stop()` for each channel instance. The SDK `connect()` task covers connection setup only; after the handshake, the SDK owns a separate receive task. Shutdown cancels an in-progress connection attempt and awaits the SDK's actual asynchronous receive-task/socket cleanup before releasing lifecycle state or allowing a restart. Cancellation of `stop()` still propagates, but only after owned cleanup finishes and lifecycle references are cleared; real connection failures remain reported by `_on_ws_task_done`. - Frontend APIs: `GET /api/channels/providers`, `GET /api/channels/connections`, `POST /api/channels/{provider}/connect`, and `DELETE /api/channels/connections/{connection_id}`. - Browser APIs remain protected by normal Gateway auth/CSRF. Provider messages arrive through the already-configured channel workers. - Provider-level `connection_status` reflects the user's newest connection row. With no binding it is `not_connected`, except in auth-disabled local mode where a configured running channel reports `connected` because all channel messages already route to the default user. diff --git a/backend/app/channels/wecom.py b/backend/app/channels/wecom.py index db0addefa..a92d3fe22 100644 --- a/backend/app/channels/wecom.py +++ b/backend/app/channels/wecom.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio import base64 import hashlib +import inspect import logging from collections.abc import Awaitable, Callable from typing import Any, cast @@ -40,6 +41,8 @@ class WeComChannel(Channel): self._bot_secret: str | None = None self._ws_client = None self._ws_task: asyncio.Task | None = None + self._ws_shutdown_task: asyncio.Future[Any] | None = None + self._lifecycle_lock = asyncio.Lock() self._ws_frames: dict[str, dict[str, Any]] = {} self._ws_stream_ids: dict[str, str] = {} self._working_message = "Working on it..." @@ -67,40 +70,41 @@ class WeComChannel(Channel): return await send_reply_async(req_id, body, cmd) async def start(self) -> None: - if self._running: - return + async with self._lifecycle_lock: + if self._running: + return - bot_id = self.config.get("bot_id") - bot_secret = self.config.get("bot_secret") - working_message = self.config.get("working_message") + bot_id = self.config.get("bot_id") + bot_secret = self.config.get("bot_secret") + working_message = self.config.get("working_message") - self._bot_id = bot_id if isinstance(bot_id, str) and bot_id else None - self._bot_secret = bot_secret if isinstance(bot_secret, str) and bot_secret else None - self._working_message = working_message if isinstance(working_message, str) and working_message else "Working on it..." + self._bot_id = bot_id if isinstance(bot_id, str) and bot_id else None + self._bot_secret = bot_secret if isinstance(bot_secret, str) and bot_secret else None + self._working_message = working_message if isinstance(working_message, str) and working_message else "Working on it..." - if not self._bot_id or not self._bot_secret: - logger.error("WeCom channel requires bot_id and bot_secret") - return + if not self._bot_id or not self._bot_secret: + logger.error("WeCom channel requires bot_id and bot_secret") + return - try: - from aibot import WSClient, WSClientOptions - except ImportError: - logger.error("wecom-aibot-python-sdk is not installed. Install it with: uv add wecom-aibot-python-sdk") - return - else: - self._ws_client = WSClient(WSClientOptions(bot_id=self._bot_id, secret=self._bot_secret, logger=logger)) - self._ws_client.on("message.text", self._on_ws_text) - self._ws_client.on("message.mixed", self._on_ws_mixed) - self._ws_client.on("message.image", self._on_ws_image) - self._ws_client.on("message.file", self._on_ws_file) - self._ws_client.on("error", self._on_ws_error) - self._ws_client.on("disconnected", self._on_ws_disconnected) - self._ws_task = asyncio.create_task(self._ws_client.connect()) - self._ws_task.add_done_callback(self._on_ws_task_done) + try: + from aibot import WSClient, WSClientOptions + except ImportError: + logger.error("wecom-aibot-python-sdk is not installed. Install it with: uv add wecom-aibot-python-sdk") + return + else: + self._ws_client = WSClient(WSClientOptions(bot_id=self._bot_id, secret=self._bot_secret, logger=logger)) + self._ws_client.on("message.text", self._on_ws_text) + self._ws_client.on("message.mixed", self._on_ws_mixed) + self._ws_client.on("message.image", self._on_ws_image) + self._ws_client.on("message.file", self._on_ws_file) + self._ws_client.on("error", self._on_ws_error) + self._ws_client.on("disconnected", self._on_ws_disconnected) + self._ws_task = asyncio.create_task(self._ws_client.connect()) + self._ws_task.add_done_callback(self._on_ws_task_done) - self._running = True - self.bus.subscribe_outbound(self._on_outbound) - logger.info("WeCom channel started") + self._running = True + self.bus.subscribe_outbound(self._on_outbound) + logger.info("WeCom channel started") def _on_ws_task_done(self, task: asyncio.Task) -> None: if task.cancelled(): @@ -120,24 +124,69 @@ class WeComChannel(Channel): detail = f" ({args[0]})" if args else "" logger.warning("WeCom WebSocket disconnected%s; SDK will attempt to reconnect", detail) + def _begin_ws_shutdown(self, ws_client: Any) -> asyncio.Future[Any] | None: + ws_manager = getattr(ws_client, "_ws_manager", None) + async_disconnect = getattr(ws_manager, "_async_disconnect", None) + stop_heartbeat = getattr(ws_manager, "_stop_heartbeat", None) + clear_pending_messages = getattr(ws_manager, "_clear_pending_messages", None) + + if inspect.iscoroutinefunction(async_disconnect) and callable(stop_heartbeat) and callable(clear_pending_messages): + # wecom-aibot-python-sdk 1.0.2 makes disconnect() synchronous and + # discards the task created for _async_disconnect(). Perform its + # synchronous bookkeeping here so DeerFlow can own and await the + # actual SDK shutdown operation without scheduling a duplicate. + try: + if hasattr(ws_client, "_started"): + ws_client._started = False + ws_manager._is_manual_close = True + stop_heartbeat() + clear_pending_messages("Connection manually closed") + except Exception: + logger.exception("Failed to prepare WeCom WebSocket shutdown") + return asyncio.create_task(async_disconnect()) + + try: + result = ws_client.disconnect() + except Exception: + logger.exception("Failed to request WeCom WebSocket disconnect") + return None + if inspect.isawaitable(result): + return asyncio.ensure_future(result) + return None + async def stop(self) -> None: - self._running = False - self.bus.unsubscribe_outbound(self._on_outbound) - if self._ws_task: + async with self._lifecycle_lock: + self._running = False + self.bus.unsubscribe_outbound(self._on_outbound) + ws_client = self._ws_client + ws_task = self._ws_task + if ws_task and not ws_task.done(): + ws_task.cancel() + + shutdown_task = None try: - self._ws_task.cancel() - except Exception: - pass - self._ws_task = None - if self._ws_client: - try: - self._ws_client.disconnect() - except Exception: - pass - self._ws_client = None - self._ws_frames.clear() - self._ws_stream_ids.clear() - logger.info("WeCom channel stopped") + shutdown_task = self._begin_ws_shutdown(ws_client) if ws_client else None + self._ws_shutdown_task = shutdown_task + tasks = [task for task in (ws_task, shutdown_task) if task is not None] + drain_future = asyncio.gather(*tasks, return_exceptions=True) if tasks else None + if drain_future is not None: + try: + await asyncio.shield(drain_future) + except asyncio.CancelledError: + # Caller cancellation still propagates, but only after + # the channel-owned tasks have completed their cleanup. + await drain_future + raise + finally: + if self._ws_task is ws_task: + self._ws_task = None + if self._ws_client is ws_client: + self._ws_client = None + if self._ws_shutdown_task is shutdown_task: + self._ws_shutdown_task = None + self._ws_frames.clear() + self._ws_stream_ids.clear() + logger.info("WeCom channel stopped") async def send(self, msg: OutboundMessage, *, _max_retries: int = 3) -> None: if self._ws_client: diff --git a/backend/tests/test_channels.py b/backend/tests/test_channels.py index 745a7aadc..9d3378fc3 100644 --- a/backend/tests/test_channels.py +++ b/backend/tests/test_channels.py @@ -6603,7 +6603,244 @@ class TestFeishuCardSuccessChecks: _run(go()) +class _ControlledWeComManager: + def __init__(self, shutdown_started: asyncio.Event, release_shutdown: asyncio.Event, shutdown_finished: asyncio.Event) -> None: + self._ws = object() + self._is_manual_close = False + self.shutdown_started = shutdown_started + self.release_shutdown = release_shutdown + self.shutdown_finished = shutdown_finished + self.shutdown_tasks: list[asyncio.Task] = [] + self.heartbeat_stopped = False + self.pending_messages_cleared = False + + def _stop_heartbeat(self) -> None: + self.heartbeat_stopped = True + + def _clear_pending_messages(self, _reason: str) -> None: + self.pending_messages_cleared = True + + def disconnect(self) -> None: + self._is_manual_close = True + self._stop_heartbeat() + self._clear_pending_messages("Connection manually closed") + if self._ws: + asyncio.ensure_future(self._async_disconnect()) + + async def _async_disconnect(self) -> None: + shutdown_task = asyncio.current_task() + assert shutdown_task is not None + self.shutdown_tasks.append(shutdown_task) + self.shutdown_started.set() + try: + await self.release_shutdown.wait() + finally: + self._ws = None + self.shutdown_finished.set() + + +class _ControlledWeComClient: + def __init__(self, manager: _ControlledWeComManager) -> None: + self._ws_manager = manager + self._started = False + self.connect_started = asyncio.Event() + + def on(self, *_args) -> None: + pass + + async def connect(self): + self._started = True + self.connect_started.set() + return self + + def disconnect(self) -> None: + if not self._started: + return + self._started = False + self._ws_manager.disconnect() + + +async def _wait_for_next_event_loop_turn() -> None: + checkpoint = asyncio.get_running_loop().create_future() + asyncio.get_running_loop().call_soon(checkpoint.set_result, None) + await checkpoint + + class TestWeComChannel: + def test_stop_waits_for_connection_task_cancellation(self): + from app.channels.wecom import WeComChannel + + async def go(): + channel = WeComChannel(MessageBus(), config={}) + connection_started = asyncio.Event() + cancellation_finished = asyncio.Event() + + async def connect(): + connection_started.set() + try: + await asyncio.Future() + finally: + cancellation_finished.set() + + connection_task = asyncio.create_task(connect()) + channel._running = True + channel._ws_client = SimpleNamespace(disconnect=MagicMock()) + channel._ws_task = connection_task + await connection_started.wait() + + try: + await channel.stop() + + assert connection_task.done() + assert cancellation_finished.is_set() + assert channel._ws_task is None + finally: + if not connection_task.done(): + connection_task.cancel() + await asyncio.gather(connection_task, return_exceptions=True) + + _run(go()) + + def test_stop_waits_for_sdk_shutdown_after_connect_returns(self): + from app.channels.wecom import WeComChannel + + async def go(): + shutdown_started = asyncio.Event() + release_shutdown = asyncio.Event() + shutdown_finished = asyncio.Event() + manager = _ControlledWeComManager(shutdown_started, release_shutdown, shutdown_finished) + client = _ControlledWeComClient(manager) + channel = WeComChannel(MessageBus(), config={}) + connect_task = asyncio.create_task(client.connect()) + await connect_task + channel._running = True + channel._ws_client = client + channel._ws_task = connect_task + + stop_task = asyncio.create_task(channel.stop()) + await shutdown_started.wait() + await _wait_for_next_event_loop_turn() + + try: + assert not stop_task.done() + finally: + release_shutdown.set() + await asyncio.gather(stop_task, *manager.shutdown_tasks, return_exceptions=True) + + assert shutdown_finished.is_set() + assert len(manager.shutdown_tasks) == 1 + assert manager.heartbeat_stopped + assert manager.pending_messages_cleared + assert not client._started + assert channel._ws_client is None + assert channel._ws_task is None + assert channel._ws_shutdown_task is None + + _run(go()) + + def test_concurrent_start_waits_for_stop_before_installing_new_client(self, monkeypatch): + from app.channels.wecom import WeComChannel + + async def go(): + old_cancellation_started = asyncio.Event() + release_old_cancellation = asyncio.Event() + start_attempted = asyncio.Event() + new_client = MagicMock() + new_client.connect_started = asyncio.Event() + + async def connect_new_client(): + new_client.connect_started.set() + return new_client + + new_client.connect = connect_new_client + monkeypatch.setitem( + __import__("sys").modules, + "aibot", + SimpleNamespace( + WSClient=lambda _options: new_client, + WSClientOptions=lambda **kwargs: SimpleNamespace(**kwargs), + ), + ) + + async def connect_old_client(): + try: + await asyncio.Future() + except asyncio.CancelledError: + old_cancellation_started.set() + await release_old_cancellation.wait() + raise + + old_task = asyncio.create_task(connect_old_client()) + channel = WeComChannel(MessageBus(), config={"bot_id": "bot", "bot_secret": "secret"}) + channel._running = True + channel._ws_client = SimpleNamespace(disconnect=MagicMock()) + channel._ws_task = old_task + + stop_task = asyncio.create_task(channel.stop()) + await old_cancellation_started.wait() + + async def start_concurrently(): + start_attempted.set() + await channel.start() + + start_task = asyncio.create_task(start_concurrently()) + await start_attempted.wait() + release_old_cancellation.set() + await asyncio.gather(stop_task, start_task) + await new_client.connect_started.wait() + + assert channel._running + assert channel._ws_client is new_client + assert channel._ws_task is not None + assert channel._ws_task.done() + + _run(go()) + + def test_cancelled_stop_finishes_sdk_shutdown_and_clears_state(self): + from app.channels.wecom import WeComChannel + + async def go(): + shutdown_started = asyncio.Event() + release_shutdown = asyncio.Event() + shutdown_finished = asyncio.Event() + manager = _ControlledWeComManager(shutdown_started, release_shutdown, shutdown_finished) + client = _ControlledWeComClient(manager) + channel = WeComChannel(MessageBus(), config={}) + connect_task = asyncio.create_task(client.connect()) + await connect_task + channel._running = True + channel._ws_client = client + channel._ws_task = connect_task + channel._ws_frames["message-1"] = {"body": {}} + channel._ws_stream_ids["message-1"] = "stream-1" + + stop_task = asyncio.create_task(channel.stop()) + await shutdown_started.wait() + await _wait_for_next_event_loop_turn() + stop_task.cancel() + await _wait_for_next_event_loop_turn() + + assert not stop_task.done() + assert not shutdown_finished.is_set() + + release_shutdown.set() + + try: + with pytest.raises(asyncio.CancelledError): + await stop_task + finally: + release_shutdown.set() + await asyncio.gather(*manager.shutdown_tasks, return_exceptions=True) + + assert shutdown_finished.is_set() + assert channel._ws_client is None + assert channel._ws_task is None + assert channel._ws_shutdown_task is None + assert channel._ws_frames == {} + assert channel._ws_stream_ids == {} + + _run(go()) + def test_publish_ws_inbound_starts_stream_and_publishes_message(self, monkeypatch): from app.channels.wecom import WeComChannel