fix: 修复 MCP worker 配置加载不一致
This commit is contained in:
parent
85d7f6354f
commit
6df070ddb6
1
.gitignore
vendored
1
.gitignore
vendored
@ -50,6 +50,7 @@ cache
|
||||
.claude
|
||||
.cursor
|
||||
.trae
|
||||
.codex
|
||||
.pytest_cache
|
||||
|
||||
*.secret*
|
||||
|
||||
@ -5,7 +5,6 @@ from dataclasses import MISSING, dataclass, field, fields
|
||||
from typing import Annotated, get_args, get_origin
|
||||
|
||||
from yuxi import config as sys_config
|
||||
from yuxi.services.mcp_service import get_mcp_server_names
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
@ -69,7 +68,7 @@ class BaseContext:
|
||||
default_factory=list,
|
||||
metadata={
|
||||
"name": "MCP服务器",
|
||||
"options": lambda: get_mcp_server_names(),
|
||||
"options": [],
|
||||
"description": (
|
||||
"MCP服务器列表,建议使用支持 SSE 的 MCP 服务器,"
|
||||
"如果需要使用 uvx 或 npx 运行的服务器,也请在项目外部启动 MCP 服务器,并在项目中配置 MCP 服务器。"
|
||||
|
||||
@ -2,12 +2,14 @@
|
||||
|
||||
Responsibilities:
|
||||
- Server configuration CRUD operations
|
||||
- Configuration synchronization (Database <-> Cache)
|
||||
- Built-in configuration synchronization (Code <-> Database)
|
||||
- Unified entry point for Agent tool retrieval (auto-filtering disabled_tools)
|
||||
- MCP Client and Tools management (formerly in agents/common/mcp.py)
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
@ -26,14 +28,12 @@ from yuxi.utils import logger
|
||||
# Global Lock for MCP state
|
||||
_mcp_lock = asyncio.Lock()
|
||||
|
||||
# Global MCP tools cache
|
||||
# 本地仅缓存工具对象。配置始终以数据库为准,每次按 server_name 现查。
|
||||
# cache key 使用 server_name:config_hash,当配置变化时会自然失效。
|
||||
_mcp_tools_cache: dict[str, list[Callable[..., Any]]] = {}
|
||||
|
||||
# MCP tools statistics (for reporting enabled/disabled counts)
|
||||
_mcp_tools_stats: dict[str, dict[str, int]] = {}
|
||||
|
||||
# MCP Server configurations (Runtime cache, loaded from DB)
|
||||
MCP_SERVERS: dict[str, dict[str, Any]] = {}
|
||||
_UNSET = object()
|
||||
|
||||
# Default MCP Server configurations (Imported to DB on first run)
|
||||
@ -74,55 +74,11 @@ _SYNCED_MCP_FIELDS = (
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def load_mcp_servers_from_db() -> None:
|
||||
"""Load all enabled MCP server configurations from database to MCP_SERVERS cache."""
|
||||
global MCP_SERVERS
|
||||
async def ensure_builtin_mcp_servers_in_db() -> None:
|
||||
"""Ensure built-in MCP server definitions exist in the database.
|
||||
|
||||
# Delayed import to avoid circular references
|
||||
from yuxi.storage.postgres.manager import pg_manager
|
||||
|
||||
try:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
result = await session.execute(select(MCPServer).filter(MCPServer.enabled == 1))
|
||||
servers = result.scalars().all()
|
||||
|
||||
async with _mcp_lock:
|
||||
MCP_SERVERS.clear()
|
||||
for server in servers:
|
||||
MCP_SERVERS[server.name] = server.to_mcp_config()
|
||||
|
||||
logger.info(f"Loaded {len(MCP_SERVERS)} MCP servers from database: {list(MCP_SERVERS.keys())}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load MCP servers from database: {e}")
|
||||
|
||||
|
||||
async def sync_mcp_server_to_cache(name: str, config: dict[str, Any] | None) -> None:
|
||||
"""Sync a single MCP server configuration to cache.
|
||||
|
||||
Args:
|
||||
name: Server name
|
||||
config: Server configuration, or None to remove from cache
|
||||
"""
|
||||
global MCP_SERVERS
|
||||
|
||||
async with _mcp_lock:
|
||||
if config is None:
|
||||
MCP_SERVERS.pop(name, None)
|
||||
logger.info(f"Removed MCP server '{name}' from cache")
|
||||
else:
|
||||
MCP_SERVERS[name] = config
|
||||
logger.info(f"Synced MCP server '{name}' to cache")
|
||||
|
||||
# Clear tools cache for this server
|
||||
_mcp_tools_cache.pop(name, None)
|
||||
|
||||
|
||||
async def init_mcp_servers() -> None:
|
||||
"""Initialize MCP server configurations.
|
||||
|
||||
On first run, if database is empty, import default configurations.
|
||||
Then load configurations from database to MCP_SERVERS cache.
|
||||
Also ensures all built-in MCP servers are present in the database.
|
||||
This function only synchronizes code-defined built-ins to the database.
|
||||
It does not preload runtime configuration into memory.
|
||||
"""
|
||||
# Delayed import to avoid circular references
|
||||
from yuxi.storage.postgres.manager import pg_manager
|
||||
@ -197,11 +153,8 @@ async def init_mcp_servers() -> None:
|
||||
elif session.dirty:
|
||||
await session.commit()
|
||||
|
||||
# Load configurations from database to cache
|
||||
await load_mcp_servers_from_db()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize MCP servers: {e}, traceback: {traceback.format_exc()}")
|
||||
logger.error(f"Failed to ensure builtin MCP servers in database: {e}, traceback: {traceback.format_exc()}")
|
||||
|
||||
|
||||
async def get_mcp_client(
|
||||
@ -227,10 +180,41 @@ def to_camel_case(s: str) -> str:
|
||||
s = s[0].lower() + s[1:]
|
||||
return s
|
||||
|
||||
async def _load_enabled_mcp_server_configs(
|
||||
*,
|
||||
names: list[str] | None = None,
|
||||
db: AsyncSession | None = None,
|
||||
) -> dict[str, dict[str, Any]]:
|
||||
"""Load enabled MCP server configs directly from the database."""
|
||||
if db is not None:
|
||||
stmt = select(MCPServer).where(MCPServer.enabled == 1)
|
||||
if names:
|
||||
stmt = stmt.where(MCPServer.name.in_(names))
|
||||
result = await db.execute(stmt)
|
||||
servers = result.scalars().all()
|
||||
return {server.name: server.to_mcp_config() for server in servers}
|
||||
|
||||
from yuxi.storage.postgres.manager import pg_manager
|
||||
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
return await _load_enabled_mcp_server_configs(names=names, db=session)
|
||||
|
||||
|
||||
async def get_enabled_mcp_server_config(server_name: str, *, db: AsyncSession | None = None) -> dict[str, Any] | None:
|
||||
"""Get the latest enabled MCP server config from the database."""
|
||||
configs = await _load_enabled_mcp_server_configs(names=[server_name], db=db)
|
||||
return configs.get(server_name)
|
||||
|
||||
|
||||
async def get_enabled_mcp_server_names(*, db: AsyncSession | None = None) -> list[str]:
|
||||
"""Get enabled MCP server names from the database."""
|
||||
configs = await _load_enabled_mcp_server_configs(db=db)
|
||||
return list(configs.keys())
|
||||
|
||||
|
||||
async def get_mcp_tools(
|
||||
server_name: str,
|
||||
additional_servers: dict[str, dict] = None,
|
||||
additional_servers: dict[str, dict[str, Any]] | None = None,
|
||||
disabled_tools: list[str] = None,
|
||||
cache: bool = True,
|
||||
force_refresh: bool = False,
|
||||
@ -249,70 +233,71 @@ async def get_mcp_tools(
|
||||
cache: Whether to use/update the cache (default: True)
|
||||
force_refresh: Whether to force a refresh from the server (default: False)
|
||||
"""
|
||||
global _mcp_tools_cache
|
||||
|
||||
# 1. Prepare Server Config
|
||||
async with _mcp_lock:
|
||||
mcp_servers = MCP_SERVERS | (additional_servers or {})
|
||||
|
||||
all_processed_tools = []
|
||||
|
||||
# 2. Check Cache / Fetch Strategy
|
||||
# If we have it in cache and don't need to force refresh, use cache.
|
||||
if not force_refresh and cache and server_name in _mcp_tools_cache:
|
||||
all_processed_tools = _mcp_tools_cache[server_name]
|
||||
if additional_servers and server_name in additional_servers:
|
||||
server_config = additional_servers[server_name]
|
||||
else:
|
||||
# Need to fetch from server
|
||||
try:
|
||||
assert server_name in mcp_servers, f"Server {server_name} not found in ({list(mcp_servers.keys())})"
|
||||
server_config = await get_enabled_mcp_server_config(server_name)
|
||||
|
||||
# Extract connection config
|
||||
server_config = mcp_servers[server_name]
|
||||
if server_config is None:
|
||||
logger.warning(f"MCP server '{server_name}' not found in database or disabled")
|
||||
return []
|
||||
|
||||
# 配置 hash 直接基于完整配置生成。只要数据库中的配置发生变化,
|
||||
# 本地工具缓存 key 就会变化,从而自然触发重建。
|
||||
config_payload = json.dumps(server_config, sort_keys=True, ensure_ascii=True, separators=(",", ":"))
|
||||
config_hash = hashlib.sha256(config_payload.encode("utf-8")).hexdigest()[:16]
|
||||
cache_key = f"{server_name}:{config_hash}"
|
||||
|
||||
all_processed_tools: list[Callable[..., Any]] = []
|
||||
|
||||
async with _mcp_lock:
|
||||
if not force_refresh and cache and cache_key in _mcp_tools_cache:
|
||||
all_processed_tools = _mcp_tools_cache[cache_key]
|
||||
|
||||
if not all_processed_tools:
|
||||
try:
|
||||
# disabled_tools 只影响返回值过滤,不参与 MCP client 建连参数。
|
||||
client_config = {k: v for k, v in server_config.items() if k not in ("disabled_tools",)}
|
||||
|
||||
client = await get_mcp_client({server_name: client_config})
|
||||
if client is None:
|
||||
return []
|
||||
|
||||
# Get ALL tools (Raw)
|
||||
raw_tools = cast(list[Any], await client.get_tools())
|
||||
|
||||
# Render IDs for ALL tools
|
||||
server_cc = to_camel_case(server_name)
|
||||
for tool in raw_tools:
|
||||
# Render unique ID rule: mcp__[camelCaseServer]__[camelCaseTool]
|
||||
original_name = tool.name
|
||||
tool_cc = to_camel_case(original_name)
|
||||
unique_id = f"mcp__{server_cc}__{tool_cc}"
|
||||
|
||||
# Use metadata to store
|
||||
if tool.metadata is None:
|
||||
tool.metadata = {}
|
||||
tool.metadata["id"] = unique_id
|
||||
|
||||
all_processed_tools.append(tool)
|
||||
|
||||
# Update Cache (Store the FULL list)
|
||||
if cache:
|
||||
_mcp_tools_cache[server_name] = all_processed_tools
|
||||
async with _mcp_lock:
|
||||
stale_keys = [
|
||||
key for key in _mcp_tools_cache if key.startswith(f"{server_name}:") and key != cache_key
|
||||
]
|
||||
for stale_key in stale_keys:
|
||||
_mcp_tools_cache.pop(stale_key, None)
|
||||
_mcp_tools_cache[cache_key] = all_processed_tools
|
||||
|
||||
# Update Stats
|
||||
# Stats should reflect the GLOBAL configuration state
|
||||
# (How many are disabled in the stored config, not the transient arg)
|
||||
global_config_disabled = mcp_servers.get(server_name, {}).get("disabled_tools") or []
|
||||
global_config_disabled = server_config.get("disabled_tools") or []
|
||||
enabled_count = len([t for t in all_processed_tools if t.name not in global_config_disabled])
|
||||
|
||||
_mcp_tools_stats[server_name] = {
|
||||
"total": len(all_processed_tools),
|
||||
"enabled": enabled_count,
|
||||
"disabled": len(all_processed_tools) - enabled_count,
|
||||
}
|
||||
|
||||
logger.info(f"Refreshed MCP tools cache for '{server_name}': {len(all_processed_tools)} tools loaded.")
|
||||
logger.info(
|
||||
f"Refreshed MCP tools cache for '{server_name}' with key '{cache_key}': "
|
||||
f"{len(all_processed_tools)} tools loaded."
|
||||
)
|
||||
|
||||
except AssertionError as e:
|
||||
logger.warning(f"[assert] Failed to load tools from MCP server '{server_name}': {e}")
|
||||
return []
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to load tools from MCP server '{server_name}': {e}, traceback: {traceback.format_exc()}"
|
||||
@ -333,20 +318,14 @@ async def get_mcp_tools(
|
||||
|
||||
async def get_tools_from_all_servers() -> list[Callable[..., Any]]:
|
||||
"""Get all tools from all configured MCP servers."""
|
||||
server_configs = await _load_enabled_mcp_server_configs()
|
||||
all_tools = []
|
||||
for server_name in MCP_SERVERS.keys():
|
||||
tools = await get_mcp_tools(server_name)
|
||||
for server_name in server_configs:
|
||||
tools = await get_mcp_tools(server_name, additional_servers=server_configs)
|
||||
all_tools.extend(tools)
|
||||
return all_tools
|
||||
|
||||
|
||||
def add_mcp_server(name: str, config: dict[str, Any]) -> None:
|
||||
"""Add a new MCP server configuration."""
|
||||
MCP_SERVERS[name] = config
|
||||
# Clear client to force reinitialization with new config
|
||||
clear_mcp_cache()
|
||||
|
||||
|
||||
def clear_mcp_cache() -> None:
|
||||
"""Clear the MCP tools cache (useful for testing)."""
|
||||
global _mcp_tools_cache, _mcp_tools_stats
|
||||
@ -357,7 +336,10 @@ def clear_mcp_cache() -> None:
|
||||
def clear_mcp_server_tools_cache(server_name: str) -> None:
|
||||
"""Clear the tools cache for a specific MCP server."""
|
||||
global _mcp_tools_cache, _mcp_tools_stats
|
||||
_mcp_tools_cache.pop(server_name, None)
|
||||
server_prefix = f"{server_name}:"
|
||||
stale_keys = [key for key in _mcp_tools_cache if key.startswith(server_prefix)]
|
||||
for stale_key in stale_keys:
|
||||
_mcp_tools_cache.pop(stale_key, None)
|
||||
_mcp_tools_stats.pop(server_name, None)
|
||||
logger.info(f"Cleared tools cache for MCP server '{server_name}'")
|
||||
|
||||
@ -431,8 +413,7 @@ async def create_mcp_server(
|
||||
await db.commit()
|
||||
await db.refresh(server)
|
||||
|
||||
# Sync to cache
|
||||
await sync_mcp_server_to_cache(name, server.to_mcp_config())
|
||||
clear_mcp_server_tools_cache(name)
|
||||
|
||||
logger.info(f"Created MCP server '{name}'")
|
||||
return server
|
||||
@ -487,9 +468,7 @@ async def update_mcp_server(
|
||||
await db.commit()
|
||||
await db.refresh(server)
|
||||
|
||||
# Sync to cache (if enabled)
|
||||
if server.enabled:
|
||||
await sync_mcp_server_to_cache(name, server.to_mcp_config())
|
||||
clear_mcp_server_tools_cache(name)
|
||||
|
||||
logger.info(f"Updated MCP server '{name}'")
|
||||
return server
|
||||
@ -504,8 +483,7 @@ async def delete_mcp_server(db: AsyncSession, name: str) -> bool:
|
||||
await db.delete(server)
|
||||
await db.commit()
|
||||
|
||||
# Remove from cache
|
||||
await sync_mcp_server_to_cache(name, None)
|
||||
clear_mcp_server_tools_cache(name)
|
||||
|
||||
logger.info(f"Deleted MCP server '{name}'")
|
||||
return True
|
||||
@ -529,10 +507,8 @@ async def set_server_enabled(
|
||||
server.updated_by = updated_by
|
||||
await db.commit()
|
||||
|
||||
# Sync to cache
|
||||
is_enabled = bool(server.enabled)
|
||||
server_config = server.to_mcp_config() if is_enabled else None
|
||||
await sync_mcp_server_to_cache(name, server_config)
|
||||
clear_mcp_server_tools_cache(name)
|
||||
|
||||
logger.info(f"Set MCP server '{name}' enabled={is_enabled}")
|
||||
return is_enabled, server
|
||||
@ -585,19 +561,11 @@ async def toggle_tool_enabled(
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def get_mcp_server_names() -> list[str]:
|
||||
"""Get list of loaded MCP server names.
|
||||
|
||||
Returns a copy of keys to avoid runtime modification issues during iteration.
|
||||
"""
|
||||
return list(MCP_SERVERS.keys())
|
||||
|
||||
|
||||
async def get_enabled_mcp_tools(server_name: str) -> list:
|
||||
"""Get MCP server tools (auto-filtering disabled_tools).
|
||||
|
||||
Unified entry point for Agents, automatically:
|
||||
1. Gets server config from cache
|
||||
1. Gets the latest server config from database
|
||||
2. Gets all tools
|
||||
3. Filters out disabled_tools
|
||||
|
||||
@ -607,13 +575,13 @@ async def get_enabled_mcp_tools(server_name: str) -> list:
|
||||
Returns:
|
||||
List of enabled tools
|
||||
"""
|
||||
config = MCP_SERVERS.get(server_name)
|
||||
if not config:
|
||||
logger.warning(f"MCP server '{server_name}' not found in cache")
|
||||
config = await get_enabled_mcp_server_config(server_name)
|
||||
if config is None:
|
||||
logger.warning(f"MCP server '{server_name}' not found in database or disabled")
|
||||
return []
|
||||
|
||||
disabled_tools = config.get("disabled_tools") or []
|
||||
return await get_mcp_tools(server_name, disabled_tools=disabled_tools)
|
||||
return await get_mcp_tools(server_name, additional_servers={server_name: config}, disabled_tools=disabled_tools)
|
||||
|
||||
|
||||
async def get_servers_config(names: list[str]) -> dict[str, dict[str, Any]]:
|
||||
@ -625,7 +593,7 @@ async def get_servers_config(names: list[str]) -> dict[str, dict[str, Any]]:
|
||||
Returns:
|
||||
{name: config} dictionary, containing only found servers
|
||||
"""
|
||||
return {name: MCP_SERVERS[name] for name in names if name in MCP_SERVERS}
|
||||
return await _load_enabled_mcp_server_configs(names=names)
|
||||
|
||||
|
||||
async def get_all_mcp_tools(server_name: str) -> list:
|
||||
@ -640,10 +608,16 @@ async def get_all_mcp_tools(server_name: str) -> list:
|
||||
Returns:
|
||||
List of all tools (unfiltered)
|
||||
"""
|
||||
config = MCP_SERVERS.get(server_name)
|
||||
if not config:
|
||||
logger.warning(f"MCP server '{server_name}' not found in cache")
|
||||
config = await get_enabled_mcp_server_config(server_name)
|
||||
if config is None:
|
||||
logger.warning(f"MCP server '{server_name}' not found in database or disabled")
|
||||
return []
|
||||
|
||||
# Get all tools (no filtering, force refresh, no cache update)
|
||||
return await get_mcp_tools(server_name, disabled_tools=[], cache=False, force_refresh=True)
|
||||
return await get_mcp_tools(
|
||||
server_name,
|
||||
additional_servers={server_name: config},
|
||||
disabled_tools=[],
|
||||
cache=False,
|
||||
force_refresh=True,
|
||||
)
|
||||
|
||||
@ -12,6 +12,7 @@ from sqlalchemy import select
|
||||
from sqlalchemy.exc import OperationalError
|
||||
from yuxi.repositories.agent_run_repository import TERMINAL_RUN_STATUSES, AgentRunRepository
|
||||
from yuxi.services.chat_service import stream_agent_chat
|
||||
from yuxi.services.mcp_service import ensure_builtin_mcp_servers_in_db
|
||||
from yuxi.services.run_queue_service import (
|
||||
append_run_stream_event,
|
||||
clear_cancel_signal,
|
||||
@ -362,9 +363,11 @@ async def process_agent_run(ctx, run_id: str):
|
||||
|
||||
|
||||
async def _worker_startup(ctx):
|
||||
del ctx
|
||||
pg_manager.initialize()
|
||||
await pg_manager.create_business_tables()
|
||||
await pg_manager.ensure_business_schema()
|
||||
await ensure_builtin_mcp_servers_in_db()
|
||||
|
||||
|
||||
async def _worker_shutdown(ctx):
|
||||
|
||||
@ -15,7 +15,7 @@ import yaml
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from yuxi import config as sys_config
|
||||
from yuxi.repositories.skill_repository import SkillRepository
|
||||
from yuxi.services.mcp_service import get_mcp_server_names
|
||||
from yuxi.services.mcp_service import get_enabled_mcp_server_names
|
||||
from yuxi.storage.postgres.models_business import Skill
|
||||
from yuxi.utils.logging_config import logger
|
||||
|
||||
@ -259,7 +259,7 @@ async def get_skill_dependency_options(db: AsyncSession) -> dict[str, list[str]
|
||||
items, tool_list, mcp_names = await asyncio.gather(
|
||||
get_skills(),
|
||||
asyncio.to_thread(get_tools),
|
||||
asyncio.to_thread(get_mcp_server_names),
|
||||
get_enabled_mcp_server_names(db=db),
|
||||
)
|
||||
|
||||
return {
|
||||
@ -282,7 +282,7 @@ def _get_all_tool_names() -> list[str]:
|
||||
return [tool["id"] for tool in all_tools]
|
||||
|
||||
|
||||
def _validate_dependencies(
|
||||
async def _validate_dependencies(
|
||||
*,
|
||||
slug: str,
|
||||
tool_dependencies: list[str],
|
||||
@ -300,7 +300,7 @@ def _validate_dependencies(
|
||||
if invalid_tools:
|
||||
raise ValueError(f"存在无效工具依赖: {', '.join(invalid_tools)}")
|
||||
|
||||
available_mcps = set(get_mcp_server_names())
|
||||
available_mcps = set(await get_enabled_mcp_server_names(db=None))
|
||||
invalid_mcps = [name for name in mcps if name not in available_mcps]
|
||||
if invalid_mcps:
|
||||
raise ValueError(f"存在无效 MCP 依赖: {', '.join(invalid_mcps)}")
|
||||
@ -328,7 +328,7 @@ async def update_skill_dependencies(
|
||||
repo = SkillRepository(db)
|
||||
skill_items = await repo.list_all()
|
||||
available_skill_slugs = {skill.slug for skill in skill_items}
|
||||
tools, mcps, skills = _validate_dependencies(
|
||||
tools, mcps, skills = await _validate_dependencies(
|
||||
slug=slug,
|
||||
tool_dependencies=tool_dependencies,
|
||||
mcp_dependencies=mcp_dependencies,
|
||||
|
||||
@ -4,7 +4,7 @@ from fastapi import FastAPI
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
|
||||
from yuxi.services.task_service import tasker
|
||||
from yuxi.services.mcp_service import init_mcp_servers
|
||||
from yuxi.services.mcp_service import ensure_builtin_mcp_servers_in_db
|
||||
from yuxi.services.subagent_service import init_builtin_subagents
|
||||
from yuxi.services.run_queue_service import close_queue_clients, get_redis_client
|
||||
from yuxi.storage.postgres.manager import pg_manager
|
||||
@ -26,11 +26,11 @@ async def lifespan(app: FastAPI):
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize database during startup: {e}")
|
||||
|
||||
# 初始化 MCP 服务器配置
|
||||
# 确保内置 MCP 服务器定义存在于数据库
|
||||
try:
|
||||
await init_mcp_servers()
|
||||
await ensure_builtin_mcp_servers_in_db()
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize MCP servers during startup: {e}")
|
||||
logger.error(f"Failed to ensure builtin MCP servers during startup: {e}")
|
||||
|
||||
# 初始化内置 SubAgent
|
||||
try:
|
||||
|
||||
114
backend/test/unit/services/test_mcp_service.py
Normal file
114
backend/test/unit/services/test_mcp_service.py
Normal file
@ -0,0 +1,114 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
from yuxi.services import mcp_service
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, tools):
|
||||
self._tools = tools
|
||||
|
||||
async def get_tools(self):
|
||||
return self._tools
|
||||
|
||||
|
||||
async def test_get_enabled_mcp_tools_loads_latest_config_from_db(monkeypatch):
|
||||
captured: list[dict] = []
|
||||
|
||||
async def fake_get_enabled_mcp_server_config(server_name: str, db=None):
|
||||
del db
|
||||
assert server_name == "demo"
|
||||
return {"transport": "stdio", "command": "demo", "disabled_tools": ["tool_b"]}
|
||||
|
||||
async def fake_get_mcp_tools(server_name: str, additional_servers=None, disabled_tools=None, **kwargs):
|
||||
del kwargs
|
||||
captured.append(
|
||||
{
|
||||
"server_name": server_name,
|
||||
"additional_servers": additional_servers,
|
||||
"disabled_tools": list(disabled_tools or []),
|
||||
}
|
||||
)
|
||||
return ["tool-a"]
|
||||
|
||||
monkeypatch.setattr(mcp_service, "get_enabled_mcp_server_config", fake_get_enabled_mcp_server_config)
|
||||
monkeypatch.setattr(mcp_service, "get_mcp_tools", fake_get_mcp_tools)
|
||||
|
||||
tools = await mcp_service.get_enabled_mcp_tools("demo")
|
||||
|
||||
assert tools == ["tool-a"]
|
||||
assert captured == [
|
||||
{
|
||||
"server_name": "demo",
|
||||
"additional_servers": {
|
||||
"demo": {"transport": "stdio", "command": "demo", "disabled_tools": ["tool_b"]}
|
||||
},
|
||||
"disabled_tools": ["tool_b"],
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
async def test_get_mcp_tools_rebuilds_cache_when_config_hash_changes(monkeypatch):
|
||||
mcp_service.clear_mcp_cache()
|
||||
|
||||
configs = [
|
||||
{"transport": "stdio", "command": "demo-v1", "disabled_tools": []},
|
||||
{"transport": "stdio", "command": "demo-v2", "disabled_tools": []},
|
||||
]
|
||||
build_calls: list[str] = []
|
||||
|
||||
async def fake_get_enabled_mcp_server_config(server_name: str, db=None):
|
||||
del db
|
||||
assert server_name == "demo"
|
||||
return configs[0]
|
||||
|
||||
async def fake_get_mcp_client(server_configs):
|
||||
config = server_configs["demo"]
|
||||
build_calls.append(config["command"])
|
||||
tool = SimpleNamespace(name=f"tool_for_{config['command']}", metadata={})
|
||||
return _FakeClient([tool])
|
||||
|
||||
monkeypatch.setattr(mcp_service, "get_enabled_mcp_server_config", fake_get_enabled_mcp_server_config)
|
||||
monkeypatch.setattr(mcp_service, "get_mcp_client", fake_get_mcp_client)
|
||||
|
||||
tools_v1_first = await mcp_service.get_mcp_tools("demo")
|
||||
tools_v1_second = await mcp_service.get_mcp_tools("demo")
|
||||
|
||||
configs[0] = configs[1]
|
||||
tools_v2 = await mcp_service.get_mcp_tools("demo")
|
||||
|
||||
assert [tool.name for tool in tools_v1_first] == ["tool_for_demo-v1"]
|
||||
assert [tool.name for tool in tools_v1_second] == ["tool_for_demo-v1"]
|
||||
assert [tool.name for tool in tools_v2] == ["tool_for_demo-v2"]
|
||||
assert build_calls == ["demo-v1", "demo-v2"]
|
||||
|
||||
mcp_service.clear_mcp_cache()
|
||||
|
||||
|
||||
async def test_get_tools_from_all_servers_loads_names_from_db_once(monkeypatch):
|
||||
server_configs = {
|
||||
"alpha": {"transport": "stdio", "command": "cmd-a", "disabled_tools": []},
|
||||
"beta": {"transport": "stdio", "command": "cmd-b", "disabled_tools": []},
|
||||
}
|
||||
calls: list[tuple[str, dict[str, dict]]] = []
|
||||
|
||||
async def fake_load_enabled_mcp_server_configs(*, names=None, db=None):
|
||||
del names, db
|
||||
return server_configs
|
||||
|
||||
async def fake_get_mcp_tools(server_name: str, additional_servers=None, **kwargs):
|
||||
del kwargs
|
||||
calls.append((server_name, additional_servers or {}))
|
||||
return [server_name]
|
||||
|
||||
monkeypatch.setattr(mcp_service, "_load_enabled_mcp_server_configs", fake_load_enabled_mcp_server_configs)
|
||||
monkeypatch.setattr(mcp_service, "get_mcp_tools", fake_get_mcp_tools)
|
||||
|
||||
tools = await mcp_service.get_tools_from_all_servers()
|
||||
|
||||
assert tools == ["alpha", "beta"]
|
||||
assert calls == [
|
||||
("alpha", server_configs),
|
||||
("beta", server_configs),
|
||||
]
|
||||
@ -151,3 +151,34 @@ async def test_process_agent_run_retryable_error_retries_then_completes(monkeypa
|
||||
|
||||
await run_worker.process_agent_run({"job_try": 2}, "run-1")
|
||||
assert terminal_statuses == ["completed"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_startup_ensures_builtin_mcp_servers(monkeypatch: pytest.MonkeyPatch):
|
||||
calls: list[str] = []
|
||||
|
||||
def fake_initialize():
|
||||
calls.append("initialize")
|
||||
|
||||
async def fake_create_business_tables():
|
||||
calls.append("create_business_tables")
|
||||
|
||||
async def fake_ensure_business_schema():
|
||||
calls.append("ensure_business_schema")
|
||||
|
||||
async def fake_ensure_builtin_mcp_servers_in_db():
|
||||
calls.append("ensure_builtin_mcp_servers_in_db")
|
||||
|
||||
monkeypatch.setattr(run_worker.pg_manager, "initialize", fake_initialize)
|
||||
monkeypatch.setattr(run_worker.pg_manager, "create_business_tables", fake_create_business_tables)
|
||||
monkeypatch.setattr(run_worker.pg_manager, "ensure_business_schema", fake_ensure_business_schema)
|
||||
monkeypatch.setattr(run_worker, "ensure_builtin_mcp_servers_in_db", fake_ensure_builtin_mcp_servers_in_db)
|
||||
|
||||
await run_worker._worker_startup({})
|
||||
|
||||
assert calls == [
|
||||
"initialize",
|
||||
"create_business_tables",
|
||||
"ensure_business_schema",
|
||||
"ensure_builtin_mcp_servers_in_db",
|
||||
]
|
||||
|
||||
@ -74,7 +74,11 @@ async def test_get_skill_dependency_options(monkeypatch: pytest.MonkeyPatch):
|
||||
]
|
||||
|
||||
monkeypatch.setattr(tool_service, "get_tool_metadata", fake_get_tool_metadata)
|
||||
monkeypatch.setattr(svc, "get_mcp_server_names", lambda: ["mcp-a", "mcp-b"])
|
||||
async def fake_get_enabled_mcp_server_names(db=None):
|
||||
del db
|
||||
return ["mcp-a", "mcp-b"]
|
||||
|
||||
monkeypatch.setattr(svc, "get_enabled_mcp_server_names", fake_get_enabled_mcp_server_names)
|
||||
|
||||
class FakeRepo:
|
||||
def __init__(self, _db):
|
||||
@ -312,7 +316,12 @@ async def test_update_skill_dependencies(monkeypatch: pytest.MonkeyPatch):
|
||||
return [{"id": "calculator", "name": "Calculator"}]
|
||||
|
||||
monkeypatch.setattr(tool_service, "get_tool_metadata", fake_get_tool_metadata)
|
||||
monkeypatch.setattr(svc, "get_mcp_server_names", lambda: ["mcp-a"])
|
||||
|
||||
async def fake_get_enabled_mcp_server_names(db=None):
|
||||
del db
|
||||
return ["mcp-a"]
|
||||
|
||||
monkeypatch.setattr(svc, "get_enabled_mcp_server_names", fake_get_enabled_mcp_server_names)
|
||||
|
||||
async def fake_get_skill_or_raise(_db, slug: str):
|
||||
assert slug == "alpha"
|
||||
|
||||
@ -39,6 +39,7 @@
|
||||
- 调整 Skills 导入能力:`/api/system/skills/import` 现在除 ZIP 外也支持直接上传单个 `SKILL.md`,前端上传入口与后端导入服务同步兼容,便于快速导入单文件技能
|
||||
- 新增 Skills 远程安装能力:Skills 管理页支持填写 `owner/repo` 或 GitHub URL,后端通过隔离的临时 `HOME` 调用 `npx skills add` 下载指定 skill,再复用现有导入链路写入 `saves/skills` 和数据库,避免将 `~/.agents/skills` 直接作为系统主存储;前端远程安装弹窗补充多选串行安装与批量进度展示,复用现有单 skill 安装接口逐个提交请求
|
||||
- 调整部门删除语义:删除部门时不再要求用户数为 0,而是将部门下用户迁移到默认部门,同时清理部门级配置和部门 API Key,保证测试部门、撤换部门等场景可直接删除,并补充对应集成测试覆盖该链路
|
||||
- 重构 MCP 运行时配置加载模型:移除 `MCP_SERVERS` 作为运行正确性前提的设计,改为每次直接从数据库读取最新 MCP 配置,并用 `server_name:config_hash` 作为本地工具缓存 key;同时将内置 MCP 初始化职责收敛为仅同步数据库默认项,前端 MCP 选项改为直接使用实时资源列表,解决 `api`/`worker` 分进程下的配置不一致与缓存失效问题
|
||||
|
||||
---
|
||||
|
||||
|
||||
@ -451,6 +451,7 @@ import ModelSelectorComponent from '@/components/ModelSelectorComponent.vue'
|
||||
import { useAgentStore } from '@/stores/agent'
|
||||
import { useUserStore } from '@/stores/user'
|
||||
import { useDatabaseStore } from '@/stores/database'
|
||||
import { mcpApi } from '@/apis/mcp_api'
|
||||
import { skillApi } from '@/apis/skill_api'
|
||||
import { subagentApi } from '@/apis/subagent_api'
|
||||
import { toolApi } from '@/apis/tool_api'
|
||||
@ -484,6 +485,7 @@ watch(
|
||||
// 强制刷新以获取最新数据
|
||||
databaseStore.loadDatabases(true).catch(() => {})
|
||||
loadLiveSkillOptions(true).catch(() => {})
|
||||
loadMcpOptions(true).catch(() => {})
|
||||
loadSubagentOptions(true).catch(() => {})
|
||||
loadToolOptions(true).catch(() => {})
|
||||
if (selectedAgentId.value) {
|
||||
@ -501,7 +503,6 @@ watch(
|
||||
)
|
||||
|
||||
const {
|
||||
availableTools,
|
||||
selectedAgent,
|
||||
selectedAgentId,
|
||||
selectedAgentConfigId,
|
||||
@ -522,6 +523,7 @@ const systemPromptModalOpen = ref(false)
|
||||
const currentSystemPromptKey = ref(null)
|
||||
const systemPromptDraft = ref('')
|
||||
const liveSkillOptions = ref([])
|
||||
const liveMcpOptions = ref([])
|
||||
const liveSubagentOptions = ref([])
|
||||
const toolOptionsFromApi = ref([])
|
||||
const createConfigModalOpen = ref(false)
|
||||
@ -616,6 +618,29 @@ const loadLiveSkillOptions = async (force = false) => {
|
||||
}
|
||||
}
|
||||
|
||||
const loadMcpOptions = async (force = false) => {
|
||||
if (!userStore.isAdmin) {
|
||||
liveMcpOptions.value = []
|
||||
return
|
||||
}
|
||||
if (!force && liveMcpOptions.value.length > 0) {
|
||||
return
|
||||
}
|
||||
try {
|
||||
const result = await mcpApi.getMcpServers()
|
||||
const rows = result?.data || []
|
||||
liveMcpOptions.value = rows
|
||||
.filter((item) => item?.enabled !== false)
|
||||
.map((item) => ({
|
||||
id: item.name,
|
||||
name: item.name,
|
||||
description: item.description || ''
|
||||
}))
|
||||
} catch (error) {
|
||||
console.warn('加载 MCP 列表失败:', error)
|
||||
}
|
||||
}
|
||||
|
||||
const loadToolOptions = async (force = false) => {
|
||||
if (!userStore.isAdmin) {
|
||||
toolOptionsFromApi.value = []
|
||||
@ -687,8 +712,8 @@ const refreshConfigOptions = async (_key, kind) => {
|
||||
message.success('Subagents 列表已刷新')
|
||||
break
|
||||
case 'mcps':
|
||||
// MCP 没有前端 store,提示用户刷新页面
|
||||
message.info('请在 MCP 管理页面刷新')
|
||||
await loadMcpOptions(true)
|
||||
message.success('MCP 列表已刷新')
|
||||
break
|
||||
}
|
||||
} catch (error) {
|
||||
@ -725,25 +750,28 @@ const navigateToConfigPage = (kind) => {
|
||||
}
|
||||
|
||||
// 通用选项获取与处理
|
||||
const resolveOptionValue = (option) => {
|
||||
if (typeof option === 'object' && option !== null) {
|
||||
return option.id || option.value || option.name || option.db_id || option.slug
|
||||
}
|
||||
return option
|
||||
}
|
||||
|
||||
const getConfigOptions = (value) => {
|
||||
if (value?.template_metadata?.kind === 'tools') {
|
||||
// 优先使用从 API 获取的工具列表,否则回退到 configurableItems 中的选项
|
||||
return toolOptionsFromApi.value.length > 0
|
||||
? toolOptionsFromApi.value
|
||||
: availableTools.value
|
||||
? Object.values(availableTools.value)
|
||||
: []
|
||||
return toolOptionsFromApi.value || []
|
||||
}
|
||||
if (value?.template_metadata?.kind === 'knowledges') {
|
||||
return databaseStore.databases || []
|
||||
}
|
||||
if (value?.template_metadata?.kind === 'mcps') {
|
||||
return liveMcpOptions.value || []
|
||||
}
|
||||
if (value?.template_metadata?.kind === 'skills') {
|
||||
return liveSkillOptions.value.length > 0 ? liveSkillOptions.value : value?.options || []
|
||||
}
|
||||
if (value?.template_metadata?.kind === 'subagents') {
|
||||
const options =
|
||||
liveSubagentOptions.value.length > 0 ? liveSubagentOptions.value : value?.options || []
|
||||
return options.filter((option) => option?.enabled !== false)
|
||||
return liveSubagentOptions.value || []
|
||||
}
|
||||
return value?.options || []
|
||||
}
|
||||
@ -755,10 +783,7 @@ const isListConfig = (key, value) => {
|
||||
}
|
||||
|
||||
const getOptionValue = (option) => {
|
||||
if (typeof option === 'object' && option !== null) {
|
||||
return option.id || option.value || option.name || option.db_id || option.slug
|
||||
}
|
||||
return option
|
||||
return resolveOptionValue(option)
|
||||
}
|
||||
|
||||
const getOptionLabel = (option) => {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user