diff --git a/backend/app/channels/telegram.py b/backend/app/channels/telegram.py index 9e6877663..d13857f06 100644 --- a/backend/app/channels/telegram.py +++ b/backend/app/channels/telegram.py @@ -724,13 +724,18 @@ class TelegramChannel(Channel): return username 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: return False state = await self._connection_repo.consume_oauth_state(provider="telegram", state=state_token) 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 owner_user_id = state["owner_user_id"] @@ -751,9 +756,30 @@ class TelegramChannel(Channel): 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"]) - 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 + 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: return await attach_connection_identity( inbound, @@ -820,6 +846,24 @@ class TelegramChannel(Channel): if reservation is not None: 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: """Forward slash commands to the channel manager.""" if not self._check_user(update.effective_user.id): @@ -856,16 +900,15 @@ class TelegramChannel(Channel): if reservation is None: return try: - inbound = await self._attach_connection_identity(inbound) scheduled = self._submit_threadsafe_coroutine( - self._process_incoming_with_reply( + self._process_incoming_with_identity( chat_id, update.message.message_id, inbound, reservation=reservation, ), self._main_loop, - name="process_incoming_with_reply", + name="process_incoming_with_identity", msg_id=update.message.message_id, reservation=reservation, ) @@ -927,16 +970,15 @@ class TelegramChannel(Channel): if reservation is None: return try: - inbound = await self._attach_connection_identity(inbound) scheduled = self._submit_threadsafe_coroutine( - self._process_incoming_with_reply( + self._process_incoming_with_identity( chat_id, update.message.message_id, inbound, reservation=reservation, ), self._main_loop, - name="process_incoming_with_reply", + name="process_incoming_with_identity", msg_id=update.message.message_id, reservation=reservation, ) diff --git a/backend/tests/test_telegram_channel_connections.py b/backend/tests/test_telegram_channel_connections.py index 1f5ef16dc..f4e96f2d1 100644 --- a/backend/tests/test_telegram_channel_connections.py +++ b/backend/tests/test_telegram_channel_connections.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio from datetime import UTC, datetime, timedelta from pathlib import Path 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 +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 async def test_start_with_deep_link_state_binds_telegram_chat(repo): state = "telegram-bind-state" @@ -54,13 +87,15 @@ async def test_start_with_deep_link_state_binds_telegram_chat(repo): bus=MessageBus(), config={"bot_token": "test-token", "connection_repo": repo}, ) + channel._main_loop = asyncio.get_running_loop() update = _telegram_update(text=f"/start {state}") context = MagicMock() context.args = [state] 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 connections[0]["provider"] == "telegram" 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 }, ) + channel._main_loop = asyncio.get_running_loop() update = _telegram_update(text=f"/start {state}", user_id=42) context = MagicMock() context.args = [state] 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 connections[0]["external_account_id"] == "42" 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, 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() 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.chat_id == "100" 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