fix(runtime): stabilize reasoning chat delivery

This commit is contained in:
fdc310
2026-08-02 23:07:57 +08:00
parent c48c345b13
commit 5461628eec
6 changed files with 220 additions and 26 deletions
+1 -3
View File
@@ -132,9 +132,7 @@ class Controller:
break break
if selected_query: # 找到了 if not selected_query: # 没有请求,或所有 query 对应的 session 都已达到并发上限
queries.remove(selected_query)
else: # 没找到 说明:没有请求 或者 所有query对应的session都已达到并发上限
await self.ap.query_pool.condition.wait() await self.ap.query_pool.condition.wait()
continue continue
@@ -5,6 +5,7 @@ import contextvars
import logging import logging
import time import time
import typing import typing
from dataclasses import dataclass
from datetime import datetime from datetime import datetime
import pydantic 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): class WebSocketMessage(pydantic.BaseModel):
"""WebSocket消息格式""" """WebSocket消息格式"""
@@ -265,6 +275,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
embed_target = self._parse_embed_target(sender_id) embed_target = self._parse_embed_target(sender_id)
if embed_target is not None: if embed_target is not None:
return embed_target 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') raise ValueError('WebSocket reply target is not bound to this adapter scope')
async def send_message( 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 # 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 = ( listeners = (
owner_bot.adapter.listeners owner_bot.adapter.listeners
if (owner_bot and hasattr(owner_bot.adapter, 'listeners') and 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: if len(listener_tasks) >= 100:
await self.logger.warning('WebSocket inbound listener capacity reached; dropping message') await self.logger.warning('WebSocket inbound listener capacity reached; dropping message')
return return
token = _current_pipeline_uuid.set(pipeline_uuid) async def run_listener() -> None:
try: token = _current_pipeline_uuid.set(pipeline_uuid)
task_manager = getattr(self.ap, 'task_mgr', None) try:
if task_manager is None or not isinstance(getattr(task_manager, 'tasks', None), list): await listeners[event.__class__](event, callback_adapter)
listener_task = asyncio.create_task(listeners[event.__class__](event, callback_adapter)) finally:
else: _current_pipeline_uuid.reset(token)
listener_task = task_manager.create_task(
listeners[event.__class__](event, callback_adapter), task_manager = getattr(self.ap, 'task_mgr', None)
kind='websocket-message', if task_manager is None or not isinstance(getattr(task_manager, 'tasks', None), list):
name=f'websocket-message-{connection.connection_id}', listener_task = asyncio.create_task(run_listener())
scopes=[ else:
core_entities.LifecycleControlScope.APPLICATION, listener_task = task_manager.create_task(
core_entities.LifecycleControlScope.PLATFORM, run_listener(),
], kind='websocket-message',
instance_uuid=connection.instance_uuid, name=f'websocket-message-{connection.connection_id}',
workspace_uuid=connection.workspace_uuid, scopes=[
placement_generation=connection.placement_generation, core_entities.LifecycleControlScope.APPLICATION,
).task core_entities.LifecycleControlScope.PLATFORM,
listener_tasks.add(listener_task) ],
listener_task.add_done_callback(self._listener_task_done) instance_uuid=connection.instance_uuid,
finally: workspace_uuid=connection.workspace_uuid,
_current_pipeline_uuid.reset(token) placement_generation=connection.placement_generation,
).task
listener_tasks.add(listener_task)
listener_task.add_done_callback(self._listener_task_done)
def get_websocket_messages( def get_websocket_messages(
self, self,
@@ -466,6 +466,8 @@ class LiteLLMRequester(requester.ProviderAPIRequester):
provider = self._reasoning_provider(model.model_entity.name) provider = self._reasoning_provider(model.model_entity.name)
if level == 'disabled': if level == 'disabled':
if provider == 'deepseek':
return {'extra_body': {'thinking': {'type': 'disabled'}}}
if provider == 'volcengine': if provider == 'volcengine':
return {'thinking': {'type': 'disabled'}} return {'thinking': {'type': 'disabled'}}
return {'reasoning_effort': 'none'} return {'reasoning_effort': 'none'}
@@ -844,7 +846,15 @@ class LiteLLMRequester(requester.ProviderAPIRequester):
raise errors.RequesterError( raise errors.RequesterError(
'reasoning_config conflicts with advanced parameters: ' + ', '.join(dict.fromkeys(conflicts)) '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': if 'reasoning_effort' in reasoning_args and self._reasoning_provider(model.model_entity.name) == 'openai':
allowed_openai_params = args.get('allowed_openai_params') or [] allowed_openai_params = args.get('allowed_openai_params') or []
if not isinstance(allowed_openai_params, (list, tuple, set)): if not isinstance(allowed_openai_params, (list, tuple, set)):
@@ -28,6 +28,57 @@ def _prepare_scheduler(mock_app):
return query_pool, session 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 @pytest.mark.asyncio
async def test_controller_drops_stale_query_before_pipeline_lookup( async def test_controller_drops_stale_query_before_pipeline_lookup(
mock_app, mock_app,
@@ -1,12 +1,15 @@
"""Regression tests for isolated embed-widget conversations.""" """Regression tests for isolated embed-widget conversations."""
import asyncio import asyncio
import contextvars
from pathlib import Path from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock from unittest.mock import AsyncMock, Mock
import pytest import pytest
import langbot_plugin.api.entities.builtin.platform.events as platform_events 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 import websocket_adapter as websocket_adapter_module
from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter, WebSocketMessage, WebSocketSession from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter, WebSocketMessage, WebSocketSession
from langbot.pkg.platform.sources.websocket_manager import ( 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}' 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 @pytest.mark.asyncio
async def test_embed_group_event_uses_stable_session_launcher(monkeypatch): async def test_embed_group_event_uses_stable_session_launcher(monkeypatch):
manager = WebSocketConnectionManager() 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(): def test_session_ids_must_be_canonical_random_uuids():
assert is_valid_session_id('31c0f2e9-b115-4ee6-8f15-3e624d6456b1') assert is_valid_session_id('31c0f2e9-b115-4ee6-8f15-3e624d6456b1')
assert not is_valid_session_id('session-a') assert not is_valid_session_id('session-a')
@@ -193,6 +193,9 @@ def test_reasoning_argument_translation(monkeypatch):
assert deepseek_request._build_reasoning_args( assert deepseek_request._build_reasoning_args(
_runtime_model(deepseek_request, 'enabled', name='deepseek-chat') _runtime_model(deepseek_request, 'enabled', name='deepseek-chat')
) == {'thinking': {'type': 'enabled'}} ) == {'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): 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 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: class _Dumpable:
def __init__(self, data: dict): def __init__(self, data: dict):
self.data = data self.data = data