feat(tenancy): implement workspace isolation

This commit is contained in:
Junyan Qin
2026-07-19 09:58:59 +08:00
parent 37099ddf7e
commit 8b7ce77cec
271 changed files with 31166 additions and 6513 deletions
+422 -305
View File
@@ -1,40 +1,55 @@
from __future__ import annotations
import asyncio
import sqlalchemy
import traceback
from typing import TypeVar
from . import requester
import sqlalchemy
from ...api.http.context import (
ExecutionContext,
PrincipalContext,
PrincipalType,
RequestContext,
)
from ...api.http.service.tenant import TenantContext, require_workspace_uuid
from ...core import app
from ...discover import engine
from . import token
from ...entity.persistence import model as persistence_model
from ...entity.errors import provider as provider_errors
from ...entity.persistence import model as persistence_model
from ...workspace.entities import WorkspaceExecutionBinding
from ...workspace.errors import WorkspaceError, WorkspaceInvariantError
from . import requester, token
_CacheKey = tuple[str, str, int, str]
_ModelEntity = TypeVar(
'_ModelEntity',
persistence_model.LLMModel,
persistence_model.EmbeddingModel,
persistence_model.RerankModel,
)
class ModelManager:
"""Model manager"""
"""Workspace-scoped runtime provider and model cache."""
ap: app.Application
provider_dict: dict[str, requester.RuntimeProvider]
"""运行时模型提供商字典, uuid -> RuntimeProvider"""
llm_models: list[requester.RuntimeLLMModel]
embedding_models: list[requester.RuntimeEmbeddingModel]
rerank_models: list[requester.RuntimeRerankModel]
provider_dict: dict[_CacheKey, requester.RuntimeProvider]
llm_model_dict: dict[_CacheKey, requester.RuntimeLLMModel]
embedding_model_dict: dict[_CacheKey, requester.RuntimeEmbeddingModel]
rerank_model_dict: dict[_CacheKey, requester.RuntimeRerankModel]
requester_components: list[engine.Component]
requester_dict: dict[str, type[requester.ProviderAPIRequester]]
def __init__(self, ap: app.Application):
self.ap = ap
self.llm_models = []
self.embedding_models = []
self.rerank_models = []
self.provider_dict = {}
self.llm_model_dict = {}
self.embedding_model_dict = {}
self.rerank_model_dict = {}
self.requester_components = []
self.requester_dict = {}
@@ -60,12 +75,75 @@ class ModelManager:
return litellm_provider
return None
async def initialize(self):
@staticmethod
def _context_from_binding(
binding: WorkspaceExecutionBinding,
*,
trigger_principal: PrincipalContext | None = None,
) -> ExecutionContext:
return ExecutionContext(
instance_uuid=binding.instance_uuid,
workspace_uuid=binding.workspace_uuid,
placement_generation=binding.placement_generation,
trigger_principal=trigger_principal,
)
@staticmethod
def _cache_key(context: ExecutionContext, resource_uuid: str) -> _CacheKey:
return (
context.instance_uuid,
context.workspace_uuid,
context.placement_generation,
resource_uuid,
)
@staticmethod
def _ensure_same_scope(
expected: ExecutionContext,
actual: ExecutionContext,
*,
resource: str,
) -> None:
if (
actual.instance_uuid != expected.instance_uuid
or actual.workspace_uuid != expected.workspace_uuid
or actual.placement_generation != expected.placement_generation
):
raise WorkspaceInvariantError(f'{resource} runtime belongs to another Workspace execution scope')
@staticmethod
def _ensure_entity_workspace(entity: object, context: ExecutionContext, *, resource: str) -> None:
workspace_uuid = getattr(entity, 'workspace_uuid', None)
if workspace_uuid != context.workspace_uuid:
raise WorkspaceInvariantError(f'{resource} belongs to another Workspace')
async def resolve_execution_context(self, context: TenantContext) -> ExecutionContext:
"""Resolve and fence-check an explicit tenant context for runtime access."""
workspace_uuid = require_workspace_uuid(context)
expected_generation = None
supplied_instance_uuid = None
trigger_principal = None
if isinstance(context, (RequestContext, ExecutionContext)):
expected_generation = context.placement_generation
supplied_instance_uuid = context.instance_uuid
trigger_principal = context.principal if isinstance(context, RequestContext) else context.trigger_principal
binding = await self.ap.workspace_service.get_execution_binding(
workspace_uuid,
expected_generation=expected_generation,
)
if supplied_instance_uuid is not None and supplied_instance_uuid != binding.instance_uuid:
raise WorkspaceInvariantError('Runtime context belongs to another LangBot instance')
return self._context_from_binding(binding, trigger_principal=trigger_principal)
async def initialize(self) -> None:
self.requester_components = self.ap.discover.get_components_by_kind('LLMAPIRequester')
requester_dict: dict[str, type[requester.ProviderAPIRequester]] = {}
for component in self.requester_components:
# Skip components that use litellm_provider (they will use litellmchat.py instead)
litellm_provider = self._get_litellm_provider_from_manifest(component)
if litellm_provider:
self.ap.logger.debug(
@@ -76,133 +154,151 @@ class ModelManager:
requester_dict[component.metadata.name] = component.get_python_component_class()
self.requester_dict = requester_dict
await self.load_models_from_db()
# Check if space models service is disabled
space_config = self.ap.instance_config.data.get('space', {})
if space_config.get('disable_models_service', False):
self.ap.logger.info('LangBot Space Models service is disabled, skipping sync.')
return
# Space model synchronization is a legacy OSS-singleton facility. A
# cloud instance must receive tenant model projections from its control
# plane and must never infer one Workspace for this global operation.
try:
binding = await self.ap.workspace_service.get_local_execution_binding()
except WorkspaceError as exc:
self.ap.logger.info(f'Skipping LangBot Space model sync outside an OSS local Workspace: {exc}')
return
sync_context = self._context_from_binding(
binding,
trigger_principal=PrincipalContext(principal_type=PrincipalType.SYSTEM),
)
sync_timeout = space_config.get('models_sync_timeout')
try:
if sync_timeout:
await asyncio.wait_for(
self.sync_new_models_from_space(),
self.sync_new_models_from_space(sync_context),
timeout=float(sync_timeout),
)
else:
await self.sync_new_models_from_space()
await self.sync_new_models_from_space(sync_context)
except asyncio.TimeoutError:
self.ap.logger.warning(f'LangBot Space model sync timed out after {sync_timeout}s, skipping startup sync.')
except Exception as e:
except Exception as exc:
self.ap.logger.warning('Failed to sync new models from LangBot Space, model list may not be updated.')
self.ap.logger.warning(f' - Error: {e}')
self.ap.logger.warning(f' - Error: {exc}')
async def load_models_from_db(self) -> None:
"""Load every active projected Workspace into isolated runtime caches."""
async def load_models_from_db(self):
"""Load models from database"""
self.ap.logger.info('Loading models from db...')
self.llm_models = []
self.embedding_models = []
self.rerank_models = []
self.provider_dict = {}
self.llm_model_dict = {}
self.embedding_model_dict = {}
self.rerank_model_dict = {}
contexts: dict[str, ExecutionContext] = {}
async def context_for(workspace_uuid: str | None) -> ExecutionContext:
if not workspace_uuid:
raise WorkspaceInvariantError('Runtime model resource has no Workspace')
cached = contexts.get(workspace_uuid)
if cached is not None:
return cached
binding = await self.ap.workspace_service.get_execution_binding(workspace_uuid)
resolved = self._context_from_binding(
binding,
trigger_principal=PrincipalContext(principal_type=PrincipalType.SYSTEM),
)
contexts[workspace_uuid] = resolved
return resolved
providers_result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_model.ModelProvider)
)
for provider in providers_result.all():
for provider_entity in providers_result.all():
try:
runtime_provider = await self.load_provider(provider)
self.provider_dict[provider.uuid] = runtime_provider
except provider_errors.RequesterNotFoundError as e:
self.ap.logger.warning(f'Requester {e.requester_name} not found, skipping provider {provider.uuid}')
continue
except Exception as e:
self.ap.logger.error(f'Failed to load provider {provider.uuid}: {e}\n{traceback.format_exc()}')
context = await context_for(provider_entity.workspace_uuid)
runtime_provider = await self._build_provider(context, provider_entity)
self.provider_dict[self._cache_key(context, provider_entity.uuid)] = runtime_provider
except provider_errors.RequesterNotFoundError as exc:
self.ap.logger.warning(
f'Requester {exc.requester_name} not found, skipping provider {provider_entity.uuid}'
)
except Exception as exc:
self.ap.logger.error(f'Failed to load provider {provider_entity.uuid}: {exc}\n{traceback.format_exc()}')
# Load LLM models
result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.LLMModel))
llm_models = result.all()
for llm_model in llm_models:
try:
provider = self.provider_dict.get(llm_model.provider_uuid)
if provider is None:
self.ap.logger.warning(f'Provider {llm_model.provider_uuid} not found for model {llm_model.uuid}')
continue
runtime_llm_model = await self.load_llm_model_with_provider(llm_model, provider)
self.llm_models.append(runtime_llm_model)
except Exception as e:
self.ap.logger.error(f'Failed to load model {llm_model.uuid}: {e}\n{traceback.format_exc()}')
await self._load_model_kind(
persistence_model.LLMModel,
self.llm_model_dict,
self._build_llm_model,
context_for,
)
await self._load_model_kind(
persistence_model.EmbeddingModel,
self.embedding_model_dict,
self._build_embedding_model,
context_for,
)
await self._load_model_kind(
persistence_model.RerankModel,
self.rerank_model_dict,
self._build_rerank_model,
context_for,
)
# Load embedding models
result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.EmbeddingModel))
embedding_models = result.all()
for embedding_model in embedding_models:
async def _load_model_kind(self, entity_type, cache: dict, builder, context_for) -> None:
result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(entity_type))
for model_entity in result.all():
try:
provider = self.provider_dict.get(embedding_model.provider_uuid)
context = await context_for(model_entity.workspace_uuid)
provider = self.provider_dict.get(self._cache_key(context, model_entity.provider_uuid))
if provider is None:
self.ap.logger.warning(
f'Provider {embedding_model.provider_uuid} not found for model {embedding_model.uuid}'
f'Provider {model_entity.provider_uuid} not found for model {model_entity.uuid}'
)
continue
runtime_embedding_model = await self.load_embedding_model_with_provider(embedding_model, provider)
self.embedding_models.append(runtime_embedding_model)
except Exception as e:
self.ap.logger.error(f'Failed to load model {embedding_model.uuid}: {e}\n{traceback.format_exc()}')
runtime_model = builder(context, model_entity, provider)
cache[self._cache_key(context, model_entity.uuid)] = runtime_model
except Exception as exc:
self.ap.logger.error(f'Failed to load model {model_entity.uuid}: {exc}\n{traceback.format_exc()}')
# Load rerank models
result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.RerankModel))
rerank_models = result.all()
for rerank_model in rerank_models:
try:
provider = self.provider_dict.get(rerank_model.provider_uuid)
if provider is None:
self.ap.logger.warning(
f'Provider {rerank_model.provider_uuid} not found for model {rerank_model.uuid}'
)
continue
runtime_rerank_model = await self.load_rerank_model_with_provider(rerank_model, provider)
self.rerank_models.append(runtime_rerank_model)
except Exception as e:
self.ap.logger.error(f'Failed to load model {rerank_model.uuid}: {e}\n{traceback.format_exc()}')
async def sync_new_models_from_space(self, context: ExecutionContext) -> None:
"""Sync legacy Space models for the explicitly selected OSS Workspace."""
async def sync_new_models_from_space(self):
"""Sync models from Space"""
space_model_provider = await self.ap.persistence_mgr.execute_async(
context = await self.resolve_execution_context(context)
await self.ap.workspace_service.get_local_execution_binding(
context.workspace_uuid,
expected_generation=context.placement_generation,
)
space_model_provider_result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_model.ModelProvider).where(
persistence_model.ModelProvider.requester == 'space-chat-completions'
persistence_model.ModelProvider.workspace_uuid == context.workspace_uuid,
persistence_model.ModelProvider.requester == 'space-chat-completions',
)
)
result = space_model_provider.first()
if result is None:
space_model_provider = space_model_provider_result.first()
if space_model_provider is None:
raise provider_errors.ProviderNotFoundError('LangBot Models')
space_model_provider = result
# get the latest models from space
space_models = await self.ap.space_service.get_models()
# Index existing models by uuid. Space reuses a model's uuid across
# renames / re-specs (e.g. the uuid that used to be ``claude-opus-4-6``
# may later become ``claude-opus-4-7``). So for Space-managed models we
# upsert: create when the uuid is new, otherwise update name/abilities/
# ranking to track Space. Models owned by other providers are never
# touched, even on an (unexpected) uuid collision.
existing_llm_models = {m['uuid']: m for m in await self.ap.llm_model_service.get_llm_models()}
existing_llm_models = {
model['uuid']: model
for model in await self.ap.llm_model_service.get_llm_models(context, include_secret=True)
}
existing_embedding_models = {
m['uuid']: m for m in await self.ap.embedding_models_service.get_embedding_models()
model['uuid']: model
for model in await self.ap.embedding_models_service.get_embedding_models(context, include_secret=True)
}
created = 0
updated = 0
for space_model in space_models:
if space_model.category == 'chat':
existing = existing_llm_models.get(space_model.uuid)
if existing is None:
# model will be automatically loaded
await self.ap.llm_model_service.create_llm_model(
context,
{
'uuid': space_model.uuid,
'name': space_model.model_id,
@@ -227,14 +323,14 @@ class ModelManager:
or list(existing.get('abilities') or []) != list(desired['abilities'])
or existing.get('prefered_ranking') != desired['prefered_ranking']
):
await self.ap.llm_model_service.update_llm_model(space_model.uuid, dict(desired))
await self.ap.llm_model_service.update_llm_model(context, space_model.uuid, dict(desired))
updated += 1
elif space_model.category == 'embedding':
existing = existing_embedding_models.get(space_model.uuid)
if existing is None:
# model will be automatically loaded
await self.ap.embedding_models_service.create_embedding_model(
context,
{
'uuid': space_model.uuid,
'name': space_model.model_id,
@@ -255,7 +351,11 @@ class ModelManager:
existing.get('name') != desired['name']
or existing.get('prefered_ranking') != desired['prefered_ranking']
):
await self.ap.embedding_models_service.update_embedding_model(space_model.uuid, dict(desired))
await self.ap.embedding_models_service.update_embedding_model(
context,
space_model.uuid,
dict(desired),
)
updated += 1
if created or updated:
@@ -263,313 +363,330 @@ class ModelManager:
async def init_temporary_runtime_llm_model(
self,
context: TenantContext,
model_info: dict,
) -> requester.RuntimeLLMModel:
"""Initialize runtime LLM model from dict (for testing)"""
provider_info = model_info.get('provider', {})
runtime_provider = await self.load_provider(provider_info)
runtime_llm_model = requester.RuntimeLLMModel(
model_entity=persistence_model.LLMModel(
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid='',
abilities=model_info.get('abilities', []),
context_length=model_info.get('context_length'),
extra_args=model_info.get('extra_args', {}),
),
provider=runtime_provider,
execution_context = await self.resolve_execution_context(context)
provider_info = {**model_info.get('provider', {}), 'workspace_uuid': execution_context.workspace_uuid}
runtime_provider = await self._build_provider(
execution_context,
persistence_model.ModelProvider(**provider_info),
)
return runtime_llm_model
model_entity = persistence_model.LLMModel(
workspace_uuid=execution_context.workspace_uuid,
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid=runtime_provider.provider_entity.uuid,
abilities=model_info.get('abilities', []),
context_length=model_info.get('context_length'),
extra_args=model_info.get('extra_args', {}),
)
return self._build_llm_model(execution_context, model_entity, runtime_provider)
async def init_temporary_runtime_embedding_model(
self,
context: TenantContext,
model_info: dict,
) -> requester.RuntimeEmbeddingModel:
"""Initialize runtime embedding model from dict (for testing)"""
provider_info = model_info.get('provider', {})
runtime_provider = await self.load_provider(provider_info)
runtime_embedding_model = requester.RuntimeEmbeddingModel(
model_entity=persistence_model.EmbeddingModel(
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid='',
extra_args=model_info.get('extra_args', {}),
),
provider=runtime_provider,
execution_context = await self.resolve_execution_context(context)
provider_info = {**model_info.get('provider', {}), 'workspace_uuid': execution_context.workspace_uuid}
runtime_provider = await self._build_provider(
execution_context,
persistence_model.ModelProvider(**provider_info),
)
return runtime_embedding_model
model_entity = persistence_model.EmbeddingModel(
workspace_uuid=execution_context.workspace_uuid,
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid=runtime_provider.provider_entity.uuid,
extra_args=model_info.get('extra_args', {}),
)
return self._build_embedding_model(execution_context, model_entity, runtime_provider)
async def init_temporary_runtime_rerank_model(
self,
context: TenantContext,
model_info: dict,
) -> requester.RuntimeRerankModel:
"""Initialize runtime rerank model from dict (for testing)"""
provider_info = model_info.get('provider', {})
runtime_provider = await self.load_provider(provider_info)
runtime_rerank_model = requester.RuntimeRerankModel(
model_entity=persistence_model.RerankModel(
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid='',
extra_args=model_info.get('extra_args', {}),
),
provider=runtime_provider,
execution_context = await self.resolve_execution_context(context)
provider_info = {**model_info.get('provider', {}), 'workspace_uuid': execution_context.workspace_uuid}
runtime_provider = await self._build_provider(
execution_context,
persistence_model.ModelProvider(**provider_info),
)
model_entity = persistence_model.RerankModel(
workspace_uuid=execution_context.workspace_uuid,
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid=runtime_provider.provider_entity.uuid,
extra_args=model_info.get('extra_args', {}),
)
return self._build_rerank_model(execution_context, model_entity, runtime_provider)
return runtime_rerank_model
async def load_provider(
self, provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict
) -> requester.RuntimeProvider:
"""Load provider from dict"""
@staticmethod
def _coerce_provider(
provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict,
context: ExecutionContext,
) -> persistence_model.ModelProvider:
if isinstance(provider_info, sqlalchemy.Row):
provider_entity = persistence_model.ModelProvider(**provider_info._mapping)
elif isinstance(provider_info, dict):
provider_entity = persistence_model.ModelProvider(**provider_info)
provider_entity = persistence_model.ModelProvider(
**{**provider_info, 'workspace_uuid': context.workspace_uuid}
)
else:
provider_entity = provider_info
ModelManager._ensure_entity_workspace(provider_entity, context, resource='Provider')
return provider_entity
# Get requester manifest to check for litellm_provider
async def _build_provider(
self,
context: ExecutionContext,
provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict,
) -> requester.RuntimeProvider:
provider_entity = self._coerce_provider(provider_info, context)
requester_manifest = self.get_available_requester_manifest_by_name(provider_entity.requester)
litellm_provider = self._get_litellm_provider_from_manifest(requester_manifest)
# Build config from base_url
config = {'base_url': provider_entity.base_url}
# Check if requester manifest specifies litellm_provider
if litellm_provider:
from .requesters import litellmchat
# Use unified LiteLLMRequester with provider prefix
# Map litellm_provider (YAML spec) to custom_llm_provider (config)
config['custom_llm_provider'] = litellm_provider
requester_inst = litellmchat.LiteLLMRequester(
ap=self.ap,
config=config,
)
requester_inst = litellmchat.LiteLLMRequester(ap=self.ap, config=config)
self.ap.logger.debug(
f'Using LiteLLMRequester for {provider_entity.requester} '
f'with custom_llm_provider={config["custom_llm_provider"]}'
)
else:
# Use original requester class (for backward compatibility)
if provider_entity.requester not in self.requester_dict:
raise provider_errors.RequesterNotFoundError(provider_entity.requester)
requester_inst = self.requester_dict[provider_entity.requester](
ap=self.ap,
config=config,
)
requester_inst = self.requester_dict[provider_entity.requester](ap=self.ap, config=config)
await requester_inst.initialize()
token_mgr = token.TokenManager(name=provider_entity.uuid, tokens=provider_entity.api_keys or [])
provider = requester.RuntimeProvider(
return requester.RuntimeProvider(
execution_context=context,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
)
async def load_provider(
self,
context: TenantContext,
provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict,
) -> requester.RuntimeProvider:
execution_context = await self.resolve_execution_context(context)
return await self._build_provider(execution_context, provider_info)
async def cache_provider(self, context: TenantContext, provider: requester.RuntimeProvider) -> None:
execution_context = await self.resolve_execution_context(context)
self._ensure_same_scope(execution_context, provider.execution_context, resource='Provider')
self._ensure_entity_workspace(provider.provider_entity, execution_context, resource='Provider')
self.provider_dict[self._cache_key(execution_context, provider.provider_entity.uuid)] = provider
async def get_provider_by_uuid(
self,
context: TenantContext,
provider_uuid: str,
) -> requester.RuntimeProvider:
execution_context = await self.resolve_execution_context(context)
provider = self.provider_dict.get(self._cache_key(execution_context, provider_uuid))
if provider is None:
raise ValueError(f'Model provider {provider_uuid} not found')
self._ensure_same_scope(execution_context, provider.execution_context, resource='Provider')
return provider
async def remove_provider(self, provider_uuid: str):
"""Remove provider
async def remove_provider(self, context: TenantContext, provider_uuid: str) -> None:
execution_context = await self.resolve_execution_context(context)
self.provider_dict.pop(self._cache_key(execution_context, provider_uuid), None)
This method will not consider the models using this provider,
because the models should be removed by the caller.
"""
del self.provider_dict[provider_uuid]
async def reload_provider(self, provider_uuid: str):
"""Reload provider"""
provider_entity = await self.ap.persistence_mgr.execute_async(
async def reload_provider(self, context: TenantContext, provider_uuid: str) -> None:
execution_context = await self.resolve_execution_context(context)
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_model.ModelProvider).where(
persistence_model.ModelProvider.uuid == provider_uuid
persistence_model.ModelProvider.workspace_uuid == execution_context.workspace_uuid,
persistence_model.ModelProvider.uuid == provider_uuid,
)
)
provider_entity = provider_entity.first()
provider_entity = result.first()
if provider_entity is None:
raise provider_errors.ProviderNotFoundError(provider_uuid)
new_runtime_provider = await self.load_provider(provider_entity)
new_provider = await self._build_provider(execution_context, provider_entity)
cache_prefix = self._cache_key(execution_context, '')[:3]
for cache in (self.llm_model_dict, self.embedding_model_dict, self.rerank_model_dict):
for key, model in cache.items():
if key[:3] == cache_prefix and model.provider.provider_entity.uuid == provider_uuid:
model.provider = new_provider
self.provider_dict[self._cache_key(execution_context, provider_uuid)] = new_provider
# update refs in runtime models
for model in self.llm_models:
if model.provider.provider_entity.uuid == provider_uuid:
model.provider = new_runtime_provider
for model in self.embedding_models:
if model.provider.provider_entity.uuid == provider_uuid:
model.provider = new_runtime_provider
for model in self.rerank_models:
if model.provider.provider_entity.uuid == provider_uuid:
model.provider = new_runtime_provider
@staticmethod
def _coerce_model(model_info: _ModelEntity | sqlalchemy.Row, entity_type: type[_ModelEntity]) -> _ModelEntity:
if isinstance(model_info, sqlalchemy.Row):
return entity_type(**model_info._mapping)
return model_info
# update ref in provider dict
self.provider_dict[provider_uuid] = new_runtime_provider
async def load_llm_model_with_provider(
def _validate_model_provider(
self,
context: ExecutionContext,
model_entity: _ModelEntity,
provider: requester.RuntimeProvider,
) -> None:
self._ensure_entity_workspace(model_entity, context, resource='Model')
self._ensure_same_scope(context, provider.execution_context, resource='Provider')
if model_entity.provider_uuid != provider.provider_entity.uuid:
raise WorkspaceInvariantError('Model references a different provider')
def _build_llm_model(
self,
context: ExecutionContext,
model_info: persistence_model.LLMModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeLLMModel:
"""Load LLM model with provider info"""
if isinstance(model_info, sqlalchemy.Row):
model_info = persistence_model.LLMModel(**model_info._mapping)
runtime_llm_model = requester.RuntimeLLMModel(
model_entity=model_info,
model_entity = self._coerce_model(model_info, persistence_model.LLMModel)
self._validate_model_provider(context, model_entity, provider)
return requester.RuntimeLLMModel(
execution_context=context,
model_entity=model_entity,
provider=provider,
)
return runtime_llm_model
async def load_embedding_model_with_provider(
def _build_embedding_model(
self,
context: ExecutionContext,
model_info: persistence_model.EmbeddingModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeEmbeddingModel:
"""Load embedding model with provider info"""
if isinstance(model_info, sqlalchemy.Row):
model_info = persistence_model.EmbeddingModel(**model_info._mapping)
runtime_embedding_model = requester.RuntimeEmbeddingModel(
model_entity=model_info,
model_entity = self._coerce_model(model_info, persistence_model.EmbeddingModel)
self._validate_model_provider(context, model_entity, provider)
return requester.RuntimeEmbeddingModel(
execution_context=context,
model_entity=model_entity,
provider=provider,
)
return runtime_embedding_model
async def load_rerank_model_with_provider(
def _build_rerank_model(
self,
context: ExecutionContext,
model_info: persistence_model.RerankModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeRerankModel:
"""Load rerank model with provider info"""
if isinstance(model_info, sqlalchemy.Row):
model_info = persistence_model.RerankModel(**model_info._mapping)
runtime_rerank_model = requester.RuntimeRerankModel(
model_entity=model_info,
model_entity = self._coerce_model(model_info, persistence_model.RerankModel)
self._validate_model_provider(context, model_entity, provider)
return requester.RuntimeRerankModel(
execution_context=context,
model_entity=model_entity,
provider=provider,
)
return runtime_rerank_model
async def load_llm_model_with_provider(
self,
context: TenantContext,
model_info: persistence_model.LLMModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeLLMModel:
execution_context = await self.resolve_execution_context(context)
return self._build_llm_model(execution_context, model_info, provider)
async def load_llm_model(self, model_info: dict):
"""Load LLM model from dict (with provider info)"""
provider_info = model_info.get('provider', {})
if not provider_info:
raise ValueError('Provider info is required')
async def load_embedding_model_with_provider(
self,
context: TenantContext,
model_info: persistence_model.EmbeddingModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeEmbeddingModel:
execution_context = await self.resolve_execution_context(context)
return self._build_embedding_model(execution_context, model_info, provider)
model_entity = persistence_model.LLMModel(
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid=model_info.get('provider_uuid', ''),
abilities=model_info.get('abilities', []),
context_length=model_info.get('context_length'),
extra_args=model_info.get('extra_args', {}),
)
async def load_rerank_model_with_provider(
self,
context: TenantContext,
model_info: persistence_model.RerankModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeRerankModel:
execution_context = await self.resolve_execution_context(context)
return self._build_rerank_model(execution_context, model_info, provider)
provider_entity = persistence_model.ModelProvider(
uuid=provider_info.get('uuid', ''),
name=provider_info.get('name', ''),
requester=provider_info.get('requester', ''),
base_url=provider_info.get('base_url', ''),
api_keys=provider_info.get('api_keys', []),
)
async def cache_llm_model(self, context: TenantContext, model: requester.RuntimeLLMModel) -> None:
execution_context = await self.resolve_execution_context(context)
self._ensure_same_scope(execution_context, model.execution_context, resource='LLM model')
self.llm_model_dict[self._cache_key(execution_context, model.model_entity.uuid)] = model
await self.load_llm_model_with_provider(model_entity, provider_entity)
async def cache_embedding_model(
self,
context: TenantContext,
model: requester.RuntimeEmbeddingModel,
) -> None:
execution_context = await self.resolve_execution_context(context)
self._ensure_same_scope(execution_context, model.execution_context, resource='Embedding model')
self.embedding_model_dict[self._cache_key(execution_context, model.model_entity.uuid)] = model
async def load_embedding_model(self, model_info: dict):
"""Load embedding model from dict (with provider info)"""
provider_info = model_info.get('provider', {})
if not provider_info:
raise ValueError('Provider info is required')
async def cache_rerank_model(self, context: TenantContext, model: requester.RuntimeRerankModel) -> None:
execution_context = await self.resolve_execution_context(context)
self._ensure_same_scope(execution_context, model.execution_context, resource='Rerank model')
self.rerank_model_dict[self._cache_key(execution_context, model.model_entity.uuid)] = model
model_entity = persistence_model.EmbeddingModel(
uuid=model_info.get('uuid', ''),
name=model_info.get('name', ''),
provider_uuid=model_info.get('provider_uuid', ''),
extra_args=model_info.get('extra_args', {}),
)
async def get_model_by_uuid(self, context: TenantContext, model_uuid: str) -> requester.RuntimeLLMModel:
execution_context = await self.resolve_execution_context(context)
model = self.llm_model_dict.get(self._cache_key(execution_context, model_uuid))
if model is None:
raise ValueError(f'LLM model {model_uuid} not found')
self._ensure_same_scope(execution_context, model.execution_context, resource='LLM model')
return model
provider_entity = persistence_model.ModelProvider(
uuid=provider_info.get('uuid', ''),
name=provider_info.get('name', ''),
requester=provider_info.get('requester', ''),
base_url=provider_info.get('base_url', ''),
api_keys=provider_info.get('api_keys', []),
)
async def get_embedding_model_by_uuid(
self,
context: TenantContext,
model_uuid: str,
) -> requester.RuntimeEmbeddingModel:
execution_context = await self.resolve_execution_context(context)
model = self.embedding_model_dict.get(self._cache_key(execution_context, model_uuid))
if model is None:
raise ValueError(f'Embedding model {model_uuid} not found')
self._ensure_same_scope(execution_context, model.execution_context, resource='Embedding model')
return model
await self.load_embedding_model_with_provider(model_entity, provider_entity)
async def get_rerank_model_by_uuid(
self,
context: TenantContext,
model_uuid: str,
) -> requester.RuntimeRerankModel:
execution_context = await self.resolve_execution_context(context)
model = self.rerank_model_dict.get(self._cache_key(execution_context, model_uuid))
if model is None:
raise ValueError(f'Rerank model {model_uuid} not found')
self._ensure_same_scope(execution_context, model.execution_context, resource='Rerank model')
return model
async def get_model_by_uuid(self, uuid: str) -> requester.RuntimeLLMModel:
"""Get LLM model by uuid"""
for model in self.llm_models:
if model.model_entity.uuid == uuid:
return model
raise ValueError(f'LLM model {uuid} not found')
async def remove_llm_model(self, context: TenantContext, model_uuid: str) -> None:
execution_context = await self.resolve_execution_context(context)
self.llm_model_dict.pop(self._cache_key(execution_context, model_uuid), None)
async def get_embedding_model_by_uuid(self, uuid: str) -> requester.RuntimeEmbeddingModel:
"""Get embedding model by uuid"""
for model in self.embedding_models:
if model.model_entity.uuid == uuid:
return model
raise ValueError(f'Embedding model {uuid} not found')
async def remove_embedding_model(self, context: TenantContext, model_uuid: str) -> None:
execution_context = await self.resolve_execution_context(context)
self.embedding_model_dict.pop(self._cache_key(execution_context, model_uuid), None)
async def get_rerank_model_by_uuid(self, uuid: str) -> requester.RuntimeRerankModel:
"""Get rerank model by uuid"""
for model in self.rerank_models:
if model.model_entity.uuid == uuid:
return model
raise ValueError(f'Rerank model {uuid} not found')
async def remove_llm_model(self, model_uuid: str):
"""Remove LLM model"""
for model in self.llm_models:
if model.model_entity.uuid == model_uuid:
self.llm_models.remove(model)
return
async def remove_embedding_model(self, model_uuid: str):
"""Remove embedding model"""
for model in self.embedding_models:
if model.model_entity.uuid == model_uuid:
self.embedding_models.remove(model)
return
async def remove_rerank_model(self, model_uuid: str):
"""Remove rerank model"""
for model in self.rerank_models:
if model.model_entity.uuid == model_uuid:
self.rerank_models.remove(model)
return
async def remove_rerank_model(self, context: TenantContext, model_uuid: str) -> None:
execution_context = await self.resolve_execution_context(context)
self.rerank_model_dict.pop(self._cache_key(execution_context, model_uuid), None)
def get_available_requesters_info(self, model_type: str) -> list[dict]:
"""Get all available requesters"""
if model_type != '':
if model_type:
return [
component.to_plain_dict()
for component in self.requester_components
if model_type in component.spec['support_type']
]
else:
return [component.to_plain_dict() for component in self.requester_components]
return [component.to_plain_dict() for component in self.requester_components]
def get_available_requester_info_by_name(self, name: str) -> dict | None:
"""Get requester info by name"""
for component in self.requester_components:
if component.metadata.name == name:
return component.to_plain_dict()
return None
def get_available_requester_manifest_by_name(self, name: str) -> engine.Component | None:
"""Get requester manifest by name"""
for component in self.requester_components:
if component.metadata.name == name:
return component
+86 -2
View File
@@ -5,7 +5,9 @@ import typing
import time
from ...core import app
from ...api.http.context import ExecutionContext
from ...entity.persistence import model as persistence_model
from ...workspace.errors import WorkspaceInvariantError
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
from . import token
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
@@ -16,6 +18,20 @@ LLM_USAGE_QUERY_VARIABLE = '_llm_usage'
STREAM_USAGE_QUERY_VARIABLE = '_stream_usage'
def _ensure_same_execution_scope(
expected: ExecutionContext,
actual: ExecutionContext,
*,
resource: str,
) -> None:
if (
actual.instance_uuid != expected.instance_uuid
or actual.workspace_uuid != expected.workspace_uuid
or actual.placement_generation != expected.placement_generation
):
raise WorkspaceInvariantError(f'{resource} belongs to another Workspace execution scope')
def _store_llm_usage(query: pipeline_query.Query | None, usage_info: dict | None) -> None:
"""Store the latest provider usage on the query for upstream action handlers."""
if query is None or not usage_info:
@@ -39,24 +55,61 @@ class RuntimeProvider:
def __init__(
self,
execution_context: ExecutionContext,
provider_entity: persistence_model.ModelProvider,
token_mgr: token.TokenManager,
requester: ProviderAPIRequester,
):
if provider_entity.workspace_uuid != execution_context.workspace_uuid:
raise WorkspaceInvariantError('Provider belongs to another Workspace')
self.execution_context = execution_context
self.provider_entity = provider_entity
self.token_mgr = token_mgr
self.requester = requester
def _validate_invocation(
self,
model: RuntimeLLMModel | RuntimeEmbeddingModel | RuntimeRerankModel,
execution_context: ExecutionContext,
) -> None:
_ensure_same_execution_scope(self.execution_context, execution_context, resource='Provider invocation')
_ensure_same_execution_scope(self.execution_context, model.execution_context, resource='Runtime model')
if model.provider is not self:
raise WorkspaceInvariantError('Runtime model is attached to another provider')
def _resolve_llm_execution_context(
self,
query: pipeline_query.Query | None,
execution_context: ExecutionContext | None,
) -> ExecutionContext:
if query is not None:
from ...pipeline.pool import get_query_execution_context
query_context = get_query_execution_context(query)
if execution_context is not None:
_ensure_same_execution_scope(
query_context,
execution_context,
resource='Explicit LLM invocation context',
)
return query_context
if execution_context is None:
raise WorkspaceInvariantError('LLM invocation requires an ExecutionContext when query is absent')
return execution_context
async def invoke_llm(
self,
query: pipeline_query.Query,
query: pipeline_query.Query | None,
model: RuntimeLLMModel,
messages: typing.List[provider_message.Message],
funcs: typing.List[resource_tool.LLMTool] = None,
extra_args: dict[str, typing.Any] = {},
remove_think: bool = False,
execution_context: ExecutionContext | None = None,
) -> provider_message.Message:
"""Bridge method for invoking LLM with monitoring"""
invocation_context = self._resolve_llm_execution_context(query, execution_context)
self._validate_invocation(model, invocation_context)
# Start timing for monitoring
start_time = time.time()
input_tokens = 0
@@ -130,14 +183,17 @@ class RuntimeProvider:
async def invoke_llm_stream(
self,
query: pipeline_query.Query,
query: pipeline_query.Query | None,
model: RuntimeLLMModel,
messages: typing.List[provider_message.Message],
funcs: typing.List[resource_tool.LLMTool] = None,
extra_args: dict[str, typing.Any] = {},
remove_think: bool = False,
execution_context: ExecutionContext | None = None,
) -> provider_message.MessageChunk:
"""Bridge method for invoking LLM stream with monitoring"""
invocation_context = self._resolve_llm_execution_context(query, execution_context)
self._validate_invocation(model, invocation_context)
# Start timing for monitoring
start_time = time.time()
status = 'success'
@@ -212,6 +268,8 @@ class RuntimeProvider:
model: RuntimeEmbeddingModel,
input_text: typing.List[str],
extra_args: dict[str, typing.Any] = {},
*,
execution_context: ExecutionContext,
knowledge_base_id: str | None = None,
query_text: str | None = None,
session_id: str | None = None,
@@ -219,6 +277,7 @@ class RuntimeProvider:
call_type: str | None = None,
) -> typing.List[typing.List[float]]:
"""Bridge method for invoking embedding with monitoring"""
self._validate_invocation(model, execution_context)
# Start timing for monitoring
start_time = time.time()
prompt_tokens = 0
@@ -254,6 +313,7 @@ class RuntimeProvider:
try:
await self.requester.ap.monitoring_service.record_embedding_call(
execution_context,
model_name=model.model_entity.name,
prompt_tokens=prompt_tokens,
total_tokens=total_tokens,
@@ -276,8 +336,11 @@ class RuntimeProvider:
query: str,
documents: typing.List[str],
extra_args: dict[str, typing.Any] = {},
*,
execution_context: ExecutionContext,
) -> typing.List[dict]:
"""Bridge method for invoking rerank with monitoring"""
self._validate_invocation(model, execution_context)
start_time = time.time()
status = 'success'
@@ -316,9 +379,16 @@ class RuntimeLLMModel:
def __init__(
self,
execution_context: ExecutionContext,
model_entity: persistence_model.LLMModel,
provider: RuntimeProvider,
):
_ensure_same_execution_scope(provider.execution_context, execution_context, resource='LLM model')
if model_entity.workspace_uuid != execution_context.workspace_uuid:
raise WorkspaceInvariantError('LLM model belongs to another Workspace')
if model_entity.provider_uuid != provider.provider_entity.uuid:
raise WorkspaceInvariantError('LLM model references another provider')
self.execution_context = execution_context
self.model_entity = model_entity
self.provider = provider
@@ -334,9 +404,16 @@ class RuntimeEmbeddingModel:
def __init__(
self,
execution_context: ExecutionContext,
model_entity: persistence_model.EmbeddingModel,
provider: RuntimeProvider,
):
_ensure_same_execution_scope(provider.execution_context, execution_context, resource='Embedding model')
if model_entity.workspace_uuid != execution_context.workspace_uuid:
raise WorkspaceInvariantError('Embedding model belongs to another Workspace')
if model_entity.provider_uuid != provider.provider_entity.uuid:
raise WorkspaceInvariantError('Embedding model references another provider')
self.execution_context = execution_context
self.model_entity = model_entity
self.provider = provider
@@ -352,9 +429,16 @@ class RuntimeRerankModel:
def __init__(
self,
execution_context: ExecutionContext,
model_entity: persistence_model.RerankModel,
provider: RuntimeProvider,
):
_ensure_same_execution_scope(provider.execution_context, execution_context, resource='Rerank model')
if model_entity.workspace_uuid != execution_context.workspace_uuid:
raise WorkspaceInvariantError('Rerank model belongs to another Workspace')
if model_entity.provider_uuid != provider.provider_entity.uuid:
raise WorkspaceInvariantError('Rerank model references another provider')
self.execution_context = execution_context
self.model_entity = model_entity
self.provider = provider
+12 -7
View File
@@ -22,10 +22,12 @@ from langbot.libs.dify_service_api.v1 import client, errors
import httpx
# Module-level store for paused-workflow form state. The key isolates the bot,
# pipeline, adapter, and launcher; each value holds an insertion-ordered map of
# form_token -> form_data so one conversation can pause multiple workflows.
PendingFormKey = tuple[str, str, str, str, str]
# Module-level store for paused-workflow form state. The key includes the full
# execution scope before the bot, pipeline, adapter, and launcher dimensions;
# each value holds an insertion-ordered map of form_token -> form_data so one
# conversation can pause multiple workflows without crossing Workspaces or
# placement generations.
PendingFormKey = tuple[str, str, int, str, str, str, str, str]
_PENDING_FORMS: dict[PendingFormKey, 'OrderedDict[str, dict[str, typing.Any]]'] = {}
_PENDING_FORM_DEFAULT_TTL = 30 * 60 # 30 minutes safety cap
_STREAM_FORM_PLACEHOLDER = '\u200b'
@@ -48,10 +50,13 @@ def _dify_user_from_query(query: pipeline_query.Query) -> str:
def _session_key_from_query(query: pipeline_query.Query) -> PendingFormKey:
"""Build a process-local pending-form key isolated by bot and pipeline."""
"""Build a process-local pending-form key isolated by execution scope."""
adapter = getattr(query, 'adapter', None)
adapter_type = f'{type(adapter).__module__}.{type(adapter).__qualname__}'
return (
str(getattr(query, 'instance_uuid', '') or ''),
str(getattr(query, 'workspace_uuid', '') or ''),
int(getattr(query, 'placement_generation', 0) or 0),
str(getattr(query, 'bot_uuid', '') or ''),
str(getattr(query, 'pipeline_uuid', '') or ''),
adapter_type,
@@ -74,8 +79,8 @@ def _prune_pending_forms(now: float | None = None) -> None:
def _set_pending_form(session_key: PendingFormKey, form_data: dict[str, typing.Any]) -> None:
_prune_pending_forms()
if isinstance(session_key, tuple) and len(session_key) > 1:
form_data['pipeline_uuid'] = session_key[1]
if isinstance(session_key, tuple) and len(session_key) == 8:
form_data['pipeline_uuid'] = session_key[4]
stored = dict(form_data)
expiration_time = stored.get('expiration_time')
try:
+17 -4
View File
@@ -11,6 +11,7 @@ import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.rag.context as rag_context
from ...pipeline.pool import get_query_execution_context
rag_combined_prompt_template = """
The following are relevant context entries retrieved from the knowledge base.
@@ -227,7 +228,10 @@ class LocalAgentRunner(runner.RequestRunner):
# Primary model
if query.use_llm_model_uuid:
try:
primary = await self.ap.model_mgr.get_model_by_uuid(query.use_llm_model_uuid)
primary = await self.ap.model_mgr.get_model_by_uuid(
get_query_execution_context(query),
query.use_llm_model_uuid,
)
candidates.append(primary)
except ValueError:
self.ap.logger.warning(f'Primary model {query.use_llm_model_uuid} not found')
@@ -236,7 +240,10 @@ class LocalAgentRunner(runner.RequestRunner):
fallback_uuids = (query.variables or {}).get('_fallback_model_uuids', [])
for fb_uuid in fallback_uuids:
try:
fb_model = await self.ap.model_mgr.get_model_by_uuid(fb_uuid)
fb_model = await self.ap.model_mgr.get_model_by_uuid(
get_query_execution_context(query),
fb_uuid,
)
candidates.append(fb_model)
except ValueError:
self.ap.logger.warning(f'Fallback model {fb_uuid} not found, skipping')
@@ -346,12 +353,13 @@ class LocalAgentRunner(runner.RequestRunner):
if kb_uuids and user_message_text:
# only support text for now
all_results: list[rag_context.RetrievalResultEntry] = []
execution_context = get_query_execution_context(query)
kb_engine_plugins: set[str] = set()
# Retrieve from each knowledge base
for kb_uuid in kb_uuids:
kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(execution_context, kb_uuid)
if not kb:
self.ap.logger.warning(f'Knowledge base {kb_uuid} not found, skipping')
@@ -364,6 +372,7 @@ class LocalAgentRunner(runner.RequestRunner):
kb_engine_plugins.add(engine_plugin_id)
result = await kb.retrieve(
execution_context,
user_message_text,
settings={
'bot_uuid': query.bot_uuid or '',
@@ -398,7 +407,10 @@ class LocalAgentRunner(runner.RequestRunner):
)
if all_results and rerank_model_uuid:
try:
rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(rerank_model_uuid)
rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(
execution_context,
rerank_model_uuid,
)
rerank_top_k = int(local_agent_config.get('rerank-top-k', 5))
doc_texts = []
@@ -411,6 +423,7 @@ class LocalAgentRunner(runner.RequestRunner):
model=rerank_model,
query=user_message_text,
documents=doc_texts_capped,
execution_context=execution_context,
)
scored = sorted(scores, key=lambda x: x.get('relevance_score', 0), reverse=True)
+76 -3
View File
@@ -1,12 +1,48 @@
from __future__ import annotations
import asyncio
import dataclasses
from ...core import app
from langbot_plugin.api.entities.builtin.provider import message as provider_message, prompt as provider_prompt
import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
from ...api.http.context import ExecutionContext
from ...core import app
from ...pipeline.pool import (
ExecutionContextMismatchError,
ExecutionContextRequiredError,
bind_execution_context,
get_query_execution_context,
)
SessionKey = tuple[
str,
str,
int,
str,
str,
int | str,
]
def _query_session_key(query: pipeline_query.Query) -> tuple[SessionKey, ExecutionContext]:
execution_context = get_query_execution_context(query)
bot_uuid = getattr(query, 'bot_uuid', None)
if not isinstance(bot_uuid, str) or not bot_uuid.strip():
raise ExecutionContextRequiredError('Query.bot_uuid is required for session lookup')
execution_context = bind_execution_context(execution_context, bot_uuid=bot_uuid)
key: SessionKey = (
execution_context.instance_uuid,
execution_context.workspace_uuid,
execution_context.placement_generation,
bot_uuid,
query.launcher_type.value,
query.launcher_id,
)
return key, execution_context
class SessionManager:
"""会话管理器"""
@@ -24,17 +60,39 @@ class SessionManager:
async def get_session(self, query: pipeline_query.Query) -> provider_session.Session:
"""获取会话"""
session_key, execution_context = _query_session_key(query)
for session in self.session_list:
if query.launcher_type == session.launcher_type and query.launcher_id == session.launcher_id:
if getattr(session, '_langbot_session_key', None) == session_key:
return session
session_concurrency = self.ap.instance_config.data['concurrency']['session']
session = provider_session.Session(
instance_uuid=execution_context.instance_uuid,
workspace_uuid=execution_context.workspace_uuid,
placement_generation=execution_context.placement_generation,
bot_uuid=query.bot_uuid,
launcher_type=query.launcher_type,
launcher_id=query.launcher_id,
sender_id=query.sender_id,
)
session_context = dataclasses.replace(
execution_context,
pipeline_uuid=None,
query_uuid=None,
)
# langbot-plugin 0.4.13 ignores Workspace fields. Preserve them until
# the Workspace-aware SDK becomes the minimum supported version.
object.__setattr__(session, 'instance_uuid', session_context.instance_uuid)
object.__setattr__(session, 'workspace_uuid', session_context.workspace_uuid)
object.__setattr__(
session,
'placement_generation',
session_context.placement_generation,
)
object.__setattr__(session, 'bot_uuid', query.bot_uuid)
object.__setattr__(session, '_execution_context', session_context)
object.__setattr__(session, '_langbot_session_key', session_key)
session._semaphore = asyncio.Semaphore(session_concurrency)
self.session_list.append(session)
return session
@@ -49,6 +107,17 @@ class SessionManager:
) -> provider_session.Conversation:
"""获取对话或创建对话"""
session_key, execution_context = _query_session_key(query)
if getattr(session, '_langbot_session_key', None) != session_key:
raise ExecutionContextMismatchError('Session does not belong to the Query execution scope')
execution_context = bind_execution_context(
execution_context,
bot_uuid=bot_uuid,
pipeline_uuid=pipeline_uuid,
)
if execution_context.bot_uuid != getattr(session, 'bot_uuid', None):
raise ExecutionContextMismatchError('Session bot_uuid does not match the Query execution scope')
if not session.conversations:
session.conversations = []
@@ -63,7 +132,11 @@ class SessionManager:
messages=prompt_messages,
)
if session.using_conversation is None or session.using_conversation.pipeline_uuid != pipeline_uuid:
if (
session.using_conversation is None
or session.using_conversation.pipeline_uuid != pipeline_uuid
or session.using_conversation.bot_uuid != bot_uuid
):
conversation = provider_session.Conversation(
prompt=prompt,
messages=[],
@@ -11,7 +11,7 @@ async def is_box_backend_available(ap: Any) -> bool:
if not getattr(box_service, 'available', False):
return False
try:
status = await box_service.get_status()
status = await box_service.get_backend_status()
backend_info = status.get('backend', {})
return bool(backend_info.get('available', False))
except Exception:
+298 -98
View File
@@ -26,6 +26,9 @@ from pydantic import AnyUrl
from .. import loader
from ....core import app
from ....api.http.context import ExecutionContext
from ....api.http.service.tenant import TenantContext, require_workspace_uuid
from ....workspace.errors import WorkspaceError, WorkspaceInvariantError
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
import langbot_plugin.api.entities.builtin.provider.message as provider_message
from ....entity.persistence import mcp as persistence_mcp
@@ -223,6 +226,8 @@ class MCPToolCallTimeoutError(TimeoutError):
class RuntimeMCPSession:
"""运行时 MCP 会话"""
_FENCE_POLL_INTERVAL = 5.0
ap: app.Application
server_name: str
@@ -262,11 +267,19 @@ class RuntimeMCPSession:
_box_stdio_runtime: BoxStdioSessionRuntime
def __init__(self, server_name: str, server_config: dict, enable: bool, ap: app.Application):
def __init__(
self,
server_name: str,
server_config: dict,
enable: bool,
ap: app.Application,
execution_context: ExecutionContext,
):
self.server_name = server_name
self.server_uuid = server_config.get('uuid', '')
self.server_config = server_config
self.ap = ap
self.execution_context = execution_context
self.enable = enable
self.session = None
self.tool_call_timeout_sec = self._parse_tool_call_timeout(
@@ -312,24 +325,46 @@ class RuntimeMCPSession:
self._box_stdio_runtime = BoxStdioSessionRuntime(self)
self.box_config = self._box_stdio_runtime.config
def _parse_tool_call_timeout(self, value: typing.Any) -> float:
"""Return a safe tool-call timeout; zero explicitly disables it."""
try:
timeout = -1 if isinstance(value, bool) else float(value)
if timeout > 0:
# Validate the exact conversion used for each call here, so a
# finite-but-enormous manual config cannot fail at invocation.
timedelta(seconds=timeout)
except (TypeError, ValueError, OverflowError):
timeout = -1
async def _assert_execution_active(self) -> None:
"""Fail closed when this long-lived session belongs to a stale placement."""
if not math.isfinite(timeout) or timeout < 0:
self.ap.logger.warning(
f'Invalid MCP tool call timeout {value!r} for {self.server_name}; '
f'using {MCP_TOOL_CALL_TIMEOUT_DEFAULT_SECONDS:g} seconds'
)
return MCP_TOOL_CALL_TIMEOUT_DEFAULT_SECONDS
return timeout
binding = await self.ap.workspace_service.get_execution_binding(
self.execution_context.workspace_uuid,
expected_generation=self.execution_context.placement_generation,
)
if binding.instance_uuid != self.execution_context.instance_uuid:
raise WorkspaceInvariantError('MCP session instance does not match the active Workspace binding')
async def _monitor_execution_fence(self) -> None:
"""Poll the placement fence while an MCP transport is idle."""
while not self._shutdown_event.is_set():
await asyncio.sleep(self._FENCE_POLL_INTERVAL)
if self._shutdown_event.is_set():
return
await self._assert_execution_active()
async def _sleep_with_execution_fence(self, delay: float) -> None:
"""Back off without reconnecting after the captured placement expires."""
await self._assert_execution_active()
try:
await asyncio.wait_for(self._shutdown_event.wait(), timeout=delay)
except asyncio.TimeoutError:
pass
if not self._shutdown_event.is_set():
await self._assert_execution_active()
def _stop_for_stale_execution(self, error: WorkspaceError) -> None:
"""Mark the session terminal without retrying a fenced placement."""
self.status = MCPSessionStatus.ERROR
self.error_message = 'Workspace execution binding is stale'
self._shutdown_event.set()
self._ready_event.set()
self.ap.logger.info(
f'MCP session {self.server_name} stopped because its Workspace execution binding is stale: {error}'
)
async def _init_stdio_python_server(self):
if self._uses_box_stdio():
@@ -458,6 +493,7 @@ class RuntimeMCPSession:
async def _lifecycle_loop(self):
"""Manage the full MCP session lifecycle in a background task."""
try:
await self._assert_execution_active()
if self.server_config['mode'] == 'stdio':
await self._init_stdio_python_server()
elif self.server_config['mode'] == 'remote':
@@ -467,9 +503,11 @@ class RuntimeMCPSession:
elif self.server_config['mode'] == 'http':
await self._init_streamable_http_server()
else:
raise ValueError(f'Unknown MCP server mode: {self.server_name}: {self.server_config}')
raise ValueError(f'Unknown MCP server mode for {self.server_name}')
await self._assert_execution_active()
await self.refresh()
await self._assert_execution_active()
self.status = MCPSessionStatus.CONNECTED
@@ -481,12 +519,16 @@ class RuntimeMCPSession:
monitor_task = asyncio.create_task(self._box_stdio_runtime.monitor_process_health())
shutdown_task = asyncio.create_task(self._shutdown_event.wait())
reconnect_task = asyncio.create_task(self._reconnect_event.wait())
fence_task = asyncio.create_task(self._monitor_execution_fence())
done, pending = await asyncio.wait(
[shutdown_task, monitor_task, reconnect_task],
[shutdown_task, monitor_task, reconnect_task, fence_task],
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
if fence_task in done and not self._shutdown_event.is_set():
fence_task.result()
if reconnect_task in done and not self._shutdown_event.is_set():
self._reconnect_event.clear()
self.ap.logger.info(
@@ -522,12 +564,16 @@ class RuntimeMCPSession:
else:
shutdown_task = asyncio.create_task(self._shutdown_event.wait())
reconnect_task = asyncio.create_task(self._reconnect_event.wait())
fence_task = asyncio.create_task(self._monitor_execution_fence())
done, pending = await asyncio.wait(
[shutdown_task, reconnect_task],
[shutdown_task, reconnect_task, fence_task],
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
if fence_task in done and not self._shutdown_event.is_set():
fence_task.result()
if reconnect_task in done and not self._shutdown_event.is_set():
self._reconnect_event.clear()
self.ap.logger.info(
@@ -590,7 +636,11 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.CONNECTING
self.error_message = None
self.error_phase = None
await asyncio.sleep(1)
try:
await self._sleep_with_execution_fence(1)
except WorkspaceError as fence_error:
self._stop_for_stale_execution(fence_error)
return
continue
except _CallerReconnect:
# A tool/resource call hit a server-expired session and asked us
@@ -607,6 +657,7 @@ class RuntimeMCPSession:
self.error_message = None
self.error_phase = None
try:
await self._assert_execution_active()
if self.server_config['mode'] == 'stdio':
await self._init_stdio_python_server()
elif self.server_config['mode'] == 'remote':
@@ -616,8 +667,12 @@ class RuntimeMCPSession:
elif self.server_config['mode'] == 'http':
await self._init_streamable_http_server()
await self.refresh()
await self._assert_execution_active()
self.status = MCPSessionStatus.CONNECTED
self.ap.logger.info(f'MCP session {self.server_name} reconnected successfully after session expiry')
except WorkspaceError as reconnect_err:
self._stop_for_stale_execution(reconnect_err)
return
except Exception as reconnect_err:
self.status = MCPSessionStatus.ERROR
self.error_message = str(reconnect_err)
@@ -645,8 +700,15 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.CONNECTING
self.error_message = None
self.error_phase = None
await asyncio.sleep(2)
try:
await self._sleep_with_execution_fence(2)
except WorkspaceError as fence_error:
self._stop_for_stale_execution(fence_error)
return
continue
except WorkspaceError as e:
self._stop_for_stale_execution(e)
return
except Exception as e:
if self._shutdown_event.is_set():
return # Shutdown requested, don't retry
@@ -686,7 +748,11 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.CONNECTING
self.error_message = None
self.error_phase = None
await asyncio.sleep(delay)
try:
await self._sleep_with_execution_fence(delay)
except WorkspaceError as fence_error:
self._stop_for_stale_execution(fence_error)
return
attempt += 1
@staticmethod
@@ -769,6 +835,7 @@ class RuntimeMCPSession:
Returns True if reconnection succeeded within the timeout.
"""
await self._assert_execution_active()
if self._shutdown_event.is_set():
return False
@@ -779,6 +846,7 @@ class RuntimeMCPSession:
try:
await asyncio.wait_for(reconnected_event.wait(), timeout=self._RECONNECT_WAIT_TIMEOUT)
await self._assert_execution_active()
return self.status == MCPSessionStatus.CONNECTED
except asyncio.TimeoutError:
self.ap.logger.warning(f'MCP session {self.server_name} reconnect timed out')
@@ -794,6 +862,7 @@ class RuntimeMCPSession:
if not self.enable:
return
await self._assert_execution_active()
# Create background task for lifecycle management with retry
self._lifecycle_task = asyncio.create_task(self._lifecycle_loop_with_retry())
@@ -805,11 +874,13 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.ERROR
raise Exception(f'Connection timeout after {startup_timeout} seconds')
await self._assert_execution_active()
# Check for errors
if self.status == MCPSessionStatus.ERROR:
raise Exception('Connection failed, please check URL')
async def refresh(self):
await self._assert_execution_active()
if not self.session:
return
@@ -825,6 +896,7 @@ class RuntimeMCPSession:
self.resource_capabilities = {}
tools = await self.session.list_tools()
await self._assert_execution_active()
self.ap.logger.debug(f'Refresh MCP tools: {tools}')
@@ -846,34 +918,44 @@ class RuntimeMCPSession:
)
await self._refresh_resources()
await self._assert_execution_active()
async def _refresh_resources(self):
await self._assert_execution_active()
if not self.session:
return
try:
cursor: str | None = None
for _ in range(MCP_RESOURCE_DISCOVERY_MAX_PAGES):
await self._assert_execution_active()
resources_result = await self.session.list_resources(cursor)
await self._assert_execution_active()
for resource in resources_result.resources:
self.resources.append(_resource_to_dict(resource))
cursor = getattr(resources_result, 'nextCursor', None)
if not cursor:
break
self.ap.logger.debug(f'Refresh MCP resources: {len(self.resources)} resources found')
except WorkspaceError:
raise
except Exception as e:
self.ap.logger.debug(f'MCP server {self.server_name} does not support resources or failed to list: {e}')
try:
cursor = None
for _ in range(MCP_RESOURCE_DISCOVERY_MAX_PAGES):
await self._assert_execution_active()
templates_result = await self.session.list_resource_templates(cursor)
await self._assert_execution_active()
for template in templates_result.resourceTemplates:
self.resource_templates.append(_resource_template_to_dict(template))
cursor = getattr(templates_result, 'nextCursor', None)
if not cursor:
break
self.ap.logger.debug(f'Refresh MCP resource templates: {len(self.resource_templates)} templates found')
except WorkspaceError:
raise
except Exception as e:
self.ap.logger.debug(
f'MCP server {self.server_name} does not support resource templates or failed to list: {e}'
@@ -992,17 +1074,15 @@ class RuntimeMCPSession:
arguments: dict,
query: pipeline_query.Query | None = None,
) -> list[provider_message.ContentElement]:
await self._assert_execution_active()
for attempt in range(2):
if not self.session:
raise Exception('MCP session is not connected')
try:
read_timeout = timedelta(seconds=self.tool_call_timeout_sec) if self.tool_call_timeout_sec > 0 else None
result = await self.session.call_tool(
tool_name,
arguments,
read_timeout_seconds=read_timeout,
)
await self._assert_execution_active()
result = await self.session.call_tool(tool_name, arguments)
await self._assert_execution_active()
except Exception as e:
if self._is_tool_call_timeout(e):
self.ap.logger.warning(
@@ -1087,6 +1167,7 @@ class RuntimeMCPSession:
query: pipeline_query.Query | None = None,
) -> dict:
"""Read a resource by URI with safety limits and audit metadata."""
await self._assert_execution_active()
if not self.session:
raise Exception('MCP session is not connected')
@@ -1113,7 +1194,9 @@ class RuntimeMCPSession:
if not self.session:
raise Exception('MCP session is not connected')
try:
await self._assert_execution_active()
result = await self.session.read_resource(AnyUrl(uri))
await self._assert_execution_active()
break
except Exception as e:
if attempt == 0 and self._is_session_terminated(e):
@@ -1194,6 +1277,7 @@ class RuntimeMCPSession:
'cache_hit': False,
'warnings': warnings,
}
await self._assert_execution_active()
self._resource_cache[cache_key] = {'cached_at': now, 'envelope': envelope}
self._record_resource_read_trace(query, envelope)
return envelope
@@ -1228,7 +1312,11 @@ class RuntimeMCPSession:
def get_runtime_info_dict(self) -> dict:
info = {
'status': self.status.value,
'error_message': self.error_message,
# Raw transport exceptions may echo command arguments, headers, or
# environment values. Detailed diagnostics belong in AUDIT_VIEW
# logs; resource-list responses expose only a stable status.
'error_message': 'MCP runtime failed' if self.error_message else None,
'error_code': 'runtime_error' if self.error_message else None,
'error_phase': self.error_phase.value if self.error_phase else None,
'retry_count': self.retry_count,
'tool_count': len(self.get_tools()),
@@ -1336,6 +1424,37 @@ class RuntimeMCPSession:
await self._box_stdio_runtime.cleanup_session()
def _execution_context_from_tenant(context: TenantContext) -> ExecutionContext:
workspace_uuid = require_workspace_uuid(context)
instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
generation = getattr(context, 'placement_generation', None)
if not instance_uuid:
raise ValueError('MCP runtime requires an explicit instance UUID')
if isinstance(generation, bool) or not isinstance(generation, int) or generation <= 0:
raise ValueError('MCP runtime requires a positive placement generation')
return ExecutionContext(
instance_uuid=instance_uuid,
workspace_uuid=workspace_uuid,
placement_generation=generation,
bot_uuid=getattr(context, 'bot_uuid', None),
pipeline_uuid=getattr(context, 'pipeline_uuid', None),
query_uuid=getattr(context, 'query_uuid', None),
)
def _execution_context_from_query(query: pipeline_query.Query) -> ExecutionContext:
return _execution_context_from_tenant(
ExecutionContext(
instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
placement_generation=getattr(query, 'placement_generation', 0) or 0,
bot_uuid=getattr(query, 'bot_uuid', None),
pipeline_uuid=getattr(query, 'pipeline_uuid', None),
query_uuid=getattr(query, 'query_uuid', None),
)
)
# @loader.loader_class('mcp')
class MCPLoader(loader.ToolLoader):
"""MCP 工具加载器。
@@ -1343,7 +1462,7 @@ class MCPLoader(loader.ToolLoader):
在此加载器中管理所有与 MCP Server 的连接。
"""
sessions: dict[str, RuntimeMCPSession]
sessions: dict[tuple[str, str, int, str], RuntimeMCPSession]
_last_listed_functions: list[resource_tool.LLMTool]
@@ -1355,6 +1474,21 @@ class MCPLoader(loader.ToolLoader):
self._last_listed_functions = []
self._hosted_mcp_tasks = []
async def _assert_execution_active(
self,
context: TenantContext,
) -> ExecutionContext:
"""Validate a caller's placement before accessing an MCP session."""
execution_context = _execution_context_from_tenant(context)
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('MCP caller instance does not match the active Workspace binding')
return execution_context
async def initialize(self):
await self.load_mcp_servers_from_db()
@@ -1368,15 +1502,51 @@ class MCPLoader(loader.ToolLoader):
for server in servers:
config = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server)
try:
binding = await self.ap.workspace_service.get_execution_binding(server.workspace_uuid)
execution_context = ExecutionContext(
instance_uuid=binding.instance_uuid,
workspace_uuid=binding.workspace_uuid,
placement_generation=binding.placement_generation,
)
except Exception as exc:
self.ap.logger.warning(
f'Skipping MCP server {server.uuid}: Workspace execution binding is unavailable: {exc}'
)
continue
task = asyncio.create_task(self.host_mcp_server(config))
task = asyncio.create_task(self.host_mcp_server(execution_context, config))
self._hosted_mcp_tasks.append(task)
async def host_mcp_server(self, server_config: dict):
@staticmethod
def _scope_key(context: TenantContext) -> tuple[str, str, int]:
execution_context = _execution_context_from_tenant(context)
return (
execution_context.instance_uuid,
execution_context.workspace_uuid,
execution_context.placement_generation,
)
@classmethod
def _session_key(cls, context: TenantContext, server_name: str) -> tuple[str, str, int, str]:
return (*cls._scope_key(context), server_name)
def _sessions_for_context(self, context: TenantContext) -> list[RuntimeMCPSession]:
scope_key = self._scope_key(context)
return [session for key, session in self.sessions.items() if key[:3] == scope_key]
async def host_mcp_server(self, context: TenantContext, server_config: dict):
execution_context = await self._assert_execution_active(context)
configured_workspace = str(server_config.get('workspace_uuid') or '').strip()
if configured_workspace and configured_workspace != execution_context.workspace_uuid:
raise ValueError('MCP server configuration belongs to another Workspace')
server_config = dict(server_config)
server_config['workspace_uuid'] = execution_context.workspace_uuid
self.ap.logger.debug(f'Loading MCP server {server_config}')
try:
session = await self.load_mcp_server(server_config)
self.sessions[server_config['name']] = session
session = await self.load_mcp_server(execution_context, server_config)
await self._assert_execution_active(execution_context)
self.sessions[self._session_key(execution_context, server_config['name'])] = session
except Exception as e:
self.ap.logger.error(
f'Failed to load MCP server from db: {server_config["name"]}({server_config["uuid"]}): {e}\n{traceback.format_exc()}'
@@ -1385,6 +1555,7 @@ class MCPLoader(loader.ToolLoader):
self.ap.logger.debug(f'Starting MCP server {server_config["name"]}({server_config["uuid"]})')
try:
await self._assert_execution_active(execution_context)
await session.start()
except Exception as e:
self.ap.logger.error(
@@ -1394,7 +1565,7 @@ class MCPLoader(loader.ToolLoader):
self.ap.logger.debug(f'Started MCP server {server_config["name"]}({server_config["uuid"]})')
async def load_mcp_server(self, server_config: dict) -> RuntimeMCPSession:
async def load_mcp_server(self, context: TenantContext, server_config: dict) -> RuntimeMCPSession:
"""加载 MCP 服务器到运行时
Args:
@@ -1404,6 +1575,13 @@ class MCPLoader(loader.ToolLoader):
- enable: 是否启用
- extra_args: 额外的配置参数 (可选)
"""
execution_context = await self._assert_execution_active(context)
server_config = dict(server_config)
configured_workspace = str(server_config.get('workspace_uuid') or '').strip()
if configured_workspace and configured_workspace != execution_context.workspace_uuid:
raise ValueError('MCP server configuration belongs to another Workspace')
server_config['workspace_uuid'] = execution_context.workspace_uuid
uuid_ = server_config.get('uuid')
is_transient = False
if not uuid_:
@@ -1429,7 +1607,7 @@ class MCPLoader(loader.ToolLoader):
**extra_args,
}
session = RuntimeMCPSession(name, mixed_config, enable, self.ap)
session = RuntimeMCPSession(name, mixed_config, enable, self.ap, execution_context)
return session
@@ -1438,9 +1616,13 @@ class MCPLoader(loader.ToolLoader):
v = getattr(query, 'variables', None) or {}
return v.get('_pipeline_bound_mcp_servers', None)
def _eligible_sessions_for_bound(self, bound_mcp_servers: list[str] | None) -> list[RuntimeMCPSession]:
def _eligible_sessions_for_bound(
self,
context: TenantContext,
bound_mcp_servers: list[str] | None,
) -> list[RuntimeMCPSession]:
out: list[RuntimeMCPSession] = []
for session in self.sessions.values():
for session in self._sessions_for_context(context):
if not session.enable:
continue
if session.status != MCPSessionStatus.CONNECTED:
@@ -1452,10 +1634,14 @@ class MCPLoader(loader.ToolLoader):
out.append(session)
return out
def _eligible_resource_sessions_for_bound(self, bound_mcp_servers: list[str] | None) -> list[RuntimeMCPSession]:
def _eligible_resource_sessions_for_bound(
self,
context: TenantContext,
bound_mcp_servers: list[str] | None,
) -> list[RuntimeMCPSession]:
return [
session
for session in self._eligible_sessions_for_bound(bound_mcp_servers)
for session in self._eligible_sessions_for_bound(context, bound_mcp_servers)
if session.has_resource_support()
]
@@ -1486,12 +1672,13 @@ class MCPLoader(loader.ToolLoader):
]
async def _invoke_mcp_list_resources(self, parameters: dict, query: pipeline_query.Query) -> typing.Any:
execution_context = _execution_context_from_query(query)
server_name = parameters.get('server_name') if parameters else None
if not server_name or not isinstance(server_name, str):
return [provider_message.ContentElement.from_text('Error: "server_name" (string) is required.')]
bound = self._get_bound_mcp_from_query(query)
allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(bound)}
allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(execution_context, bound)}
if server_name not in allowed:
return [
provider_message.ContentElement.from_text(
@@ -1501,7 +1688,7 @@ class MCPLoader(loader.ToolLoader):
)
]
session = self.get_session(server_name)
session = self.get_session(execution_context, server_name)
if session is None or session.status != MCPSessionStatus.CONNECTED:
return [provider_message.ContentElement.from_text(f'Error: MCP server not connected: {server_name!r}')]
@@ -1518,6 +1705,7 @@ class MCPLoader(loader.ToolLoader):
return [provider_message.ContentElement.from_text(json.dumps(body, ensure_ascii=False, indent=2))]
async def _invoke_mcp_read_resource(self, parameters: dict, query: pipeline_query.Query) -> typing.Any:
execution_context = _execution_context_from_query(query)
server_name = parameters.get('server_name') if parameters else None
uri = parameters.get('uri') if parameters else None
if not server_name or not isinstance(server_name, str):
@@ -1526,7 +1714,7 @@ class MCPLoader(loader.ToolLoader):
return [provider_message.ContentElement.from_text('Error: "uri" (string) is required.')]
bound = self._get_bound_mcp_from_query(query)
allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(bound)}
allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(execution_context, bound)}
if server_name not in allowed:
return [
provider_message.ContentElement.from_text(
@@ -1535,7 +1723,7 @@ class MCPLoader(loader.ToolLoader):
)
]
session = self.get_session(server_name)
session = self.get_session(execution_context, server_name)
if session is None or session.status != MCPSessionStatus.CONNECTED:
return [provider_message.ContentElement.from_text(f'Error: MCP server not connected: {server_name!r}')]
@@ -1586,13 +1774,15 @@ class MCPLoader(loader.ToolLoader):
async def get_tools(
self,
context: TenantContext,
bound_mcp_servers: list[str] | None = None,
*,
include_resource_tools: bool = True,
) -> list[resource_tool.LLMTool]:
await self._assert_execution_active(context)
all_functions: list[resource_tool.LLMTool] = []
for session in self.sessions.values():
for session in self._sessions_for_context(context):
# If bound_mcp_servers is specified, only include tools from those servers
if bound_mcp_servers is not None:
if session.server_uuid in bound_mcp_servers:
@@ -1601,7 +1791,7 @@ class MCPLoader(loader.ToolLoader):
# If no bound servers specified, include all tools
all_functions.extend(session.get_tools())
if include_resource_tools and self._eligible_resource_sessions_for_bound(bound_mcp_servers):
if include_resource_tools and self._eligible_resource_sessions_for_bound(context, bound_mcp_servers):
all_functions.extend(self._mcp_synthetic_resource_tools())
self._last_listed_functions = all_functions
@@ -1610,13 +1800,15 @@ class MCPLoader(loader.ToolLoader):
async def get_tool_catalog(
self,
context: TenantContext,
bound_mcp_servers: list[str] | None = None,
*,
include_resource_tools: bool = False,
) -> list[dict[str, typing.Any]]:
await self._assert_execution_active(context)
items: list[dict[str, typing.Any]] = []
for session in self.sessions.values():
for session in self._sessions_for_context(context):
if bound_mcp_servers is not None and session.server_uuid not in bound_mcp_servers:
continue
for tool in session.get_tools():
@@ -1632,7 +1824,7 @@ class MCPLoader(loader.ToolLoader):
}
)
if include_resource_tools and self._eligible_resource_sessions_for_bound(bound_mcp_servers):
if include_resource_tools and self._eligible_resource_sessions_for_bound(context, bound_mcp_servers):
for tool in self._mcp_synthetic_resource_tools():
items.append(
{
@@ -1648,18 +1840,20 @@ class MCPLoader(loader.ToolLoader):
return items
async def has_tool(self, name: str) -> bool:
async def has_tool(self, context: TenantContext, name: str) -> bool:
"""检查工具是否存在"""
await self._assert_execution_active(context)
if name in (MCP_TOOL_LIST_RESOURCES, MCP_TOOL_READ_RESOURCE):
return bool(self._eligible_resource_sessions_for_bound(None))
for session in self.sessions.values():
return bool(self._eligible_resource_sessions_for_bound(context, None))
for session in self._sessions_for_context(context):
for function in session.get_tools():
if function.name == name:
return True
return False
async def get_tool(self, name: str) -> resource_tool.LLMTool | None:
for session in self.sessions.values():
async def get_tool(self, context: TenantContext, name: str) -> resource_tool.LLMTool | None:
await self._assert_execution_active(context)
for session in self._sessions_for_context(context):
for function in session.get_tools():
if function.name == name:
return function
@@ -1667,6 +1861,7 @@ class MCPLoader(loader.ToolLoader):
async def invoke_tool(self, name: str, parameters: dict, query: pipeline_query.Query) -> typing.Any:
"""执行工具调用"""
execution_context = await self._assert_execution_active(_execution_context_from_query(query))
if name == MCP_TOOL_LIST_RESOURCES:
if getattr(query, 'variables', {}).get('_pipeline_mcp_resource_agent_read_enabled', True) is False:
return [provider_message.ContentElement.from_text('Error: MCP resource agent reads are disabled.')]
@@ -1676,7 +1871,7 @@ class MCPLoader(loader.ToolLoader):
return [provider_message.ContentElement.from_text('Error: MCP resource agent reads are disabled.')]
return await self._invoke_mcp_read_resource(parameters, query)
for session in self.sessions.values():
for session in self._sessions_for_context(execution_context):
for function in session.get_tools():
if function.name == name:
self.ap.logger.debug(f'Invoking MCP tool: {name} with parameters: {parameters}')
@@ -1690,22 +1885,25 @@ class MCPLoader(loader.ToolLoader):
raise ValueError(f'Tool not found: {name}')
async def get_resources(self, server_name: str) -> list[dict]:
async def get_resources(self, context: TenantContext, server_name: str) -> list[dict]:
"""Get resources from a specific MCP server."""
session = self.get_session(server_name)
await self._assert_execution_active(context)
session = self.get_session(context, server_name)
if session is None:
raise ValueError(f'MCP server not found: {server_name}')
return session.get_resources()
async def get_resource_templates(self, server_name: str) -> list[dict]:
async def get_resource_templates(self, context: TenantContext, server_name: str) -> list[dict]:
"""Get resource templates from a specific MCP server."""
session = self.get_session(server_name)
await self._assert_execution_active(context)
session = self.get_session(context, server_name)
if session is None:
raise ValueError(f'MCP server not found: {server_name}')
return session.get_resource_templates()
async def read_resource_envelope(
self,
context: TenantContext,
server_name: str,
uri: str,
*,
@@ -1716,7 +1914,8 @@ class MCPLoader(loader.ToolLoader):
query: pipeline_query.Query | None = None,
) -> dict:
"""Read a resource from a specific MCP server and return metadata plus contents."""
session = self.get_session(server_name)
await self._assert_execution_active(context)
session = self.get_session(context, server_name)
if session is None:
raise ValueError(f'MCP server not found: {server_name}')
return await session.read_resource_envelope(
@@ -1728,24 +1927,28 @@ class MCPLoader(loader.ToolLoader):
query=query,
)
async def read_resource(self, server_name: str, uri: str) -> list[dict]:
async def read_resource(self, context: TenantContext, server_name: str, uri: str) -> list[dict]:
"""Read a resource from a specific MCP server."""
envelope = await self.read_resource_envelope(server_name, uri)
envelope = await self.read_resource_envelope(context, server_name, uri)
return envelope['contents']
def get_session_by_uuid(self, server_uuid: str) -> RuntimeMCPSession | None:
for session in self.sessions.values():
def get_session_by_uuid(self, context: TenantContext, server_uuid: str) -> RuntimeMCPSession | None:
for session in self._sessions_for_context(context):
if session.server_uuid == server_uuid:
return session
return None
def _resolve_attachment_session(self, attachment: dict) -> RuntimeMCPSession | None:
def _resolve_attachment_session(
self,
context: TenantContext,
attachment: dict,
) -> RuntimeMCPSession | None:
server_uuid = attachment.get('server_uuid') or attachment.get('server_id')
server_name = attachment.get('server_name')
if server_uuid:
return self.get_session_by_uuid(server_uuid)
return self.get_session_by_uuid(context, server_uuid)
if server_name:
return self.get_session(server_name)
return self.get_session(context, server_name)
return None
async def build_resource_context_for_query(
@@ -1756,6 +1959,7 @@ class MCPLoader(loader.ToolLoader):
default_max_bytes: int = MCP_RESOURCE_CONTEXT_MAX_BYTES,
) -> str:
"""Build host-controlled MCP resource context for the current query."""
execution_context = await self._assert_execution_active(_execution_context_from_query(query))
if getattr(query, 'variables', {}).get('_pipeline_mcp_resource_agent_read_enabled', True) is False:
return ''
@@ -1764,7 +1968,7 @@ class MCPLoader(loader.ToolLoader):
return ''
bound = self._get_bound_mcp_from_query(query)
eligible = self._eligible_resource_sessions_for_bound(bound)
eligible = self._eligible_resource_sessions_for_bound(execution_context, bound)
eligible_by_uuid = {session.server_uuid: session for session in eligible}
eligible_by_name = {session.server_name: session for session in eligible}
@@ -1772,6 +1976,7 @@ class MCPLoader(loader.ToolLoader):
remaining_tokens = default_max_tokens
for raw_attachment in attachments:
await self._assert_execution_active(execution_context)
if remaining_tokens <= 0:
break
if not isinstance(raw_attachment, dict) or raw_attachment.get('enabled') is False:
@@ -1786,7 +1991,7 @@ class MCPLoader(loader.ToolLoader):
if not uri or not isinstance(uri, str):
continue
session = self._resolve_attachment_session(attachment)
session = self._resolve_attachment_session(execution_context, attachment)
if session is None:
continue
if session.server_uuid not in eligible_by_uuid and session.server_name not in eligible_by_name:
@@ -1804,6 +2009,8 @@ class MCPLoader(loader.ToolLoader):
source='preloaded',
query=query,
)
except WorkspaceError:
raise
except Exception as e:
self.ap.logger.warning(f'Failed to preload MCP resource {uri!r} from {session.server_name!r}: {e}')
continue
@@ -1843,37 +2050,40 @@ class MCPLoader(loader.ToolLoader):
pass
return context
async def remove_mcp_server(self, server_name: str):
async def remove_mcp_server(self, context: TenantContext, server_name: str):
"""移除 MCP 服务器"""
if server_name not in self.sessions:
await self._assert_execution_active(context)
key = self._session_key(context, server_name)
if key not in self.sessions:
self.ap.logger.warning(f'MCP server {server_name} not found in sessions, skipping removal')
return
session = self.sessions.pop(server_name)
session = self.sessions.pop(key)
await session.shutdown()
self.ap.logger.info(f'Removed MCP server: {server_name}')
def get_session(self, server_name: str) -> RuntimeMCPSession | None:
def get_session(self, context: TenantContext, server_name: str) -> RuntimeMCPSession | None:
"""获取指定名称的 MCP 会话"""
return self.sessions.get(server_name)
return self.sessions.get(self._session_key(context, server_name))
def has_session(self, server_name: str) -> bool:
def has_session(self, context: TenantContext, server_name: str) -> bool:
"""检查是否存在指定名称的 MCP 会话"""
return server_name in self.sessions
return self._session_key(context, server_name) in self.sessions
def get_all_server_names(self) -> list[str]:
def get_all_server_names(self, context: TenantContext) -> list[str]:
"""获取所有已加载的 MCP 服务器名称"""
return list(self.sessions.keys())
return [session.server_name for session in self._sessions_for_context(context)]
def get_server_tool_count(self, server_name: str) -> int:
def get_server_tool_count(self, context: TenantContext, server_name: str) -> int:
"""获取指定服务器的工具数量"""
session = self.get_session(server_name)
session = self.get_session(context, server_name)
return len(session.get_tools()) if session else 0
def get_all_servers_info(self) -> dict[str, dict]:
def get_all_servers_info(self, context: TenantContext) -> dict[str, dict]:
"""获取所有服务器的信息"""
info = {}
for server_name, session in self.sessions.items():
for session in self._sessions_for_context(context):
server_name = session.server_name
tools = session.get_tools()
info[server_name] = {
'name': server_name,
@@ -1887,23 +2097,13 @@ class MCPLoader(loader.ToolLoader):
async def shutdown(self):
"""关闭所有工具"""
self.ap.logger.info('Shutting down all MCP sessions...')
hosted_tasks = [task for task in self._hosted_mcp_tasks if not task.done()]
for task in hosted_tasks:
task.cancel()
if hosted_tasks:
await asyncio.gather(*hosted_tasks, return_exceptions=True)
self._hosted_mcp_tasks.clear()
async def shutdown_session(server_name: str, session: RuntimeMCPSession) -> None:
for key, session in list(self.sessions.items()):
try:
await session.shutdown()
self.ap.logger.debug(f'Shutdown MCP session: {server_name}')
self.ap.logger.debug(f'Shutdown MCP session: {session.server_name}')
except Exception as e:
self.ap.logger.error(f'Error shutting down MCP session {server_name}: {e}\n{traceback.format_exc()}')
await asyncio.gather(
*(shutdown_session(server_name, session) for server_name, session in list(self.sessions.items()))
)
self.ap.logger.error(
f'Error shutting down MCP session {session.server_name}: {e}\n{traceback.format_exc()}'
)
self.sessions.clear()
self.ap.logger.info('All MCP sessions shutdown complete')
@@ -6,7 +6,7 @@ import os
import shutil
import shlex
import threading
from contextlib import suppress, AsyncExitStack
from contextlib import suppress, AsyncExitStack, asynccontextmanager
from typing import TYPE_CHECKING, Any
import pydantic
@@ -94,6 +94,60 @@ class MCPServerBoxConfig(pydantic.BaseModel):
_HANDSHAKE_ATTEMPT_TIMEOUT_SEC = 10.0
@asynccontextmanager
async def authenticated_websocket_client(url: str, headers: dict[str, str]):
"""MCP WebSocket transport with host-only Box relay headers.
The upstream MCP helper does not expose WebSocket handshake headers. This
mirrors that transport while keeping the Box control token out of the URL,
JSON-RPC payloads, and logs.
"""
import json
import anyio
import mcp.types as mcp_types
from mcp.shared.message import SessionMessage
from pydantic import ValidationError
from websockets.asyncio.client import connect as ws_connect
from websockets.typing import Subprotocol
read_stream_writer, read_stream = anyio.create_memory_object_stream(0)
write_stream, write_stream_reader = anyio.create_memory_object_stream(0)
async with ws_connect(
url,
subprotocols=[Subprotocol('mcp')],
additional_headers=dict(headers),
proxy=None,
) as websocket:
async def ws_reader():
async with read_stream_writer:
async for raw_text in websocket:
try:
message = mcp_types.JSONRPCMessage.model_validate_json(raw_text)
await read_stream_writer.send(SessionMessage(message))
except ValidationError as exc: # pragma: no cover - upstream parity
await read_stream_writer.send(exc)
async def ws_writer():
async with write_stream_reader:
async for session_message in write_stream_reader:
payload = session_message.message.model_dump(
by_alias=True,
mode='json',
exclude_none=True,
)
await websocket.send(json.dumps(payload))
async with anyio.create_task_group() as task_group:
task_group.start_soon(ws_reader)
task_group.start_soon(ws_writer)
yield read_stream, write_stream
task_group.cancel_scope.cancel()
class _TransferredStack:
"""Adapts an already-populated AsyncExitStack into an async context manager
so ownership of its resources can be transferred into another exit stack.
@@ -149,6 +203,7 @@ class BoxStdioSessionRuntime:
resolved_host_path = self.resolve_host_path() if host_path is ... else host_path
return BoxWorkspaceSession(
self.ap.box_service,
self.owner.execution_context,
self.owner._build_box_session_id(),
host_path=resolved_host_path,
host_path_mode=self.config.host_path_mode,
@@ -249,7 +304,11 @@ class BoxStdioSessionRuntime:
if install_cmd:
payload = self._wrap_process_payload_with_python_env(payload, process_cwd)
payload['process_id'] = self.process_id
await workspace.box_service.start_managed_process(workspace.session_id, payload)
await workspace.box_service.start_managed_process(
workspace.execution_context,
workspace.session_id,
payload,
)
except Exception:
self.owner.error_phase = MCPSessionErrorPhase.PROCESS_START
raise
@@ -259,7 +318,10 @@ class BoxStdioSessionRuntime:
f'process_id={self.process_id} (transport reconnect)'
)
websocket_url = workspace.get_managed_process_websocket_url(self.process_id)
(
websocket_url,
websocket_headers,
) = await workspace.get_managed_process_websocket_connection(self.process_id)
# Attach the WS transport + MCP session ONCE, on the owner's exit stack,
# in the same task as the serve loop that follows. websocket_client and
@@ -277,7 +339,12 @@ class BoxStdioSessionRuntime:
# attempt re-attaches to the same live process; once it has finished
# cold start the handshake succeeds and stays healthy.
try:
transport = await self.owner.exit_stack.enter_async_context(websocket_client(websocket_url))
transport_context = (
authenticated_websocket_client(websocket_url, websocket_headers)
if websocket_headers
else websocket_client(websocket_url)
)
transport = await self.owner.exit_stack.enter_async_context(transport_context)
read_stream, write_stream = transport
self.owner.session = await self.owner.exit_stack.enter_async_context(
ClientSession(read_stream, write_stream)
@@ -11,6 +11,7 @@ from .. import loader
from ..errors import ToolNotFoundError
from .availability import is_box_backend_available
from . import skill as skill_loader
from ....api.http.context import ExecutionContext
EXEC_TOOL_NAME = 'exec'
READ_TOOL_NAME = 'read'
@@ -56,6 +57,17 @@ class NativeToolLoader(loader.ToolLoader):
"""Check if the box backend is truly available (not just the runtime)."""
return await is_box_backend_available(self.ap)
@staticmethod
def _execution_context(query: pipeline_query.Query) -> ExecutionContext:
return ExecutionContext(
instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
placement_generation=getattr(query, 'placement_generation', 0) or 0,
bot_uuid=getattr(query, 'bot_uuid', None),
pipeline_uuid=getattr(query, 'pipeline_uuid', None),
query_uuid=getattr(query, 'query_uuid', None),
)
async def get_tools(self, bound_plugins: list[str] | None = None) -> list[resource_tool.LLMTool]:
if not await self._is_sandbox_available():
return []
@@ -142,7 +154,7 @@ class NativeToolLoader(loader.ToolLoader):
result = self._normalize_exec_result(result)
if selected_skill is not None:
self._refresh_skill_from_disk(selected_skill)
self._refresh_skill_from_disk(query, selected_skill)
return result
def _resolve_host_path(
@@ -162,7 +174,11 @@ class NativeToolLoader(loader.ToolLoader):
)
box_service = self.ap.box_service
host_root = selected_skill.get('package_root') if selected_skill is not None else box_service.default_workspace
host_root = (
selected_skill.get('package_root')
if selected_skill is not None
else box_service._tenant_workspace(self._execution_context(query))
)
if not host_root:
raise ValueError('No host workspace configured for file operations.')
@@ -522,11 +538,19 @@ else:
return self._read_text_file_preview(host_path, parameters)
try:
result = await self.ap.box_service.read_skill_file(selected_skill['name'], relative)
result = await self.ap.box_service.read_skill_file(
self._execution_context(query),
selected_skill['name'],
relative,
)
return self._build_read_result_from_text(str(result.get('content', '')), parameters)
except Exception:
try:
result = await self.ap.box_service.list_skill_files(selected_skill['name'], relative)
result = await self.ap.box_service.list_skill_files(
self._execution_context(query),
selected_skill['name'],
relative,
)
entries = [entry['name'] for entry in result.get('entries', [])]
return self._build_directory_result(entries)
except Exception as exc:
@@ -562,8 +586,9 @@ else:
if encoding != 'text':
return {'ok': False, 'error': 'base64 writes to skill packages are not supported.'}
selected_skill, relative = skill_request
await self.ap.box_service.write_skill_file(selected_skill['name'], relative, content)
await self.ap.skill_mgr.reload_skills()
execution_context = self._execution_context(query)
await self.ap.box_service.write_skill_file(execution_context, selected_skill['name'], relative, content)
await self.ap.skill_mgr.reload_skills(execution_context)
return {'ok': True, 'path': path}
host_path, selected_skill = self._resolve_host_path(
@@ -579,7 +604,7 @@ else:
self._write_host_file(host_path, content, parameters)
except ValueError as exc:
return {'ok': False, 'error': str(exc)}
self._refresh_skill_from_disk(selected_skill)
self._refresh_skill_from_disk(query, selected_skill)
return {'ok': True, 'path': path}
async def _invoke_edit(self, parameters: dict, query: pipeline_query.Query) -> dict:
@@ -603,7 +628,11 @@ else:
):
selected_skill, relative = skill_request
try:
result = await self.ap.box_service.read_skill_file(selected_skill['name'], relative)
result = await self.ap.box_service.read_skill_file(
self._execution_context(query),
selected_skill['name'],
relative,
)
except Exception:
return {'ok': False, 'error': f'File not found: {path}'}
content = result.get('content', '')
@@ -613,8 +642,14 @@ else:
if count > 1:
return {'ok': False, 'error': f'old_string matches {count} locations; provide a more unique string.'}
new_content = content.replace(old_string, new_string, 1)
await self.ap.box_service.write_skill_file(selected_skill['name'], relative, new_content)
await self.ap.skill_mgr.reload_skills()
execution_context = self._execution_context(query)
await self.ap.box_service.write_skill_file(
execution_context,
selected_skill['name'],
relative,
new_content,
)
await self.ap.skill_mgr.reload_skills(execution_context)
return {'ok': True, 'path': path}
host_path, selected_skill = self._resolve_host_path(
@@ -637,10 +672,10 @@ else:
new_content = content.replace(old_string, new_string, 1)
with open(host_path, 'w', encoding='utf-8') as f:
f.write(new_content)
self._refresh_skill_from_disk(selected_skill)
self._refresh_skill_from_disk(query, selected_skill)
return {'ok': True, 'path': path}
def _refresh_skill_from_disk(self, selected_skill: dict | None) -> None:
def _refresh_skill_from_disk(self, query: pipeline_query.Query, selected_skill: dict | None) -> None:
if selected_skill is None:
return
@@ -650,7 +685,7 @@ else:
refresh_skill = getattr(skill_mgr, 'refresh_skill_from_disk', None)
if callable(refresh_skill):
refresh_skill(selected_skill.get('name', ''))
refresh_skill(self._execution_context(query), selected_skill.get('name', ''))
async def _is_sandbox_available(self) -> bool:
"""Refresh backend availability so Box reconnects restore tool exposure."""
@@ -67,7 +67,11 @@ class PluginToolLoader(loader.ToolLoader):
async def invoke_tool(self, name: str, parameters: dict, query: pipeline_query.Query) -> typing.Any:
try:
return await self.ap.plugin_connector.call_tool(
name, parameters, session=query.session, query_id=query.query_id
name,
parameters,
session=query.session,
query_id=query.query_id,
query_uuid=query.query_uuid,
)
except Exception as e:
self.ap.logger.error(f'执行函数 {name} 时发生错误: {e}')
@@ -4,6 +4,7 @@ import re
import typing
from ....box import workspace as box_workspace
from ....api.http.context import ExecutionContext
if typing.TYPE_CHECKING:
from ....core import app
@@ -36,7 +37,15 @@ def get_visible_skills(ap: app.Application, query: pipeline_query.Query) -> dict
if skill_mgr is None:
return {}
visible_skills = getattr(skill_mgr, 'skills', {})
execution_context = ExecutionContext(
instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
placement_generation=getattr(query, 'placement_generation', 0) or 0,
bot_uuid=getattr(query, 'bot_uuid', None),
pipeline_uuid=getattr(query, 'pipeline_uuid', None),
query_uuid=getattr(query, 'query_uuid', None),
)
visible_skills = skill_mgr.get_skills(execution_context)
bound_skills = get_bound_skill_names(query)
if bound_skills is None:
return visible_skills
@@ -7,6 +7,7 @@ import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
from .. import loader
from .availability import is_box_backend_available
from ....api.http.context import ExecutionContext
# Align with Claude Code's Skill tool design:
# - activate: Activate a skill via Tool Call, returns SKILL.md content
@@ -75,7 +76,7 @@ class SkillToolLoader(loader.ToolLoader):
if name == ACTIVATE_SKILL_TOOL_NAME:
return await self._invoke_activate_skill(parameters, query)
if name == REGISTER_SKILL_TOOL_NAME:
return await self._invoke_register_skill(parameters)
return await self._invoke_register_skill(parameters, query)
raise ValueError(f'Unknown skill tool: {name}')
async def shutdown(self):
@@ -128,7 +129,7 @@ class SkillToolLoader(loader.ToolLoader):
'content': result_content,
}
async def _invoke_register_skill(self, parameters: dict) -> typing.Any:
async def _invoke_register_skill(self, parameters: dict, query) -> typing.Any:
"""Register a skill from sandbox directory to data/skills/."""
sandbox_path = str(parameters.get('path', '') or '').strip()
if not sandbox_path:
@@ -143,7 +144,15 @@ class SkillToolLoader(loader.ToolLoader):
raise ValueError('Skill service not available')
# Scan and register the skill
scanned = await skill_service.scan_directory_async(host_path)
execution_context = ExecutionContext(
instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
placement_generation=getattr(query, 'placement_generation', 0) or 0,
bot_uuid=getattr(query, 'bot_uuid', None),
pipeline_uuid=getattr(query, 'pipeline_uuid', None),
query_uuid=getattr(query, 'query_uuid', None),
)
scanned = await skill_service.scan_directory_async(execution_context, host_path)
# Override name if provided
skill_name = str(parameters.get('name') or scanned['name']).strip()
@@ -152,13 +161,14 @@ class SkillToolLoader(loader.ToolLoader):
# Create the skill
created = await skill_service.create_skill(
execution_context,
{
'name': skill_name,
'display_name': str(parameters.get('display_name') or scanned.get('display_name', '')).strip(),
'description': str(parameters.get('description') or scanned.get('description', '')).strip(),
'instructions': str(parameters.get('instructions') or scanned.get('instructions', '')),
'package_root': host_path,
}
},
)
return {
+11 -4
View File
@@ -9,6 +9,8 @@ from langbot_plugin.api.entities.events import pipeline_query
from . import loader as tool_loader
from .errors import ToolNotFoundError
from ...pipeline.pool import get_query_execution_context
from ...api.http.service.tenant import TenantContext
if TYPE_CHECKING:
from ...core import app
@@ -57,6 +59,7 @@ class ToolManager:
async def get_all_tools(
self,
context: TenantContext,
bound_plugins: list[str] | None = None,
bound_mcp_servers: list[str] | None = None,
include_skill_authoring: bool = False,
@@ -70,6 +73,7 @@ class ToolManager:
all_functions.extend(await self.plugin_tool_loader.get_tools(bound_plugins))
all_functions.extend(
await self.mcp_tool_loader.get_tools(
context,
bound_mcp_servers,
include_resource_tools=include_mcp_resource_tools,
)
@@ -79,6 +83,7 @@ class ToolManager:
async def get_tool_catalog(
self,
context: TenantContext,
bound_plugins: list[str] | None = None,
bound_mcp_servers: list[str] | None = None,
include_skill_authoring: bool = False,
@@ -106,6 +111,7 @@ class ToolManager:
if self.mcp_tool_loader:
for item in await self.mcp_tool_loader.get_tool_catalog(
context,
bound_mcp_servers,
include_resource_tools=include_mcp_resource_tools,
):
@@ -113,19 +119,18 @@ class ToolManager:
return catalog
async def get_tool_by_name(self, name: str) -> tool_loader.ToolLookupResult | None:
async def get_tool_by_name(self, context: TenantContext, name: str) -> tool_loader.ToolLookupResult | None:
"""Get tool by name from any active loader."""
for active_loader in (
self.native_tool_loader,
self.plugin_tool_loader,
self.mcp_tool_loader,
self.skill_tool_loader,
):
tool = await active_loader.get_tool(name)
if tool:
return tool
return None
return await self.mcp_tool_loader.get_tool(context, name)
async def generate_tools_for_openai(self, use_funcs: list[resource_tool.LLMTool]) -> list:
tools = []
@@ -175,6 +180,7 @@ class ToolManager:
try:
await monitoring_service.record_tool_call(
get_query_execution_context(query),
tool_name=name,
tool_source=source,
duration=duration_ms,
@@ -249,7 +255,8 @@ class ToolManager:
query=query,
invoke=lambda: self.plugin_tool_loader.invoke_tool(name, parameters, query),
)
if await self.mcp_tool_loader.has_tool(name):
execution_context = get_query_execution_context(query)
if await self.mcp_tool_loader.has_tool(execution_context, name):
telemetry_features.increment(query, 'tool_calls', 'mcp')
return await self._invoke_tool_with_monitoring(
source='mcp',