fix(send_message): cross-loop dispatch to live WeCom adapter
When send_message is invoked from the agent's worker thread (a different event loop than the gateway's), awaiting the WeCom adapter directly can hang because the adapter enqueues onto the gateway loop. Dispatch via run_coroutine_threadsafe onto the gateway loop when the caller loop differs, with caller-cancellation shielded so an already-enqueued send is not cancelled mid-flight (which would otherwise cause a false-failure retry -> duplicate). Recognizes WeCom native chat IDs as explicit send targets and whitelists WeCom for media delivery. Part of the async queue design this branch introduces.
This commit is contained in:
@@ -0,0 +1,201 @@
|
||||
"""Regression tests for the cross-event-loop deadlock fix in send_message.
|
||||
|
||||
When the agent's tool worker thread calls _send_via_adapter() while the
|
||||
adapter's queues live on the gateway's main event loop, the send must be
|
||||
dispatched via run_coroutine_threadsafe to the gateway loop — NOT awaited
|
||||
directly on the worker loop (which would deadlock due to the selector never
|
||||
being woken by cross-thread future.set_result).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
import threading
|
||||
from types import ModuleType, SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from gateway.config import Platform
|
||||
|
||||
|
||||
class TestSendViaAdapterCrossLoopDispatch:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cross_loop_dispatches_to_gateway_loop(self, monkeypatch):
|
||||
"""adapter.send() runs on gateway loop, not the caller's loop."""
|
||||
from tools.send_message_tool import _send_via_adapter
|
||||
|
||||
send_loop_id = {}
|
||||
platform = Platform("wecom")
|
||||
|
||||
class FakeAdapter:
|
||||
async def send(self, *, chat_id, content, metadata=None):
|
||||
send_loop_id["loop"] = id(asyncio.get_running_loop())
|
||||
return SimpleNamespace(success=True, message_id="cross-ok")
|
||||
|
||||
gateway_loop = asyncio.new_event_loop()
|
||||
started = threading.Event()
|
||||
|
||||
def run_gateway():
|
||||
asyncio.set_event_loop(gateway_loop)
|
||||
started.set()
|
||||
gateway_loop.run_forever()
|
||||
|
||||
t = threading.Thread(target=run_gateway, daemon=True)
|
||||
t.start()
|
||||
started.wait(timeout=2)
|
||||
|
||||
try:
|
||||
runner = SimpleNamespace(
|
||||
adapters={platform: FakeAdapter()},
|
||||
_gateway_loop=gateway_loop,
|
||||
)
|
||||
fake_gateway_run = ModuleType("gateway.run")
|
||||
fake_gateway_run._gateway_runner_ref = lambda: runner
|
||||
monkeypatch.setitem(sys.modules, "gateway.run", fake_gateway_run)
|
||||
|
||||
result = await _send_via_adapter(
|
||||
platform,
|
||||
SimpleNamespace(extra={}),
|
||||
"wr_group_123",
|
||||
"hello from worker",
|
||||
)
|
||||
|
||||
assert result == {"success": True, "message_id": "cross-ok"}
|
||||
# Verify send() ran on the gateway loop, not our current loop
|
||||
assert send_loop_id["loop"] == id(gateway_loop)
|
||||
finally:
|
||||
gateway_loop.call_soon_threadsafe(gateway_loop.stop)
|
||||
t.join(timeout=2)
|
||||
gateway_loop.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_loop_uses_direct_await(self, monkeypatch):
|
||||
"""When current loop IS the gateway loop, adapter.send() is awaited
|
||||
directly — no run_coroutine_threadsafe (which would self-lock)."""
|
||||
from tools.send_message_tool import _send_via_adapter
|
||||
|
||||
current_loop = asyncio.get_running_loop()
|
||||
platform = Platform("wecom")
|
||||
called_directly = {}
|
||||
|
||||
class FakeAdapter:
|
||||
async def send(self, *, chat_id, content, metadata=None):
|
||||
called_directly["loop"] = id(asyncio.get_running_loop())
|
||||
return SimpleNamespace(success=True, message_id="direct-ok")
|
||||
|
||||
runner = SimpleNamespace(
|
||||
adapters={platform: FakeAdapter()},
|
||||
_gateway_loop=current_loop,
|
||||
)
|
||||
fake_gateway_run = ModuleType("gateway.run")
|
||||
fake_gateway_run._gateway_runner_ref = lambda: runner
|
||||
monkeypatch.setitem(sys.modules, "gateway.run", fake_gateway_run)
|
||||
|
||||
result = await _send_via_adapter(
|
||||
platform,
|
||||
SimpleNamespace(extra={}),
|
||||
"wr_group_456",
|
||||
"direct send",
|
||||
)
|
||||
|
||||
assert result == {"success": True, "message_id": "direct-ok"}
|
||||
assert called_directly["loop"] == id(current_loop)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gateway_loop_not_running_returns_error(self, monkeypatch):
|
||||
"""When gateway loop exists but is stopped, return an error rather
|
||||
than attempting direct await on a loop-bound adapter."""
|
||||
from tools.send_message_tool import _send_via_adapter
|
||||
|
||||
stopped_loop = asyncio.new_event_loop()
|
||||
stopped_loop.close()
|
||||
platform = Platform("wecom")
|
||||
|
||||
class FakeAdapter:
|
||||
async def send(self, *, chat_id, content, metadata=None):
|
||||
raise AssertionError("should not be called")
|
||||
|
||||
runner = SimpleNamespace(
|
||||
adapters={platform: FakeAdapter()},
|
||||
_gateway_loop=stopped_loop,
|
||||
)
|
||||
fake_gateway_run = ModuleType("gateway.run")
|
||||
fake_gateway_run._gateway_runner_ref = lambda: runner
|
||||
monkeypatch.setitem(sys.modules, "gateway.run", fake_gateway_run)
|
||||
|
||||
result = await _send_via_adapter(
|
||||
platform,
|
||||
SimpleNamespace(extra={}),
|
||||
"wr_group_789",
|
||||
"should fail",
|
||||
)
|
||||
|
||||
assert "error" in result
|
||||
assert "not running" in result["error"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shield_prevents_cancel_of_enqueued_send(self, monkeypatch):
|
||||
"""asyncio.shield ensures that cancelling the caller does NOT cancel
|
||||
the already-dispatched send on the gateway loop."""
|
||||
from tools.send_message_tool import _send_via_adapter
|
||||
|
||||
send_completed = asyncio.Event()
|
||||
send_result_holder = {}
|
||||
platform = Platform("wecom")
|
||||
|
||||
class FakeAdapter:
|
||||
async def send(self, *, chat_id, content, metadata=None):
|
||||
# Simulate a slow send (token bucket wait)
|
||||
await asyncio.sleep(0.3)
|
||||
send_result_holder["sent"] = True
|
||||
send_completed.set()
|
||||
return SimpleNamespace(success=True, message_id="shielded")
|
||||
|
||||
gateway_loop = asyncio.new_event_loop()
|
||||
started = threading.Event()
|
||||
|
||||
def run_gateway():
|
||||
asyncio.set_event_loop(gateway_loop)
|
||||
started.set()
|
||||
gateway_loop.run_forever()
|
||||
|
||||
t = threading.Thread(target=run_gateway, daemon=True)
|
||||
t.start()
|
||||
started.wait(timeout=2)
|
||||
|
||||
try:
|
||||
runner = SimpleNamespace(
|
||||
adapters={platform: FakeAdapter()},
|
||||
_gateway_loop=gateway_loop,
|
||||
)
|
||||
fake_gateway_run = ModuleType("gateway.run")
|
||||
fake_gateway_run._gateway_runner_ref = lambda: runner
|
||||
monkeypatch.setitem(sys.modules, "gateway.run", fake_gateway_run)
|
||||
|
||||
# Start the send, then cancel the caller task after a short delay
|
||||
async def do_send():
|
||||
return await _send_via_adapter(
|
||||
platform,
|
||||
SimpleNamespace(extra={}),
|
||||
"wr_group_shield",
|
||||
"shielded msg",
|
||||
)
|
||||
|
||||
task = asyncio.create_task(do_send())
|
||||
await asyncio.sleep(0.1) # let it dispatch to gateway loop
|
||||
task.cancel()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
# The send on the gateway loop should still complete despite cancel
|
||||
fut = asyncio.run_coroutine_threadsafe(
|
||||
asyncio.wait_for(send_completed.wait(), timeout=1.0),
|
||||
gateway_loop,
|
||||
)
|
||||
fut.result(timeout=2)
|
||||
assert send_result_holder.get("sent") is True
|
||||
finally:
|
||||
gateway_loop.call_soon_threadsafe(gateway_loop.stop)
|
||||
t.join(timeout=2)
|
||||
gateway_loop.close()
|
||||
@@ -600,6 +600,14 @@ def _parse_target_ref(platform_name: str, target_ref: str):
|
||||
if group_id:
|
||||
return f"group:{group_id}", None, True
|
||||
return None, None, False
|
||||
# WeCom: group IDs start with "wr" or "wc", user IDs start with "wo" or
|
||||
# are bare alphanumeric strings. Treat any non-empty WeCom target_ref as
|
||||
# an explicit chat_id — the adapter resolves whether to use APP_CMD_RESPONSE
|
||||
# (groups) or APP_CMD_SEND (DMs) internally.
|
||||
if platform_name == "wecom":
|
||||
stripped = target_ref.strip()
|
||||
if stripped:
|
||||
return stripped, None, True
|
||||
if platform_name in _PHONE_PLATFORMS:
|
||||
match = _E164_TARGET_RE.fullmatch(target_ref)
|
||||
if match:
|
||||
@@ -870,7 +878,48 @@ async def _send_via_adapter(
|
||||
metadata["publish_topic"] = chat_id
|
||||
if not metadata:
|
||||
metadata = None
|
||||
result = await adapter.send(chat_id=chat_id, content=chunk, metadata=metadata)
|
||||
|
||||
# The adapter's send() uses asyncio.Queue + worker tasks bound
|
||||
# to the gateway's main event loop. Calling send() from a
|
||||
# different thread/loop (the agent's tool worker thread) causes
|
||||
# a cross-loop Future deadlock: the worker loop's selector never
|
||||
# gets woken when the gateway loop resolves the future.
|
||||
# When on a different loop, dispatch onto the gateway loop via
|
||||
# run_coroutine_threadsafe and await the wrapped future.
|
||||
gateway_loop = getattr(runner, "_gateway_loop", None)
|
||||
try:
|
||||
_current_loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
_current_loop = None
|
||||
|
||||
_need_cross_loop = (
|
||||
gateway_loop is not None
|
||||
and _current_loop is not gateway_loop
|
||||
)
|
||||
|
||||
if _need_cross_loop:
|
||||
if not gateway_loop.is_running():
|
||||
return {"error": "Gateway loop is not running; cannot dispatch adapter send"}
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
fut = safe_schedule_threadsafe(
|
||||
adapter.send(chat_id=chat_id, content=chunk, metadata=metadata),
|
||||
gateway_loop,
|
||||
logger=logger,
|
||||
log_message="send_message: failed to schedule on gateway loop",
|
||||
)
|
||||
if fut is None:
|
||||
return {"error": "Gateway loop unavailable for send dispatch"}
|
||||
# Use shield so that if the caller's task is cancelled (e.g.
|
||||
# agent interrupt), the already-enqueued send on the gateway
|
||||
# loop is NOT cancelled — preventing "tool failed but message
|
||||
# still sent later" followed by agent retry causing duplicates.
|
||||
# No explicit timeout here: the adapter's internal request
|
||||
# timeout (15s) and the upper-layer _run_async 300s timeout
|
||||
# provide sufficient protection against hangs.
|
||||
result = await asyncio.shield(asyncio.wrap_future(fut))
|
||||
else:
|
||||
# Same loop or no gateway loop (CLI, tests) — direct await.
|
||||
result = await adapter.send(chat_id=chat_id, content=chunk, metadata=metadata)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e:
|
||||
@@ -1259,6 +1308,25 @@ async def _send_to_platform(platform, pconfig, chat_id, message, thread_id=None,
|
||||
last_result = result
|
||||
return last_result
|
||||
|
||||
# --- WeCom: native media attachment support via live gateway adapter ---
|
||||
if platform == Platform.WECOM and media_files:
|
||||
last_result = None
|
||||
for i, chunk in enumerate(chunks):
|
||||
is_last = (i == len(chunks) - 1)
|
||||
result = await _send_via_adapter(
|
||||
platform,
|
||||
pconfig,
|
||||
chat_id,
|
||||
chunk,
|
||||
thread_id=thread_id,
|
||||
media_files=media_files if is_last else None,
|
||||
force_document=force_document,
|
||||
)
|
||||
if isinstance(result, dict) and result.get("error"):
|
||||
return result
|
||||
last_result = result
|
||||
return last_result
|
||||
|
||||
# --- Non-media platforms ---
|
||||
if media_files and not message.strip():
|
||||
return {
|
||||
|
||||
Reference in New Issue
Block a user