From c772e5ca3a94fb7373e210b054ab4b571a0ed8a1 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 21 May 2026 19:32:10 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20Agent=20=E8=BF=90=E8=A1=8C=E6=97=B6?= =?UTF-8?q?=E6=9E=B6=E6=9E=84=E9=87=8D=E6=9E=84=20-=20=E7=A7=BB=E9=99=A4?= =?UTF-8?q?=20RuntimeConfigMiddleware=EF=BC=8C=E7=BB=9F=E4=B8=80=E4=B8=8A?= =?UTF-8?q?=E4=B8=8B=E6=96=87=E5=87=86=E5=A4=87=E4=B8=8E=E5=B7=A5=E5=85=B7?= =?UTF-8?q?=E8=A7=A3=E6=9E=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 删除 runtime_config_middleware.py,将功能合并到: - prepare_agent_runtime_context(): 统一上下文准备入口 - resolve_configured_runtime_tools(): 运行时工具解析 - context.py: 规范化函数重命名 (_names → _keys),重构 config 加载流程 - chatbot/deep_agent graph: 集成新上下文准备流程,移除 RuntimeConfigMiddleware 引用 - skills_middleware: 抽取 normalize_string_list 为共享工具函数 - subagent_service: get_subagents_from_names → get_subagents_from_slugs,支持 slug 查询 - 对应更新 repositories/services/routers 适配新接口 --- .../package/yuxi/agents/backends/composite.py | 17 +- .../agents/backends/knowledge_base_backend.py | 4 +- .../yuxi/agents/backends/sandbox/backend.py | 8 +- backend/package/yuxi/agents/base.py | 20 +- .../yuxi/agents/buildin/chatbot/graph.py | 17 +- .../yuxi/agents/buildin/chatbot/prompt.py | 7 +- .../yuxi/agents/buildin/deep_agent/graph.py | 21 ++- backend/package/yuxi/agents/context.py | 171 ++++++++++++++---- .../yuxi/agents/middlewares/__init__.py | 2 - .../middlewares/knowledge_base_middleware.py | 13 +- .../middlewares/runtime_config_middleware.py | 161 ----------------- .../agents/middlewares/skills_middleware.py | 84 +++------ .../agents/toolkits/buildin/install_skill.py | 11 +- .../package/yuxi/agents/toolkits/kbs/tools.py | 90 ++++----- .../repositories/evaluation_repository.py | 8 +- .../repositories/knowledge_base_repository.py | 12 +- .../knowledge_chunk_repository.py | 34 ++-- .../repositories/knowledge_file_repository.py | 8 +- .../knowledge_graph_repository.py | 22 +-- .../repositories/mcp_server_repository.py | 24 +-- .../yuxi/repositories/subagent_repository.py | 16 +- .../yuxi/services/filesystem_service.py | 13 +- backend/package/yuxi/services/mcp_service.py | 164 +++++++++-------- .../package/yuxi/services/skill_service.py | 50 ++--- .../package/yuxi/services/subagent_service.py | 51 +++--- backend/package/yuxi/services/tool_service.py | 56 +++++- .../services/viewer_filesystem_service.py | 10 +- .../package/yuxi/storage/postgres/manager.py | 52 +++--- .../yuxi/storage/postgres/models_business.py | 16 +- backend/server/routers/chat_router.py | 8 +- backend/server/routers/dashboard_router.py | 2 +- .../server/routers/knowledge_eval_router.py | 48 ++--- backend/server/routers/mcp_router.py | 78 ++++---- backend/server/routers/subagent_router.py | 50 +++-- backend/server/routers/tool_router.py | 2 +- 35 files changed, 677 insertions(+), 673 deletions(-) delete mode 100644 backend/package/yuxi/agents/middlewares/runtime_config_middleware.py diff --git a/backend/package/yuxi/agents/backends/composite.py b/backend/package/yuxi/agents/backends/composite.py index ded2a3b6..0ac1e37d 100644 --- a/backend/package/yuxi/agents/backends/composite.py +++ b/backend/package/yuxi/agents/backends/composite.py @@ -8,7 +8,7 @@ from deepagents.backends.composite import ( ) from deepagents.backends.protocol import FileInfo -from yuxi.agents.middlewares.skills_middleware import normalize_selected_skills +from yuxi.services.skill_service import normalize_string_list from .sandbox import ProvisionerSandboxBackend from .skills_backend import SelectedSkillsReadonlyBackend @@ -67,13 +67,10 @@ class CustomCompositeBackend(CompositeBackend): return await self.default.aglob_info(pattern, path) -def _get_visible_skills_from_runtime(runtime) -> list[str]: - """获取运行时可见的 skills 列表""" +def _get_readable_skills_from_runtime(runtime) -> list[str]: context = getattr(runtime, "context", None) - selected = getattr(context, "_visible_skills", None) - if not isinstance(selected, list): - selected = getattr(context, "skills", None) or [] - return normalize_selected_skills(selected) + selected = getattr(context, "_readable_skills", []) + return normalize_string_list(selected if isinstance(selected, list) else []) def _extract_thread_id(runtime) -> str: @@ -111,12 +108,12 @@ def _extract_uid(runtime) -> str: def create_agent_composite_backend(runtime) -> CompositeBackend: - visible_skills = _get_visible_skills_from_runtime(runtime) + readable_skills = _get_readable_skills_from_runtime(runtime) thread_id = _extract_thread_id(runtime) uid = _extract_uid(runtime) return CustomCompositeBackend( - default=ProvisionerSandboxBackend(thread_id=thread_id, uid=uid, visible_skills=visible_skills), + default=ProvisionerSandboxBackend(thread_id=thread_id, uid=uid, readable_skills=readable_skills), routes={ - "/skills/": SelectedSkillsReadonlyBackend(selected_slugs=visible_skills), + "/skills/": SelectedSkillsReadonlyBackend(selected_slugs=readable_skills), }, ) diff --git a/backend/package/yuxi/agents/backends/knowledge_base_backend.py b/backend/package/yuxi/agents/backends/knowledge_base_backend.py index b2aedb0b..4825f820 100644 --- a/backend/package/yuxi/agents/backends/knowledge_base_backend.py +++ b/backend/package/yuxi/agents/backends/knowledge_base_backend.py @@ -15,8 +15,8 @@ async def resolve_visible_knowledge_bases_for_context(context) -> list[dict[str, databases = result.get("databases") or [] enabled_knowledges = getattr(context, "knowledges", None) if enabled_knowledges is not None: - enabled_names = {str(name).strip() for name in enabled_knowledges if str(name).strip()} - databases = [db for db in databases if str(db.get("name") or "").strip() in enabled_names] + enabled_ids = {str(value).strip() for value in enabled_knowledges if str(value).strip()} + databases = [db for db in databases if str(db.get("kb_id") or "").strip() in enabled_ids] setattr(context, "_visible_knowledge_bases", databases) return databases diff --git a/backend/package/yuxi/agents/backends/sandbox/backend.py b/backend/package/yuxi/agents/backends/sandbox/backend.py index b1e45a24..f0d92987 100644 --- a/backend/package/yuxi/agents/backends/sandbox/backend.py +++ b/backend/package/yuxi/agents/backends/sandbox/backend.py @@ -17,7 +17,7 @@ from deepagents.backends.protocol import ( from deepagents.backends.sandbox import BaseSandbox from yuxi import config as conf -from yuxi.services.skill_service import sync_thread_visible_skills +from yuxi.services.skill_service import sync_thread_readable_skills from yuxi.utils.logging_config import logger from .provider import get_sandbox_provider, sandbox_id_for_thread @@ -62,7 +62,7 @@ def _looks_like_binary(content: bytes) -> bool: class ProvisionerSandboxBackend(BaseSandbox): - def __init__(self, thread_id: str, *, uid: str, visible_skills: list[str] | None = None): + def __init__(self, thread_id: str, *, uid: str, readable_skills: list[str] | None = None): self._thread_id = str(thread_id or "").strip() if not self._thread_id: raise ValueError("thread_id is required for ProvisionerSandboxBackend") @@ -70,7 +70,7 @@ class ProvisionerSandboxBackend(BaseSandbox): if not self._uid: raise ValueError("uid is required for ProvisionerSandboxBackend") - self._visible_skills = list(visible_skills or []) + self._readable_skills = list(readable_skills or []) self._provider = get_sandbox_provider() self._id = sandbox_id_for_thread(self._thread_id) self._client: Any | None = None @@ -93,7 +93,7 @@ class ProvisionerSandboxBackend(BaseSandbox): return AgentSandboxClient(base_url=sandbox_url, timeout=self._command_timeout_seconds) def _get_client(self) -> Any: - sync_thread_visible_skills(self._thread_id, self._visible_skills) + sync_thread_readable_skills(self._thread_id, self._readable_skills) connection = self._provider.get(self._thread_id, uid=self._uid, create_if_missing=True) if connection is None: raise RuntimeError(f"sandbox is unavailable for thread {self._thread_id}") diff --git a/backend/package/yuxi/agents/base.py b/backend/package/yuxi/agents/base.py index 30a01b36..d71ca2d9 100644 --- a/backend/package/yuxi/agents/base.py +++ b/backend/package/yuxi/agents/base.py @@ -9,7 +9,7 @@ from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite from langgraph.graph.state import CompiledStateGraph from yuxi import config as sys_config -from yuxi.agents.context import BaseContext +from yuxi.agents.context import BaseContext, resolve_agent_resource_options from yuxi.storage.postgres.manager import pg_manager from yuxi.utils import logger @@ -41,12 +41,28 @@ class BaseAgent: """Get the agent's class name.""" return self.__class__.__name__ - async def get_info(self, include_configurable_items: bool = True, user_role: str | None = None): + async def get_info( + self, + include_configurable_items: bool = True, + user_role: str | None = None, + db=None, + user=None, + ): # metadata 固定在代码中,由各 Agent 的类属性提供 metadata = self.load_metadata() configurable_items = {} if include_configurable_items: configurable_items = self.context_schema.get_configurable_items(user_role=user_role) + if db is not None and user is not None: + resource_fields = { + item["kind"] + for item in configurable_items.values() + if item.get("kind") in {"tools", "knowledges", "mcps", "skills", "subagents"} + } + resource_options = await resolve_agent_resource_options(resource_fields, db=db, user=user) + for item in configurable_items.values(): + if item.get("kind") in resource_options: + item["options"] = resource_options[item["kind"]] # Merge metadata with class attributes, metadata takes precedence return { diff --git a/backend/package/yuxi/agents/buildin/chatbot/graph.py b/backend/package/yuxi/agents/buildin/chatbot/graph.py index d7dcdc2d..158ebf6f 100644 --- a/backend/package/yuxi/agents/buildin/chatbot/graph.py +++ b/backend/package/yuxi/agents/buildin/chatbot/graph.py @@ -6,23 +6,21 @@ from langchain.agents.middleware import ModelRetryMiddleware, TodoListMiddleware from yuxi.agents import BaseAgent, BaseState, load_chat_model from yuxi.agents.backends import create_agent_composite_backend +from yuxi.agents.context import prepare_agent_runtime_context from yuxi.agents.middlewares import ( - RuntimeConfigMiddleware, SummaryOffloadMiddleware, save_attachments_to_fs, ) from yuxi.agents.middlewares.knowledge_base_middleware import KnowledgeBaseMiddleware from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware -from yuxi.services.mcp_service import get_tools_from_all_servers -from yuxi.services.subagent_service import get_subagents_from_names +from yuxi.services.subagent_service import get_subagents_from_slugs +from yuxi.services.tool_service import resolve_configured_runtime_tools from .prompt import TODO_MID_PROMPT, build_prompt_with_context async def _build_middlewares(context): """构建中间件列表""" - all_mcp_tools = await get_tools_from_all_servers() # 因为异步加载,无法放在 RuntimeConfigMiddleware 的 __init__ 中 - # summary middleware # 主 Agent 上下文优化:90k tokens 触发压缩(128k context window 的 70%) summary_middleware = SummaryOffloadMiddleware( @@ -34,7 +32,7 @@ async def _build_middlewares(context): ) # subagents - subagents = await get_subagents_from_names(context.subagents) + subagents = await get_subagents_from_slugs(context.subagents) subagents_middleware = SubAgentMiddleware( default_model=load_chat_model(fully_specified_name=context.subagents_model), subagents=subagents, @@ -50,7 +48,6 @@ async def _build_middlewares(context): FilesystemMiddleware(backend=create_agent_composite_backend), # 文件系统后端 save_attachments_to_fs, # 附件注入提示词 KnowledgeBaseMiddleware(), # 知识库工具 - RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/MCP/提示词) SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活) subagents_middleware, summary_middleware, @@ -80,11 +77,15 @@ class ChatbotAgent(BaseAgent): async def get_graph(self, context=None, **kwargs): - context = context or self.context_schema() # 获取上下文配置 + context = await prepare_agent_runtime_context( + context or self.context_schema(), + context_schema=self.context_schema, + ) # 使用 create_agent 创建智能体 graph = create_agent( model=load_chat_model(fully_specified_name=context.model), + tools=await resolve_configured_runtime_tools(context), system_prompt=build_prompt_with_context(context), middleware=await _build_middlewares(context), state_schema=BaseState, diff --git a/backend/package/yuxi/agents/buildin/chatbot/prompt.py b/backend/package/yuxi/agents/buildin/chatbot/prompt.py index 213970f6..24931012 100644 --- a/backend/package/yuxi/agents/buildin/chatbot/prompt.py +++ b/backend/package/yuxi/agents/buildin/chatbot/prompt.py @@ -1,3 +1,4 @@ +from yuxi.utils.datetime_utils import shanghai_now from yuxi.utils.paths import ( VIRTUAL_PATH_OUTPUTS, VIRTUAL_PATH_PREFIX, @@ -23,7 +24,8 @@ PROMPT = f""" - {VIRTUAL_PATH_WORKSPACE}:用于存放用户文件(用户私人目录,除非用户要求,否则不得写入) - 其他路径:非必要不写入其他路径 -非必要不写入其他路径 +<| 风格规范 |> +保持专业严谨,减少使用 Emoji """ # 效果不好,暂时不启用 @@ -48,5 +50,6 @@ TODO_MID_PROMPT = """ def build_prompt_with_context(context): - system_prompt = f"{PROMPT.strip()}\n\n{context.system_prompt or ''}" + current_date = f"当前日期:{shanghai_now().strftime('%Y-%m-%d')}" + system_prompt = f"{current_date}\n\n{PROMPT.strip()}\n\n{context.system_prompt or ''}" return system_prompt.strip() diff --git a/backend/package/yuxi/agents/buildin/deep_agent/graph.py b/backend/package/yuxi/agents/buildin/deep_agent/graph.py index ca3c6ed4..a2bdf0af 100644 --- a/backend/package/yuxi/agents/buildin/deep_agent/graph.py +++ b/backend/package/yuxi/agents/buildin/deep_agent/graph.py @@ -11,17 +11,18 @@ from langchain.agents.middleware import ( from yuxi.agents import BaseAgent, BaseState, load_chat_model from yuxi.agents.backends import create_agent_composite_backend +from yuxi.agents.context import prepare_agent_runtime_context from yuxi.agents.middlewares import ( - RuntimeConfigMiddleware, SummaryOffloadMiddleware, save_attachments_to_fs, ) from yuxi.agents.middlewares.knowledge_base_middleware import KnowledgeBaseMiddleware from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware from yuxi.agents.toolkits.buildin.tools import _create_tavily_search -from yuxi.services.mcp_service import get_tools_from_all_servers -from yuxi.services.subagent_service import get_subagents_from_names +from yuxi.services.subagent_service import get_subagents_from_slugs +from yuxi.services.tool_service import resolve_configured_runtime_tools from yuxi.utils import logger +from yuxi.utils.datetime_utils import shanghai_now from .prompt import DEEP_PROMPT @@ -51,17 +52,19 @@ class DeepAgent(BaseAgent): async def get_graph(self, context=None, **kwargs): - context = context or self.context_schema() # 获取上下文配置 - system_prompt = f"{DEEP_PROMPT.strip()}\n\n{context.system_prompt or ''}" + context = await prepare_agent_runtime_context( + context or self.context_schema(), + context_schema=self.context_schema, + ) + current_date = f"当前日期:{shanghai_now().strftime('%Y-%m-%d')}" + system_prompt = f"{current_date}\n\n{DEEP_PROMPT.strip()}\n\n{context.system_prompt or ''}" model = load_chat_model(context.model) sub_model = load_chat_model(context.subagents_model) search_tools = await self.get_tools() - all_mcp_tools = await get_tools_from_all_servers() - # 合并搜索工具和 MCP 工具 # 从数据库加载 subagent specs(工具名称已解析) - user_subagents = await get_subagents_from_names(context.subagents) + user_subagents = await get_subagents_from_slugs(context.subagents) # 主 Agent 上下文优化:90k tokens 触发压缩(128k context window 的 70%) summary_middleware = SummaryOffloadMiddleware( @@ -93,10 +96,10 @@ class DeepAgent(BaseAgent): # 使用 create_deep_agent 创建深度智能体 graph = create_agent( model=model, + tools=await resolve_configured_runtime_tools(context), system_prompt=system_prompt, middleware=[ FilesystemMiddleware(backend=create_agent_composite_backend), # 文件系统后端 - RuntimeConfigMiddleware(extra_tools=all_mcp_tools), SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活) save_attachments_to_fs, # 附件注入提示词 TodoListMiddleware(system_prompt="任务结束前,应该检查维护的待办事项列表是否结束。"), diff --git a/backend/package/yuxi/agents/context.py b/backend/package/yuxi/agents/context.py index 44248d22..4aceb8b5 100644 --- a/backend/package/yuxi/agents/context.py +++ b/backend/package/yuxi/agents/context.py @@ -215,7 +215,7 @@ class BaseContext: _DEFAULT_ALL_CONTEXT_FIELDS = frozenset({"tools", "knowledges", "mcps", "skills", "subagents"}) -def _normalize_selected_resource_names(value: Any, available: list[str]) -> list[str]: +def _normalize_selected_resource_keys(value: Any, available: list[str]) -> list[str]: if not isinstance(value, list): return [] @@ -225,15 +225,15 @@ def _normalize_selected_resource_names(value: Any, available: list[str]) -> list for item in value: if not isinstance(item, str): continue - name = item.strip() - if not name or name in seen or name not in allowed: + key = item.strip() + if not key or key in seen or key not in allowed: continue - seen.add(name) - normalized.append(name) + seen.add(key) + normalized.append(key) return normalized -def _resource_fields_requiring_available_names(normalized: dict, resource_fields: set[str]) -> set[str]: +def _resource_fields_requiring_available_keys(normalized: dict, resource_fields: set[str]) -> set[str]: fields_to_load: set[str] = set() for field_name in resource_fields: current = normalized.get(field_name) @@ -246,6 +246,73 @@ def _resource_fields_requiring_available_names(normalized: dict, resource_fields return fields_to_load +def _resource_option(key: Any, name: Any = None, description: Any = None) -> dict[str, str]: + key_value = str(key) + return { + "key": key_value, + "name": str(name or key_value), + "description": str(description or ""), + } + + +async def resolve_agent_resource_options( + resource_fields: set[str] | None = None, + *, + db, + user, +) -> dict[str, list[dict[str, str]]]: + fields_to_load = _DEFAULT_ALL_CONTEXT_FIELDS if resource_fields is None else resource_fields + if not fields_to_load: + return {} + + options: dict[str, list[dict[str, str]]] = {} + + if "tools" in fields_to_load: + from yuxi.services.tool_service import get_tool_metadata + + options["tools"] = [ + _resource_option(tool["slug"], tool.get("name"), tool.get("description")) + for tool in get_tool_metadata(category="buildin") + if tool.get("slug") + ] + if "knowledges" in fields_to_load: + from yuxi.knowledge import knowledge_base + + databases = (await knowledge_base.get_databases_by_user(user)).get("databases", []) + options["knowledges"] = [ + _resource_option(item.get("kb_id"), item.get("name"), item.get("description")) + for item in databases + if isinstance(item, dict) and item.get("kb_id") + ] + if "mcps" in fields_to_load: + from yuxi.services.mcp_service import get_all_mcp_servers + + servers = await get_all_mcp_servers(db) + options["mcps"] = [ + _resource_option(server.slug, server.name, server.description) + for server in servers + if server.enabled and server.slug + ] + if "skills" in fields_to_load: + from yuxi.services.skill_service import list_skills + + skills = await list_skills(db) + options["skills"] = [ + _resource_option(skill.slug, skill.name, skill.description) for skill in skills if skill.slug + ] + if "subagents" in fields_to_load: + from yuxi.services.subagent_service import get_all_subagents + + subagents = await get_all_subagents(db) + options["subagents"] = [ + _resource_option(item.get("slug"), item.get("name"), item.get("description")) + for item in subagents + if item.get("enabled") and item.get("slug") + ] + + return options + + async def normalize_agent_context_config( context: dict | None, *, @@ -262,44 +329,72 @@ async def normalize_agent_context_config( if not resource_fields: return normalized - fields_to_load = _resource_fields_requiring_available_names(normalized, resource_fields) + fields_to_load = _resource_fields_requiring_available_keys(normalized, resource_fields) if not fields_to_load: return normalized - available: dict[str, list[str]] = {} - if "tools" in fields_to_load: - from yuxi.agents.toolkits import get_all_tool_instances + resource_options = await resolve_agent_resource_options(fields_to_load, db=db, user=user) + available = { + field_name: [option["key"] for option in field_options] + for field_name, field_options in resource_options.items() + } - available["tools"] = [ - tool.name for tool in get_all_tool_instances() if isinstance(getattr(tool, "name", None), str) - ] - if "knowledges" in fields_to_load: - from yuxi.knowledge import knowledge_base - - databases = (await knowledge_base.get_databases_by_user(user)).get("databases", []) - available["knowledges"] = [ - str(db_item.get("db_id") or db_item.get("id")) - for db_item in databases - if isinstance(db_item, dict) and (db_item.get("db_id") or db_item.get("id")) - ] - if "mcps" in fields_to_load: - from yuxi.services.mcp_service import get_enabled_mcp_server_names - - available["mcps"] = await get_enabled_mcp_server_names(db=db) - if "skills" in fields_to_load: - from yuxi.services.skill_service import list_skill_slugs - - available["skills"] = await list_skill_slugs(db) - if "subagents" in fields_to_load: - from yuxi.services.subagent_service import get_enabled_subagent_names - - available["subagents"] = await get_enabled_subagent_names(db) - - for field_name, available_names in available.items(): + for field_name, available_keys in available.items(): current = normalized.get(field_name) if current is None: - normalized[field_name] = available_names + normalized[field_name] = available_keys else: - normalized[field_name] = _normalize_selected_resource_names(current, available_names) + normalized[field_name] = _normalize_selected_resource_keys(current, available_keys) return normalized + + +async def prepare_agent_runtime_context( + context: BaseContext, + *, + context_schema: type[BaseContext] | None = None, +) -> BaseContext: + """准备 Agent 运行时上下文,主要是根据 context 中的 uid 加载用户可访问的资源列表,并进行规范化处理。""" + uid = str(getattr(context, "uid", "") or "").strip() + if not uid: + return context + + from yuxi.agents.backends.knowledge_base_backend import resolve_visible_knowledge_bases_for_context + from yuxi.agents.middlewares.skills_middleware import resolve_runtime_skills_for_context + from yuxi.repositories.user_repository import UserRepository + from yuxi.storage.postgres.manager import pg_manager + + resource_fields = _DEFAULT_ALL_CONTEXT_FIELDS + async with pg_manager.get_async_session_context() as db: + user = await UserRepository().get_by_uid_with_db(db, uid) + if user is None: + for field_name in resource_fields: + if hasattr(context, field_name): + setattr(context, field_name, []) + setattr(context, "_visible_knowledge_bases", []) + setattr(context, "_prompt_skills", []) + setattr(context, "_readable_skills", []) + return context + + raw_resources = { + field_name: getattr(context, field_name, None) + for field_name in resource_fields + if hasattr(context, field_name) + } + normalized = await normalize_agent_context_config( + raw_resources, + db=db, + user=user, + context_schema=context_schema, + ) + for field_name in resource_fields: + if hasattr(context, field_name): + setattr(context, field_name, normalized.get(field_name, [])) + + await resolve_visible_knowledge_bases_for_context(context) + skill_scope = await resolve_runtime_skills_for_context(context, db=db) + context.skills = skill_scope["context_skills"] + setattr(context, "_prompt_skills", skill_scope["prompt_skills"]) + setattr(context, "_readable_skills", skill_scope["readable_skills"]) + + return context diff --git a/backend/package/yuxi/agents/middlewares/__init__.py b/backend/package/yuxi/agents/middlewares/__init__.py index a59c7d69..e77515c8 100644 --- a/backend/package/yuxi/agents/middlewares/__init__.py +++ b/backend/package/yuxi/agents/middlewares/__init__.py @@ -1,12 +1,10 @@ from .attachment_middleware import inject_attachment_context, save_attachments_to_fs from .context_middlewares import context_aware_prompt, context_based_model from .dynamic_tool_middleware import DynamicToolMiddleware -from .runtime_config_middleware import RuntimeConfigMiddleware from .summary_middleware import SummaryOffloadMiddleware, create_summary_offload_middleware __all__ = [ "DynamicToolMiddleware", - "RuntimeConfigMiddleware", "SummaryOffloadMiddleware", "context_aware_prompt", "context_based_model", diff --git a/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py b/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py index 285e94b2..b3c3ad8e 100644 --- a/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py +++ b/backend/package/yuxi/agents/middlewares/knowledge_base_middleware.py @@ -1,16 +1,13 @@ """知识库中间件 - 提供通用知识库工具""" -from collections.abc import Callable +from langchain.agents.middleware import AgentMiddleware -from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse - -from yuxi.agents.backends.knowledge_base_backend import resolve_visible_knowledge_bases_for_context from yuxi.agents.toolkits.kbs import get_common_kb_tools from yuxi.utils.logging_config import logger class KnowledgeBaseMiddleware(AgentMiddleware): - """知识库中间件 - 提供通用知识库工具 + """知识库中间件 - 提供通用知识库工具,其他没有任何作用 提供通用知识库工具: - list_kbs: 列出用户可访问的知识库 @@ -26,9 +23,3 @@ class KnowledgeBaseMiddleware(AgentMiddleware): self.kb_tools = get_common_kb_tools() self.tools = self.kb_tools logger.debug(f"Initialized KnowledgeBaseMiddleware with {len(self.kb_tools)} tools") - - async def awrap_model_call( - self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] - ) -> ModelResponse: - await resolve_visible_knowledge_bases_for_context(request.runtime.context) - return await handler(request) diff --git a/backend/package/yuxi/agents/middlewares/runtime_config_middleware.py b/backend/package/yuxi/agents/middlewares/runtime_config_middleware.py deleted file mode 100644 index a1b6e98f..00000000 --- a/backend/package/yuxi/agents/middlewares/runtime_config_middleware.py +++ /dev/null @@ -1,161 +0,0 @@ -from __future__ import annotations - -from collections.abc import Callable -from typing import Any - -from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse -from langchain_core.messages import SystemMessage - -from yuxi.agents import load_chat_model -from yuxi.agents.toolkits import get_all_tool_instances -from yuxi.services.mcp_service import get_enabled_mcp_tools -from yuxi.utils.datetime_utils import shanghai_now -from yuxi.utils.logging_config import logger - - -class RuntimeConfigMiddleware(AgentMiddleware): - """运行时配置中间件 - 应用模型/工具/MCP/提示词配置 - - 知识库工具已移至独立的 KnowledgeBaseMiddleware - Skills 功能已移至独立的 SkillsMiddleware - - 支持自定义上下文字段名称,以便在不同场景(如主智能体/子智能体)使用不同的配置字段 - """ - - def __init__( - self, - *, - extra_tools: list[Any] | None = None, - model_context_name: str = "model", - system_prompt_context_name: str = "system_prompt", - tools_context_name: str = "tools", - knowledges_context_name: str = "knowledges", - mcps_context_name: str = "mcps", - enable_model_override: bool = True, - enable_system_prompt_override: bool = True, - enable_tools_override: bool = True, - ): - """初始化中间件 - - Args: - extra_tools: 额外工具列表(从 create_agent 的 tools 参数传入) - model_context_name: 上下文中的模型字段名称(默认 "model") - system_prompt_context_name: 上下文中的系统提示词字段名称(默认 "system_prompt") - tools_context_name: 上下文中的工具列表字段名称(默认 "tools") - knowledges_context_name: 上下文中的知识库列表字段名称(默认 "knowledges") - mcps_context_name: 上下文中的 MCP 服务器列表字段名称(默认 "mcps") - enable_model_override: 是否允许覆盖模型配置(默认 True) - enable_system_prompt_override: 是否允许覆盖系统提示词(默认 True) - enable_tools_override: 是否允许覆盖工具列表(默认 True) - """ - super().__init__() - # 存储自定义字段名称 - self.model_context_name = model_context_name - self.system_prompt_context_name = system_prompt_context_name - self.tools_context_name = tools_context_name - self.knowledges_context_name = knowledges_context_name - self.mcps_context_name = mcps_context_name - # 存储覆盖配置 - self.enable_model_override = enable_model_override - self.enable_system_prompt_override = enable_system_prompt_override - self.enable_tools_override = enable_tools_override - - self.tools: list[Any] = [] - # 预加载工具列表(仅当启用工具覆盖时) - # 注意:知识库工具已移至独立的 KnowledgeBaseMiddleware - if self.enable_tools_override: - self.base_tools = get_all_tool_instances() - self.tools = self.base_tools + (extra_tools or []) - elif extra_tools: - logger.warning( - "RuntimeConfigMiddleware: extra_tools 参数已提供,但 enable_tools_override=False," - "将忽略 extra_tools 并不会应用任何工具覆盖。" - ) - - async def awrap_model_call( - self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] - ) -> ModelResponse: - runtime_context = request.runtime.context - overrides: dict[str, Any] = {} - - # 1. 模型覆盖(可选) - if self.enable_model_override: - model = load_chat_model(getattr(runtime_context, self.model_context_name, None)) - overrides["model"] = model - - # 2. 工具覆盖(可选) - # 注意:Skills 依赖的工具加载已移至 SkillsMiddleware - if self.enable_tools_override: - # 获取上下文配置的工具 - enabled_tools = await self.get_tools_from_context(runtime_context) - existing_tools = list(request.tools or []) - enabled_tool_names = {t.name for t in enabled_tools} - managed_tool_names = {t.name for t in self.tools} - merged_tools = [] - for t_bind in existing_tools: - # (1) 已启用的工具保留 - # (2) 非本中间件管理的工具保留 - if t_bind.name in enabled_tool_names or t_bind.name not in managed_tool_names: - merged_tools.append(t_bind) - overrides["tools"] = merged_tools - logger.debug(f"RuntimeConfigMiddleware selected tools: {[t.name for t in merged_tools]}") - - # 3. 系统提示词覆盖(可选) - if self.enable_system_prompt_override: - cur_datetime = f"当前时间:{shanghai_now().strftime('%Y-%m-%d %H:%M:%S')} UTC" - system_prompt = getattr(runtime_context, self.system_prompt_context_name, "") or "" - merged_system_prompt = f"{cur_datetime}\n\n{system_prompt}" - - content_blocks = list(request.system_message.content_blocks) if request.system_message else [] - new_content = content_blocks + [{"type": "text", "text": merged_system_prompt}] - new_system_message = SystemMessage(content=new_content) - overrides["system_message"] = new_system_message - - if overrides: - request = request.override(**overrides) - - return await handler(request) - - async def get_tools_from_context(self, context) -> list: - """从上下文配置中获取工具列表""" - selected_tools = [] - selected_tool_names: set[str] = set() - - # 1. 基础工具 (从 context.tools 中筛选) - tools = getattr(context, self.tools_context_name, None) or [] - all_tool_names = [] - for tool_name in tools: - if isinstance(tool_name, str): - all_tool_names.append(tool_name) - - tools_map = {t.name: t for t in self.tools} - for tool_name in all_tool_names: - if tool_name in selected_tool_names: - continue - if tool_name in tools_map: - selected_tools.append(tools_map[tool_name]) - selected_tool_names.add(tool_name) - continue - logger.warning(f"RuntimeConfigMiddleware: tool dependency not found, skip: {tool_name}") - - # 2. MCP 工具(使用统一入口,自动过滤 disabled_tools) - mcps = getattr(context, self.mcps_context_name, None) or [] - all_mcp_names: list[str] = [] - for server_name in mcps: - if isinstance(server_name, str): - all_mcp_names.append(server_name) - - selected_mcp_servers: set[str] = set() - for server_name in all_mcp_names: - if server_name in selected_mcp_servers: - continue - selected_mcp_servers.add(server_name) - try: - mcp_tools = await get_enabled_mcp_tools(server_name) - if not mcp_tools: - logger.warning(f"RuntimeConfigMiddleware: mcp dependency unavailable, skip: {server_name}") - selected_tools.extend(mcp_tools) - except Exception as e: - logger.warning(f"RuntimeConfigMiddleware: failed to load mcp dependency '{server_name}': {e}") - - return selected_tools diff --git a/backend/package/yuxi/agents/middlewares/skills_middleware.py b/backend/package/yuxi/agents/middlewares/skills_middleware.py index 21a1d754..2cc5f29f 100644 --- a/backend/package/yuxi/agents/middlewares/skills_middleware.py +++ b/backend/package/yuxi/agents/middlewares/skills_middleware.py @@ -16,7 +16,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from yuxi.agents.toolkits import get_all_tool_instances from yuxi.repositories.skill_repository import SkillRepository from yuxi.services.mcp_service import get_enabled_mcp_tools -from yuxi.services.skill_service import _normalize_string_list, is_valid_skill_slug +from yuxi.services.skill_service import is_valid_skill_slug, normalize_string_list from yuxi.storage.postgres.manager import pg_manager from yuxi.utils.logging_config import logger @@ -72,24 +72,19 @@ async def get_dependency_map(db: AsyncSession | None = None) -> dict[str, SkillD result: dict[str, SkillDependencyNode] = {} for item in skills: result[item.slug] = { - "tools": normalize_selected_skills(item.tool_dependencies or []), - "mcps": normalize_selected_skills(item.mcp_dependencies or []), - "skills": normalize_selected_skills(item.skill_dependencies or []), + "tools": normalize_string_list(item.tool_dependencies or []), + "mcps": normalize_string_list(item.mcp_dependencies or []), + "skills": normalize_string_list(item.skill_dependencies or []), } return result -def normalize_selected_skills(selected_skills: list[str] | None) -> list[str]: - """规范化 skills 列表,去重并过滤无效值""" - return _normalize_string_list(selected_skills) - - def expand_skill_closure( slugs: list[str] | None, dependency_map: dict[str, SkillDependencyNode], ) -> list[str]: """展开 skills 依赖闭包,返回包含所有依赖的列表""" - ordered_roots = normalize_selected_skills(slugs) + ordered_roots = normalize_string_list(slugs) if not ordered_roots: return [] @@ -120,6 +115,19 @@ def expand_skill_closure( return result +async def resolve_runtime_skills_for_context(context, *, db: AsyncSession | None = None) -> dict[str, list[str]]: + dependency_map = await get_dependency_map(db) + installed = set(dependency_map) + selected = normalize_string_list(getattr(context, "skills", None)) + context_skills = [slug for slug in selected if slug in installed] + prompt_skills = expand_skill_closure(context_skills, dependency_map) + return { + "context_skills": context_skills, + "prompt_skills": prompt_skills, + "readable_skills": prompt_skills, + } + + def _activated_skills_reducer(left: list[str] | None, right: list[str] | None) -> list[str]: """合并 activated_skills 列表""" merged: list[str] = [] @@ -182,24 +190,16 @@ class SkillsMiddleware(AgentMiddleware): if getattr(runtime_context, "_skills_prompt_injected", False): return None - # 从数据库加载 skills 数据(使用缓存) - dependency_map = await get_dependency_map() - - # 获取配置的 skills - configured_skills = getattr(runtime_context, self.skills_context_name, None) or [] - selected_skills = normalize_selected_skills(configured_skills) - - if not selected_skills: + prompt_skills = getattr(runtime_context, "_prompt_skills", None) + if not isinstance(prompt_skills, list): return None - # 计算 visible_skills - visible_skills = expand_skill_closure(selected_skills, dependency_map) - - if not visible_skills: + prompt_skills = normalize_string_list(prompt_skills) + if not prompt_skills: return None # 收集提示词元数据并构建提示段 - skills_meta = await self._collect_prompt_metadata(visible_skills) + skills_meta = await self._collect_prompt_metadata(prompt_skills) skills_section = self._build_skills_section(skills_meta) # 注入提示词 @@ -208,9 +208,6 @@ class SkillsMiddleware(AgentMiddleware): setattr(runtime_context, "system_prompt", merged_prompt) setattr(runtime_context, "_skills_prompt_injected", True) - # 存储 visible_skills 供后续使用 - setattr(runtime_context, "_visible_skills", visible_skills) - return None async def awrap_model_call( @@ -219,39 +216,23 @@ class SkillsMiddleware(AgentMiddleware): """包装模型调用,处理动态激活和依赖展开""" runtime_context = request.runtime.context - # 从缓存加载 skills 数据 - dependency_map = await get_dependency_map() - - # 1. 获取配置的 skills - configured_skills = getattr(runtime_context, self.skills_context_name, None) or [] - configured = normalize_selected_skills(configured_skills) - - # 2. 获取运行时动态激活的 skills state = request.state if isinstance(request.state, dict) else {} activated = state.get("activated_skills", []) or [] if not isinstance(activated, list): activated = [] - # 3. 合并并展开闭包 - all_skills = normalize_selected_skills(configured + activated) - visible_skills = expand_skill_closure(all_skills, dependency_map) + readable_skills = self._get_readable_skills(runtime_context) + activated = [slug for slug in normalize_string_list(activated) if slug in readable_skills] - # 4. 更新 runtime_context 中的 visible_skills - setattr(runtime_context, "_visible_skills", visible_skills) - - # 5. 构建依赖包(只从直接激活的 skills 获取依赖,不包含闭包展开的依赖) deps_bundle = await self._build_dependency_bundle(activated) - # 6. 加载依赖的工具(普通工具 + MCP 工具) enabled_tools = [] - # 6.1 从 toolkits 获取普通工具 if deps_bundle["tools"]: all_tools = get_all_tool_instances() required_tool_names = set(deps_bundle["tools"]) enabled_tools = [t for t in all_tools if t.name in required_tool_names] - # 6.2 获取 MCP 工具 if deps_bundle["mcps"]: mcp_tools = await self._get_mcp_tools_from_context( runtime_context, @@ -418,18 +399,13 @@ class SkillsMiddleware(AgentMiddleware): return None return slug + def _get_readable_skills(self, runtime_context) -> set[str]: + selected = getattr(runtime_context, "_readable_skills", []) + return set(normalize_string_list(selected if isinstance(selected, list) else [])) + def _is_visible_skill_slug(self, request: ToolCallRequest, slug: str) -> bool: """检查 slug 是否可见""" - runtime_context = request.runtime.context - visible_skills = getattr(runtime_context, "_visible_skills", None) - - if isinstance(visible_skills, list): - return slug in visible_skills - - # 后备:从配置的 skills 检查 - configured_skills = getattr(runtime_context, self.skills_context_name, None) or [] - normalized = normalize_selected_skills(configured_skills) - return slug in normalized + return slug in self._get_readable_skills(request.runtime.context) def _merge_activated_skill_update(self, result: Any, slug: str): """合并动态激活的 skill 更新""" diff --git a/backend/package/yuxi/agents/toolkits/buildin/install_skill.py b/backend/package/yuxi/agents/toolkits/buildin/install_skill.py index 6ad0b448..98662bd7 100644 --- a/backend/package/yuxi/agents/toolkits/buildin/install_skill.py +++ b/backend/package/yuxi/agents/toolkits/buildin/install_skill.py @@ -145,9 +145,12 @@ async def _run_install_task( skill_names: list[str] | None = None, ) -> Command: """执行异步安装任务的核心逻辑""" - from yuxi.agents.middlewares.skills_middleware import normalize_selected_skills from yuxi.services.remote_skill_install_service import prepare_remote_skills_batch - from yuxi.services.skill_service import import_skill_dir, sync_thread_visible_skills + from yuxi.services.skill_service import ( + import_skill_dir, + normalize_string_list, + sync_thread_readable_skills, + ) uid = getattr(runtime.context, "uid", None) thread_id = getattr(runtime.context, "thread_id", None) @@ -215,8 +218,8 @@ async def _run_install_task( preparation.cleanup() # 文件同步 - current_skills = normalize_selected_skills(getattr(runtime.context, "skills", None)) - sync_thread_visible_skills(thread_id, normalize_selected_skills(current_skills + installed_slugs)) + current_skills = normalize_string_list(getattr(runtime.context, "skills", None)) + sync_thread_readable_skills(thread_id, normalize_string_list(current_skills + installed_slugs)) # 响应 lines = [] diff --git a/backend/package/yuxi/agents/toolkits/kbs/tools.py b/backend/package/yuxi/agents/toolkits/kbs/tools.py index b11d5655..0375c1ae 100644 --- a/backend/package/yuxi/agents/toolkits/kbs/tools.py +++ b/backend/package/yuxi/agents/toolkits/kbs/tools.py @@ -72,7 +72,7 @@ async def list_kbs(dummy: str, runtime: ToolRuntime) -> str: # Now has 2 params for kb in available_kbs: name = kb.get("name", "") desc = kb.get("description") or "无描述" - kb_list.append({"resource_id": kb.get("db_id"), "name": name, "description": desc}) + kb_list.append({"kb_id": kb.get("kb_id"), "name": name, "description": desc}) return kb_list @@ -103,22 +103,22 @@ async def get_mindmap(kb_name: str, runtime: ToolRuntime) -> str: retrievers = knowledge_base.get_retrievers() # 查找对应的知识库 - target_db_id = None + target_kb_id = None target_info = None - for db_id, info in retrievers.items(): + for kb_id, info in retrievers.items(): if info["name"] == kb_name: - target_db_id = db_id + target_kb_id = kb_id target_info = info break - if not target_db_id: + if not target_kb_id: return f"知识库 '{kb_name}' 不存在" try: from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository kb_repo = KnowledgeBaseRepository() - kb = await kb_repo.get_by_id(target_db_id) + kb = await kb_repo.get_by_kb_id(target_kb_id) if kb is None: return f"知识库 {target_info['name']} 不存在" @@ -175,40 +175,40 @@ async def _resolve_visible_knowledge_bases_for_query(runtime: ToolRuntime | None def _find_query_target( *, - resource_id: str, + kb_id: str, retrievers: dict[str, Any], visible_kbs: list[dict[str, Any]], ) -> tuple[dict[str, Any] | None, str | None, str | None]: if not visible_kbs: return None, None, "无法获取当前会话可访问的知识库" - normalized_resource_id = str(resource_id or "").strip() - visible_resource_ids = {str(kb.get("db_id") or "").strip() for kb in visible_kbs} - if normalized_resource_id not in visible_resource_ids: - return None, None, f"知识库资源 '{normalized_resource_id}' 不存在或当前会话未启用" + normalized_kb_id = str(kb_id or "").strip() + visible_kb_ids = {str(kb.get("kb_id") or "").strip() for kb in visible_kbs} + if normalized_kb_id not in visible_kb_ids: + return None, None, f"知识库资源 '{normalized_kb_id}' 不存在或当前会话未启用" - target_info = retrievers.get(normalized_resource_id) + target_info = retrievers.get(normalized_kb_id) if target_info is None: - return None, None, f"知识库资源 '{normalized_resource_id}' 不存在" - return target_info, normalized_resource_id, None + return None, None, f"知识库资源 '{normalized_kb_id}' 不存在" + return target_info, normalized_kb_id, None @tool(category="knowledge", tags=["知识库"], args_schema=QueryKBInput) -async def query_kb(resource_id: str, query_text: str, file_name: str | None = None, runtime: ToolRuntime = None) -> Any: +async def query_kb(kb_id: str, query_text: str, file_name: str | None = None, runtime: ToolRuntime = None) -> Any: """在指定知识库中检索内容 - 当用户需要查询具体内容时使用此工具。resource_id 是知识库资源 ID,也就是 kb_id;返回结果中的 + 当用户需要查询具体内容时使用此工具。kb_id 是知识库资源 ID,也就是 kb_id;返回结果中的 file_id 可继续用于 find_kb_document 或 open_kb_document。 """ - if not resource_id: - return "请提供 resource_id" + if not kb_id: + return "请提供 kb_id" if not query_text: return "请提供查询内容" retrievers = knowledge_base.get_retrievers() visible_kbs = await _resolve_visible_knowledge_bases_for_query(runtime) - target_info, target_db_id, target_error = _find_query_target( - resource_id=resource_id, + target_info, target_kb_id, target_error = _find_query_target( + kb_id=kb_id, retrievers=retrievers, visible_kbs=visible_kbs, ) @@ -228,11 +228,11 @@ async def query_kb(resource_id: str, query_text: str, file_name: str | None = No if ( isinstance(result, dict) - and result.get("resource_id") == target_db_id + and result.get("kb_id") == target_kb_id and isinstance(result.get("results"), list) ): return SearchOutputSchema(**result).model_dump() - return KnowledgeBase.build_search_output(target_db_id, result) + return KnowledgeBase.build_search_output(target_kb_id, result) except Exception as e: logger.error(f"检索失败: {e}") @@ -241,7 +241,7 @@ async def query_kb(resource_id: str, query_text: str, file_name: str | None = No @tool(category="knowledge", tags=["知识库"], args_schema=OpenKBDocumentInput) async def open_kb_document( - resource_id: str, + kb_id: str, file_id: str, line: int | None = None, offset: int | None = None, @@ -251,12 +251,12 @@ async def open_kb_document( """按行窗口打开知识库文档原文 当 query_kb 返回的片段不足以回答问题,或需要查看某个文档的上下文时使用。 - resource_id 是知识库资源 ID,也就是 kb_id;file_id 是知识库文件 ID。 + kb_id 是知识库资源 ID,也就是 kb_id;file_id 是知识库文件 ID。 """ - normalized_resource_id = str(resource_id or "").strip() + normalized_kb_id = str(kb_id or "").strip() normalized_file_id = str(file_id or "").strip() - if not normalized_resource_id: - return "请提供 resource_id" + if not normalized_kb_id: + return "请提供 kb_id" if not normalized_file_id: return "请提供 file_id" @@ -264,14 +264,14 @@ async def open_kb_document( if not visible_kbs: return "无法获取当前会话可访问的知识库" - visible_resource_ids = {str(kb.get("db_id") or "").strip() for kb in visible_kbs} - if normalized_resource_id not in visible_resource_ids: - return f"知识库资源 '{normalized_resource_id}' 不存在或当前会话未启用" + visible_kb_ids = {str(kb.get("kb_id") or "").strip() for kb in visible_kbs} + if normalized_kb_id not in visible_kb_ids: + return f"知识库资源 '{normalized_kb_id}' 不存在或当前会话未启用" retrievers = knowledge_base.get_retrievers() - target_info = retrievers.get(normalized_resource_id) + target_info = retrievers.get(normalized_kb_id) if target_info is None: - return f"知识库资源 '{normalized_resource_id}' 不存在" + return f"知识库资源 '{normalized_kb_id}' 不存在" metadata = target_info.get("metadata") if isinstance(target_info, dict) else None kb_type = str((metadata or {}).get("kb_type") or "").strip().lower() @@ -281,12 +281,12 @@ async def open_kb_document( try: start_offset = int(line) - 1 if line is not None else int(offset or 0) window = await knowledge_base.open_file_content( - normalized_resource_id, + normalized_kb_id, normalized_file_id, offset=start_offset, limit=window_size, ) - return OpenOutputSchema(resource_id=normalized_resource_id, file_id=normalized_file_id, **window).model_dump() + return OpenOutputSchema(kb_id=normalized_kb_id, file_id=normalized_file_id, **window).model_dump() except Exception as e: logger.error(f"打开知识库文档失败: {e}") @@ -295,7 +295,7 @@ async def open_kb_document( @tool(category="knowledge", tags=["知识库"], args_schema=FindKBDocumentInput) async def find_kb_document( - resource_id: str, + kb_id: str, file_id: str, patterns: list[str], use_regex: bool = False, @@ -308,10 +308,10 @@ async def find_kb_document( 当 query_kb 已找到候选文件,但需要在该文件内定位术语、指标、章节或实体时使用。 """ - normalized_resource_id = str(resource_id or "").strip() + normalized_kb_id = str(kb_id or "").strip() normalized_file_id = str(file_id or "").strip() - if not normalized_resource_id: - return "请提供 resource_id" + if not normalized_kb_id: + return "请提供 kb_id" if not normalized_file_id: return "请提供 file_id" if not patterns: @@ -321,14 +321,14 @@ async def find_kb_document( if not visible_kbs: return "无法获取当前会话可访问的知识库" - visible_resource_ids = {str(kb.get("db_id") or "").strip() for kb in visible_kbs} - if normalized_resource_id not in visible_resource_ids: - return f"知识库资源 '{normalized_resource_id}' 不存在或当前会话未启用" + visible_kb_ids = {str(kb.get("kb_id") or "").strip() for kb in visible_kbs} + if normalized_kb_id not in visible_kb_ids: + return f"知识库资源 '{normalized_kb_id}' 不存在或当前会话未启用" retrievers = knowledge_base.get_retrievers() - target_info = retrievers.get(normalized_resource_id) + target_info = retrievers.get(normalized_kb_id) if target_info is None: - return f"知识库资源 '{normalized_resource_id}' 不存在" + return f"知识库资源 '{normalized_kb_id}' 不存在" metadata = target_info.get("metadata") if isinstance(target_info, dict) else None kb_type = str((metadata or {}).get("kb_type") or "").strip().lower() @@ -337,7 +337,7 @@ async def find_kb_document( try: result = await knowledge_base.find_file_content( - normalized_resource_id, + normalized_kb_id, normalized_file_id, patterns, use_regex=use_regex, @@ -345,7 +345,7 @@ async def find_kb_document( max_windows=max_windows, window_size=window_size, ) - return FindOutputSchema(resource_id=normalized_resource_id, file_id=normalized_file_id, **result).model_dump() + return FindOutputSchema(kb_id=normalized_kb_id, file_id=normalized_file_id, **result).model_dump() except Exception as e: logger.error(f"知识库文档内检索失败: {e}") return f"知识库文档内检索失败: {str(e)}" diff --git a/backend/package/yuxi/repositories/evaluation_repository.py b/backend/package/yuxi/repositories/evaluation_repository.py index f9c74d65..cd02d224 100644 --- a/backend/package/yuxi/repositories/evaluation_repository.py +++ b/backend/package/yuxi/repositories/evaluation_repository.py @@ -50,11 +50,11 @@ class EvaluationRepository: result = await session.execute(select(EvaluationDataset).where(EvaluationDataset.dataset_id == dataset_id)) return result.scalar_one_or_none() - async def list_datasets(self, db_id: str) -> list[EvaluationDataset]: + async def list_datasets(self, kb_id: str) -> list[EvaluationDataset]: async with pg_manager.get_async_session_context() as session: result = await session.execute( select(EvaluationDataset) - .where(EvaluationDataset.db_id == db_id) + .where(EvaluationDataset.kb_id == kb_id) .order_by(EvaluationDataset.created_at.desc()) ) return list(result.scalars().all()) @@ -106,10 +106,10 @@ class EvaluationRepository: result = await session.execute(select(EvaluationRun).where(EvaluationRun.run_id == run_id)) return result.scalar_one_or_none() - async def list_runs(self, db_id: str) -> list[EvaluationRun]: + async def list_runs(self, kb_id: str) -> list[EvaluationRun]: async with pg_manager.get_async_session_context() as session: result = await session.execute( - select(EvaluationRun).where(EvaluationRun.db_id == db_id).order_by(EvaluationRun.started_at.desc()) + select(EvaluationRun).where(EvaluationRun.kb_id == kb_id).order_by(EvaluationRun.started_at.desc()) ) return list(result.scalars().all()) diff --git a/backend/package/yuxi/repositories/knowledge_base_repository.py b/backend/package/yuxi/repositories/knowledge_base_repository.py index 4c11a74f..62d58ca2 100644 --- a/backend/package/yuxi/repositories/knowledge_base_repository.py +++ b/backend/package/yuxi/repositories/knowledge_base_repository.py @@ -14,9 +14,9 @@ class KnowledgeBaseRepository: result = await session.execute(select(KnowledgeBase)) return list(result.scalars().all()) - async def get_by_id(self, db_id: str) -> KnowledgeBase | None: + async def get_by_kb_id(self, kb_id: str) -> KnowledgeBase | None: async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(KnowledgeBase).where(KnowledgeBase.db_id == db_id)) + result = await session.execute(select(KnowledgeBase).where(KnowledgeBase.kb_id == kb_id)) return result.scalar_one_or_none() async def create(self, data: dict[str, Any]) -> KnowledgeBase: @@ -25,9 +25,9 @@ class KnowledgeBaseRepository: session.add(kb) return kb - async def update(self, db_id: str, data: dict[str, Any]) -> KnowledgeBase | None: + async def update(self, kb_id: str, data: dict[str, Any]) -> KnowledgeBase | None: async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(KnowledgeBase).where(KnowledgeBase.db_id == db_id)) + result = await session.execute(select(KnowledgeBase).where(KnowledgeBase.kb_id == kb_id)) kb = result.scalar_one_or_none() if kb is None: return None @@ -35,9 +35,9 @@ class KnowledgeBaseRepository: setattr(kb, key, value) return kb - async def delete(self, db_id: str) -> None: + async def delete(self, kb_id: str) -> None: async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(KnowledgeBase).where(KnowledgeBase.db_id == db_id)) + result = await session.execute(select(KnowledgeBase).where(KnowledgeBase.kb_id == kb_id)) kb = result.scalar_one_or_none() if kb is not None: await session.delete(kb) diff --git a/backend/package/yuxi/repositories/knowledge_chunk_repository.py b/backend/package/yuxi/repositories/knowledge_chunk_repository.py index 4a828d34..fc7113bd 100644 --- a/backend/package/yuxi/repositories/knowledge_chunk_repository.py +++ b/backend/package/yuxi/repositories/knowledge_chunk_repository.py @@ -12,7 +12,7 @@ class KnowledgeChunkRepository: _writable_fields = { "chunk_id", "file_id", - "db_id", + "kb_id", "chunk_index", "content", "start_char_pos", @@ -39,10 +39,10 @@ class KnowledgeChunkRepository: ) return list(result.scalars().all()) - async def list_by_db_id(self, db_id: str) -> list[KnowledgeChunk]: + async def list_by_kb_id(self, kb_id: str) -> list[KnowledgeChunk]: async with pg_manager.get_async_session_context() as session: result = await session.execute( - select(KnowledgeChunk).where(KnowledgeChunk.db_id == db_id).order_by(KnowledgeChunk.id.asc()) + select(KnowledgeChunk).where(KnowledgeChunk.kb_id == kb_id).order_by(KnowledgeChunk.id.asc()) ) return list(result.scalars().all()) @@ -86,16 +86,16 @@ class KnowledgeChunkRepository: result = await session.execute(delete(KnowledgeChunk).where(KnowledgeChunk.file_id == file_id)) return int(result.rowcount or 0) - async def delete_by_db_id(self, db_id: str) -> int: + async def delete_by_kb_id(self, kb_id: str) -> int: async with pg_manager.get_async_session_context() as session: - result = await session.execute(delete(KnowledgeChunk).where(KnowledgeChunk.db_id == db_id)) + result = await session.execute(delete(KnowledgeChunk).where(KnowledgeChunk.kb_id == kb_id)) return int(result.rowcount or 0) - async def count_by_db_id(self, db_id: str) -> int: - return await self._count_by_db_id(db_id) + async def count_by_kb_id(self, kb_id: str) -> int: + return await self._count_by_kb_id(kb_id) - async def count_graph_indexed_by_db_id(self, db_id: str) -> int: - return await self._count_by_db_id(db_id, KnowledgeChunk.graph_indexed.is_(True)) + async def count_graph_indexed_by_kb_id(self, kb_id: str) -> int: + return await self._count_by_kb_id(kb_id, KnowledgeChunk.graph_indexed.is_(True)) async def count_graph_indexed_by_file_id(self, file_id: str) -> int: async with pg_manager.get_async_session_context() as session: @@ -106,21 +106,21 @@ class KnowledgeChunkRepository: ) return int(result.scalar() or 0) - async def count_graph_pending_by_db_id(self, db_id: str) -> int: - return await self._count_by_db_id(db_id, KnowledgeChunk.graph_indexed.is_not(True)) + async def count_graph_pending_by_kb_id(self, kb_id: str) -> int: + return await self._count_by_kb_id(kb_id, KnowledgeChunk.graph_indexed.is_not(True)) - async def _count_by_db_id(self, db_id: str, *conditions: Any) -> int: + async def _count_by_kb_id(self, kb_id: str, *conditions: Any) -> int: async with pg_manager.get_async_session_context() as session: result = await session.execute( - select(func.count()).select_from(KnowledgeChunk).where(KnowledgeChunk.db_id == db_id, *conditions) + select(func.count()).select_from(KnowledgeChunk).where(KnowledgeChunk.kb_id == kb_id, *conditions) ) return int(result.scalar() or 0) - async def list_graph_pending_by_db_id(self, db_id: str, limit: int) -> list[KnowledgeChunk]: + async def list_graph_pending_by_kb_id(self, kb_id: str, limit: int) -> list[KnowledgeChunk]: async with pg_manager.get_async_session_context() as session: result = await session.execute( select(KnowledgeChunk) - .where(KnowledgeChunk.db_id == db_id, KnowledgeChunk.graph_indexed.is_not(True)) + .where(KnowledgeChunk.kb_id == kb_id, KnowledgeChunk.graph_indexed.is_not(True)) .order_by(KnowledgeChunk.id.asc()) .limit(max(limit, 1)) ) @@ -149,11 +149,11 @@ class KnowledgeChunkRepository: async with pg_manager.get_async_session_context() as session: await session.execute(update(KnowledgeChunk).where(KnowledgeChunk.chunk_id == chunk_id).values(**values)) - async def reset_graph_state_by_db_id(self, db_id: str, clear_extraction_result: bool) -> int: + async def reset_graph_state_by_kb_id(self, kb_id: str, clear_extraction_result: bool) -> int: values: dict[str, Any] = {"graph_indexed": False} if clear_extraction_result: values.update({"extraction_result": None, "ent_ids": None, "tags": None}) async with pg_manager.get_async_session_context() as session: - result = await session.execute(update(KnowledgeChunk).where(KnowledgeChunk.db_id == db_id).values(**values)) + result = await session.execute(update(KnowledgeChunk).where(KnowledgeChunk.kb_id == kb_id).values(**values)) return int(result.rowcount or 0) diff --git a/backend/package/yuxi/repositories/knowledge_file_repository.py b/backend/package/yuxi/repositories/knowledge_file_repository.py index 6f7883c0..00d7c70e 100644 --- a/backend/package/yuxi/repositories/knowledge_file_repository.py +++ b/backend/package/yuxi/repositories/knowledge_file_repository.py @@ -20,9 +20,9 @@ class KnowledgeFileRepository: result = await session.execute(select(KnowledgeFile).where(KnowledgeFile.file_id == file_id)) return result.scalar_one_or_none() - async def list_by_db_id(self, db_id: str) -> list[KnowledgeFile]: + async def list_by_kb_id(self, kb_id: str) -> list[KnowledgeFile]: async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(KnowledgeFile).where(KnowledgeFile.db_id == db_id)) + result = await session.execute(select(KnowledgeFile).where(KnowledgeFile.kb_id == kb_id)) return list(result.scalars().all()) async def upsert(self, file_id: str, data: dict[str, Any]) -> KnowledgeFile: @@ -44,8 +44,8 @@ class KnowledgeFileRepository: if record is not None: await session.delete(record) - async def delete_by_db_id(self, db_id: str) -> None: + async def delete_by_kb_id(self, kb_id: str) -> None: async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(KnowledgeFile).where(KnowledgeFile.db_id == db_id)) + result = await session.execute(select(KnowledgeFile).where(KnowledgeFile.kb_id == kb_id)) for record in result.scalars().all(): await session.delete(record) diff --git a/backend/package/yuxi/repositories/knowledge_graph_repository.py b/backend/package/yuxi/repositories/knowledge_graph_repository.py index e1fadfd8..2fa72351 100644 --- a/backend/package/yuxi/repositories/knowledge_graph_repository.py +++ b/backend/package/yuxi/repositories/knowledge_graph_repository.py @@ -15,20 +15,20 @@ from yuxi.storage.postgres.models_knowledge import ( class KnowledgeGraphRepository: - async def count_by_db_id(self, db_id: str) -> tuple[int, int]: + async def count_by_kb_id(self, kb_id: str) -> tuple[int, int]: async with pg_manager.get_async_session_context() as session: entity_count = await session.scalar( - select(func.count()).select_from(KnowledgeGraphEntity).where(KnowledgeGraphEntity.db_id == db_id) + select(func.count()).select_from(KnowledgeGraphEntity).where(KnowledgeGraphEntity.kb_id == kb_id) ) triple_count = await session.scalar( - select(func.count()).select_from(KnowledgeGraphTriple).where(KnowledgeGraphTriple.db_id == db_id) + select(func.count()).select_from(KnowledgeGraphTriple).where(KnowledgeGraphTriple.kb_id == kb_id) ) return int(entity_count or 0), int(triple_count or 0) async def upsert_chunk_graph( self, *, - db_id: str, + kb_id: str, file_id: str, chunk_id: str, entities: list[dict[str, Any]], @@ -54,7 +54,7 @@ class KnowledgeGraphRepository: [ { "entity_id": entity["entity_id"], - "db_id": db_id, + "kb_id": kb_id, "file_id": file_id, "chunk_id": chunk_id, } @@ -86,7 +86,7 @@ class KnowledgeGraphRepository: [ { "triple_id": triple["triple_id"], - "db_id": db_id, + "kb_id": kb_id, "file_id": file_id, "chunk_id": chunk_id, "text": triple.get("text"), @@ -183,9 +183,9 @@ class KnowledgeGraphRepository: return orphan_entity_ids, orphan_triple_ids - async def delete_by_db_id(self, db_id: str) -> None: + async def delete_by_kb_id(self, kb_id: str) -> None: async with pg_manager.get_async_session_context() as session: - await session.execute(delete(KnowledgeGraphTripleMention).where(KnowledgeGraphTripleMention.db_id == db_id)) - await session.execute(delete(KnowledgeGraphEntityMention).where(KnowledgeGraphEntityMention.db_id == db_id)) - await session.execute(delete(KnowledgeGraphTriple).where(KnowledgeGraphTriple.db_id == db_id)) - await session.execute(delete(KnowledgeGraphEntity).where(KnowledgeGraphEntity.db_id == db_id)) + await session.execute(delete(KnowledgeGraphTripleMention).where(KnowledgeGraphTripleMention.kb_id == kb_id)) + await session.execute(delete(KnowledgeGraphEntityMention).where(KnowledgeGraphEntityMention.kb_id == kb_id)) + await session.execute(delete(KnowledgeGraphTriple).where(KnowledgeGraphTriple.kb_id == kb_id)) + await session.execute(delete(KnowledgeGraphEntity).where(KnowledgeGraphEntity.kb_id == kb_id)) diff --git a/backend/package/yuxi/repositories/mcp_server_repository.py b/backend/package/yuxi/repositories/mcp_server_repository.py index e2e14566..6ada7695 100644 --- a/backend/package/yuxi/repositories/mcp_server_repository.py +++ b/backend/package/yuxi/repositories/mcp_server_repository.py @@ -11,10 +11,10 @@ from yuxi.storage.postgres.models_business import MCPServer class MCPServerRepository: """MCP 服务器数据访问层""" - async def get_by_name(self, name: str) -> MCPServer | None: + async def get_by_slug(self, slug: str) -> MCPServer | None: """根据名称获取 MCP 服务器""" async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(MCPServer).where(MCPServer.name == name)) + result = await session.execute(select(MCPServer).where(MCPServer.slug == slug)) return result.scalar_one_or_none() async def list(self) -> list[MCPServer]: @@ -36,22 +36,22 @@ class MCPServerRepository: session.add(server) return server - async def update(self, name: str, data: dict[str, Any]) -> MCPServer | None: + async def update(self, slug: str, data: dict[str, Any]) -> MCPServer | None: """更新 MCP 服务器""" async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(MCPServer).where(MCPServer.name == name)) + result = await session.execute(select(MCPServer).where(MCPServer.slug == slug)) server = result.scalar_one_or_none() if server is None: return None for key, value in data.items(): - if key != "name": + if key != "slug": setattr(server, key, value) return server - async def delete(self, name: str) -> bool: + async def delete(self, slug: str) -> bool: """删除 MCP 服务器""" async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(MCPServer).where(MCPServer.name == name)) + result = await session.execute(select(MCPServer).where(MCPServer.slug == slug)) server = result.scalar_one_or_none() if server is None: return False @@ -60,22 +60,22 @@ class MCPServerRepository: async def upsert(self, data: dict[str, Any]) -> MCPServer: """插入或更新 MCP 服务器""" - name = data.get("name") + slug = data.get("slug") async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(MCPServer).where(MCPServer.name == name)) + result = await session.execute(select(MCPServer).where(MCPServer.slug == slug)) existing = result.scalar_one_or_none() if existing is None: server = MCPServer(**data) session.add(server) else: for key, value in data.items(): - if key != "name": + if key != "slug": setattr(existing, key, value) server = existing return server - async def exists_by_name(self, name: str) -> bool: + async def exists_by_slug(self, slug: str) -> bool: """检查 MCP 服务器是否存在""" async with pg_manager.get_async_session_context() as session: - result = await session.execute(select(MCPServer.id).where(MCPServer.name == name)) + result = await session.execute(select(MCPServer.id).where(MCPServer.slug == slug)) return result.scalar_one_or_none() is not None diff --git a/backend/package/yuxi/repositories/subagent_repository.py b/backend/package/yuxi/repositories/subagent_repository.py index 9ca994c8..dcebdae9 100644 --- a/backend/package/yuxi/repositories/subagent_repository.py +++ b/backend/package/yuxi/repositories/subagent_repository.py @@ -30,21 +30,22 @@ class SubAgentRepository: items = await self.list_enabled() return [item.to_subagent_spec() for item in items] - async def get_by_name(self, name: str) -> SubAgent | None: - """根据名称获取 SubAgent""" - result = await self.db.execute(select(SubAgent).where(SubAgent.name == name)) + async def get_by_slug(self, slug: str) -> SubAgent | None: + """根据 slug 获取 SubAgent""" + result = await self.db.execute(select(SubAgent).where(SubAgent.slug == slug)) return result.scalar_one_or_none() - async def exists_name(self, name: str) -> bool: - """检查名称是否存在(仅查询计数,不获取完整数据)""" + async def exists_slug(self, slug: str) -> bool: + """检查 slug 是否存在(仅查询计数,不获取完整数据)""" from sqlalchemy import func, select - result = await self.db.execute(select(func.count()).select_from(SubAgent).where(SubAgent.name == name)) + result = await self.db.execute(select(func.count()).select_from(SubAgent).where(SubAgent.slug == slug)) return result.scalar() > 0 async def create( self, *, + slug: str, name: str, description: str, system_prompt: str, @@ -55,6 +56,7 @@ class SubAgentRepository: ) -> SubAgent: now = utc_now_naive() item = SubAgent( + slug=slug, name=name, description=description, system_prompt=system_prompt, @@ -76,6 +78,7 @@ class SubAgentRepository: self, item: SubAgent, *, + name: str | None, description: str | None, system_prompt: str | None, tools: list[str] | None, @@ -85,6 +88,7 @@ class SubAgentRepository: ) -> SubAgent: # 批量更新非空字段 updates = { + "name": name, "description": description, "system_prompt": system_prompt, "tools": tools, diff --git a/backend/package/yuxi/services/filesystem_service.py b/backend/package/yuxi/services/filesystem_service.py index 9f928c42..0085a6c9 100644 --- a/backend/package/yuxi/services/filesystem_service.py +++ b/backend/package/yuxi/services/filesystem_service.py @@ -4,13 +4,9 @@ import asyncio from fastapi import HTTPException from sqlalchemy.ext.asyncio import AsyncSession -from yuxi.agents.backends import ( - ProvisionerSandboxBackend, - create_agent_composite_backend, - resolve_visible_knowledge_bases_for_context, -) +from yuxi.agents.backends import create_agent_composite_backend from yuxi.agents.backends.sandbox.backend import _looks_like_binary -from yuxi.agents.context import BaseContext, normalize_agent_context_config +from yuxi.agents.context import BaseContext, normalize_agent_context_config, prepare_agent_runtime_context from yuxi.repositories.agent_config_repository import AgentConfigRepository from yuxi.repositories.conversation_repository import ConversationRepository from yuxi.services.conversation_service import require_user_conversation @@ -68,10 +64,9 @@ async def _resolve_filesystem_state( ) runtime_context.thread_id = thread_id runtime_context.uid = str(user.uid) - await resolve_visible_knowledge_bases_for_context(runtime_context) + await prepare_agent_runtime_context(runtime_context) - sandbox_backend = ProvisionerSandboxBackend(thread_id=thread_id, uid=str(user.uid)) - return conversation, runtime_context, sandbox_backend + return conversation, runtime_context async def list_filesystem_entries_view( diff --git a/backend/package/yuxi/services/mcp_service.py b/backend/package/yuxi/services/mcp_service.py index ea9bae61..be050bea 100644 --- a/backend/package/yuxi/services/mcp_service.py +++ b/backend/package/yuxi/services/mcp_service.py @@ -28,8 +28,8 @@ from yuxi.utils import logger # Global Lock for MCP state _mcp_lock = asyncio.Lock() -# 本地仅缓存工具对象。配置始终以数据库为准,每次按 server_name 现查。 -# cache key 使用 server_name:config_hash,当配置变化时会自然失效。 +# 本地仅缓存工具对象。配置始终以数据库为准,每次按 server_slug 现查。 +# cache key 使用 server_slug:config_hash,当配置变化时会自然失效。 _mcp_tools_cache: dict[str, list[Callable[..., Any]]] = {} # MCP tools statistics (for reporting enabled/disabled counts) @@ -86,15 +86,16 @@ async def ensure_builtin_mcp_servers_in_db() -> None: try: async with pg_manager.get_async_session_context() as session: # Check if database has MCP configurations - result = await session.execute(select(func.count(MCPServer.name))) + result = await session.execute(select(func.count(MCPServer.slug))) count = result.scalar() if count == 0: # Database is empty, import default configurations logger.info("No MCP servers in database, importing default configurations...") - for name, config in _DEFAULT_MCP_SERVERS.items(): + for slug, config in _DEFAULT_MCP_SERVERS.items(): server = MCPServer( - name=name, + slug=slug, + name=config.get("name", slug), description=config.get("description"), transport=config["transport"], url=config.get("url"), @@ -115,12 +116,13 @@ async def ensure_builtin_mcp_servers_in_db() -> None: logger.info(f"Imported {len(_DEFAULT_MCP_SERVERS)} default MCP servers to database") else: # Ensure all built-in MCP servers exist in database - for name, config in _DEFAULT_MCP_SERVERS.items(): - result = await session.execute(select(MCPServer).filter(MCPServer.name == name)) + for slug, config in _DEFAULT_MCP_SERVERS.items(): + result = await session.execute(select(MCPServer).filter(MCPServer.slug == slug)) existing = result.scalar_one_or_none() if not existing: server = MCPServer( - name=name, + slug=slug, + name=config.get("name", slug), description=config.get("description"), transport=config["transport"], url=config.get("url"), @@ -137,7 +139,7 @@ async def ensure_builtin_mcp_servers_in_db() -> None: updated_by="system", ) session.add(server) - logger.info(f"Added built-in MCP server '{name}' to database") + logger.info(f"Added built-in MCP server '{slug}' to database") else: changed = False for field in _SYNCED_MCP_FIELDS: @@ -190,10 +192,10 @@ async def _load_enabled_mcp_server_configs( if db is not None: stmt = select(MCPServer).where(MCPServer.enabled == 1) if names: - stmt = stmt.where(MCPServer.name.in_(names)) + stmt = stmt.where(MCPServer.slug.in_(names)) result = await db.execute(stmt) servers = result.scalars().all() - return {server.name: server.to_mcp_config() for server in servers} + return {server.slug: server.to_mcp_config() for server in servers} from yuxi.storage.postgres.manager import pg_manager @@ -201,26 +203,26 @@ async def _load_enabled_mcp_server_configs( 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: +async def get_enabled_mcp_server_config(server_slug: 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) + configs = await _load_enabled_mcp_server_configs(names=[server_slug], db=db) + return configs.get(server_slug) -async def get_enabled_mcp_server_names(*, db: AsyncSession | None = None) -> list[str]: - """Get enabled MCP server names from the database.""" +async def get_enabled_mcp_server_slugs(*, db: AsyncSession | None = None) -> list[str]: + """Get enabled MCP server slugs from the database.""" if db is not None: - result = await db.execute(select(MCPServer.name).where(MCPServer.enabled == 1)) + result = await db.execute(select(MCPServer.slug).where(MCPServer.enabled == 1)) return [name for name in result.scalars().all() if isinstance(name, str)] from yuxi.storage.postgres.manager import pg_manager async with pg_manager.get_async_session_context() as session: - return await get_enabled_mcp_server_names(db=session) + return await get_enabled_mcp_server_slugs(db=session) async def get_mcp_tools( - server_name: str, + server_slug: str, additional_servers: dict[str, dict[str, Any]] | None = None, disabled_tools: list[str] = None, cache: bool = True, @@ -234,26 +236,26 @@ async def get_mcp_tools( 3. Filtering: Filters the return value based on `disabled_tools` argument. Args: - server_name: Server name + server_slug: Server slug additional_servers: Additional server configurations disabled_tools: List of tool names to filter out from the RETURN value (does not affect cache) cache: Whether to use/update the cache (default: True) force_refresh: Whether to force a refresh from the server (default: False) """ - if additional_servers and server_name in additional_servers: - server_config = additional_servers[server_name] + if additional_servers and server_slug in additional_servers: + server_config = additional_servers[server_slug] else: - server_config = await get_enabled_mcp_server_config(server_name) + server_config = await get_enabled_mcp_server_config(server_slug) if server_config is None: - logger.warning(f"MCP server '{server_name}' not found in database or disabled") + logger.warning(f"MCP server '{server_slug}' 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}" + cache_key = f"{server_slug}:{config_hash}" all_processed_tools: list[Callable[..., Any]] = [] @@ -266,13 +268,13 @@ async def get_mcp_tools( # 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}) + client = await get_mcp_client({server_slug: client_config}) if client is None: return [] raw_tools = cast(list[Any], await client.get_tools()) - server_cc = to_camel_case(server_name) + server_cc = to_camel_case(server_slug) for tool in raw_tools: original_name = tool.name tool_cc = to_camel_case(original_name) @@ -288,7 +290,7 @@ async def get_mcp_tools( if cache: async with _mcp_lock: stale_keys = [ - key for key in _mcp_tools_cache if key.startswith(f"{server_name}:") and key != cache_key + key for key in _mcp_tools_cache if key.startswith(f"{server_slug}:") and key != cache_key ] for stale_key in stale_keys: _mcp_tools_cache.pop(stale_key, None) @@ -296,20 +298,20 @@ async def get_mcp_tools( 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] = { + _mcp_tools_stats[server_slug] = { "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}' with key '{cache_key}': " + f"Refreshed MCP tools cache for '{server_slug}' with key '{cache_key}': " f"{len(all_processed_tools)} tools loaded." ) except Exception as e: logger.error( - f"Failed to load tools from MCP server '{server_name}': {e}, traceback: {traceback.format_exc()}" + f"Failed to load tools from MCP server '{server_slug}': {e}, traceback: {traceback.format_exc()}" ) return [] @@ -317,7 +319,7 @@ async def get_mcp_tools( if disabled_tools: filtered_tools = [t for t in all_processed_tools if t.name not in disabled_tools] logger.debug( - f"Returning {len(filtered_tools)}/{len(all_processed_tools)} tools for '{server_name}' " + f"Returning {len(filtered_tools)}/{len(all_processed_tools)} tools for '{server_slug}' " f"(filtered {len(disabled_tools)} by argument)" ) return filtered_tools @@ -329,8 +331,8 @@ 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 server_configs: - tools = await get_mcp_tools(server_name, additional_servers=server_configs) + for server_slug in server_configs: + tools = await get_mcp_tools(server_slug, additional_servers=server_configs) all_tools.extend(tools) return all_tools @@ -342,24 +344,24 @@ def clear_mcp_cache() -> None: _mcp_tools_stats = {} -def clear_mcp_server_tools_cache(server_name: str) -> None: +def clear_mcp_server_tools_cache(server_slug: str) -> None: """Clear the tools cache for a specific MCP server.""" global _mcp_tools_cache, _mcp_tools_stats - server_prefix = f"{server_name}:" + server_prefix = f"{server_slug}:" 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}'") + _mcp_tools_stats.pop(server_slug, None) + logger.info(f"Cleared tools cache for MCP server '{server_slug}'") -def get_mcp_tools_stats(server_name: str) -> dict[str, int] | None: +def get_mcp_tools_stats(server_slug: str) -> dict[str, int] | None: """Get tools statistics for a MCP server. Returns: dict with 'total', 'enabled', 'disabled' counts, or None if not available """ - return _mcp_tools_stats.get(server_name) + return _mcp_tools_stats.get(server_slug) # ============================================================================= @@ -367,9 +369,9 @@ def get_mcp_tools_stats(server_name: str) -> dict[str, int] | None: # ============================================================================= -async def get_mcp_server(db: AsyncSession, name: str) -> MCPServer | None: - """Get single server configuration.""" - result = await db.execute(select(MCPServer).filter(MCPServer.name == name)) +async def get_mcp_server(db: AsyncSession, slug: str) -> MCPServer | None: + """Get single server configuration by slug.""" + result = await db.execute(select(MCPServer).filter(MCPServer.slug == slug)) return result.scalar_one_or_none() @@ -381,6 +383,7 @@ async def get_all_mcp_servers(db: AsyncSession) -> list[MCPServer]: async def create_mcp_server( db: AsyncSession, + slug: str, name: str, transport: str, url: str = None, @@ -396,12 +399,12 @@ async def create_mcp_server( created_by: str = None, ) -> MCPServer: """Create server.""" - # Check if name exists - existing = await get_mcp_server(db, name) + existing = await get_mcp_server(db, slug) if existing: - raise ValueError(f"Server name '{name}' already exists") + raise ValueError(f"Server slug '{slug}' already exists") server = MCPServer( + slug=slug, name=name, description=description, transport=transport, @@ -422,15 +425,16 @@ async def create_mcp_server( await db.commit() await db.refresh(server) - clear_mcp_server_tools_cache(name) + clear_mcp_server_tools_cache(slug) - logger.info(f"Created MCP server '{name}'") + logger.info(f"Created MCP server '{slug}'") return server async def update_mcp_server( db: AsyncSession, - name: str, + slug: str, + name: str = None, description: str = None, transport: str = None, url: str = None, @@ -445,10 +449,12 @@ async def update_mcp_server( updated_by: str = None, ) -> MCPServer: """Update server configuration.""" - server = await get_mcp_server(db, name) + server = await get_mcp_server(db, slug) if not server: - raise ValueError(f"Server '{name}' does not exist") + raise ValueError(f"Server '{slug}' does not exist") + if name is not None: + server.name = name if description is not None: server.description = description if transport is not None: @@ -477,24 +483,24 @@ async def update_mcp_server( await db.commit() await db.refresh(server) - clear_mcp_server_tools_cache(name) + clear_mcp_server_tools_cache(slug) - logger.info(f"Updated MCP server '{name}'") + logger.info(f"Updated MCP server '{slug}'") return server -async def delete_mcp_server(db: AsyncSession, name: str) -> bool: +async def delete_mcp_server(db: AsyncSession, slug: str) -> bool: """Delete server.""" - server = await get_mcp_server(db, name) + server = await get_mcp_server(db, slug) if not server: return False await db.delete(server) await db.commit() - clear_mcp_server_tools_cache(name) + clear_mcp_server_tools_cache(slug) - logger.info(f"Deleted MCP server '{name}'") + logger.info(f"Deleted MCP server '{slug}'") return True @@ -504,12 +510,12 @@ async def delete_mcp_server(db: AsyncSession, name: str) -> bool: async def set_server_enabled( - db: AsyncSession, name: str, enabled: bool, updated_by: str = None + db: AsyncSession, slug: str, enabled: bool, updated_by: str = None ) -> tuple[bool, MCPServer]: """Set server enabled status.""" - server = await get_mcp_server(db, name) + server = await get_mcp_server(db, slug) if not server: - raise ValueError(f"Server '{name}' does not exist") + raise ValueError(f"Server '{slug}' does not exist") server.enabled = 1 if enabled else 0 if updated_by is not None: @@ -517,15 +523,15 @@ async def set_server_enabled( await db.commit() is_enabled = bool(server.enabled) - clear_mcp_server_tools_cache(name) + clear_mcp_server_tools_cache(slug) - logger.info(f"Set MCP server '{name}' enabled={is_enabled}") + logger.info(f"Set MCP server '{slug}' enabled={is_enabled}") return is_enabled, server async def toggle_tool_enabled( db: AsyncSession, - server_name: str, + server_slug: str, tool_name: str, updated_by: str = None, ) -> tuple[bool, MCPServer]: @@ -533,16 +539,16 @@ async def toggle_tool_enabled( Args: db: Database session - server_name: Server name + server_slug: Server slug tool_name: Tool name updated_by: Updater Returns: (enabled, server): Tool enabled status and updated server object """ - server = await get_mcp_server(db, server_name) + server = await get_mcp_server(db, server_slug) if not server: - raise ValueError(f"Server '{server_name}' does not exist") + raise ValueError(f"Server '{server_slug}' does not exist") disabled_tools = list(server.disabled_tools or []) @@ -559,9 +565,9 @@ async def toggle_tool_enabled( await db.commit() # Clear tool cache (re-filtered on next fetch) - clear_mcp_server_tools_cache(server_name) + clear_mcp_server_tools_cache(server_slug) - logger.info(f"Toggled tool '{tool_name}' for server '{server_name}' enabled={enabled}") + logger.info(f"Toggled tool '{tool_name}' for server '{server_slug}' enabled={enabled}") return enabled, server @@ -570,7 +576,7 @@ async def toggle_tool_enabled( # ============================================================================= -async def get_enabled_mcp_tools(server_name: str) -> list: +async def get_enabled_mcp_tools(server_slug: str) -> list: """Get MCP server tools (auto-filtering disabled_tools). Unified entry point for Agents, automatically: @@ -579,18 +585,18 @@ async def get_enabled_mcp_tools(server_name: str) -> list: 3. Filters out disabled_tools Args: - server_name: Server name + server_slug: Server slug Returns: List of enabled tools """ - config = await get_enabled_mcp_server_config(server_name) + config = await get_enabled_mcp_server_config(server_slug) if config is None: - logger.warning(f"MCP server '{server_name}' not found in database or disabled") + logger.warning(f"MCP server '{server_slug}' not found in database or disabled") return [] disabled_tools = config.get("disabled_tools") or [] - return await get_mcp_tools(server_name, additional_servers={server_name: config}, disabled_tools=disabled_tools) + return await get_mcp_tools(server_slug, additional_servers={server_slug: config}, disabled_tools=disabled_tools) async def get_servers_config(names: list[str]) -> dict[str, dict[str, Any]]: @@ -605,27 +611,27 @@ async def get_servers_config(names: list[str]) -> dict[str, dict[str, Any]]: return await _load_enabled_mcp_server_configs(names=names) -async def get_all_mcp_tools(server_name: str) -> list: +async def get_all_mcp_tools(server_slug: str) -> list: """Get all tools of an MCP server (no filtering). For management UI to display tool list, supports viewing all tools and their enabled status. Does NOT update the global tools cache to avoid polluting agent's filtered view. Args: - server_name: Server name + server_slug: Server slug Returns: List of all tools (unfiltered) """ - config = await get_enabled_mcp_server_config(server_name) + config = await get_enabled_mcp_server_config(server_slug) if config is None: - logger.warning(f"MCP server '{server_name}' not found in database or disabled") + logger.warning(f"MCP server '{server_slug}' not found in database or disabled") return [] # Get all tools (no filtering, force refresh, no cache update) return await get_mcp_tools( - server_name, - additional_servers={server_name: config}, + server_slug, + additional_servers={server_slug: config}, disabled_tools=[], cache=False, force_refresh=True, diff --git a/backend/package/yuxi/services/skill_service.py b/backend/package/yuxi/services/skill_service.py index 2f8521d1..b3513fc9 100644 --- a/backend/package/yuxi/services/skill_service.py +++ b/backend/package/yuxi/services/skill_service.py @@ -16,7 +16,7 @@ from sqlalchemy import select 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_enabled_mcp_server_names +from yuxi.services.mcp_service import get_enabled_mcp_server_slugs from yuxi.storage.postgres.models_business import Skill from yuxi.utils.logging_config import logger @@ -73,7 +73,7 @@ def _get_thread_skills_lock(thread_id: str) -> threading.Lock: return lock -def _normalize_string_list(values: list[str] | None) -> list[str]: +def normalize_string_list(values: list[str] | None) -> list[str]: if not values: return [] normalized: list[str] = [] @@ -113,14 +113,14 @@ def get_thread_skills_root_dir(thread_id: str) -> Path: return root -def sync_thread_visible_skills(thread_id: str, selected_slugs: list[str] | None) -> Path: +def sync_thread_readable_skills(thread_id: str, selected_slugs: list[str] | None) -> Path: skills_root = get_skills_root_dir().resolve() thread_skills_root = get_thread_skills_root_dir(thread_id) - normalized_slugs = [slug for slug in _normalize_string_list(selected_slugs) if is_valid_skill_slug(slug)] - visible_slugs = set(normalized_slugs) + normalized_slugs = [slug for slug in normalize_string_list(selected_slugs) if is_valid_skill_slug(slug)] + readable_slugs = set(normalized_slugs) with _get_thread_skills_lock(thread_id): for entry in thread_skills_root.iterdir(): - if entry.name in visible_slugs: + if entry.name in readable_slugs: continue if entry.is_dir() and not entry.is_symlink(): shutil.rmtree(entry) @@ -251,12 +251,12 @@ async def get_skill_dependency_options(db: AsyncSession) -> dict[str, list[str] def get_tools(): all_tools = get_tool_metadata() - return [{"id": tool["id"], "name": tool.get("name", tool["id"])} for tool in all_tools] + return [{"slug": tool["slug"], "name": tool.get("name", tool["slug"])} for tool in all_tools] skill_slugs, tool_list, mcp_names = await asyncio.gather( list_skill_slugs(db), asyncio.to_thread(get_tools), - get_enabled_mcp_server_names(db=db), + get_enabled_mcp_server_slugs(db=db), ) return { @@ -281,7 +281,7 @@ def _get_all_tool_names() -> list[str]: from yuxi.services.tool_service import get_tool_metadata all_tools = get_tool_metadata() - return [tool["id"] for tool in all_tools] + return [tool["slug"] for tool in all_tools] async def _validate_dependencies( @@ -292,9 +292,9 @@ async def _validate_dependencies( skill_dependencies: list[str], available_skill_slugs: set[str], ) -> tuple[list[str], list[str], list[str]]: - tools = _normalize_string_list(tool_dependencies) - mcps = _normalize_string_list(mcp_dependencies) - skills = _normalize_string_list(skill_dependencies) + tools = normalize_string_list(tool_dependencies) + mcps = normalize_string_list(mcp_dependencies) + skills = normalize_string_list(skill_dependencies) # 验证所有工具(不仅仅是 buildin) available_tools = set(_get_all_tool_names()) @@ -302,7 +302,7 @@ async def _validate_dependencies( if invalid_tools: raise ValueError(f"存在无效工具依赖: {', '.join(invalid_tools)}") - available_mcps = set(await get_enabled_mcp_server_names(db=None)) + available_mcps = set(await get_enabled_mcp_server_slugs(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)}") @@ -817,9 +817,9 @@ async def init_builtin_skills(db: AsyncSession, *, created_by: str = "system") - parsed_name, _, meta = _parse_skill_markdown(content) if parsed_name != slug: raise ValueError(f"内置 skill frontmatter.name 必须等于 slug: {slug}") - _normalize_string_list(meta.get("tool_dependencies")) - _normalize_string_list(meta.get("mcp_dependencies")) - _normalize_string_list(meta.get("skill_dependencies")) + normalize_string_list(meta.get("tool_dependencies")) + normalize_string_list(meta.get("mcp_dependencies")) + normalize_string_list(meta.get("skill_dependencies")) _compute_dir_hash(source_dir) @@ -830,9 +830,9 @@ def list_builtin_skill_specs() -> list[dict[str, Any]]: source_dir = Path(str(getattr(raw_spec, "source_dir", ""))).resolve() configured_description = str(getattr(raw_spec, "description", "")).strip() version = str(getattr(raw_spec, "version", "1.0.0")).strip() or "1.0.0" - configured_tools = _normalize_string_list(getattr(raw_spec, "tool_dependencies", None)) - configured_mcps = _normalize_string_list(getattr(raw_spec, "mcp_dependencies", None)) - configured_skills = _normalize_string_list(getattr(raw_spec, "skill_dependencies", None)) + configured_tools = normalize_string_list(getattr(raw_spec, "tool_dependencies", None)) + configured_mcps = normalize_string_list(getattr(raw_spec, "mcp_dependencies", None)) + configured_skills = normalize_string_list(getattr(raw_spec, "skill_dependencies", None)) if not is_valid_skill_slug(slug): raise ValueError(f"内置 skill slug 非法: {slug}") @@ -855,9 +855,9 @@ def list_builtin_skill_specs() -> list[dict[str, Any]]: "name": slug, "description": configured_description or parsed_desc, "version": version, - "tool_dependencies": configured_tools or _normalize_string_list(meta.get("tool_dependencies")), - "mcp_dependencies": configured_mcps or _normalize_string_list(meta.get("mcp_dependencies")), - "skill_dependencies": configured_skills or _normalize_string_list(meta.get("skill_dependencies")), + "tool_dependencies": configured_tools or normalize_string_list(meta.get("tool_dependencies")), + "mcp_dependencies": configured_mcps or normalize_string_list(meta.get("mcp_dependencies")), + "skill_dependencies": configured_skills or normalize_string_list(meta.get("skill_dependencies")), "content_hash": _compute_dir_hash(source_dir), "source_dir": source_dir, } @@ -930,9 +930,9 @@ async def update_builtin_skill( ) if ( - _normalize_string_list(item.tool_dependencies or []) != spec["tool_dependencies"] - or _normalize_string_list(item.mcp_dependencies or []) != spec["mcp_dependencies"] - or _normalize_string_list(item.skill_dependencies or []) != spec["skill_dependencies"] + normalize_string_list(item.tool_dependencies or []) != spec["tool_dependencies"] + or normalize_string_list(item.mcp_dependencies or []) != spec["mcp_dependencies"] + or normalize_string_list(item.skill_dependencies or []) != spec["skill_dependencies"] ): await repo.update_dependencies( item, diff --git a/backend/package/yuxi/services/subagent_service.py b/backend/package/yuxi/services/subagent_service.py index 1ee515d7..0d5d7943 100644 --- a/backend/package/yuxi/services/subagent_service.py +++ b/backend/package/yuxi/services/subagent_service.py @@ -31,7 +31,8 @@ async def _get_session(db: AsyncSession | None = None): # 内置 SubAgent 配置 _DEFAULT_SUBAGENTS = [ { - "name": "research-agent", + "slug": "research-agent", + "name": "研究员", "description": "利用搜索工具,用于研究更深入的问题。将调研结果写入到主题研究文件中。", "system_prompt": ( "你是一位专注的研究员。你的工作是根据用户的问题进行研究。" @@ -43,7 +44,8 @@ _DEFAULT_SUBAGENTS = [ "is_builtin": True, }, { - "name": "critique-agent", + "slug": "critique-agent", + "name": "评论员", "description": "用于评论最终报告。给这个代理一些关于你希望它如何评论报告的信息。", "system_prompt": ( "你是一位专注的编辑。你的任务是评论一份报告。\n\n" @@ -66,7 +68,7 @@ _DEFAULT_SUBAGENTS = [ }, ] -_SYNCED_SUBAGENT_FIELDS = ("description", "system_prompt", "tools", "model", "is_builtin") +_SYNCED_SUBAGENT_FIELDS = ("name", "description", "system_prompt", "tools", "model", "is_builtin") async def init_builtin_subagents() -> None: @@ -74,9 +76,10 @@ async def init_builtin_subagents() -> None: async with pg_manager.get_async_session_context() as session: repo = SubAgentRepository(session) for data in _DEFAULT_SUBAGENTS: - item = await repo.get_by_name(data["name"]) + item = await repo.get_by_slug(data["slug"]) if item is None: await repo.create( + slug=data["slug"], name=data["name"], description=data["description"], system_prompt=data["system_prompt"], @@ -114,15 +117,15 @@ async def get_subagent_specs(db: AsyncSession | None = None) -> list[dict[str, A return deepcopy(_subagent_specs_cache) -async def get_enabled_subagent_names(db: AsyncSession | None = None) -> list[str]: +async def get_enabled_subagent_slugs(db: AsyncSession | None = None) -> list[str]: if _subagent_specs_cache is not None: - return [spec["name"] for spec in _subagent_specs_cache if isinstance(spec.get("name"), str)] + return [spec["slug"] for spec in _subagent_specs_cache if isinstance(spec.get("slug"), str)] async with _get_session(db) as session: result = await session.execute( - select(SubAgent.name).where(SubAgent.enabled.is_(True)).order_by(SubAgent.updated_at.desc()) + select(SubAgent.slug).where(SubAgent.enabled.is_(True)).order_by(SubAgent.updated_at.desc()) ) - return [name for name in result.scalars().all() if isinstance(name, str)] + return [slug for slug in result.scalars().all() if isinstance(slug, str)] def clear_specs_cache() -> None: @@ -131,17 +134,17 @@ def clear_specs_cache() -> None: _subagent_specs_cache = None -async def get_subagents_from_names(selected_names: Any, *, db: AsyncSession | None = None) -> list[dict[str, Any]]: - """根据名称获取 subagent specs(含工具解析)。""" +async def get_subagents_from_slugs(selected_slugs: Any, *, db: AsyncSession | None = None) -> list[dict[str, Any]]: + """根据 slug 获取 subagent specs(含工具解析)。""" specs = await get_subagent_specs(db) - if not selected_names: + if not selected_slugs: return [] - selected_set = set(selected_names) - available = {spec["name"] for spec in specs if isinstance(spec.get("name"), str)} - matched = [spec for spec in specs if spec.get("name") in selected_set] - missing = [n for n in selected_names if n not in available] + selected_set = set(selected_slugs) + available = {spec["slug"] for spec in specs if isinstance(spec.get("slug"), str)} + matched = [spec for spec in specs if spec.get("slug") in selected_set] + missing = [slug for slug in selected_slugs if slug not in available] if missing: logger.warning(f"Configured subagents not found, skip: {missing}") @@ -169,11 +172,11 @@ async def get_all_subagents(db: AsyncSession | None = None) -> list[dict[str, An return [item.to_dict() for item in items] -async def get_subagent(name: str, db: AsyncSession | None = None) -> dict[str, Any] | None: +async def get_subagent(slug: str, db: AsyncSession | None = None) -> dict[str, Any] | None: """获取单个 SubAgent""" async with _get_session(db) as session: repo = SubAgentRepository(session) - item = await repo.get_by_name(name) + item = await repo.get_by_slug(slug) return item.to_dict() if item else None @@ -186,6 +189,7 @@ async def create_subagent( async with _get_session(db) as session: repo = SubAgentRepository(session) item = await repo.create( + slug=data["slug"], name=data["name"], description=data["description"], system_prompt=data["system_prompt"], @@ -199,7 +203,7 @@ async def create_subagent( async def update_subagent( - name: str, + slug: str, data: dict[str, Any], updated_by: str | None, db: AsyncSession | None = None, @@ -207,13 +211,14 @@ async def update_subagent( """更新 SubAgent""" async with _get_session(db) as session: repo = SubAgentRepository(session) - item = await repo.get_by_name(name) + item = await repo.get_by_slug(slug) if not item: return None if item.is_builtin: raise ValueError("内置 SubAgent 不可编辑") item = await repo.update( item, + name=data.get("name"), description=data.get("description"), system_prompt=data.get("system_prompt"), tools=data.get("tools"), @@ -225,11 +230,11 @@ async def update_subagent( return item.to_dict() -async def delete_subagent(name: str, db: AsyncSession | None = None) -> bool: +async def delete_subagent(slug: str, db: AsyncSession | None = None) -> bool: """删除 SubAgent""" async with _get_session(db) as session: repo = SubAgentRepository(session) - item = await repo.get_by_name(name) + item = await repo.get_by_slug(slug) if not item: return False if item.is_builtin: @@ -240,7 +245,7 @@ async def delete_subagent(name: str, db: AsyncSession | None = None) -> bool: async def set_subagent_enabled( - name: str, + slug: str, enabled: bool, *, updated_by: str | None, @@ -249,7 +254,7 @@ async def set_subagent_enabled( """更新 SubAgent 启用状态。""" async with _get_session(db) as session: repo = SubAgentRepository(session) - item = await repo.get_by_name(name) + item = await repo.get_by_slug(slug) if not item: return None item.enabled = enabled diff --git a/backend/package/yuxi/services/tool_service.py b/backend/package/yuxi/services/tool_service.py index 05f3ba70..3b14cecb 100644 --- a/backend/package/yuxi/services/tool_service.py +++ b/backend/package/yuxi/services/tool_service.py @@ -1,3 +1,5 @@ +from typing import Any + from yuxi.utils import logger # 工具元数据缓存 @@ -8,7 +10,7 @@ def _extract_tool_info(tool_obj) -> dict: """从 tool_obj 提取基础信息""" metadata = getattr(tool_obj, "metadata", {}) or {} info = { - "id": tool_obj.name, + "slug": tool_obj.name, "name": metadata.get("name", tool_obj.name), # 显示名称优先从 metadata 获取 "description": tool_obj.description, "metadata": metadata, @@ -77,3 +79,55 @@ def get_tool_metadata(category: str = None) -> list[dict]: if category: return [t for t in _metadata_cache if t.get("category") == category] return _metadata_cache + + +def get_tool_instances_by_category(category: str) -> list[Any]: + from yuxi.agents.toolkits.registry import get_all_extra_metadata, get_all_tool_instances + + extra_meta = get_all_extra_metadata() + tools = [] + for tool in get_all_tool_instances(): + tool_meta = extra_meta.get(tool.name) + tool_category = tool_meta.category if tool_meta else "buildin" + if tool_category == category: + tools.append(tool) + return tools + + +async def resolve_configured_runtime_tools(context) -> list[Any]: + from yuxi.services.mcp_service import get_enabled_mcp_tools + + selected_tools = [] + selected_tool_names: set[str] = set() + buildin_tools = {tool.name: tool for tool in get_tool_instances_by_category("buildin")} + + for tool_name in getattr(context, "tools", None) or []: + if not isinstance(tool_name, str) or tool_name in selected_tool_names: + continue + tool = buildin_tools.get(tool_name) + if tool is None: + logger.warning(f"Configured buildin tool not found, skip: {tool_name}") + continue + selected_tools.append(tool) + selected_tool_names.add(tool_name) + + selected_mcp_servers: set[str] = set() + for server_name in getattr(context, "mcps", None) or []: + if not isinstance(server_name, str) or server_name in selected_mcp_servers: + continue + selected_mcp_servers.add(server_name) + try: + mcp_tools = await get_enabled_mcp_tools(server_name) + except Exception as e: + logger.warning(f"Failed to load configured MCP tools '{server_name}': {e}") + continue + if not mcp_tools: + logger.warning(f"Configured MCP unavailable, skip: {server_name}") + continue + for tool in mcp_tools: + if tool.name in selected_tool_names: + continue + selected_tools.append(tool) + selected_tool_names.add(tool.name) + + return selected_tools diff --git a/backend/package/yuxi/services/viewer_filesystem_service.py b/backend/package/yuxi/services/viewer_filesystem_service.py index 538ef921..21e4b240 100644 --- a/backend/package/yuxi/services/viewer_filesystem_service.py +++ b/backend/package/yuxi/services/viewer_filesystem_service.py @@ -12,6 +12,7 @@ import aiofiles from fastapi import HTTPException, UploadFile from fastapi.responses import FileResponse, StreamingResponse from sqlalchemy.ext.asyncio import AsyncSession +from yuxi.agents.backends import create_agent_composite_backend from yuxi.agents.backends.sandbox import ( SKILLS_PATH, USER_DATA_PATH, @@ -22,7 +23,7 @@ from yuxi.agents.backends.sandbox import ( virtual_path_for_thread_file, ) from yuxi.agents.backends.skills_backend import SelectedSkillsReadonlyBackend -from yuxi.agents.middlewares.skills_middleware import normalize_selected_skills +from yuxi.services.skill_service import normalize_string_list from yuxi.services.filesystem_service import _resolve_filesystem_state from yuxi.storage.postgres.models_business import User from yuxi.utils.datetime_utils import utc_isoformat_from_timestamp @@ -338,14 +339,17 @@ async def _resolve_viewer_state( current_user: User, db: AsyncSession, ): - _conversation, runtime_context, sandbox_backend = await _resolve_filesystem_state( + _conversation, runtime_context = await _resolve_filesystem_state( thread_id=thread_id, user=current_user, db=db, agent_id=agent_id, agent_config_id=agent_config_id, ) - selected_skills = normalize_selected_skills(getattr(runtime_context, "skills", None) or []) + selected_skills = getattr(runtime_context, "_readable_skills", []) + selected_skills = normalize_string_list(selected_skills if isinstance(selected_skills, list) else []) + runtime_stub = type("RuntimeStub", (), {"context": runtime_context})() + sandbox_backend = create_agent_composite_backend(runtime_stub) skills_backend = SelectedSkillsReadonlyBackend(selected_slugs=selected_skills) return sandbox_backend, skills_backend, selected_skills diff --git a/backend/package/yuxi/storage/postgres/manager.py b/backend/package/yuxi/storage/postgres/manager.py index 31d0badf..4d566af3 100644 --- a/backend/package/yuxi/storage/postgres/manager.py +++ b/backend/package/yuxi/storage/postgres/manager.py @@ -168,7 +168,7 @@ class PostgresManager(metaclass=SingletonMeta): CREATE TABLE IF NOT EXISTS evaluation_datasets ( id SERIAL PRIMARY KEY, dataset_id VARCHAR(64) NOT NULL UNIQUE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, name VARCHAR(255) NOT NULL, description TEXT, item_count INTEGER DEFAULT 0, @@ -185,7 +185,7 @@ class PostgresManager(metaclass=SingletonMeta): id SERIAL PRIMARY KEY, item_id VARCHAR(64) NOT NULL UNIQUE, dataset_id VARCHAR(64) NOT NULL REFERENCES evaluation_datasets(dataset_id) ON DELETE CASCADE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, item_index INTEGER NOT NULL, query_text TEXT NOT NULL, gold_chunk_ids JSONB, @@ -198,7 +198,7 @@ class PostgresManager(metaclass=SingletonMeta): CREATE TABLE IF NOT EXISTS evaluation_runs ( id SERIAL PRIMARY KEY, run_id VARCHAR(64) NOT NULL UNIQUE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, dataset_id VARCHAR(64) REFERENCES evaluation_datasets(dataset_id) ON DELETE SET NULL, status VARCHAR(32) DEFAULT 'running', retrieval_config JSONB, @@ -232,7 +232,7 @@ class PostgresManager(metaclass=SingletonMeta): id SERIAL PRIMARY KEY, chunk_id VARCHAR(128) NOT NULL UNIQUE, file_id VARCHAR(64) NOT NULL REFERENCES knowledge_files(file_id) ON DELETE CASCADE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, chunk_index INTEGER NOT NULL, content TEXT NOT NULL, start_char_pos INTEGER, @@ -252,21 +252,21 @@ class PostgresManager(metaclass=SingletonMeta): CREATE TABLE IF NOT EXISTS knowledge_graph_entities ( id SERIAL PRIMARY KEY, entity_id VARCHAR(64) NOT NULL UNIQUE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, normalized_name VARCHAR(512) NOT NULL, label VARCHAR(128) NOT NULL, name VARCHAR(512) NOT NULL, attributes JSONB, created_at TIMESTAMPTZ DEFAULT NOW(), updated_at TIMESTAMPTZ DEFAULT NOW(), - CONSTRAINT uq_knowledge_graph_entities_identity UNIQUE (db_id, normalized_name, label) + CONSTRAINT uq_knowledge_graph_entities_identity UNIQUE (slug, normalized_name, label) ) """, """ CREATE TABLE IF NOT EXISTS knowledge_graph_entity_mentions ( id SERIAL PRIMARY KEY, entity_id VARCHAR(64) NOT NULL REFERENCES knowledge_graph_entities(entity_id) ON DELETE CASCADE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, file_id VARCHAR(64) NOT NULL REFERENCES knowledge_files(file_id) ON DELETE CASCADE, chunk_id VARCHAR(128) NOT NULL REFERENCES knowledge_chunks(chunk_id) ON DELETE CASCADE, created_at TIMESTAMPTZ DEFAULT NOW(), @@ -277,7 +277,7 @@ class PostgresManager(metaclass=SingletonMeta): CREATE TABLE IF NOT EXISTS knowledge_graph_triples ( id SERIAL PRIMARY KEY, triple_id VARCHAR(64) NOT NULL UNIQUE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, source_entity_id VARCHAR(64) NOT NULL REFERENCES knowledge_graph_entities(entity_id) ON DELETE CASCADE, target_entity_id VARCHAR(64) NOT NULL REFERENCES knowledge_graph_entities(entity_id) ON DELETE CASCADE, relation_type VARCHAR(256) NOT NULL, @@ -290,7 +290,7 @@ class PostgresManager(metaclass=SingletonMeta): CREATE TABLE IF NOT EXISTS knowledge_graph_triple_mentions ( id SERIAL PRIMARY KEY, triple_id VARCHAR(64) NOT NULL REFERENCES knowledge_graph_triples(triple_id) ON DELETE CASCADE, - db_id VARCHAR(80) NOT NULL REFERENCES knowledge_bases(db_id) ON DELETE CASCADE, + slug VARCHAR(80) NOT NULL REFERENCES knowledge_bases(slug) ON DELETE CASCADE, file_id VARCHAR(64) NOT NULL REFERENCES knowledge_files(file_id) ON DELETE CASCADE, chunk_id VARCHAR(128) NOT NULL REFERENCES knowledge_chunks(chunk_id) ON DELETE CASCADE, text TEXT, @@ -299,40 +299,40 @@ class PostgresManager(metaclass=SingletonMeta): CONSTRAINT uq_knowledge_graph_triple_mentions_triple_chunk UNIQUE (triple_id, chunk_id) ) """, - # 扩展 db_id 字段长度以支持最长 75 字符的 ID(kb_private_ + 64字符hash) - "ALTER TABLE IF EXISTS knowledge_bases ALTER COLUMN db_id TYPE VARCHAR(80)", - "ALTER TABLE IF EXISTS knowledge_files ALTER COLUMN db_id TYPE VARCHAR(80)", - "ALTER TABLE IF EXISTS evaluation_datasets ALTER COLUMN db_id TYPE VARCHAR(80)", - "ALTER TABLE IF EXISTS evaluation_dataset_items ALTER COLUMN db_id TYPE VARCHAR(80)", - "ALTER TABLE IF EXISTS evaluation_runs ALTER COLUMN db_id TYPE VARCHAR(80)", + # 扩展 slug 字段长度以支持最长 75 字符的 ID(kb_private_ + 64字符hash) + "ALTER TABLE IF EXISTS knowledge_bases ALTER COLUMN slug TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS knowledge_files ALTER COLUMN slug TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS evaluation_datasets ALTER COLUMN slug TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS evaluation_dataset_items ALTER COLUMN slug TYPE VARCHAR(80)", + "ALTER TABLE IF EXISTS evaluation_runs ALTER COLUMN slug TYPE VARCHAR(80)", "CREATE INDEX IF NOT EXISTS idx_kb_type ON knowledge_bases(kb_type)", "CREATE INDEX IF NOT EXISTS idx_kb_name ON knowledge_bases(name)", - "CREATE INDEX IF NOT EXISTS idx_kf_db_id ON knowledge_files(db_id)", + "CREATE INDEX IF NOT EXISTS idx_kf_slug ON knowledge_files(slug)", "CREATE INDEX IF NOT EXISTS idx_kf_parent ON knowledge_files(parent_id)", "CREATE INDEX IF NOT EXISTS idx_kf_status ON knowledge_files(status)", "CREATE INDEX IF NOT EXISTS idx_kf_hash ON knowledge_files(content_hash)", - "CREATE INDEX IF NOT EXISTS ix_evaluation_datasets_db_id ON evaluation_datasets(db_id)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_datasets_slug ON evaluation_datasets(slug)", ( "CREATE INDEX IF NOT EXISTS ix_evaluation_dataset_items_dataset_index " "ON evaluation_dataset_items(dataset_id, item_index)" ), - "CREATE INDEX IF NOT EXISTS ix_evaluation_dataset_items_db_id ON evaluation_dataset_items(db_id)", - "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_db_id ON evaluation_runs(db_id)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_dataset_items_slug ON evaluation_dataset_items(slug)", + "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_slug ON evaluation_runs(slug)", "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_status ON evaluation_runs(status)", "CREATE INDEX IF NOT EXISTS ix_evaluation_runs_started ON evaluation_runs(started_at DESC)", "CREATE INDEX IF NOT EXISTS ix_evaluation_run_items_run_index ON evaluation_run_items(run_id, item_index)", "CREATE UNIQUE INDEX IF NOT EXISTS uq_knowledge_chunks_chunk_id ON knowledge_chunks(chunk_id)", "CREATE INDEX IF NOT EXISTS ix_knowledge_chunks_file_id ON knowledge_chunks(file_id)", - "CREATE INDEX IF NOT EXISTS ix_knowledge_chunks_db_id ON knowledge_chunks(db_id)", + "CREATE INDEX IF NOT EXISTS ix_knowledge_chunks_slug ON knowledge_chunks(slug)", "CREATE INDEX IF NOT EXISTS ix_knowledge_chunks_graph_indexed ON knowledge_chunks(graph_indexed)", ( "CREATE UNIQUE INDEX IF NOT EXISTS uq_knowledge_graph_entities_entity_id " "ON knowledge_graph_entities(entity_id)" ), - "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_entities_db_id ON knowledge_graph_entities(db_id)", + "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_entities_slug ON knowledge_graph_entities(slug)", ( - "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_entity_mentions_db_id " - "ON knowledge_graph_entity_mentions(db_id)" + "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_entity_mentions_slug " + "ON knowledge_graph_entity_mentions(slug)" ), ( "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_entity_mentions_file_id " @@ -346,10 +346,10 @@ class PostgresManager(metaclass=SingletonMeta): "CREATE UNIQUE INDEX IF NOT EXISTS uq_knowledge_graph_triples_triple_id " "ON knowledge_graph_triples(triple_id)" ), - "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_triples_db_id ON knowledge_graph_triples(db_id)", + "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_triples_slug ON knowledge_graph_triples(slug)", ( - "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_triple_mentions_db_id " - "ON knowledge_graph_triple_mentions(db_id)" + "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_triple_mentions_slug " + "ON knowledge_graph_triple_mentions(slug)" ), ( "CREATE INDEX IF NOT EXISTS ix_knowledge_graph_triple_mentions_file_id " diff --git a/backend/package/yuxi/storage/postgres/models_business.py b/backend/package/yuxi/storage/postgres/models_business.py index 96ba528d..5b2d30ca 100644 --- a/backend/package/yuxi/storage/postgres/models_business.py +++ b/backend/package/yuxi/storage/postgres/models_business.py @@ -429,8 +429,9 @@ class MCPServer(Base): __tablename__ = "mcp_servers" - # 核心字段 - name 作为主键 - name = Column(String(100), primary_key=True, comment="服务器名称(唯一标识)") + id = Column(Integer, primary_key=True, autoincrement=True) + slug = Column(String(100), nullable=False, unique=True, index=True, comment="稳定标识") + name = Column(String(100), nullable=False, comment="展示名称") description = Column(String(500), nullable=True, comment="描述") # 连接配置 @@ -461,6 +462,8 @@ class MCPServer(Base): def to_dict(self) -> dict[str, Any]: return { + "id": self.id, + "slug": self.slug, "name": self.name, "description": self.description, "transport": self.transport, @@ -634,7 +637,9 @@ class SubAgent(Base): __tablename__ = "subagents" - name = Column(String(128), primary_key=True, comment="唯一标识") + id = Column(Integer, primary_key=True, autoincrement=True) + slug = Column(String(128), nullable=False, unique=True, index=True, comment="稳定标识") + name = Column(String(128), nullable=False, comment="展示名称") description = Column(Text, nullable=False, comment="描述") system_prompt = Column(Text, nullable=False, comment="系统提示词") tools = Column(JSON, nullable=False, default=list, comment="工具名称列表") @@ -650,6 +655,8 @@ class SubAgent(Base): def to_dict(self) -> dict[str, Any]: return { + "id": self.id, + "slug": self.slug, "name": self.name, "description": self.description, "system_prompt": self.system_prompt, @@ -666,7 +673,8 @@ class SubAgent(Base): def to_subagent_spec(self) -> dict[str, Any]: """转换为 SubAgentMiddleware 需要的 spec 格式""" spec = { - "name": self.name, + "slug": self.slug, + "name": self.slug, "description": self.description, "system_prompt": self.system_prompt, "tools": self.tools or [], diff --git a/backend/server/routers/chat_router.py b/backend/server/routers/chat_router.py index dce6ec30..d12d5576 100644 --- a/backend/server/routers/chat_router.py +++ b/backend/server/routers/chat_router.py @@ -193,7 +193,11 @@ async def get_agent(current_user: User = Depends(get_required_user)): @chat.get("/agent/{agent_id}") -async def get_single_agent(agent_id: str, current_user: User = Depends(get_config_user)): +async def get_single_agent( + agent_id: str, + current_user: User = Depends(get_config_user), + db: AsyncSession = Depends(get_db), +): """获取指定智能体的完整信息(包含配置选项)(需要登录)""" try: # 检查智能体是否存在 @@ -201,7 +205,7 @@ async def get_single_agent(agent_id: str, current_user: User = Depends(get_confi raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在") # 获取智能体的完整信息(包含 configurable_items) - agent_info = await agent.get_info(user_role=current_user.role) + agent_info = await agent.get_info(user_role=current_user.role, db=db, user=current_user) return agent_info diff --git a/backend/server/routers/dashboard_router.py b/backend/server/routers/dashboard_router.py index 0c1f2328..98cb034a 100644 --- a/backend/server/routers/dashboard_router.py +++ b/backend/server/routers/dashboard_router.py @@ -452,7 +452,7 @@ async def get_knowledge_stats( }.get(kb_type, kb.kb_type or "未知类型") databases_by_type[display_type] = databases_by_type.get(display_type, 0) + 1 - files = await file_repo.list_by_db_id(kb.db_id) + files = await file_repo.list_by_kb_id(kb.kb_id) total_files += len(files) for record in files: file_ext = (record.file_type or "").lower() diff --git a/backend/server/routers/knowledge_eval_router.py b/backend/server/routers/knowledge_eval_router.py index 7ac1abda..e244d5de 100644 --- a/backend/server/routers/knowledge_eval_router.py +++ b/backend/server/routers/knowledge_eval_router.py @@ -37,9 +37,9 @@ class RunEvaluationRequest(BaseModel): retrieval_config: dict[str, Any] = Field(default_factory=dict, alias="model_config") -@evaluation.post("/databases/{db_id}/datasets/upload") +@evaluation.post("/databases/{kb_id}/datasets/upload") async def upload_evaluation_dataset( - db_id: str, + kb_id: str, file: UploadFile = File(...), name: str = Form(...), description: str = Form(""), @@ -54,7 +54,7 @@ async def upload_evaluation_dataset( service = EvaluationService() result = await service.upload_dataset( - db_id=db_id, + kb_id=kb_id, file_content=await file.read(), filename=file.filename, name=name, @@ -69,23 +69,23 @@ async def upload_evaluation_dataset( raise HTTPException(status_code=500, detail=f"上传评估数据集失败: {str(e)}") -@evaluation.get("/databases/{db_id}/datasets") -async def list_evaluation_datasets(db_id: str, current_user: User = Depends(get_admin_user)): +@evaluation.get("/databases/{kb_id}/datasets") +async def list_evaluation_datasets(kb_id: str, current_user: User = Depends(get_admin_user)): """获取知识库的评估数据集列表""" from yuxi.knowledge.eval.service import EvaluationService try: service = EvaluationService() - datasets = await service.list_datasets(db_id) + datasets = await service.list_datasets(kb_id) return {"message": "success", "data": datasets} except Exception as e: logger.error(f"获取评估数据集列表失败: {e}, {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"获取评估数据集列表失败: {str(e)}") -@evaluation.get("/databases/{db_id}/datasets/{dataset_id}") +@evaluation.get("/databases/{kb_id}/datasets/{dataset_id}") async def get_evaluation_dataset( - db_id: str, dataset_id: str, page: int = 1, page_size: int = 10, current_user: User = Depends(get_admin_user) + kb_id: str, dataset_id: str, page: int = 1, page_size: int = 10, current_user: User = Depends(get_admin_user) ): """获取评估数据集详情""" from yuxi.knowledge.eval.service import EvaluationService @@ -97,7 +97,7 @@ async def get_evaluation_dataset( raise HTTPException(status_code=400, detail="每页大小必须在1-100之间") service = EvaluationService() - dataset = await service.get_dataset_detail(db_id, dataset_id, page, page_size) + dataset = await service.get_dataset_detail(kb_id, dataset_id, page, page_size) return {"message": "success", "data": dataset} except HTTPException: raise @@ -151,9 +151,9 @@ async def delete_evaluation_dataset(dataset_id: str, current_user: User = Depend raise HTTPException(status_code=500, detail=f"删除评估数据集失败: {str(e)}") -@evaluation.post("/databases/{db_id}/datasets/generate") +@evaluation.post("/databases/{kb_id}/datasets/generate") async def generate_evaluation_dataset( - db_id: str, request: GenerateDatasetRequest, current_user: User = Depends(get_admin_user) + kb_id: str, request: GenerateDatasetRequest, current_user: User = Depends(get_admin_user) ): """自动生成评估数据集""" from yuxi.knowledge.eval.service import EvaluationService @@ -161,7 +161,7 @@ async def generate_evaluation_dataset( try: service = EvaluationService() result = await service.generate_dataset( - db_id=db_id, + kb_id=kb_id, name=request.name, description=request.description, count=request.count, @@ -180,15 +180,15 @@ async def generate_evaluation_dataset( raise HTTPException(status_code=500, detail=f"生成评估数据集失败: {str(e)}") -@evaluation.post("/databases/{db_id}/runs") -async def run_evaluation(db_id: str, request: RunEvaluationRequest, current_user: User = Depends(get_admin_user)): +@evaluation.post("/databases/{kb_id}/runs") +async def run_evaluation(kb_id: str, request: RunEvaluationRequest, current_user: User = Depends(get_admin_user)): """运行RAG评估""" from yuxi.knowledge.eval.service import EvaluationService try: service = EvaluationService() run_id = await service.run_evaluation( - db_id=db_id, + kb_id=kb_id, dataset_id=request.dataset_id, model_config=request.retrieval_config, created_by=current_user.uid, @@ -203,23 +203,23 @@ async def run_evaluation(db_id: str, request: RunEvaluationRequest, current_user raise HTTPException(status_code=500, detail=f"启动评估失败: {str(e)}") -@evaluation.get("/databases/{db_id}/runs") -async def list_evaluation_runs(db_id: str, current_user: User = Depends(get_admin_user)): +@evaluation.get("/databases/{kb_id}/runs") +async def list_evaluation_runs(kb_id: str, current_user: User = Depends(get_admin_user)): """获取知识库评估运行历史""" from yuxi.knowledge.eval.service import EvaluationService try: service = EvaluationService() - runs = await service.list_runs(db_id) + runs = await service.list_runs(kb_id) return {"message": "success", "data": runs} except Exception as e: logger.error(f"获取评估运行历史失败: {e}, {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"获取评估运行历史失败: {str(e)}") -@evaluation.get("/databases/{db_id}/runs/{run_id}") +@evaluation.get("/databases/{kb_id}/runs/{run_id}") async def get_evaluation_run_results( - db_id: str, + kb_id: str, run_id: str, page: int = 1, page_size: int = 20, @@ -236,7 +236,7 @@ async def get_evaluation_run_results( raise HTTPException(status_code=400, detail="每页大小必须在1-100之间") service = EvaluationService() - results = await service.get_run_results(db_id, run_id, page=page, page_size=page_size, error_only=error_only) + results = await service.get_run_results(kb_id, run_id, page=page, page_size=page_size, error_only=error_only) return {"message": "success", "data": results} except HTTPException: raise @@ -249,14 +249,14 @@ async def get_evaluation_run_results( raise HTTPException(status_code=500, detail=f"获取评估运行结果失败: {str(e)}") -@evaluation.delete("/databases/{db_id}/runs/{run_id}") -async def delete_evaluation_run(db_id: str, run_id: str, current_user: User = Depends(get_admin_user)): +@evaluation.delete("/databases/{kb_id}/runs/{run_id}") +async def delete_evaluation_run(kb_id: str, run_id: str, current_user: User = Depends(get_admin_user)): """删除评估运行""" from yuxi.knowledge.eval.service import EvaluationService try: service = EvaluationService() - await service.delete_run(db_id, run_id) + await service.delete_run(kb_id, run_id) return {"message": "success", "data": None} except ValueError as e: if "not found" in str(e).lower(): diff --git a/backend/server/routers/mcp_router.py b/backend/server/routers/mcp_router.py index 76fdbd96..8c253f8f 100644 --- a/backend/server/routers/mcp_router.py +++ b/backend/server/routers/mcp_router.py @@ -28,7 +28,8 @@ mcp = APIRouter(prefix="/system/mcp-servers", tags=["mcp"]) class CreateMcpServerRequest(BaseModel): - name: str = Field(..., description="服务器名称") + slug: str = Field(..., description="稳定标识") + name: str = Field(..., description="展示名称") transport: str = Field(..., description="传输类型:sse/streamable_http/stdio") url: str | None = Field(None, description="服务器 URL(sse/streamable_http)") command: str | None = Field(None, description="命令(stdio)") @@ -43,6 +44,7 @@ class CreateMcpServerRequest(BaseModel): class UpdateMcpServerRequest(BaseModel): + name: str | None = Field(None, description="展示名称") transport: str | None = Field(None, description="传输类型") url: str | None = Field(None, description="服务器 URL") command: str | None = Field(None, description="命令(stdio)") @@ -65,11 +67,11 @@ class UpdateMcpServerStatusRequest(BaseModel): # ============================================================================= -async def get_server_or_404(db: AsyncSession, name: str): +async def get_server_or_404(db: AsyncSession, slug: str): """Helper to get server or raise 404.""" - server = await get_mcp_server(db, name) + server = await get_mcp_server(db, slug) if not server: - raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在") + raise HTTPException(status_code=404, detail=f"服务器 '{slug}' 不存在") return server @@ -129,6 +131,7 @@ async def create_mcp_server_route( try: server = await create_mcp_server( db, + slug=request.slug, name=request.name, transport=request.transport, url=request.url, @@ -151,15 +154,15 @@ async def create_mcp_server_route( raise HTTPException(status_code=500, detail=str(e)) -@mcp.get("/{name}") +@mcp.get("/{slug}") async def get_mcp_server_route( - name: str, + slug: str, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """获取单个 MCP 服务器配置""" try: - server = await get_server_or_404(db, name) + server = await get_server_or_404(db, slug) return {"success": True, "data": server.to_dict()} except HTTPException: raise @@ -168,9 +171,9 @@ async def get_mcp_server_route( raise HTTPException(status_code=500, detail=str(e)) -@mcp.put("/{name}") +@mcp.put("/{slug}") async def update_mcp_server_route( - name: str, + slug: str, request: UpdateMcpServerRequest, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), @@ -189,7 +192,8 @@ async def update_mcp_server_route( server = await update_mcp_server( db, - name=name, + slug=slug, + name=request.name, description=request.description, transport=request.transport, url=request.url, @@ -211,23 +215,23 @@ async def update_mcp_server_route( raise HTTPException(status_code=500, detail=str(e)) -@mcp.delete("/{name}") +@mcp.delete("/{slug}") async def delete_mcp_server_route( - name: str, + slug: str, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """删除 MCP 服务器""" try: # 检查是否为系统内置服务器 - server = await get_mcp_server(db, name) + server = await get_mcp_server(db, slug) if server and server.created_by == "system": raise HTTPException(status_code=403, detail="系统内置的 MCP 服务器无法删除") - deleted = await delete_mcp_server(db, name) + deleted = await delete_mcp_server(db, slug) if not deleted: - raise HTTPException(status_code=404, detail=f"服务器 '{name}' 不存在") - return {"success": True, "message": f"服务器 '{name}' 已删除"} + raise HTTPException(status_code=404, detail=f"服务器 '{slug}' 不存在") + return {"success": True, "message": f"服务器 '{slug}' 已删除"} except HTTPException: raise except Exception as e: @@ -240,18 +244,18 @@ async def delete_mcp_server_route( # ============================================================================= -@mcp.post("/{name}/test") +@mcp.post("/{slug}/test") async def test_mcp_server( - name: str, + slug: str, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """测试 MCP 服务器连接""" try: - await get_server_or_404(db, name) + await get_server_or_404(db, slug) try: - tools = await get_all_mcp_tools(name) + tools = await get_all_mcp_tools(slug) return { "success": True, "message": f"连接成功,共发现 {len(tools)} 个工具", @@ -266,21 +270,21 @@ async def test_mcp_server( raise HTTPException(status_code=500, detail=str(e)) -@mcp.put("/{name}/status") +@mcp.put("/{slug}/status") async def update_mcp_server_status_route( - name: str, + slug: str, request: UpdateMcpServerStatusRequest, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """更新 MCP 服务器启用状态""" try: - is_enabled, server = await set_server_enabled(db, name, request.enabled, current_user.username) + is_enabled, server = await set_server_enabled(db, slug, request.enabled, current_user.username) return { "success": True, "enabled": is_enabled, "data": server.to_dict(), - "message": f"MCP '{name}' 已{'添加' if is_enabled else '移除'}", + "message": f"MCP '{slug}' 已{'添加' if is_enabled else '移除'}", } except ValueError as ve: raise HTTPException(status_code=404, detail=str(ve)) @@ -294,20 +298,20 @@ async def update_mcp_server_status_route( # ============================================================================= -@mcp.get("/{name}/tools") +@mcp.get("/{slug}/tools") async def get_mcp_server_tools( - name: str, + slug: str, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """获取 MCP 服务器的工具列表""" try: - server = await get_server_or_404(db, name) + server = await get_server_or_404(db, slug) disabled_tools = server.disabled_tools or [] try: # 获取所有工具(不过滤 disabled_tools) - tools = await get_all_mcp_tools(name) + tools = await get_all_mcp_tools(slug) tool_list = [] for tool in tools: @@ -336,7 +340,7 @@ async def get_mcp_server_tools( "total": len(tool_list), } except Exception as tool_error: - logger.error(f"Failed to get tools from MCP server '{name}': {tool_error}") + logger.error(f"Failed to get tools from MCP server '{slug}': {tool_error}") raise HTTPException(status_code=500, detail=f"获取工具失败: {str(tool_error)}") except HTTPException: raise @@ -345,22 +349,22 @@ async def get_mcp_server_tools( raise HTTPException(status_code=500, detail=str(e)) -@mcp.post("/{name}/tools/refresh") +@mcp.post("/{slug}/tools/refresh") async def refresh_mcp_server_tools( - name: str, + slug: str, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """刷新 MCP 服务器的工具列表(清除缓存重新获取)""" try: - await get_server_or_404(db, name) + await get_server_or_404(db, slug) try: # 获取所有工具(不过滤 disabled_tools) - tools = await get_all_mcp_tools(name) + tools = await get_all_mcp_tools(slug) # 获取统计信息 - stats = get_mcp_tools_stats(name) + stats = get_mcp_tools_stats(slug) enabled_count = stats.get("enabled", len(tools)) if stats else len(tools) disabled_count = stats.get("disabled", 0) if stats else 0 @@ -386,16 +390,16 @@ async def refresh_mcp_server_tools( raise HTTPException(status_code=500, detail=str(e)) -@mcp.put("/{name}/tools/{tool_name}/toggle") +@mcp.put("/{slug}/tools/{tool_name}/toggle") async def toggle_mcp_server_tool_route( - name: str, + slug: str, tool_name: str, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """切换单个工具的启用状态""" try: - enabled, server = await toggle_tool_enabled(db, name, tool_name, current_user.username) + enabled, server = await toggle_tool_enabled(db, slug, tool_name, current_user.username) return { "success": True, "tool_name": tool_name, diff --git a/backend/server/routers/subagent_router.py b/backend/server/routers/subagent_router.py index f535006c..7b6c0c03 100644 --- a/backend/server/routers/subagent_router.py +++ b/backend/server/routers/subagent_router.py @@ -17,7 +17,8 @@ subagents_router = APIRouter(prefix="/system/subagents", tags=["subagents"]) class SubAgentCreateRequest(BaseModel): - name: str = Field(..., description="唯一标识") + slug: str = Field(..., description="稳定标识") + name: str = Field(..., description="展示名称") description: str = Field(..., description="描述") system_prompt: str = Field(..., description="系统提示词") tools: list[str] = Field(default_factory=list, description="工具名称列表") @@ -25,6 +26,7 @@ class SubAgentCreateRequest(BaseModel): class SubAgentUpdateRequest(BaseModel): + name: str | None = Field(None, description="展示名称") description: str | None = Field(None, description="描述") system_prompt: str | None = Field(None, description="系统提示词") tools: list[str] | None = Field(None, description="工具名称列表") @@ -46,13 +48,9 @@ def _raise_internal_error(action: str, error: Exception) -> None: raise HTTPException(status_code=500, detail=f"{action}失败") -def _is_subagent_name_duplicate_error(error: IntegrityError) -> bool: +def _is_subagent_slug_duplicate_error(error: IntegrityError) -> bool: raw_message = str(getattr(error, "orig", error)).lower() - return ( - "duplicate key" in raw_message - and "subagents" in raw_message - and ("(name)" in raw_message or "subagents_pkey" in raw_message) - ) + return "duplicate key" in raw_message and "subagents" in raw_message and "(slug)" in raw_message @subagents_router.get("") @@ -68,17 +66,17 @@ async def list_subagents_route( _raise_internal_error("获取列表", e) -@subagents_router.get("/{name}") +@subagents_router.get("/{slug}") async def get_subagent_route( - name: str, + slug: str, _current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """获取单个 SubAgent(管理员可读)""" try: - item = await service.get_subagent(name, db) + item = await service.get_subagent(slug, db) if not item: - raise HTTPException(status_code=404, detail=f"SubAgent '{name}' 不存在") + raise HTTPException(status_code=404, detail=f"SubAgent '{slug}' 不存在") return {"success": True, "data": item} except HTTPException: raise @@ -98,8 +96,8 @@ async def create_subagent_route( item = await service.create_subagent(data, created_by=current_user.username, db=db) return {"success": True, "data": item} except IntegrityError as e: - if _is_subagent_name_duplicate_error(e): - raise HTTPException(status_code=409, detail=f"SubAgent '{payload.name}' 已存在") + if _is_subagent_slug_duplicate_error(e): + raise HTTPException(status_code=409, detail=f"SubAgent '{payload.slug}' 已存在") _raise_internal_error("创建", e) except HTTPException: raise @@ -109,9 +107,9 @@ async def create_subagent_route( _raise_internal_error("创建", e) -@subagents_router.put("/{name}") +@subagents_router.put("/{slug}") async def update_subagent_route( - name: str, + slug: str, payload: SubAgentUpdateRequest, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), @@ -119,9 +117,9 @@ async def update_subagent_route( """更新 SubAgent(管理员)""" try: data = payload.model_dump(exclude_unset=True) - item = await service.update_subagent(name, data, updated_by=current_user.username, db=db) + item = await service.update_subagent(slug, data, updated_by=current_user.username, db=db) if not item: - raise HTTPException(status_code=404, detail=f"SubAgent '{name}' 不存在") + raise HTTPException(status_code=404, detail=f"SubAgent '{slug}' 不存在") return {"success": True, "data": item} except ValueError as e: _raise_from_value_error(e) @@ -131,17 +129,17 @@ async def update_subagent_route( _raise_internal_error("更新", e) -@subagents_router.delete("/{name}") +@subagents_router.delete("/{slug}") async def delete_subagent_route( - name: str, + slug: str, _current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """删除 SubAgent(管理员)""" try: - deleted = await service.delete_subagent(name, db=db) + deleted = await service.delete_subagent(slug, db=db) if not deleted: - raise HTTPException(status_code=404, detail=f"SubAgent '{name}' 不存在") + raise HTTPException(status_code=404, detail=f"SubAgent '{slug}' 不存在") return {"success": True} except ValueError as e: _raise_from_value_error(e) @@ -151,22 +149,22 @@ async def delete_subagent_route( _raise_internal_error("删除", e) -@subagents_router.put("/{name}/status") +@subagents_router.put("/{slug}/status") async def update_subagent_status_route( - name: str, + slug: str, payload: SubAgentStatusRequest, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db), ): """更新 SubAgent 启用状态(管理员)。""" try: - item = await service.set_subagent_enabled(name, payload.enabled, updated_by=current_user.username, db=db) + item = await service.set_subagent_enabled(slug, payload.enabled, updated_by=current_user.username, db=db) if not item: - raise HTTPException(status_code=404, detail=f"SubAgent '{name}' 不存在") + raise HTTPException(status_code=404, detail=f"SubAgent '{slug}' 不存在") return { "success": True, "data": item, - "message": f"SubAgent '{name}' 已{'添加' if payload.enabled else '移除'}", + "message": f"SubAgent '{slug}' 已{'添加' if payload.enabled else '移除'}", } except HTTPException: raise diff --git a/backend/server/routers/tool_router.py b/backend/server/routers/tool_router.py index f19df552..cee32c42 100644 --- a/backend/server/routers/tool_router.py +++ b/backend/server/routers/tool_router.py @@ -22,4 +22,4 @@ async def get_tool_options( ): """获取工具选项(前端下拉框用)""" all_tools = get_tool_metadata() - return {"success": True, "data": [{"label": t["name"], "value": t["id"]} for t in all_tools]} + return {"success": True, "data": [{"label": t["name"], "value": t["slug"]} for t in all_tools]}