mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-09 21:49:37 +00:00
fix(channels): move Telegram _attach_connection_identity to main event loop (#4815)
Move SQL-dependent connection identity and /start bind work onto the Gateway main loop via _submit_threadsafe_coroutine, and send PTB replies through _run_on_telegram_loop so the Telegram worker never blocks or touches SQLAlchemy/HTTP across event loops. When the main loop is not running (e.g. during gateway shutdown), the bind path logs a warning and returns False instead of running SQLAlchemy on the wrong loop, matching the Feishu bind pattern.
This commit is contained in:
parent
13f0a7f263
commit
74d9e6c2e0
@ -724,13 +724,18 @@ class TelegramChannel(Channel):
|
|||||||
return username
|
return username
|
||||||
return str(getattr(user, "id", ""))
|
return str(getattr(user, "id", ""))
|
||||||
|
|
||||||
async def _bind_connection_from_start_token(self, update, state_token: str) -> bool:
|
async def _bind_connection_from_start_token_on_main(
|
||||||
|
self,
|
||||||
|
update,
|
||||||
|
state_token: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Run the bind flow on the main event loop (SQLAlchemy session is bound there)."""
|
||||||
if self._connection_repo is None or not state_token:
|
if self._connection_repo is None or not state_token:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
state = await self._connection_repo.consume_oauth_state(provider="telegram", state=state_token)
|
state = await self._connection_repo.consume_oauth_state(provider="telegram", state=state_token)
|
||||||
if state is None:
|
if state is None:
|
||||||
await update.message.reply_text("Telegram connection link is invalid or expired.")
|
await self._run_on_telegram_loop(update.message.reply_text("Telegram connection link is invalid or expired."))
|
||||||
return True
|
return True
|
||||||
|
|
||||||
owner_user_id = state["owner_user_id"]
|
owner_user_id = state["owner_user_id"]
|
||||||
@ -751,9 +756,30 @@ class TelegramChannel(Channel):
|
|||||||
status="connected",
|
status="connected",
|
||||||
)
|
)
|
||||||
logger.info("[Telegram] bound chat=%s user=%s to DeerFlow user=%s connection=%s", chat_id, user_id, owner_user_id, connection["id"])
|
logger.info("[Telegram] bound chat=%s user=%s to DeerFlow user=%s connection=%s", chat_id, user_id, owner_user_id, connection["id"])
|
||||||
await update.message.reply_text("Telegram connected to DeerFlow.")
|
await self._run_on_telegram_loop(update.message.reply_text("Telegram connected to DeerFlow."))
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def _bind_connection_from_start_token(self, update, state_token: str) -> bool:
|
||||||
|
"""Dispatch the bind flow to the main loop if repo is configured."""
|
||||||
|
if self._connection_repo is None or not state_token:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if self._main_loop and self._main_loop.is_running():
|
||||||
|
msg_id = getattr(getattr(update, "message", None), "message_id", None)
|
||||||
|
scheduled = self._submit_threadsafe_coroutine(
|
||||||
|
self._bind_connection_from_start_token_on_main(update, state_token),
|
||||||
|
self._main_loop,
|
||||||
|
name="bind_connection",
|
||||||
|
msg_id=msg_id,
|
||||||
|
)
|
||||||
|
if not scheduled:
|
||||||
|
logger.info("[Telegram] main loop stopped before channel connection bind could be scheduled")
|
||||||
|
# Schedule counts as handled so /start does not fall through to help
|
||||||
|
# while SQL runs on the main loop.
|
||||||
|
return scheduled
|
||||||
|
logger.warning("[Telegram] main loop not running, cannot bind channel connection")
|
||||||
|
return False
|
||||||
|
|
||||||
async def _attach_connection_identity(self, inbound: InboundMessage) -> InboundMessage:
|
async def _attach_connection_identity(self, inbound: InboundMessage) -> InboundMessage:
|
||||||
return await attach_connection_identity(
|
return await attach_connection_identity(
|
||||||
inbound,
|
inbound,
|
||||||
@ -820,6 +846,24 @@ class TelegramChannel(Channel):
|
|||||||
if reservation is not None:
|
if reservation is not None:
|
||||||
reservation.release()
|
reservation.release()
|
||||||
|
|
||||||
|
async def _process_incoming_with_identity(
|
||||||
|
self,
|
||||||
|
chat_id: str,
|
||||||
|
msg_id: int,
|
||||||
|
inbound: InboundMessage,
|
||||||
|
*,
|
||||||
|
reservation: InboundReservation | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Attach connection identity and dispatch on the main event loop.
|
||||||
|
|
||||||
|
This must run on the main loop because _attach_connection_identity
|
||||||
|
uses a SQLAlchemy session factory bound to that loop. Calling it
|
||||||
|
from the Telegram polling thread's own loop raises
|
||||||
|
``RuntimeError: got Future attached to a different loop``.
|
||||||
|
"""
|
||||||
|
inbound = await self._attach_connection_identity(inbound)
|
||||||
|
await self._process_incoming_with_reply(chat_id, msg_id, inbound, reservation=reservation)
|
||||||
|
|
||||||
async def _cmd_generic(self, update, context) -> None:
|
async def _cmd_generic(self, update, context) -> None:
|
||||||
"""Forward slash commands to the channel manager."""
|
"""Forward slash commands to the channel manager."""
|
||||||
if not self._check_user(update.effective_user.id):
|
if not self._check_user(update.effective_user.id):
|
||||||
@ -856,16 +900,15 @@ class TelegramChannel(Channel):
|
|||||||
if reservation is None:
|
if reservation is None:
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
inbound = await self._attach_connection_identity(inbound)
|
|
||||||
scheduled = self._submit_threadsafe_coroutine(
|
scheduled = self._submit_threadsafe_coroutine(
|
||||||
self._process_incoming_with_reply(
|
self._process_incoming_with_identity(
|
||||||
chat_id,
|
chat_id,
|
||||||
update.message.message_id,
|
update.message.message_id,
|
||||||
inbound,
|
inbound,
|
||||||
reservation=reservation,
|
reservation=reservation,
|
||||||
),
|
),
|
||||||
self._main_loop,
|
self._main_loop,
|
||||||
name="process_incoming_with_reply",
|
name="process_incoming_with_identity",
|
||||||
msg_id=update.message.message_id,
|
msg_id=update.message.message_id,
|
||||||
reservation=reservation,
|
reservation=reservation,
|
||||||
)
|
)
|
||||||
@ -927,16 +970,15 @@ class TelegramChannel(Channel):
|
|||||||
if reservation is None:
|
if reservation is None:
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
inbound = await self._attach_connection_identity(inbound)
|
|
||||||
scheduled = self._submit_threadsafe_coroutine(
|
scheduled = self._submit_threadsafe_coroutine(
|
||||||
self._process_incoming_with_reply(
|
self._process_incoming_with_identity(
|
||||||
chat_id,
|
chat_id,
|
||||||
update.message.message_id,
|
update.message.message_id,
|
||||||
inbound,
|
inbound,
|
||||||
reservation=reservation,
|
reservation=reservation,
|
||||||
),
|
),
|
||||||
self._main_loop,
|
self._main_loop,
|
||||||
name="process_incoming_with_reply",
|
name="process_incoming_with_identity",
|
||||||
msg_id=update.message.message_id,
|
msg_id=update.message.message_id,
|
||||||
reservation=reservation,
|
reservation=reservation,
|
||||||
)
|
)
|
||||||
|
|||||||
@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
from datetime import UTC, datetime, timedelta
|
from datetime import UTC, datetime, timedelta
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
@ -41,6 +42,38 @@ def _telegram_update(*, text: str = "/start", user_id: int = 42, chat_id: int =
|
|||||||
return update
|
return update
|
||||||
|
|
||||||
|
|
||||||
|
async def _await_connections(repo, owner_user_id: str, *, timeout: float = 2.0) -> list:
|
||||||
|
"""Poll the connection repo until rows for the owner appear or the deadline passes.
|
||||||
|
|
||||||
|
The bind runs as a scheduled task on the (test) main loop, so the test must
|
||||||
|
yield until it completes. A deadline gives a loaded CI runner headroom for the
|
||||||
|
aiosqlite thread hops in consume_oauth_state / upsert_connection instead of a
|
||||||
|
fixed iteration count that can time out and mask the real bind failure.
|
||||||
|
"""
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
deadline = loop.time() + timeout
|
||||||
|
connections: list = []
|
||||||
|
while loop.time() < deadline:
|
||||||
|
connections = await repo.list_connections(owner_user_id)
|
||||||
|
if connections:
|
||||||
|
return connections
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
return connections
|
||||||
|
|
||||||
|
|
||||||
|
async def _await_reply(reply_text, *, timeout: float = 2.0) -> None:
|
||||||
|
"""Yield until the bind task has dispatched its reply (the task's final step).
|
||||||
|
|
||||||
|
The reply is sent right after the upsert, so the connection row can be visible
|
||||||
|
before the reply has been awaited; wait for the reply explicitly instead of
|
||||||
|
relying on an incidental extra await.
|
||||||
|
"""
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
deadline = loop.time() + timeout
|
||||||
|
while reply_text.await_count == 0 and loop.time() < deadline:
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_start_with_deep_link_state_binds_telegram_chat(repo):
|
async def test_start_with_deep_link_state_binds_telegram_chat(repo):
|
||||||
state = "telegram-bind-state"
|
state = "telegram-bind-state"
|
||||||
@ -54,13 +87,15 @@ async def test_start_with_deep_link_state_binds_telegram_chat(repo):
|
|||||||
bus=MessageBus(),
|
bus=MessageBus(),
|
||||||
config={"bot_token": "test-token", "connection_repo": repo},
|
config={"bot_token": "test-token", "connection_repo": repo},
|
||||||
)
|
)
|
||||||
|
channel._main_loop = asyncio.get_running_loop()
|
||||||
update = _telegram_update(text=f"/start {state}")
|
update = _telegram_update(text=f"/start {state}")
|
||||||
context = MagicMock()
|
context = MagicMock()
|
||||||
context.args = [state]
|
context.args = [state]
|
||||||
|
|
||||||
await channel._cmd_start(update, context)
|
await channel._cmd_start(update, context)
|
||||||
|
connections = await _await_connections(repo, "deerflow-user-1")
|
||||||
|
await _await_reply(update.message.reply_text)
|
||||||
|
|
||||||
connections = await repo.list_connections("deerflow-user-1")
|
|
||||||
assert len(connections) == 1
|
assert len(connections) == 1
|
||||||
assert connections[0]["provider"] == "telegram"
|
assert connections[0]["provider"] == "telegram"
|
||||||
assert connections[0]["external_account_id"] == "42"
|
assert connections[0]["external_account_id"] == "42"
|
||||||
@ -91,13 +126,15 @@ async def test_start_token_bypasses_allowed_users_filter(repo):
|
|||||||
"allowed_users": [999], # newcomer (42) is not whitelisted
|
"allowed_users": [999], # newcomer (42) is not whitelisted
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
channel._main_loop = asyncio.get_running_loop()
|
||||||
update = _telegram_update(text=f"/start {state}", user_id=42)
|
update = _telegram_update(text=f"/start {state}", user_id=42)
|
||||||
context = MagicMock()
|
context = MagicMock()
|
||||||
context.args = [state]
|
context.args = [state]
|
||||||
|
|
||||||
await channel._cmd_start(update, context)
|
await channel._cmd_start(update, context)
|
||||||
|
connections = await _await_connections(repo, "deerflow-user-1")
|
||||||
|
await _await_reply(update.message.reply_text)
|
||||||
|
|
||||||
connections = await repo.list_connections("deerflow-user-1")
|
|
||||||
assert len(connections) == 1
|
assert len(connections) == 1
|
||||||
assert connections[0]["external_account_id"] == "42"
|
assert connections[0]["external_account_id"] == "42"
|
||||||
assert "connected" in update.message.reply_text.await_args.args[0].lower()
|
assert "connected" in update.message.reply_text.await_args.args[0].lower()
|
||||||
@ -118,7 +155,7 @@ async def test_bound_telegram_message_publishes_connection_identity(repo):
|
|||||||
bus=bus,
|
bus=bus,
|
||||||
config={"bot_token": "test-token", "connection_repo": repo},
|
config={"bot_token": "test-token", "connection_repo": repo},
|
||||||
)
|
)
|
||||||
channel._main_loop = __import__("asyncio").get_event_loop()
|
channel._main_loop = asyncio.get_running_loop()
|
||||||
channel._send_running_reply = AsyncMock()
|
channel._send_running_reply = AsyncMock()
|
||||||
|
|
||||||
await channel._on_text(_telegram_update(text="hello"), None)
|
await channel._on_text(_telegram_update(text="hello"), None)
|
||||||
@ -130,3 +167,56 @@ async def test_bound_telegram_message_publishes_connection_identity(repo):
|
|||||||
assert inbound.user_id == "42"
|
assert inbound.user_id == "42"
|
||||||
assert inbound.chat_id == "100"
|
assert inbound.chat_id == "100"
|
||||||
assert inbound.text == "hello"
|
assert inbound.text == "hello"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_bind_dispatcher_uses_submit_threadsafe_when_main_loop_running(repo):
|
||||||
|
channel = TelegramChannel(
|
||||||
|
bus=MessageBus(),
|
||||||
|
config={"bot_token": "test-token", "connection_repo": repo},
|
||||||
|
)
|
||||||
|
channel._main_loop = asyncio.get_running_loop()
|
||||||
|
|
||||||
|
def _fake_submit(coroutine, *args, **kwargs):
|
||||||
|
coroutine.close()
|
||||||
|
return True
|
||||||
|
|
||||||
|
channel._submit_threadsafe_coroutine = MagicMock(side_effect=_fake_submit)
|
||||||
|
channel._bind_connection_from_start_token_on_main = AsyncMock(return_value=True)
|
||||||
|
|
||||||
|
handled = await channel._bind_connection_from_start_token(_telegram_update(), "bind-token")
|
||||||
|
|
||||||
|
assert handled is True
|
||||||
|
channel._submit_threadsafe_coroutine.assert_called_once()
|
||||||
|
assert channel._submit_threadsafe_coroutine.call_args.kwargs["name"] == "bind_connection"
|
||||||
|
channel._bind_connection_from_start_token_on_main.assert_called_once()
|
||||||
|
channel._bind_connection_from_start_token_on_main.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_bind_on_main_replies_via_telegram_loop(repo):
|
||||||
|
state = "telegram-bind-state"
|
||||||
|
await repo.create_oauth_state(
|
||||||
|
owner_user_id="deerflow-user-1",
|
||||||
|
provider="telegram",
|
||||||
|
state=state,
|
||||||
|
expires_at=datetime.now(UTC) + timedelta(minutes=5),
|
||||||
|
)
|
||||||
|
channel = TelegramChannel(
|
||||||
|
bus=MessageBus(),
|
||||||
|
config={"bot_token": "test-token", "connection_repo": repo},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _passthrough(coro):
|
||||||
|
return await coro
|
||||||
|
|
||||||
|
channel._run_on_telegram_loop = AsyncMock(side_effect=_passthrough)
|
||||||
|
update = _telegram_update(text=f"/start {state}")
|
||||||
|
|
||||||
|
assert await channel._bind_connection_from_start_token_on_main(update, state) is True
|
||||||
|
|
||||||
|
channel._run_on_telegram_loop.assert_awaited_once()
|
||||||
|
update.message.reply_text.assert_called_once()
|
||||||
|
assert "connected" in update.message.reply_text.call_args.args[0].lower()
|
||||||
|
connections = await repo.list_connections("deerflow-user-1")
|
||||||
|
assert len(connections) == 1
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user