refactor: Agent 运行时架构重构 - 移除 RuntimeConfigMiddleware,统一上下文准备与工具解析
- 删除 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 适配新接口
This commit is contained in:
parent
82fd5846ea
commit
c772e5ca3a
@ -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),
|
||||
},
|
||||
)
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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}")
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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="任务结束前,应该检查维护的待办事项列表是否结束。"),
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
@ -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 更新"""
|
||||
|
||||
@ -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 = []
|
||||
|
||||
@ -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)}"
|
||||
|
||||
@ -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())
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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 "
|
||||
|
||||
@ -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 [],
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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():
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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]}
|
||||
|
||||
Loading…
Reference in New Issue
Block a user