From e9c9e896c6cd460863e06b98e94d5bbc42a2e526 Mon Sep 17 00:00:00 2001 From: Hyu Date: Sun, 2 Aug 2026 01:49:17 +0800 Subject: [PATCH] fix(cloud): preserve pipeline routing in debug chat (#2380) Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com> --- .../pkg/platform/sources/websocket_adapter.py | 53 +++++++++++-------- .../test_websocket_session_isolation.py | 43 +++++++++++++++ 2 files changed, 74 insertions(+), 22 deletions(-) diff --git a/src/langbot/pkg/platform/sources/websocket_adapter.py b/src/langbot/pkg/platform/sources/websocket_adapter.py index 6176e2d93..0752225c6 100644 --- a/src/langbot/pkg/platform/sources/websocket_adapter.py +++ b/src/langbot/pkg/platform/sources/websocket_adapter.py @@ -707,28 +707,37 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter) if len(listener_tasks) >= 100: await self.logger.warning('WebSocket inbound listener capacity reached; dropping message') return - token = _current_pipeline_uuid.set(pipeline_uuid) - try: - task_manager = getattr(self.ap, 'task_mgr', None) - if task_manager is None or not isinstance(getattr(task_manager, 'tasks', None), list): - listener_task = asyncio.create_task(listeners[event.__class__](event, callback_adapter)) - else: - listener_task = task_manager.create_task( - listeners[event.__class__](event, callback_adapter), - kind='websocket-message', - name=f'websocket-message-{connection.connection_id}', - scopes=[ - core_entities.LifecycleControlScope.APPLICATION, - core_entities.LifecycleControlScope.PLATFORM, - ], - instance_uuid=connection.instance_uuid, - workspace_uuid=connection.workspace_uuid, - placement_generation=connection.placement_generation, - ).task - listener_tasks.add(listener_task) - listener_task.add_done_callback(self._listener_task_done) - finally: - _current_pipeline_uuid.reset(token) + listener = typing.cast( + typing.Callable[[typing.Any, typing.Any], typing.Awaitable[None]], + listeners[event.__class__], + ) + + async def run_listener(): + token = _current_pipeline_uuid.set(pipeline_uuid) + try: + await listener(event, callback_adapter) + finally: + _current_pipeline_uuid.reset(token) + + listener_coro = run_listener() + task_manager = getattr(self.ap, 'task_mgr', None) + if task_manager is None or not isinstance(getattr(task_manager, 'tasks', None), list): + listener_task = asyncio.create_task(listener_coro) + else: + listener_task = task_manager.create_task( + listener_coro, + kind='websocket-message', + name=f'websocket-message-{connection.connection_id}', + scopes=[ + core_entities.LifecycleControlScope.APPLICATION, + core_entities.LifecycleControlScope.PLATFORM, + ], + instance_uuid=connection.instance_uuid, + workspace_uuid=connection.workspace_uuid, + placement_generation=connection.placement_generation, + ).task + listener_tasks.add(listener_task) + listener_task.add_done_callback(self._listener_task_done) def get_websocket_messages( self, diff --git a/tests/unit_tests/platform/test_websocket_session_isolation.py b/tests/unit_tests/platform/test_websocket_session_isolation.py index 682a42184..958ba5888 100644 --- a/tests/unit_tests/platform/test_websocket_session_isolation.py +++ b/tests/unit_tests/platform/test_websocket_session_isolation.py @@ -1,6 +1,7 @@ """Regression tests for isolated embed-widget conversations.""" import asyncio +import contextvars from pathlib import Path from unittest.mock import AsyncMock, Mock @@ -204,6 +205,48 @@ async def test_embed_event_uses_stable_session_launcher(monkeypatch): assert received[0].sender.id == f'websocket_pipeline-1:{session_id}' +@pytest.mark.asyncio +async def test_pipeline_override_survives_detached_listener_task(monkeypatch): + manager = WebSocketConnectionManager() + connection = await manager.add_connection( + websocket=Mock(), + scope=SCOPE_A, + pipeline_uuid='pipeline-1', + session_type='person', + ) + monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager) + + class DetachedTaskManager: + def __init__(self): + self.tasks = [] + + def create_task(self, coro, **_kwargs): + task = asyncio.create_task(coro, context=contextvars.Context()) + self.tasks.append(task) + return Mock(task=task) + + task_manager = DetachedTaskManager() + adapter = WebSocketAdapter.model_construct( + ap=Mock(task_mgr=task_manager), + logger=_adapter_logger(), + ) + adapter.websocket_person_session = WebSocketSession(id='person') + adapter.websocket_group_session = WebSocketSession(id='group') + pipeline_overrides = [] + + async def listener(_event, callback_adapter): + pipeline_overrides.append(callback_adapter.get_pipeline_uuid_override()) + + adapter.listeners = {platform_events.FriendMessage: listener} + await adapter.handle_websocket_message( + connection, + {'message': [{'type': 'Plain', 'text': 'hello'}], 'stream': False}, + ) + await asyncio.gather(*task_manager.tasks) + + assert pipeline_overrides == ['pipeline-1'] + + @pytest.mark.asyncio async def test_embed_group_event_uses_stable_session_launcher(monkeypatch): manager = WebSocketConnectionManager()