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