mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-09 21:49:37 +00:00
fix(wecom): serialize websocket shutdown (#4762)
* fix(wecom): await connection task shutdown * fix(wecom): serialize websocket shutdown
This commit is contained in:
parent
baaf2bad47
commit
38ff44778a
@ -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 <code>` flow over the existing long-polling worker. Slack, Discord, Feishu/Lark, DingTalk, WeChat, and WeCom use `/connect <code>` 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.
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user