mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-08 20:30:59 +00:00
fix(runtime): stabilize reasoning chat delivery
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user