mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-09 04:40:57 +00:00
feat(tenancy): add Workspace multi-tenant foundation (#2353)
* Document multi-tenant workspace architecture * Add OSS and commercial workspace boundaries * docs: redesign multi-tenant workspace architecture * feat(tenancy): implement workspace isolation * docs(tenancy): record verification evidence * docs(tenancy): revise single-instance SaaS topology * docs(tenancy): refine architecture options * docs: finalize cloud v2 multi-tenant decisions * feat(tenancy): establish cloud isolation foundations * feat(tenancy): harden shared cloud runtime boundaries * docs(tenancy): record final isolation verification * fix(tenancy): close isolation and permission gaps * docs(tenancy): record final isolation verification * feat(tenancy): connect cloud workspace control plane * fix(build): install git for pinned SDK * docs(cloud): update control plane verification * chore: update multi-tenant SDK pin * fix(cloud): skip legacy model sync during startup * test(cloud): preserve minimal model manager fixtures * fix(cloud): preserve authenticated account context * fix(cloud): reuse authenticated account for user info * feat(cloud): complete Workspace settings navigation * test(web): cover Workspace dropdown menu * feat(web): place workspace controls in sidebar * refactor(web): streamline workspace controls * style(web): format workspace layout test * fix(cloud): surface runtime and workspace plan status * fix(plugin): keep runtime identity stable across restarts * fix(ui): widen and center workspace switcher * fix(ui): hide roles from workspace switcher * fix(ui): align workspace switcher with sidebar entries * feat(workspace): add in-product collaboration and direct Cloud launch * style: format collaboration changes * fix(workspace): bind collaboration APIs to tenant UoW * fix(cloud): preserve Core-owned collaboration state * test(cloud): require Space identity for invite registration * feat(cloud): complete secure invitation experience * style(web): format invitation flows * fix(cloud): recover box runtime without unscoped skill reload * feat(oss): enforce invitation account and owner billing flows * style: format OSS account service * test(oss): cover invitation logout handoff * fix(oss): resolve workspace owner in scoped session * feat(cloud): harden multi-tenant runtime resources * fix(cloud): bound runtime restart storms * fix(cloud): eliminate periodic runtime CPU spikes * fix(cloud): enforce instance capacity ceilings * fix(cloud): scope public login capability discovery * fix(cloud): bound tenant maintenance and monitoring work * fix(runtime): bound tenant resource amplification * fix(deps): pin green multi-tenant plugin SDK * fix(cloud): handle unavailable skill capability * fix(security): require authentication for image file endpoint (H-2) - Changed /api/v1/files/image from AuthType.NONE to USER_TOKEN_OR_API_KEY - Added Permission.RESOURCE_VIEW requirement - Prevents unauthenticated cross-tenant file access via leaked keys - Fixes HIGH severity finding from multi-tenant security review docs: add comprehensive database migration guide - Complete migration steps for OSS → multi-tenant - Backup, execution, verification procedures - Rollback scenarios and recovery plans - Performance tuning recommendations * test: add comprehensive cross-tenant isolation tests Added 7 critical test scenarios for multi-tenant boundaries: - Cross-tenant bot access prevention - Viewer role read-only enforcement - Removed member immediate access revocation - Model provider credential isolation - WebSocket message isolation - Invitation token workspace scoping - Multi-workspace context validation These tests address P0-2 coverage gaps for: - workspaces.py (membership & invitation flows) - user.py (authentication & authorization) - websocket_chat.py (real-time isolation) - plugins.py (resource access control) docs: finalize database migration guide * fix(security): resolve M-1, M-2, M-3 security findings M-1: WebSocket authorization TOCTOU race (FIXED) - Changed _revalidate_websocket_authorization to return RequestContext - Ensures validated context is used immediately without race window - Prevents removed members from sending messages during revalidation gap M-2: Model Manager cache workspace isolation (VERIFIED) - Confirmed _CacheKey already uses 4-tuple: (instance, workspace, generation, resource) - Cache is properly scoped per workspace, no cross-tenant leakage possible - No code change needed, documented as working correctly M-3: Invitation lock workspace scoping (FIXED) - Changed lock key from token_digest to workspace_uuid:token_digest - Prevents DoS where attacker locks token in Workspace A to block Workspace B - Locks now isolated per workspace All MEDIUM severity findings from security review now resolved. * fix(cloud): unblock tenant CI and enforce knowledge quotas * fix(tenancy): scope rerank model sync --------- Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
@@ -1,7 +1,9 @@
|
||||
"""WebSocket适配器 - 支持双向通信的IM系统"""
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import logging
|
||||
import time
|
||||
import typing
|
||||
from datetime import datetime
|
||||
|
||||
@@ -13,9 +15,14 @@ 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 .websocket_manager import WebSocketConnection, is_valid_session_id, ws_connection_manager
|
||||
from ...core import entities as core_entities
|
||||
from .websocket_manager import WebSocketConnection, WebSocketScope, is_valid_session_id, ws_connection_manager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_current_pipeline_uuid: contextvars.ContextVar[str | None] = contextvars.ContextVar(
|
||||
'websocket_pipeline_uuid',
|
||||
default=None,
|
||||
)
|
||||
|
||||
|
||||
class WebSocketMessage(pydantic.BaseModel):
|
||||
@@ -40,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适配器 - 支持双向实时通信"""
|
||||
@@ -70,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,
|
||||
)
|
||||
"""后端主动推送消息的队列"""
|
||||
|
||||
# 流式输出开关
|
||||
@@ -84,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
|
||||
@@ -113,9 +203,38 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
return None
|
||||
return pipeline_uuid, session_id
|
||||
|
||||
@classmethod
|
||||
async def _get_connection_from_target(cls, target_id: str):
|
||||
def _scope(self) -> WebSocketScope:
|
||||
"""Return this adapter's immutable runtime placement."""
|
||||
|
||||
return WebSocketScope.from_context(self.logger.execution_context)
|
||||
|
||||
def get_pipeline_uuid_override(self) -> str | None:
|
||||
"""Return the connection pipeline propagated into the listener task."""
|
||||
|
||||
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()
|
||||
target_value = str(target_id)
|
||||
for prefix in ('websocket_', 'websocketgroup_'):
|
||||
if target_value.startswith(prefix):
|
||||
@@ -123,14 +242,18 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
break
|
||||
else:
|
||||
return None
|
||||
connection = await ws_connection_manager.get_connection(target)
|
||||
connection = await ws_connection_manager.get_connection(target, scope=scope)
|
||||
if connection is not None:
|
||||
return connection
|
||||
embed_target = cls._parse_embed_target(target_id)
|
||||
embed_target = self._parse_embed_target(target_id)
|
||||
if embed_target is not None:
|
||||
pipeline_uuid, session_id = embed_target
|
||||
return await ws_connection_manager.get_connection_by_session_id(session_id, pipeline_uuid)
|
||||
return await ws_connection_manager.get_connection_by_session_id(target)
|
||||
return await ws_connection_manager.get_connection_by_session_id(
|
||||
session_id,
|
||||
scope=scope,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
)
|
||||
return await ws_connection_manager.get_connection_by_session_id(target, scope=scope)
|
||||
|
||||
async def _get_message_context(self, message_source) -> tuple[str, str | None]:
|
||||
"""Resolve the originating pipeline and browser session for a reply."""
|
||||
@@ -142,7 +265,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
embed_target = self._parse_embed_target(sender_id)
|
||||
if embed_target is not None:
|
||||
return embed_target
|
||||
return typing.cast(str, self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid), None
|
||||
raise ValueError('WebSocket reply target is not bound to this adapter scope')
|
||||
|
||||
async def send_message(
|
||||
self,
|
||||
@@ -160,22 +283,23 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
if connection is not None:
|
||||
pipeline_uuid = connection.pipeline_uuid
|
||||
session_id = connection.session_id
|
||||
scope = connection.scope
|
||||
else:
|
||||
embed_target = self._parse_embed_target(target_id)
|
||||
if embed_target is not None:
|
||||
pipeline_uuid, session_id = embed_target
|
||||
else:
|
||||
pipeline_uuid = typing.cast(
|
||||
str,
|
||||
self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid,
|
||||
)
|
||||
pipeline_uuid = str(target_id).strip()
|
||||
if not pipeline_uuid:
|
||||
raise ValueError('WebSocket target pipeline is required')
|
||||
session_id = None
|
||||
scope = self._scope()
|
||||
session_type = 'group' if target_type == 'group' else 'person'
|
||||
conversation_key = self._conversation_key(pipeline_uuid, session_id)
|
||||
|
||||
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,
|
||||
@@ -186,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,
|
||||
@@ -195,6 +319,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
'session_type': session_type,
|
||||
'data': message_data.model_dump(),
|
||||
},
|
||||
scope=scope,
|
||||
session_type=session_type,
|
||||
session_id=session_id,
|
||||
)
|
||||
@@ -216,10 +341,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
)
|
||||
|
||||
pipeline_uuid, session_id = await self._get_message_context(message_source)
|
||||
scope = self._scope()
|
||||
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,
|
||||
@@ -230,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,
|
||||
@@ -239,6 +365,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
'session_type': session_type,
|
||||
'data': message_data.model_dump(),
|
||||
},
|
||||
scope=scope,
|
||||
session_type=session_type,
|
||||
session_id=session_id,
|
||||
)
|
||||
@@ -262,6 +389,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
)
|
||||
|
||||
pipeline_uuid, session_id = await self._get_message_context(message_source)
|
||||
scope = self._scope()
|
||||
session_type = 'group' if isinstance(message_source, platform_events.GroupMessage) else 'person'
|
||||
conversation_key = self._conversation_key(pipeline_uuid, session_id)
|
||||
message_list = session.get_message_list(conversation_key)
|
||||
@@ -276,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',
|
||||
@@ -287,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:
|
||||
@@ -316,6 +445,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
'session_type': session_type,
|
||||
'data': message_data.model_dump(),
|
||||
},
|
||||
scope=scope,
|
||||
session_type=session_type,
|
||||
session_id=session_id,
|
||||
)
|
||||
@@ -360,7 +490,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
message = await asyncio.wait_for(self.outbound_message_queue.get(), timeout=0.1)
|
||||
# 广播到所有相关连接
|
||||
target_id = message.get('target_id', '')
|
||||
await ws_connection_manager.broadcast_to_pipeline(target_id, message)
|
||||
await ws_connection_manager.broadcast_to_pipeline(
|
||||
target_id,
|
||||
message,
|
||||
scope=self._scope(),
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
pass
|
||||
|
||||
@@ -370,9 +504,28 @@ 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, message_chain_obj: list):
|
||||
async def _process_image_components(
|
||||
self,
|
||||
connection: WebSocketConnection,
|
||||
message_chain_obj: list,
|
||||
):
|
||||
"""
|
||||
处理消息链中的图片、语音和文件组件,将 path 转换为 base64
|
||||
|
||||
@@ -387,18 +540,36 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
import base64
|
||||
import mimetypes
|
||||
|
||||
storage_mgr = self.ap.storage_mgr
|
||||
attachments = [
|
||||
component
|
||||
for component in message_chain_obj
|
||||
if component.get('path') and component.get('type') in ('Image', 'Voice', 'File')
|
||||
]
|
||||
if not attachments:
|
||||
return
|
||||
|
||||
for component in message_chain_obj:
|
||||
storage_mgr = self.ap.storage_mgr
|
||||
execution_context = connection.execution_context
|
||||
expected_prefix = storage_mgr.scoped_prefix(execution_context, owner_type='upload_image')
|
||||
|
||||
for component in attachments:
|
||||
comp_type = component.get('type', '')
|
||||
comp_path = component.get('path', '')
|
||||
|
||||
if not comp_path or comp_type not in ('Image', 'Voice', 'File'):
|
||||
continue
|
||||
if not comp_path.startswith(expected_prefix) or not storage_mgr.is_scoped_object_key(
|
||||
comp_path,
|
||||
expected_owner_type='upload_image',
|
||||
):
|
||||
await self.logger.warning(f'Rejected {comp_type} attachment outside the WebSocket connection scope')
|
||||
raise ValueError('Attachment key does not belong to this WebSocket connection')
|
||||
|
||||
try:
|
||||
file_content = await storage_mgr.storage_provider.load(comp_path)
|
||||
base64_str = base64.b64encode(file_content).decode('utf-8')
|
||||
file_content = await storage_mgr.load_scoped_object_key(
|
||||
execution_context,
|
||||
comp_path,
|
||||
expected_owner_type='upload_image',
|
||||
)
|
||||
base64_str = (await asyncio.to_thread(base64.b64encode, file_content)).decode('utf-8')
|
||||
|
||||
lowered = comp_path.lower()
|
||||
if comp_type == 'Image':
|
||||
@@ -416,10 +587,15 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
mime_type = mimetypes.guess_type(comp_path)[0] or 'application/octet-stream'
|
||||
|
||||
component['base64'] = f'data:{mime_type};base64,{base64_str}'
|
||||
await storage_mgr.storage_provider.delete(comp_path)
|
||||
await storage_mgr.delete_scoped_object_key(
|
||||
execution_context,
|
||||
comp_path,
|
||||
expected_owner_type='upload_image',
|
||||
)
|
||||
component['path'] = ''
|
||||
except Exception as e:
|
||||
await self.logger.error(f'Failed to load {comp_type} file {comp_path}: {e}')
|
||||
raise
|
||||
|
||||
async def handle_websocket_message(
|
||||
self,
|
||||
@@ -451,23 +627,23 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
|
||||
message_chain_obj = message_data.get('message', [])
|
||||
|
||||
await self._process_image_components(message_chain_obj)
|
||||
await self._process_image_components(connection, message_chain_obj)
|
||||
|
||||
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,
|
||||
@@ -476,6 +652,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
'session_type': session_type,
|
||||
'data': user_message.model_dump(),
|
||||
},
|
||||
scope=connection.scope,
|
||||
session_type=session_type,
|
||||
session_id=connection.session_id,
|
||||
)
|
||||
@@ -506,11 +683,6 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
sender=sender, message_chain=message_chain, time=datetime.now().timestamp()
|
||||
)
|
||||
|
||||
# 设置流水线UUID (proxy bot always needs it for reply_message routing)
|
||||
self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid = pipeline_uuid
|
||||
if owner_bot is not None:
|
||||
owner_bot.bot_entity.use_pipeline_uuid = pipeline_uuid
|
||||
|
||||
# 异步触发事件处理
|
||||
# Use owner_bot's listeners if available, otherwise fall back to proxy bot
|
||||
listeners = (
|
||||
@@ -525,7 +697,38 @@ 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:
|
||||
asyncio.create_task(listeners[event.__class__](event, callback_adapter))
|
||||
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:
|
||||
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)
|
||||
|
||||
def get_websocket_messages(
|
||||
self,
|
||||
@@ -547,10 +750,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 = (
|
||||
@@ -558,11 +765,15 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
||||
if session_type == 'group'
|
||||
else f'websocket_{pipeline_uuid}:{session_id}'
|
||||
)
|
||||
scope = self._scope()
|
||||
self.ap.sess_mgr.session_list = [
|
||||
candidate_session
|
||||
for candidate_session in self.ap.sess_mgr.session_list
|
||||
if not (
|
||||
str(
|
||||
getattr(candidate_session, 'instance_uuid', None) == scope.instance_uuid
|
||||
and getattr(candidate_session, 'workspace_uuid', None) == scope.workspace_uuid
|
||||
and getattr(candidate_session, 'placement_generation', None) == scope.placement_generation
|
||||
and str(
|
||||
candidate_session.launcher_type.value
|
||||
if hasattr(candidate_session.launcher_type, 'value')
|
||||
else candidate_session.launcher_type
|
||||
|
||||
Reference in New Issue
Block a user