diff --git a/src/langbot/pkg/pipeline/controller.py b/src/langbot/pkg/pipeline/controller.py index d14288c29..37d1379e6 100644 --- a/src/langbot/pkg/pipeline/controller.py +++ b/src/langbot/pkg/pipeline/controller.py @@ -132,9 +132,7 @@ class Controller: break - if selected_query: # 找到了 - queries.remove(selected_query) - else: # 没找到 说明:没有请求 或者 所有query对应的session都已达到并发上限 + if not selected_query: # 没有请求,或所有 query 对应的 session 都已达到并发上限 await self.ap.query_pool.condition.wait() continue diff --git a/src/langbot/pkg/platform/sources/websocket_adapter.py b/src/langbot/pkg/platform/sources/websocket_adapter.py index 6176e2d93..2dc29ad06 100644 --- a/src/langbot/pkg/platform/sources/websocket_adapter.py +++ b/src/langbot/pkg/platform/sources/websocket_adapter.py @@ -5,6 +5,7 @@ import contextvars import logging import time import typing +from dataclasses import dataclass from datetime import datetime import pydantic @@ -25,6 +26,15 @@ _current_pipeline_uuid: contextvars.ContextVar[str | None] = contextvars.Context ) +@dataclass(frozen=True) +class WebSocketReplyContext: + """Trusted routing context retained when the originating socket reconnects.""" + + scope: WebSocketScope + pipeline_uuid: str + session_id: str | None + + class WebSocketMessage(pydantic.BaseModel): """WebSocket消息格式""" @@ -265,6 +275,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter) embed_target = self._parse_embed_target(sender_id) if embed_target is not None: return embed_target + reply_context = getattr(message_source, '_websocket_reply_context', None) + if isinstance(reply_context, WebSocketReplyContext): + if reply_context.scope != self._scope(): + raise ValueError('WebSocket reply context does not match this adapter scope') + return reply_context.pipeline_uuid, reply_context.session_id raise ValueError('WebSocket reply target is not bound to this adapter scope') async def send_message( @@ -685,6 +700,16 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter) # 异步触发事件处理 # Use owner_bot's listeners if available, otherwise fall back to proxy bot + object.__setattr__( + event, + '_websocket_reply_context', + WebSocketReplyContext( + scope=connection.scope, + pipeline_uuid=pipeline_uuid, + session_id=connection.session_id, + ), + ) + listeners = ( owner_bot.adapter.listeners if (owner_bot and hasattr(owner_bot.adapter, 'listeners') and owner_bot.adapter.listeners) @@ -707,28 +732,31 @@ 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) + async def run_listener() -> None: + token = _current_pipeline_uuid.set(pipeline_uuid) + try: + await listeners[event.__class__](event, callback_adapter) + finally: + _current_pipeline_uuid.reset(token) + + 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(run_listener()) + else: + listener_task = task_manager.create_task( + run_listener(), + 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/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py b/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py index 38cb5c3bd..68dd50aa1 100644 --- a/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py +++ b/src/langbot/pkg/provider/modelmgr/requesters/litellmchat.py @@ -466,6 +466,8 @@ class LiteLLMRequester(requester.ProviderAPIRequester): provider = self._reasoning_provider(model.model_entity.name) if level == 'disabled': + if provider == 'deepseek': + return {'extra_body': {'thinking': {'type': 'disabled'}}} if provider == 'volcengine': return {'thinking': {'type': 'disabled'}} return {'reasoning_effort': 'none'} @@ -844,7 +846,15 @@ class LiteLLMRequester(requester.ProviderAPIRequester): raise errors.RequesterError( 'reasoning_config conflicts with advanced parameters: ' + ', '.join(dict.fromkeys(conflicts)) ) - args.update(reasoning_args) + reasoning_extra_body = reasoning_args.get('extra_body') + if isinstance(reasoning_extra_body, dict): + existing_extra_body = args.get('extra_body') or {} + if not isinstance(existing_extra_body, dict): + raise errors.RequesterError('extra_body must be an object') + args.update({key: value for key, value in reasoning_args.items() if key != 'extra_body'}) + args['extra_body'] = {**existing_extra_body, **reasoning_extra_body} + else: + args.update(reasoning_args) if 'reasoning_effort' in reasoning_args and self._reasoning_provider(model.model_entity.name) == 'openai': allowed_openai_params = args.get('allowed_openai_params') or [] if not isinstance(allowed_openai_params, (list, tuple, set)): diff --git a/tests/unit_tests/pipeline/test_controller_tenancy.py b/tests/unit_tests/pipeline/test_controller_tenancy.py index b99e54bb2..ea24dd3f5 100644 --- a/tests/unit_tests/pipeline/test_controller_tenancy.py +++ b/tests/unit_tests/pipeline/test_controller_tenancy.py @@ -28,6 +28,57 @@ def _prepare_scheduler(mock_app): return query_pool, session +@pytest.mark.asyncio +async def test_consumer_schedules_query_after_running_transition( + mock_app, + sample_query, +): + query_pool = MagicMock() + query_pool.queries = [sample_query] + query_pool.__aenter__ = AsyncMock(return_value=query_pool) + query_pool.__aexit__ = AsyncMock(return_value=None) + query_pool.remove_query = AsyncMock(return_value=True) + wait_for_query = asyncio.Event() + query_pool.condition = SimpleNamespace( + wait=AsyncMock(side_effect=wait_for_query.wait), + notify_all=Mock(), + ) + query_pool.mark_query_running_locked = Mock(side_effect=query_pool.queries.remove) + mock_app.query_pool = query_pool + + session = SimpleNamespace(_semaphore=asyncio.Semaphore(1)) + mock_app.sess_mgr.get_session = AsyncMock(return_value=session) + runtime_pipeline = SimpleNamespace(run=AsyncMock()) + mock_app.pipeline_mgr = SimpleNamespace(get_pipeline_by_uuid=AsyncMock(return_value=runtime_pipeline)) + + task_created = asyncio.Event() + process_tasks = [] + + def create_process_task(coro, **_kwargs): + process_tasks.append(asyncio.create_task(coro)) + task_created.set() + + mock_app.task_mgr.create_task = Mock(side_effect=create_process_task) + controller = Controller(mock_app) + initial_slots = controller.semaphore._value + consumer_task = asyncio.create_task(controller.consumer()) + + try: + await asyncio.wait_for(task_created.wait(), timeout=2) + finally: + consumer_task.cancel() + with pytest.raises(asyncio.CancelledError): + await consumer_task + await asyncio.gather(*process_tasks) + + query_pool.mark_query_running_locked.assert_called_once_with(sample_query) + runtime_pipeline.run.assert_awaited_once_with(sample_query) + query_pool.remove_query.assert_awaited_once_with(sample_query) + assert query_pool.queries == [] + assert session._semaphore._value == 1 + assert controller.semaphore._value == initial_slots + + @pytest.mark.asyncio async def test_controller_drops_stale_query_before_pipeline_lookup( mock_app, diff --git a/tests/unit_tests/platform/test_websocket_session_isolation.py b/tests/unit_tests/platform/test_websocket_session_isolation.py index 682a42184..c26119eab 100644 --- a/tests/unit_tests/platform/test_websocket_session_isolation.py +++ b/tests/unit_tests/platform/test_websocket_session_isolation.py @@ -1,12 +1,15 @@ """Regression tests for isolated embed-widget conversations.""" import asyncio +import contextvars from pathlib import Path +from types import SimpleNamespace from unittest.mock import AsyncMock, Mock import pytest import langbot_plugin.api.entities.builtin.platform.events as platform_events +import langbot_plugin.api.entities.builtin.platform.message as platform_message from langbot.pkg.platform.sources import websocket_adapter as websocket_adapter_module from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter, WebSocketMessage, WebSocketSession from langbot.pkg.platform.sources.websocket_manager import ( @@ -204,6 +207,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_is_set_inside_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.get_running_loop().create_task(coro, context=contextvars.Context()) + self.tasks.append(task) + return SimpleNamespace(task=task) + + task_manager = DetachedTaskManager() + adapter = WebSocketAdapter.model_construct( + ap=SimpleNamespace(task_mgr=task_manager), + logger=_adapter_logger(), + ) + adapter.websocket_person_session = WebSocketSession(id='person') + adapter.websocket_group_session = WebSocketSession(id='group') + received_pipeline_uuids = [] + + async def listener(_event, callback_adapter): + received_pipeline_uuids.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 received_pipeline_uuids == ['pipeline-1'] + + @pytest.mark.asyncio async def test_embed_group_event_uses_stable_session_launcher(monkeypatch): manager = WebSocketConnectionManager() @@ -300,6 +345,49 @@ async def test_stable_session_launcher_resolves_to_active_connection(monkeypatch ) +@pytest.mark.asyncio +async def test_dashboard_reply_survives_connection_replacement(monkeypatch): + manager = WebSocketConnectionManager() + original = 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) + + adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger()) + adapter.websocket_person_session = WebSocketSession(id='person') + adapter.websocket_group_session = WebSocketSession(id='group') + received = [] + + async def listener(event, _callback_adapter): + received.append(event) + + adapter.listeners = {platform_events.FriendMessage: listener} + await adapter.handle_websocket_message( + original, + {'message': [{'type': 'Plain', 'text': 'hello'}], 'stream': False}, + ) + await asyncio.sleep(0) + await manager.remove_connection(original.connection_id) + replacement = await manager.add_connection( + websocket=Mock(), + scope=SCOPE_A, + pipeline_uuid='pipeline-1', + session_type='person', + ) + + await adapter.reply_message( + received[0], + platform_message.MessageChain([platform_message.Plain(text='done')]), + ) + + response = await replacement.send_queue.get() + assert response['type'] == 'response' + assert response['data']['content'] == 'done' + + def test_session_ids_must_be_canonical_random_uuids(): assert is_valid_session_id('31c0f2e9-b115-4ee6-8f15-3e624d6456b1') assert not is_valid_session_id('session-a') diff --git a/tests/unit_tests/provider/test_reasoning_control.py b/tests/unit_tests/provider/test_reasoning_control.py index e6d860db3..6e7e769cc 100644 --- a/tests/unit_tests/provider/test_reasoning_control.py +++ b/tests/unit_tests/provider/test_reasoning_control.py @@ -193,6 +193,9 @@ def test_reasoning_argument_translation(monkeypatch): assert deepseek_request._build_reasoning_args( _runtime_model(deepseek_request, 'enabled', name='deepseek-chat') ) == {'thinking': {'type': 'enabled'}} + assert deepseek_request._build_reasoning_args( + _runtime_model(deepseek_request, 'disabled', name='deepseek-chat') + ) == {'extra_body': {'thinking': {'type': 'disabled'}}} def test_pipeline_reasoning_override_takes_precedence(monkeypatch): @@ -332,6 +335,22 @@ async def test_provider_default_does_not_allow_or_send_reasoning_effort(): assert 'allowed_openai_params' not in args +@pytest.mark.asyncio +async def test_deepseek_disabled_thinking_is_merged_into_extra_body(monkeypatch): + request = _requester('deepseek') + monkeypatch.setattr(request, '_supports_reasoning', lambda _: False) + model = _runtime_model(request, 'disabled', name='deepseek-chat') + model.model_entity.extra_args = {'extra_body': {'custom_extension': True}} + model.provider.token_mgr.get_token = lambda: 'test-token' + + args = await request._build_completion_args(model, []) + + assert args['extra_body'] == { + 'custom_extension': True, + 'thinking': {'type': 'disabled'}, + } + + class _Dumpable: def __init__(self, data: dict): self.data = data