mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-14 14:31:00 +00:00
feat(tenancy): implement workspace isolation
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import typing
|
||||
import traceback
|
||||
|
||||
@@ -13,7 +14,11 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
||||
import langbot_plugin.api.entities.builtin.platform.events as platform_events
|
||||
import langbot_plugin.api.entities.events as events
|
||||
from ..utils import importutil
|
||||
from ..api.http.authz import WorkspaceRequiredError
|
||||
from ..api.http.context import ExecutionContext, PrincipalContext, PrincipalType, RequestContext
|
||||
from ..workspace.errors import WorkspaceError, WorkspaceInvariantError
|
||||
from .config_coercion import coerce_pipeline_config
|
||||
from .pool import get_query_execution_context
|
||||
|
||||
import langbot_plugin.api.entities.builtin.provider.session as provider_session
|
||||
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
|
||||
@@ -82,15 +87,39 @@ class RuntimePipeline:
|
||||
enable_all_mcp_servers: bool
|
||||
"""是否启用所有MCP服务器"""
|
||||
|
||||
execution_context: ExecutionContext
|
||||
|
||||
workspace_uuid: str
|
||||
|
||||
placement_generation: int
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ap: app.Application,
|
||||
pipeline_entity: persistence_pipeline.LegacyPipeline,
|
||||
stage_containers: list[StageInstContainer],
|
||||
execution_context: ExecutionContext,
|
||||
):
|
||||
if not isinstance(execution_context, ExecutionContext):
|
||||
raise WorkspaceRequiredError('RuntimePipeline requires an ExecutionContext')
|
||||
if not execution_context.instance_uuid.strip() or not execution_context.workspace_uuid.strip():
|
||||
raise WorkspaceRequiredError('RuntimePipeline requires an instance and Workspace')
|
||||
if execution_context.placement_generation <= 0:
|
||||
raise WorkspaceRequiredError('RuntimePipeline requires a positive placement generation')
|
||||
if pipeline_entity.workspace_uuid != execution_context.workspace_uuid:
|
||||
raise WorkspaceRequiredError('RuntimePipeline entity Workspace does not match its ExecutionContext')
|
||||
if execution_context.pipeline_uuid not in (None, pipeline_entity.uuid):
|
||||
raise WorkspaceRequiredError('RuntimePipeline UUID does not match its ExecutionContext')
|
||||
|
||||
self.ap = ap
|
||||
self.pipeline_entity = pipeline_entity
|
||||
self.stage_containers = stage_containers
|
||||
self.execution_context = dataclasses.replace(
|
||||
execution_context,
|
||||
pipeline_uuid=pipeline_entity.uuid,
|
||||
)
|
||||
self.workspace_uuid = self.execution_context.workspace_uuid
|
||||
self.placement_generation = self.execution_context.placement_generation
|
||||
|
||||
# Extract bound plugins and MCP servers from extensions_preferences
|
||||
extensions_prefs = pipeline_entity.extensions_preferences or {}
|
||||
@@ -120,7 +149,37 @@ class RuntimePipeline:
|
||||
mcp_server_list = extensions_prefs.get('mcp_servers', [])
|
||||
self.bound_mcp_servers = mcp_server_list if mcp_server_list else []
|
||||
|
||||
async def _assert_execution_active(
|
||||
self,
|
||||
query: pipeline_query.Query | None = None,
|
||||
) -> ExecutionContext:
|
||||
"""Fail closed when this runtime or query belongs to a stale placement."""
|
||||
|
||||
execution_context = self.execution_context if query is None else get_query_execution_context(query)
|
||||
if (
|
||||
execution_context.instance_uuid != self.execution_context.instance_uuid
|
||||
or execution_context.workspace_uuid != self.workspace_uuid
|
||||
or execution_context.placement_generation != self.placement_generation
|
||||
or execution_context.pipeline_uuid != self.pipeline_entity.uuid
|
||||
):
|
||||
raise WorkspaceInvariantError('Query execution scope does not match RuntimePipeline')
|
||||
binding = await self.ap.workspace_service.get_execution_binding(
|
||||
execution_context.workspace_uuid,
|
||||
expected_generation=execution_context.placement_generation,
|
||||
)
|
||||
if binding.instance_uuid != execution_context.instance_uuid:
|
||||
raise WorkspaceInvariantError('RuntimePipeline instance does not match the active Workspace binding')
|
||||
return execution_context
|
||||
|
||||
async def run(self, query: pipeline_query.Query):
|
||||
if (
|
||||
query.instance_uuid != self.execution_context.instance_uuid
|
||||
or query.workspace_uuid != self.workspace_uuid
|
||||
or query.placement_generation != self.placement_generation
|
||||
or query.pipeline_uuid != self.pipeline_entity.uuid
|
||||
):
|
||||
raise WorkspaceRequiredError('Query execution scope does not match RuntimePipeline')
|
||||
await self._assert_execution_active(query)
|
||||
query.pipeline_config = self.pipeline_entity.config
|
||||
# Store bound plugins and MCP servers in query for filtering
|
||||
query.variables['_pipeline_bound_plugins'] = self.bound_plugins
|
||||
@@ -134,7 +193,11 @@ class RuntimePipeline:
|
||||
bot_name = 'WebChat'
|
||||
if query.bot_uuid:
|
||||
try:
|
||||
bot = await self.ap.bot_service.get_bot(query.bot_uuid, include_secret=False)
|
||||
bot = await self.ap.bot_service.get_bot(
|
||||
query.workspace_uuid,
|
||||
query.bot_uuid,
|
||||
include_secret=False,
|
||||
)
|
||||
if bot:
|
||||
bot_name = bot.get('name', 'Unknown')
|
||||
except Exception:
|
||||
@@ -150,6 +213,7 @@ class RuntimePipeline:
|
||||
|
||||
async def _check_output(self, query: pipeline_query.Query, result: pipeline_entities.StageProcessResult):
|
||||
"""检查输出"""
|
||||
await self._assert_execution_active(query)
|
||||
if result.user_notice:
|
||||
# 处理str类型
|
||||
|
||||
@@ -162,7 +226,9 @@ class RuntimePipeline:
|
||||
query.message_event, platform_events.GroupMessage
|
||||
):
|
||||
result.user_notice.insert(0, platform_message.At(target=query.message_event.sender.id))
|
||||
if await query.adapter.is_stream_output_supported() and query.resp_messages:
|
||||
stream_output_supported = await query.adapter.is_stream_output_supported()
|
||||
await self._assert_execution_active(query)
|
||||
if stream_output_supported and query.resp_messages:
|
||||
await query.adapter.reply_message_chunk(
|
||||
message_source=query.message_event,
|
||||
bot_message=query.resp_messages[-1],
|
||||
@@ -186,6 +252,7 @@ class RuntimePipeline:
|
||||
query.variables['_monitoring_has_error'] = True
|
||||
# Record error to monitoring system
|
||||
try:
|
||||
await self._assert_execution_active(query)
|
||||
bot_name = query.variables.get('_monitoring_bot_name', 'Unknown')
|
||||
pipeline_name = query.variables.get('_monitoring_pipeline_name', 'Unknown')
|
||||
message_id = query.variables.get('_monitoring_message_id', '')
|
||||
@@ -194,6 +261,7 @@ class RuntimePipeline:
|
||||
# Update message status to error
|
||||
if message_id:
|
||||
await self.ap.monitoring_service.update_message_status(
|
||||
get_query_execution_context(query),
|
||||
message_id=message_id,
|
||||
status='error',
|
||||
level='error',
|
||||
@@ -201,6 +269,7 @@ class RuntimePipeline:
|
||||
|
||||
# Record error log
|
||||
await self.ap.monitoring_service.record_error(
|
||||
get_query_execution_context(query),
|
||||
bot_id=query.bot_uuid or 'unknown',
|
||||
bot_name=bot_name,
|
||||
pipeline_id=self.pipeline_entity.uuid,
|
||||
@@ -242,6 +311,7 @@ class RuntimePipeline:
|
||||
i = stage_index
|
||||
|
||||
while i < len(self.stage_containers):
|
||||
await self._assert_execution_active(query)
|
||||
stage_container = self.stage_containers[i]
|
||||
|
||||
query.current_stage_name = stage_container.inst_name # 标记到 Query 对象里
|
||||
@@ -250,6 +320,7 @@ class RuntimePipeline:
|
||||
|
||||
if isinstance(result, typing.Coroutine):
|
||||
result = await result
|
||||
await self._assert_execution_active(query)
|
||||
|
||||
if isinstance(result, pipeline_entities.StageProcessResult): # 直接返回结果
|
||||
self.ap.logger.debug(
|
||||
@@ -265,7 +336,14 @@ class RuntimePipeline:
|
||||
elif isinstance(result, typing.AsyncGenerator): # 生成器
|
||||
self.ap.logger.debug(f'Stage {stage_container.inst_name} processed query {query.query_id} gen')
|
||||
|
||||
async for sub_result in result:
|
||||
iterator = result.__aiter__()
|
||||
while True:
|
||||
await self._assert_execution_active(query)
|
||||
try:
|
||||
sub_result = await anext(iterator)
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
await self._assert_execution_active(query)
|
||||
self.ap.logger.debug(
|
||||
f'Stage {stage_container.inst_name} processed query {query.query_id} res {sub_result.result_type}'
|
||||
)
|
||||
@@ -283,6 +361,7 @@ class RuntimePipeline:
|
||||
|
||||
async def process_query(self, query: pipeline_query.Query):
|
||||
"""处理请求"""
|
||||
await self._assert_execution_active(query)
|
||||
# Get monitoring metadata
|
||||
bot_name = query.variables.get('_monitoring_bot_name', 'Unknown')
|
||||
pipeline_name = query.variables.get('_monitoring_pipeline_name', 'Unknown')
|
||||
@@ -310,6 +389,7 @@ class RuntimePipeline:
|
||||
query.variables['_monitoring_message_id'] = message_id
|
||||
# Notify adapter so it can map platform-specific IDs to monitoring message ID
|
||||
if hasattr(query.adapter, 'on_monitoring_message_created'):
|
||||
await self._assert_execution_active(query)
|
||||
await query.adapter.on_monitoring_message_created(query, message_id)
|
||||
except Exception as e:
|
||||
self.ap.logger.error(f'Failed to record query start: {e}')
|
||||
@@ -334,7 +414,9 @@ class RuntimePipeline:
|
||||
message_chain=query.message_chain,
|
||||
)
|
||||
|
||||
await self._assert_execution_active(query)
|
||||
event_ctx = await self.ap.plugin_connector.emit_event(event_obj, bound_plugins)
|
||||
await self._assert_execution_active(query)
|
||||
|
||||
if event_ctx.is_prevented_default():
|
||||
self.ap.logger.debug(
|
||||
@@ -349,6 +431,7 @@ class RuntimePipeline:
|
||||
# Record query success only if no error occurred during processing
|
||||
if not query.variables.get('_monitoring_has_error', False):
|
||||
try:
|
||||
await self._assert_execution_active(query)
|
||||
await monitoring_helper.MonitoringHelper.record_query_success(
|
||||
ap=self.ap,
|
||||
message_id=message_id,
|
||||
@@ -359,6 +442,7 @@ class RuntimePipeline:
|
||||
|
||||
# Record bot response message
|
||||
try:
|
||||
await self._assert_execution_active(query)
|
||||
await monitoring_helper.MonitoringHelper.record_query_response(
|
||||
ap=self.ap,
|
||||
query=query,
|
||||
@@ -371,6 +455,8 @@ class RuntimePipeline:
|
||||
except Exception as e:
|
||||
self.ap.logger.error(f'Failed to record query response: {e}')
|
||||
|
||||
except WorkspaceError as e:
|
||||
self.ap.logger.info(f'Dropped query {query.query_id} because its Workspace execution binding is stale: {e}')
|
||||
except Exception as e:
|
||||
inst_name = query.current_stage_name if query.current_stage_name else 'unknown'
|
||||
self.ap.logger.error(f'Error processing query {query.query_id} stage={inst_name} : {e}')
|
||||
@@ -380,6 +466,7 @@ class RuntimePipeline:
|
||||
try:
|
||||
from . import monitoring_helper
|
||||
|
||||
await self._assert_execution_active(query)
|
||||
await monitoring_helper.MonitoringHelper.record_query_error(
|
||||
ap=self.ap,
|
||||
query=query,
|
||||
@@ -395,7 +482,7 @@ class RuntimePipeline:
|
||||
|
||||
finally:
|
||||
self.ap.logger.debug(f'Query {query.query_id} processed')
|
||||
del self.ap.query_pool.cached_queries[query.query_id]
|
||||
await self.ap.query_pool.remove_query(query)
|
||||
|
||||
|
||||
class PipelineManager:
|
||||
@@ -425,10 +512,38 @@ class PipelineManager:
|
||||
|
||||
# load pipelines
|
||||
for pipeline in pipelines:
|
||||
await self.load_pipeline(pipeline)
|
||||
binding = await self.ap.workspace_service.get_execution_binding(pipeline.workspace_uuid)
|
||||
await self.load_pipeline(
|
||||
ExecutionContext(
|
||||
instance_uuid=binding.instance_uuid,
|
||||
workspace_uuid=binding.workspace_uuid,
|
||||
placement_generation=binding.placement_generation,
|
||||
pipeline_uuid=pipeline.uuid,
|
||||
trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
|
||||
),
|
||||
pipeline,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_execution_context(
|
||||
context: ExecutionContext | RequestContext,
|
||||
pipeline_uuid: str,
|
||||
) -> ExecutionContext:
|
||||
if isinstance(context, RequestContext):
|
||||
return ExecutionContext.from_request(context, pipeline_uuid=pipeline_uuid)
|
||||
if not isinstance(context, ExecutionContext):
|
||||
raise WorkspaceRequiredError('Pipeline runtime operations require an ExecutionContext')
|
||||
if not context.instance_uuid.strip() or not context.workspace_uuid.strip():
|
||||
raise WorkspaceRequiredError('Pipeline runtime operations require an instance and Workspace')
|
||||
if context.placement_generation <= 0:
|
||||
raise WorkspaceRequiredError('Pipeline runtime operations require a positive placement generation')
|
||||
if context.pipeline_uuid not in (None, pipeline_uuid):
|
||||
raise WorkspaceRequiredError('Pipeline UUID does not match its ExecutionContext')
|
||||
return dataclasses.replace(context, pipeline_uuid=pipeline_uuid)
|
||||
|
||||
async def load_pipeline(
|
||||
self,
|
||||
context: ExecutionContext | RequestContext,
|
||||
pipeline_entity: persistence_pipeline.LegacyPipeline
|
||||
| sqlalchemy.Row[persistence_pipeline.LegacyPipeline]
|
||||
| dict,
|
||||
@@ -438,6 +553,14 @@ class PipelineManager:
|
||||
elif isinstance(pipeline_entity, dict):
|
||||
pipeline_entity = persistence_pipeline.LegacyPipeline(**pipeline_entity)
|
||||
|
||||
execution_context = self._normalize_execution_context(context, pipeline_entity.uuid)
|
||||
if pipeline_entity.workspace_uuid != execution_context.workspace_uuid:
|
||||
raise WorkspaceRequiredError('Pipeline entity Workspace does not match its runtime context')
|
||||
await self.ap.workspace_service.get_execution_binding(
|
||||
execution_context.workspace_uuid,
|
||||
expected_generation=execution_context.placement_generation,
|
||||
)
|
||||
|
||||
coerce_pipeline_config(
|
||||
pipeline_entity.config,
|
||||
getattr(self.ap, 'pipeline_config_meta_trigger', {'name': 'trigger', 'stages': []}),
|
||||
@@ -454,17 +577,36 @@ class PipelineManager:
|
||||
for stage_container in stage_containers:
|
||||
await stage_container.inst.initialize(pipeline_entity.config)
|
||||
|
||||
runtime_pipeline = RuntimePipeline(self.ap, pipeline_entity, stage_containers)
|
||||
runtime_pipeline = RuntimePipeline(
|
||||
self.ap,
|
||||
pipeline_entity,
|
||||
stage_containers,
|
||||
execution_context,
|
||||
)
|
||||
self.pipelines.append(runtime_pipeline)
|
||||
|
||||
async def get_pipeline_by_uuid(self, uuid: str) -> RuntimePipeline | None:
|
||||
async def get_pipeline_by_uuid(
|
||||
self,
|
||||
context: ExecutionContext | RequestContext,
|
||||
uuid: str,
|
||||
) -> RuntimePipeline | None:
|
||||
execution_context = self._normalize_execution_context(context, uuid)
|
||||
for pipeline in self.pipelines:
|
||||
if pipeline.pipeline_entity.uuid == uuid:
|
||||
if (
|
||||
pipeline.workspace_uuid == execution_context.workspace_uuid
|
||||
and pipeline.placement_generation == execution_context.placement_generation
|
||||
and pipeline.pipeline_entity.uuid == uuid
|
||||
):
|
||||
return pipeline
|
||||
return None
|
||||
|
||||
async def remove_pipeline(self, uuid: str):
|
||||
async def remove_pipeline(
|
||||
self,
|
||||
context: ExecutionContext | RequestContext,
|
||||
uuid: str,
|
||||
) -> None:
|
||||
execution_context = self._normalize_execution_context(context, uuid)
|
||||
for pipeline in self.pipelines:
|
||||
if pipeline.pipeline_entity.uuid == uuid:
|
||||
if pipeline.workspace_uuid == execution_context.workspace_uuid and pipeline.pipeline_entity.uuid == uuid:
|
||||
self.pipelines.remove(pipeline)
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user