feat(cloud): harden multi-tenant runtime resources

This commit is contained in:
Junyan Qin
2026-07-29 11:32:26 +08:00
parent 32abbb636f
commit ae85ac2b16
211 changed files with 14963 additions and 1968 deletions
@@ -3,6 +3,7 @@
import asyncio
import contextvars
import logging
import time
import typing
from datetime import datetime
@@ -14,6 +15,7 @@ import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_logger
from ...core import app
from ...core import entities as core_entities
from .websocket_manager import WebSocketConnection, WebSocketScope, is_valid_session_id, ws_connection_manager
logger = logging.getLogger(__name__)
@@ -45,21 +47,82 @@ class WebSocketSession:
stream_message_indexes: dict[str, dict[str, int]] = {}
"""流式消息索引 {pipeline_uuid: {resp_message_id: message_index}}"""
def __init__(self, id: str):
def __init__(
self,
id: str = '',
*,
max_conversations: int = 200,
max_messages: int = 100,
idle_ttl_seconds: int = 86400,
):
self.id = id
self.message_lists = {}
self.stream_message_indexes = {}
self.message_counters: dict[str, int] = {}
self.last_accessed: dict[str, float] = {}
self.max_conversations = max(int(max_conversations), 1)
self.max_messages = max(int(max_messages), 1)
self.idle_ttl_seconds = max(int(idle_ttl_seconds), 1)
def _prune(self, now: float) -> None:
expired = [
key for key, last_accessed in self.last_accessed.items() if now - last_accessed >= self.idle_ttl_seconds
]
for key in expired:
self.reset(key)
overflow = len(self.message_lists) - self.max_conversations + 1
if overflow <= 0:
return
oldest = sorted(self.last_accessed, key=self.last_accessed.get)
for key in oldest[:overflow]:
self.reset(key)
def get_message_list(self, pipeline_uuid: str) -> list[WebSocketMessage]:
now = time.monotonic()
self._prune(now)
if pipeline_uuid not in self.message_lists:
self.message_lists[pipeline_uuid] = []
self.last_accessed[pipeline_uuid] = now
return self.message_lists[pipeline_uuid]
def get_stream_message_indexes(self, pipeline_uuid: str) -> dict[str, int]:
if pipeline_uuid not in self.stream_message_indexes:
self.stream_message_indexes[pipeline_uuid] = {}
self.last_accessed[pipeline_uuid] = time.monotonic()
return self.stream_message_indexes[pipeline_uuid]
def next_message_id(self, conversation_key: str) -> int:
next_id = self.message_counters.get(conversation_key, 0) + 1
self.message_counters[conversation_key] = next_id
return next_id
def append_message(self, conversation_key: str, message: WebSocketMessage) -> None:
messages = self.get_message_list(conversation_key)
messages.append(message)
overflow = len(messages) - self.max_messages
if overflow <= 0:
return
del messages[:overflow]
indexes = self.stream_message_indexes.get(conversation_key, {})
adjusted_indexes = {
response_id: index - overflow for response_id, index in indexes.items() if index >= overflow
}
indexes.clear()
indexes.update(adjusted_indexes)
def reset(self, conversation_key: str) -> None:
self.message_lists.pop(conversation_key, None)
self.stream_message_indexes.pop(conversation_key, None)
self.message_counters.pop(conversation_key, None)
self.last_accessed.pop(conversation_key, None)
def clear(self) -> None:
self.message_lists.clear()
self.stream_message_indexes.clear()
self.message_counters.clear()
self.last_accessed.clear()
class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
"""WebSocket适配器 - 支持双向实时通信"""
@@ -75,7 +138,14 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
ap: app.Application = pydantic.Field(exclude=True)
# 主动推送消息的队列
outbound_message_queue: asyncio.Queue = pydantic.Field(default_factory=asyncio.Queue, exclude=True)
outbound_message_queue: asyncio.Queue = pydantic.Field(
default_factory=lambda: asyncio.Queue(maxsize=100),
exclude=True,
)
inbound_listener_tasks: set[asyncio.Task] = pydantic.Field(
default_factory=set,
exclude=True,
)
"""后端主动推送消息的队列"""
# 流式输出开关
@@ -89,11 +159,26 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
**kwargs,
)
self.websocket_person_session = WebSocketSession(id='websocketperson')
self.websocket_group_session = WebSocketSession(id='websocketgroup')
application = kwargs.get('ap')
instance_data = getattr(getattr(application, 'instance_config', None), 'data', {})
retention = (
instance_data.get('system', {}).get('websocket_retention', {}) if isinstance(instance_data, dict) else {}
)
session_options = {
'max_conversations': retention.get('max_conversations_per_workspace', 200),
'max_messages': retention.get('max_messages_per_conversation', 100),
'idle_ttl_seconds': retention.get('conversation_idle_ttl_seconds', 86400),
}
self.websocket_person_session = WebSocketSession(id='websocketperson', **session_options)
self.websocket_group_session = WebSocketSession(id='websocketgroup', **session_options)
self.bot_account_id = 'websocketbot'
self.outbound_message_queue = asyncio.Queue()
try:
outbound_queue_size = max(int(retention.get('send_queue_size', 100)), 1)
except (TypeError, ValueError):
outbound_queue_size = 100
self.outbound_message_queue = asyncio.Queue(maxsize=outbound_queue_size)
self.inbound_listener_tasks = set()
self.stream_enabled = True
@staticmethod
@@ -128,6 +213,25 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
return _current_pipeline_uuid.get()
def _listener_task_done(self, task: asyncio.Task) -> None:
listener_tasks = getattr(self, 'inbound_listener_tasks', None)
if listener_tasks is not None:
listener_tasks.discard(task)
if not task.cancelled():
task.exception()
@staticmethod
def _history_message_chain(message_chain: list[dict]) -> list[dict]:
"""Remove large transient payloads before retaining browser history."""
history = []
for component in message_chain:
copied = dict(component)
if copied.get('base64'):
copied['base64'] = ''
history.append(copied)
return history
async def _get_connection_from_target(self, target_id: str):
"""Resolve a person or group WebSocket launcher to its connection."""
scope = self._scope()
@@ -195,7 +299,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
session = self.websocket_group_session if session_type == 'group' else self.websocket_person_session
msg_id = len(session.get_message_list(conversation_key)) + 1
msg_id = session.next_message_id(conversation_key)
message_data = WebSocketMessage(
id=msg_id,
@@ -206,7 +310,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
is_final=True,
)
session.get_message_list(conversation_key).append(message_data)
session.append_message(conversation_key, message_data)
await ws_connection_manager.broadcast_to_pipeline(
pipeline_uuid,
@@ -241,7 +345,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
session_type = 'group' if isinstance(message_source, platform_events.GroupMessage) else 'person'
conversation_key = self._conversation_key(pipeline_uuid, session_id)
msg_id = len(session.get_message_list(conversation_key)) + 1
msg_id = session.next_message_id(conversation_key)
message_data = WebSocketMessage(
id=msg_id,
@@ -252,7 +356,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
is_final=True,
)
session.get_message_list(conversation_key).append(message_data)
session.append_message(conversation_key, message_data)
await ws_connection_manager.broadcast_to_pipeline(
pipeline_uuid,
@@ -300,7 +404,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
if existing_index is None or existing_index >= len(message_list):
# 创建新消息
msg_id = len(message_list) + 1
msg_id = session.next_message_id(conversation_key)
message_data = WebSocketMessage(
id=msg_id,
role='assistant',
@@ -311,7 +415,8 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
)
# 立即添加到历史记录(即使is_final=False),以便后续块可以更新它
message_list.append(message_data)
session.append_message(conversation_key, message_data)
message_list = session.get_message_list(conversation_key)
if resp_message_id:
stream_message_indexes[resp_message_id] = len(message_list) - 1
else:
@@ -399,7 +504,22 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
async def kill(self):
"""停止适配器"""
pass
await ws_connection_manager.close_scope(self._scope())
listener_tasks = getattr(self, 'inbound_listener_tasks', set())
inbound_tasks = list(listener_tasks)
for task in inbound_tasks:
if not task.done():
task.cancel()
if inbound_tasks:
await asyncio.gather(*inbound_tasks, return_exceptions=True)
listener_tasks.clear()
self.websocket_person_session.clear()
self.websocket_group_session.clear()
while not self.outbound_message_queue.empty():
try:
self.outbound_message_queue.get_nowait()
except asyncio.QueueEmpty:
break
async def _process_image_components(
self,
@@ -445,7 +565,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
try:
file_content = await storage_mgr.storage_provider.load(comp_path)
base64_str = base64.b64encode(file_content).decode('utf-8')
base64_str = (await asyncio.to_thread(base64.b64encode, file_content)).decode('utf-8')
lowered = comp_path.lower()
if comp_type == 'Image':
@@ -507,19 +627,19 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
message_chain = platform_message.MessageChain.model_validate(message_chain_obj)
message_id = len(use_session.get_message_list(conversation_key)) + 1
message_id = use_session.next_message_id(conversation_key)
# 保存用户消息
user_message = WebSocketMessage(
id=message_id,
role='user',
content=str(message_chain),
message_chain=message_chain_obj,
message_chain=self._history_message_chain(message_chain_obj),
timestamp=datetime.now().isoformat(),
connection_id=connection.connection_id,
is_final=True, # 用户消息始终是完整的,非流式
)
use_session.get_message_list(conversation_key).append(user_message)
use_session.append_message(conversation_key, user_message)
await ws_connection_manager.broadcast_to_pipeline(
pipeline_uuid,
@@ -573,9 +693,36 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
owner_bot.adapter.set_ws_adapter(self)
callback_adapter = owner_bot.adapter if (owner_bot and hasattr(owner_bot, 'adapter')) else self
if event.__class__ in listeners:
listener_tasks = getattr(self, 'inbound_listener_tasks', None)
if listener_tasks is None:
listener_tasks = set()
object.__setattr__(self, 'inbound_listener_tasks', listener_tasks)
for task in tuple(listener_tasks):
if task.done():
listener_tasks.discard(task)
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:
asyncio.create_task(listeners[event.__class__](event, callback_adapter))
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)
@@ -599,10 +746,14 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
"""Reset one pipeline/client conversation."""
conversation_key = self._conversation_key(pipeline_uuid, session_id)
session = self.websocket_person_session if session_type == 'person' else self.websocket_group_session
if conversation_key in session.message_lists:
session.message_lists[conversation_key] = []
if conversation_key in session.stream_message_indexes:
session.stream_message_indexes[conversation_key] = {}
if isinstance(session, WebSocketSession):
session.reset(conversation_key)
else:
# Compatibility for lightweight adapter doubles.
if conversation_key in session.message_lists:
session.message_lists[conversation_key] = []
if conversation_key in session.stream_message_indexes:
session.stream_message_indexes[conversation_key] = {}
if session_id:
launcher_id = (