refactor(skills): 通过引入SkillsMiddleware重构技能处理
- 添加了SkillsMiddleware来管理技能提示注入、依赖解析和动态激活。 - 从RuntimeConfigMiddleware中移除了与技能相关的逻辑并将其整合到SkillsMiddleware中。 - 删除了skill_resolver.py,因为其功能已集成到SkillsMiddleware中。 - 更新了各个组件以使用新的SkillsMiddleware进行技能管理。 - 改进了skill_service.py中技能和工具加载的异步处理。
This commit is contained in:
parent
24cbb41f4f
commit
393d3b6a32
@ -13,6 +13,7 @@ from src.agents.common.middlewares import (
|
||||
save_attachments_to_fs,
|
||||
)
|
||||
from src.agents.common.middlewares.knowledge_base_middleware import KnowledgeBaseMiddleware
|
||||
from src.agents.common.middlewares.skills_middleware import SkillsMiddleware
|
||||
from src.services.mcp_service import get_tools_from_all_servers
|
||||
|
||||
|
||||
@ -46,6 +47,7 @@ class ChatbotAgent(BaseAgent):
|
||||
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
|
||||
KnowledgeBaseMiddleware(), # 知识库工具
|
||||
RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/MCP/提示词)
|
||||
SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活)
|
||||
ModelRetryMiddleware(), # 模型重试中间件
|
||||
TodoListMiddleware(),
|
||||
PatchToolCallsMiddleware(),
|
||||
|
||||
@ -1,19 +1,13 @@
|
||||
from deepagents.backends import CompositeBackend, StateBackend
|
||||
|
||||
from src.services.skill_resolver import normalize_selected_skills
|
||||
from src.services.skill_service import is_valid_skill_slug
|
||||
from src.agents.common.middlewares.skills_middleware import normalize_selected_skills
|
||||
|
||||
from .skills_backend import SelectedSkillsReadonlyBackend
|
||||
|
||||
|
||||
def _get_visible_skills_from_runtime(runtime) -> list[str]:
|
||||
"""获取运行时可见的 skills 列表"""
|
||||
context = getattr(runtime, "context", None)
|
||||
snapshot = getattr(context, "skill_session_snapshot", None)
|
||||
if isinstance(snapshot, dict):
|
||||
visible = snapshot.get("visible_skills")
|
||||
if isinstance(visible, list):
|
||||
return [slug for slug in visible if isinstance(slug, str) and is_valid_skill_slug(slug)]
|
||||
|
||||
selected = getattr(context, "skills", None) or []
|
||||
return normalize_selected_skills(selected)
|
||||
|
||||
|
||||
@ -12,7 +12,6 @@ from langgraph.graph.state import CompiledStateGraph
|
||||
|
||||
from src import config as sys_config
|
||||
from src.agents.common.context import BaseContext
|
||||
from src.services.skill_resolver import get_skill_options_from_db
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
@ -50,9 +49,6 @@ class BaseAgent:
|
||||
configurable_items = {}
|
||||
if include_configurable_items:
|
||||
configurable_items = self.context_schema.get_configurable_items()
|
||||
if "skills" in configurable_items:
|
||||
configurable_items["skills"] = dict(configurable_items["skills"])
|
||||
configurable_items["skills"]["options"] = await get_skill_options_from_db()
|
||||
|
||||
# Merge metadata with class attributes, metadata takes precedence
|
||||
return {
|
||||
|
||||
@ -1,61 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Annotated, Any, NotRequired
|
||||
from typing import Any
|
||||
|
||||
from deepagents.middleware.skills import SKILLS_SYSTEM_PROMPT
|
||||
from langchain.agents import AgentState
|
||||
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse
|
||||
from langchain.tools.tool_node import ToolCallRequest
|
||||
from langchain_core.messages import SystemMessage, ToolMessage
|
||||
from langgraph.types import Command
|
||||
from langchain_core.messages import SystemMessage
|
||||
|
||||
from src.agents.common import load_chat_model
|
||||
from src.agents.common.toolkits import get_all_tool_instances
|
||||
from src.services.mcp_service import get_enabled_mcp_tools
|
||||
from src.services.skill_resolver import (
|
||||
SkillSessionSnapshot,
|
||||
build_dependency_bundle,
|
||||
collect_prompt_metadata,
|
||||
normalize_selected_skills,
|
||||
resolve_session_snapshot,
|
||||
)
|
||||
from src.services.skill_service import is_valid_skill_slug
|
||||
from src.utils.datetime_utils import shanghai_now
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
|
||||
def _activated_skills_reducer(left: list[str] | None, right: list[str] | None) -> list[str]:
|
||||
merged: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for group in (left or [], right or []):
|
||||
for value in group:
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
slug = value.strip()
|
||||
if not slug or slug in seen:
|
||||
continue
|
||||
seen.add(slug)
|
||||
merged.append(slug)
|
||||
return merged
|
||||
|
||||
|
||||
class RuntimeConfigState(AgentState):
|
||||
activated_skills: NotRequired[Annotated[list[str], _activated_skills_reducer]]
|
||||
skill_session_snapshot: NotRequired[SkillSessionSnapshot]
|
||||
|
||||
|
||||
class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
"""运行时配置中间件 - 应用模型/工具/MCP/提示词配置
|
||||
|
||||
知识库工具已移至独立的 KnowledgeBaseMiddleware
|
||||
Skills 功能已移至独立的 SkillsMiddleware
|
||||
|
||||
支持自定义上下文字段名称,以便在不同场景(如主智能体/子智能体)使用不同的配置字段
|
||||
"""
|
||||
|
||||
state_schema = RuntimeConfigState
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@ -65,12 +31,9 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
tools_context_name: str = "tools",
|
||||
knowledges_context_name: str = "knowledges",
|
||||
mcps_context_name: str = "mcps",
|
||||
skills_context_name: str = "skills",
|
||||
enable_model_override: bool = True,
|
||||
enable_system_prompt_override: bool = True,
|
||||
enable_tools_override: bool = True,
|
||||
enable_skills_prompt_override: bool = True,
|
||||
skills_sources_for_prompt: list[str] | None = None,
|
||||
):
|
||||
"""初始化中间件
|
||||
|
||||
@ -81,12 +44,9 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
tools_context_name: 上下文中的工具列表字段名称(默认 "tools")
|
||||
knowledges_context_name: 上下文中的知识库列表字段名称(默认 "knowledges")
|
||||
mcps_context_name: 上下文中的 MCP 服务器列表字段名称(默认 "mcps")
|
||||
skills_context_name: 上下文中的 skills 列表字段名称(默认 "skills")
|
||||
enable_model_override: 是否允许覆盖模型配置(默认 True)
|
||||
enable_system_prompt_override: 是否允许覆盖系统提示词(默认 True)
|
||||
enable_tools_override: 是否允许覆盖工具列表(默认 True)
|
||||
enable_skills_prompt_override: 是否启用 skills 提示段注入(默认 True)
|
||||
skills_sources_for_prompt: skills 来源路径(用于提示词展示,默认 ["/skills/"])
|
||||
"""
|
||||
super().__init__()
|
||||
# 存储自定义字段名称
|
||||
@ -95,13 +55,10 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
self.tools_context_name = tools_context_name
|
||||
self.knowledges_context_name = knowledges_context_name
|
||||
self.mcps_context_name = mcps_context_name
|
||||
self.skills_context_name = skills_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.enable_skills_prompt_override = enable_skills_prompt_override
|
||||
self.skills_sources_for_prompt = skills_sources_for_prompt or ["/skills/"]
|
||||
|
||||
self.tools: list[Any] = []
|
||||
# 预加载工具列表(仅当启用工具覆盖时)
|
||||
@ -118,48 +75,18 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
logger.debug(
|
||||
f"Initialized RuntimeConfigMiddleware with custom field names: model={model_context_name}, "
|
||||
f"system_prompt={system_prompt_context_name}, tools={tools_context_name}, "
|
||||
f"knowledges={knowledges_context_name}, mcps={mcps_context_name}, "
|
||||
f"skills={skills_context_name}"
|
||||
f"knowledges={knowledges_context_name}, mcps={mcps_context_name}"
|
||||
)
|
||||
|
||||
async def abefore_agent(self, state: RuntimeConfigState, runtime) -> dict[str, Any] | None:
|
||||
runtime_context = runtime.context
|
||||
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
||||
selected_skills = normalize_selected_skills(configured_skills)
|
||||
|
||||
try:
|
||||
snapshot = await resolve_session_snapshot(selected_skills)
|
||||
except Exception as e:
|
||||
logger.warning(f"RuntimeConfigMiddleware: failed to resolve skill snapshot in abefore_agent: {e}")
|
||||
snapshot = {
|
||||
"selected_skills": selected_skills,
|
||||
"visible_skills": [],
|
||||
"prompt_metadata": {},
|
||||
"dependency_map": {},
|
||||
}
|
||||
|
||||
setattr(runtime_context, "skill_session_snapshot", snapshot)
|
||||
|
||||
if not self.enable_system_prompt_override or not self.enable_skills_prompt_override:
|
||||
return None
|
||||
if getattr(runtime_context, "_skills_prompt_injected", False):
|
||||
return None
|
||||
if not snapshot.get("visible_skills"):
|
||||
return None
|
||||
|
||||
skills_meta = collect_prompt_metadata(snapshot, snapshot.get("visible_skills") or [])
|
||||
skills_section = self._build_skills_section(skills_meta)
|
||||
base_prompt = getattr(runtime_context, self.system_prompt_context_name, "") or ""
|
||||
merged_prompt = f"{base_prompt}\n\n{skills_section}" if base_prompt else skills_section
|
||||
setattr(runtime_context, self.system_prompt_context_name, merged_prompt)
|
||||
setattr(runtime_context, "_skills_prompt_injected", True)
|
||||
async def abefore_agent(self, state, runtime) -> dict[str, Any] | None:
|
||||
# abefore_agent 在 RuntimeConfigMiddleware 中暂无额外逻辑
|
||||
# Skills 相关逻辑已移至 SkillsMiddleware
|
||||
return None
|
||||
|
||||
async def awrap_model_call(
|
||||
self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
|
||||
) -> ModelResponse:
|
||||
runtime_context = request.runtime.context
|
||||
snapshot = self._get_skill_snapshot_from_context(runtime_context)
|
||||
overrides: dict[str, Any] = {}
|
||||
|
||||
# 1. 模型覆盖(可选)
|
||||
@ -168,17 +95,10 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
overrides["model"] = model
|
||||
|
||||
# 2. 工具覆盖(可选)
|
||||
# 注意:Skills 依赖的工具加载已移至 SkillsMiddleware
|
||||
if self.enable_tools_override:
|
||||
state = request.state if isinstance(request.state, dict) else {}
|
||||
activated_skills = state.get("activated_skills", [])
|
||||
if not isinstance(activated_skills, list):
|
||||
activated_skills = []
|
||||
deps_bundle = build_dependency_bundle(snapshot, activated_skills)
|
||||
enabled_tools = await self.get_tools_from_context(
|
||||
runtime_context,
|
||||
extra_tool_names=deps_bundle["tools"],
|
||||
extra_mcps=deps_bundle["mcps"],
|
||||
)
|
||||
# 获取上下文配置的工具
|
||||
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}
|
||||
@ -207,13 +127,7 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
|
||||
return await handler(request)
|
||||
|
||||
async def get_tools_from_context(
|
||||
self,
|
||||
context,
|
||||
*,
|
||||
extra_tool_names: list[str] | None = None,
|
||||
extra_mcps: list[str] | None = None,
|
||||
) -> list:
|
||||
async def get_tools_from_context(self, context) -> list:
|
||||
"""从上下文配置中获取工具列表"""
|
||||
selected_tools = []
|
||||
selected_tool_names: set[str] = set()
|
||||
@ -224,9 +138,6 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
for tool_name in tools:
|
||||
if isinstance(tool_name, str):
|
||||
all_tool_names.append(tool_name)
|
||||
for tool_name in extra_tool_names or []:
|
||||
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:
|
||||
@ -244,9 +155,6 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
for server_name in mcps:
|
||||
if isinstance(server_name, str):
|
||||
all_mcp_names.append(server_name)
|
||||
for server_name in extra_mcps or []:
|
||||
if isinstance(server_name, str):
|
||||
all_mcp_names.append(server_name)
|
||||
|
||||
selected_mcp_servers: set[str] = set()
|
||||
for server_name in all_mcp_names:
|
||||
@ -262,119 +170,3 @@ class RuntimeConfigMiddleware(AgentMiddleware):
|
||||
logger.warning(f"RuntimeConfigMiddleware: failed to load mcp dependency '{server_name}': {e}")
|
||||
|
||||
return selected_tools
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Any],
|
||||
):
|
||||
result = await handler(request)
|
||||
if request.tool_call.get("name") != "read_file":
|
||||
return result
|
||||
|
||||
args = request.tool_call.get("args") or {}
|
||||
file_path = args.get("file_path") if isinstance(args, dict) else None
|
||||
slug = self._extract_skill_slug_from_skill_md_path(file_path)
|
||||
if not slug:
|
||||
return result
|
||||
if not self._is_visible_skill_slug(request, slug):
|
||||
logger.warning(f"RuntimeConfigMiddleware: deny skill activation for invisible slug: {slug}")
|
||||
return result
|
||||
logger.debug(f"RuntimeConfigMiddleware: activated skill by read_file: {slug}")
|
||||
return self._merge_activated_skill_update(result, slug)
|
||||
|
||||
def wrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Any],
|
||||
):
|
||||
result = handler(request)
|
||||
if request.tool_call.get("name") != "read_file":
|
||||
return result
|
||||
|
||||
args = request.tool_call.get("args") or {}
|
||||
file_path = args.get("file_path") if isinstance(args, dict) else None
|
||||
slug = self._extract_skill_slug_from_skill_md_path(file_path)
|
||||
if not slug:
|
||||
return result
|
||||
if not self._is_visible_skill_slug(request, slug):
|
||||
logger.warning(f"RuntimeConfigMiddleware: deny skill activation for invisible slug: {slug}")
|
||||
return result
|
||||
logger.debug(f"RuntimeConfigMiddleware: activated skill by read_file: {slug}")
|
||||
return self._merge_activated_skill_update(result, slug)
|
||||
|
||||
def _extract_skill_slug_from_skill_md_path(self, file_path: Any) -> str | None:
|
||||
if not isinstance(file_path, str):
|
||||
return None
|
||||
raw = file_path.strip()
|
||||
if not raw:
|
||||
return None
|
||||
pure = PurePosixPath(raw if raw.startswith("/") else f"/{raw}")
|
||||
parts = [p for p in pure.parts if p not in ("/", "")]
|
||||
if len(parts) != 3:
|
||||
return None
|
||||
if parts[0] != "skills" or parts[2] != "SKILL.md":
|
||||
return None
|
||||
slug = parts[1]
|
||||
if not is_valid_skill_slug(slug):
|
||||
return None
|
||||
return slug
|
||||
|
||||
def _merge_activated_skill_update(self, result: Any, slug: str):
|
||||
if isinstance(result, Command):
|
||||
update = dict(result.update or {})
|
||||
current = update.get("activated_skills") or []
|
||||
update["activated_skills"] = _activated_skills_reducer(current, [slug])
|
||||
return Command(graph=result.graph, update=update, resume=result.resume, goto=result.goto)
|
||||
|
||||
if isinstance(result, ToolMessage):
|
||||
return Command(update={"messages": [result], "activated_skills": [slug]})
|
||||
|
||||
return result
|
||||
|
||||
def _is_visible_skill_slug(self, request: ToolCallRequest, slug: str) -> bool:
|
||||
snapshot = self._get_skill_snapshot_from_context(request.runtime.context)
|
||||
if snapshot:
|
||||
visible_skills = snapshot.get("visible_skills")
|
||||
if isinstance(visible_skills, list):
|
||||
return slug in visible_skills
|
||||
|
||||
configured_skills = getattr(request.runtime.context, self.skills_context_name, None) or []
|
||||
normalized = normalize_selected_skills(configured_skills)
|
||||
return slug in normalized
|
||||
|
||||
def _get_skill_snapshot_from_context(self, context: Any) -> SkillSessionSnapshot | None:
|
||||
snapshot = getattr(context, "skill_session_snapshot", None)
|
||||
if not isinstance(snapshot, dict):
|
||||
return None
|
||||
visible_skills = snapshot.get("visible_skills")
|
||||
selected_skills = snapshot.get("selected_skills")
|
||||
if not isinstance(visible_skills, list) or not isinstance(selected_skills, list):
|
||||
return None
|
||||
return snapshot
|
||||
|
||||
def _format_skills_locations(self, sources: list[str]) -> str:
|
||||
locations = []
|
||||
for i, source_path in enumerate(sources):
|
||||
name = PurePosixPath(source_path.rstrip("/")).name.capitalize()
|
||||
suffix = " (higher priority)" if i == len(sources) - 1 else ""
|
||||
locations.append(f"**{name} Skills**: `{source_path}`{suffix}")
|
||||
return "\n".join(locations)
|
||||
|
||||
def _format_skills_list(self, skills_meta: list[dict[str, str]]) -> str:
|
||||
if not skills_meta:
|
||||
return f"(No skills available yet. You can create skills in {' or '.join(self.skills_sources_for_prompt)})"
|
||||
|
||||
lines = []
|
||||
for skill in skills_meta:
|
||||
lines.append(f"- **{skill['name']}**: {skill['description']}")
|
||||
lines.append(f" -> Read `{skill['path']}` for full instructions")
|
||||
return "\n".join(lines)
|
||||
|
||||
def _build_skills_section(self, skills_meta: list[dict[str, str]]) -> str:
|
||||
skills_locations = self._format_skills_locations(self.skills_sources_for_prompt)
|
||||
skills_list = self._format_skills_list(skills_meta)
|
||||
return SKILLS_SYSTEM_PROMPT.format(
|
||||
skills_locations=skills_locations,
|
||||
skills_list=skills_list,
|
||||
)
|
||||
|
||||
475
src/agents/common/middlewares/skills_middleware.py
Normal file
475
src/agents/common/middlewares/skills_middleware.py
Normal file
@ -0,0 +1,475 @@
|
||||
"""Skills 中间件 - 处理 skills 提示词注入、依赖展开、动态激活"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Annotated, Any, NotRequired, TypedDict
|
||||
|
||||
from deepagents.middleware.skills import SKILLS_SYSTEM_PROMPT
|
||||
from langchain.agents import AgentState
|
||||
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse
|
||||
from langchain.tools.tool_node import ToolCallRequest
|
||||
from langgraph.types import Command
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.repositories.skill_repository import SkillRepository
|
||||
from src.services.mcp_service import get_enabled_mcp_tools
|
||||
from src.services.skill_service import _normalize_string_list, is_valid_skill_slug
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# 类型定义
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class SkillPromptMetadata(TypedDict):
|
||||
name: str
|
||||
description: str
|
||||
path: str
|
||||
|
||||
|
||||
class SkillDependencyNode(TypedDict):
|
||||
tools: list[str]
|
||||
mcps: list[str]
|
||||
skills: list[str]
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# 运行时数据加载函数
|
||||
# =============================================================================
|
||||
|
||||
|
||||
async def _list_skills_from_db(db: AsyncSession | None = None) -> list:
|
||||
"""从数据库加载 skills 列表"""
|
||||
if db is not None:
|
||||
repo = SkillRepository(db)
|
||||
return await repo.list_all()
|
||||
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
repo = SkillRepository(session)
|
||||
return await repo.list_all()
|
||||
|
||||
|
||||
async def get_prompt_metadata(db: AsyncSession | None = None) -> dict[str, SkillPromptMetadata]:
|
||||
"""获取提示词元数据(直接从数据库加载)"""
|
||||
skills = await _list_skills_from_db(db)
|
||||
return {
|
||||
item.slug: {
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
"path": f"/skills/{item.slug}/SKILL.md",
|
||||
}
|
||||
for item in skills
|
||||
}
|
||||
|
||||
|
||||
async def get_dependency_map(db: AsyncSession | None = None) -> dict[str, SkillDependencyNode]:
|
||||
"""获取依赖关系映射(直接从数据库加载)"""
|
||||
skills = await _list_skills_from_db(db)
|
||||
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 []),
|
||||
}
|
||||
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)
|
||||
if not ordered_roots:
|
||||
return []
|
||||
|
||||
result: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
def dfs(slug: str, stack: set[str]) -> None:
|
||||
if slug in stack:
|
||||
logger.warning(f"Cycle detected in skill dependencies, skip: {' -> '.join([*stack, slug])}")
|
||||
return
|
||||
if slug in seen:
|
||||
return
|
||||
|
||||
node = dependency_map.get(slug)
|
||||
if not node:
|
||||
logger.warning(f"Skill dependency target not found in DB, skip: {slug}")
|
||||
return
|
||||
|
||||
seen.add(slug)
|
||||
result.append(slug)
|
||||
next_stack = set(stack)
|
||||
next_stack.add(slug)
|
||||
for dep in node.get("skills", []):
|
||||
dfs(dep, next_stack)
|
||||
|
||||
for root in ordered_roots:
|
||||
dfs(root, set())
|
||||
return result
|
||||
|
||||
|
||||
def _activated_skills_reducer(left: list[str] | None, right: list[str] | None) -> list[str]:
|
||||
"""合并 activated_skills 列表"""
|
||||
merged: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for group in (left or [], right or []):
|
||||
for value in group:
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
slug = value.strip()
|
||||
if not slug or slug in seen:
|
||||
continue
|
||||
seen.add(slug)
|
||||
merged.append(slug)
|
||||
return merged
|
||||
|
||||
|
||||
class SkillsState(AgentState):
|
||||
"""Skills 状态定义"""
|
||||
|
||||
activated_skills: NotRequired[Annotated[list[str], _activated_skills_reducer]]
|
||||
|
||||
|
||||
class SkillsMiddleware(AgentMiddleware):
|
||||
"""Skills 中间件 - 处理 skills 提示词注入、依赖展开、动态激活
|
||||
|
||||
职责:
|
||||
- Skills 提示词注入(直接从数据库加载)
|
||||
- 依赖展开(用户配置 + 动态激活)
|
||||
- 工具/MCP 动态加载
|
||||
"""
|
||||
|
||||
state_schema = SkillsState
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
skills_context_name: str = "skills",
|
||||
enable_skills_prompt: bool = True,
|
||||
skills_sources_for_prompt: list[str] | None = None,
|
||||
):
|
||||
"""初始化中间件
|
||||
|
||||
Args:
|
||||
skills_context_name: 上下文中的 skills 列表字段名称(默认 "skills")
|
||||
enable_skills_prompt: 是否启用 skills 提示段注入(默认 True)
|
||||
skills_sources_for_prompt: skills 来源路径(用于提示词展示,默认 ["/skills/"])
|
||||
"""
|
||||
super().__init__()
|
||||
self.skills_context_name = skills_context_name
|
||||
self.enable_skills_prompt = enable_skills_prompt
|
||||
self.skills_sources_for_prompt = skills_sources_for_prompt or ["/skills/"]
|
||||
|
||||
async def abefore_agent(self, state: SkillsState, runtime) -> dict[str, Any] | None:
|
||||
"""在 agent 执行前注入 skills 提示词"""
|
||||
runtime_context = runtime.context
|
||||
|
||||
# 检查是否需要注入
|
||||
if not self.enable_skills_prompt:
|
||||
return None
|
||||
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:
|
||||
return None
|
||||
|
||||
# 计算 visible_skills
|
||||
visible_skills = expand_skill_closure(selected_skills, dependency_map)
|
||||
|
||||
if not visible_skills:
|
||||
return None
|
||||
|
||||
# 收集提示词元数据并构建提示段
|
||||
skills_meta = await self._collect_prompt_metadata(visible_skills)
|
||||
skills_section = self._build_skills_section(skills_meta)
|
||||
|
||||
# 注入提示词
|
||||
base_prompt = getattr(runtime_context, "system_prompt", "") or ""
|
||||
merged_prompt = f"{base_prompt}\n\n{skills_section}" if base_prompt else skills_section
|
||||
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(
|
||||
self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
|
||||
) -> ModelResponse:
|
||||
"""包装模型调用,处理动态激活和依赖展开"""
|
||||
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)
|
||||
|
||||
# 4. 更新 runtime_context 中的 visible_skills
|
||||
setattr(runtime_context, "_visible_skills", visible_skills)
|
||||
|
||||
# 5. 构建依赖包
|
||||
deps_bundle = await self._build_dependency_bundle(visible_skills)
|
||||
|
||||
# 6. 加载依赖的工具
|
||||
if deps_bundle["tools"] or deps_bundle["mcps"]:
|
||||
enabled_tools = await self._get_tools_from_context(
|
||||
runtime_context,
|
||||
extra_tool_names=deps_bundle["tools"],
|
||||
extra_mcps=deps_bundle["mcps"],
|
||||
)
|
||||
|
||||
# 合并工具
|
||||
if enabled_tools:
|
||||
existing_tools = list(request.tools or [])
|
||||
enabled_tool_names = {t.name for t in enabled_tools}
|
||||
merged_tools = []
|
||||
for t_bind in existing_tools:
|
||||
if t_bind.name in enabled_tool_names:
|
||||
merged_tools.append(t_bind)
|
||||
if merged_tools:
|
||||
request = request.override(tools=merged_tools)
|
||||
|
||||
return await handler(request)
|
||||
|
||||
async def _build_dependency_bundle(self, visible_skills: list[str]) -> dict[str, list[str]]:
|
||||
"""根据 visible_skills 构建依赖包"""
|
||||
dependency_map = await get_dependency_map()
|
||||
|
||||
tools: list[str] = []
|
||||
mcps: list[str] = []
|
||||
seen_tools: set[str] = set()
|
||||
seen_mcps: set[str] = set()
|
||||
|
||||
for slug in visible_skills:
|
||||
dep = dependency_map.get(slug, {})
|
||||
for tool_name in dep.get("tools", []):
|
||||
if tool_name in seen_tools:
|
||||
continue
|
||||
seen_tools.add(tool_name)
|
||||
tools.append(tool_name)
|
||||
for mcp_name in dep.get("mcps", []):
|
||||
if mcp_name in seen_mcps:
|
||||
continue
|
||||
seen_mcps.add(mcp_name)
|
||||
mcps.append(mcp_name)
|
||||
|
||||
return {"tools": tools, "mcps": mcps, "skills": visible_skills}
|
||||
|
||||
async def _collect_prompt_metadata(self, slugs: list[str]) -> list[SkillPromptMetadata]:
|
||||
"""收集指定 slugs 的提示词元数据"""
|
||||
prompt_metadata = await get_prompt_metadata()
|
||||
|
||||
result: list[SkillPromptMetadata] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
for slug in slugs:
|
||||
if not isinstance(slug, str):
|
||||
continue
|
||||
normalized = slug.strip()
|
||||
if not normalized or normalized in seen:
|
||||
continue
|
||||
seen.add(normalized)
|
||||
|
||||
item = prompt_metadata.get(normalized)
|
||||
if not item:
|
||||
logger.debug(f"Skill slug not found in prompt metadata, skip: {normalized}")
|
||||
continue
|
||||
result.append(dict(item))
|
||||
|
||||
return result
|
||||
|
||||
async def _get_tools_from_context(
|
||||
self,
|
||||
context,
|
||||
*,
|
||||
extra_tool_names: list[str] | None = None,
|
||||
extra_mcps: list[str] | None = None,
|
||||
) -> list:
|
||||
"""从上下文配置中获取工具列表"""
|
||||
import asyncio
|
||||
|
||||
selected_tools = []
|
||||
|
||||
# 1. 工具(从 extra_tool_names 获取)
|
||||
all_tool_names: list[str] = []
|
||||
for tool_name in extra_tool_names or []:
|
||||
if isinstance(tool_name, str):
|
||||
all_tool_names.append(tool_name)
|
||||
|
||||
# 这里简化处理:假设工具已经在其他 middleware 中加载
|
||||
# SkillsMiddleware 主要负责 MCP 工具的加载
|
||||
|
||||
# 2. MCP 工具(并行加载)
|
||||
mcps = getattr(context, "mcps", None) or []
|
||||
all_mcp_names: list[str] = []
|
||||
for server_name in mcps:
|
||||
if isinstance(server_name, str):
|
||||
all_mcp_names.append(server_name)
|
||||
for server_name in extra_mcps or []:
|
||||
if isinstance(server_name, str):
|
||||
all_mcp_names.append(server_name)
|
||||
|
||||
# 去重
|
||||
unique_mcp_names = list(dict.fromkeys(all_mcp_names))
|
||||
|
||||
async def load_mcp_tools(server_name: str) -> list:
|
||||
"""加载单个 MCP 服务器的工具"""
|
||||
try:
|
||||
mcp_tools = await get_enabled_mcp_tools(server_name)
|
||||
if not mcp_tools:
|
||||
logger.warning(f"SkillsMiddleware: mcp dependency unavailable, skip: {server_name}")
|
||||
return mcp_tools
|
||||
except Exception as e:
|
||||
logger.warning(f"SkillsMiddleware: failed to load mcp dependency '{server_name}': {e}")
|
||||
return []
|
||||
|
||||
# 并行加载所有 MCP 工具
|
||||
results = await asyncio.gather(*[load_mcp_tools(name) for name in unique_mcp_names])
|
||||
for tools in results:
|
||||
selected_tools.extend(tools)
|
||||
|
||||
return selected_tools
|
||||
|
||||
def _process_tool_call_result(self, result: Any, request: ToolCallRequest) -> Any:
|
||||
"""处理工具调用结果,检查并处理 skill 动态激活"""
|
||||
if request.tool_call.get("name") != "read_file":
|
||||
return result
|
||||
|
||||
args = request.tool_call.get("args") or {}
|
||||
file_path = args.get("file_path") if isinstance(args, dict) else None
|
||||
slug = self._extract_skill_slug_from_skill_md_path(file_path)
|
||||
|
||||
if not slug:
|
||||
return result
|
||||
|
||||
if not self._is_visible_skill_slug(request, slug):
|
||||
logger.warning(f"SkillsMiddleware: deny skill activation for invisible slug: {slug}")
|
||||
return result
|
||||
|
||||
logger.debug(f"SkillsMiddleware: activated skill by read_file: {slug}")
|
||||
return self._merge_activated_skill_update(result, slug)
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Any],
|
||||
):
|
||||
"""包装工具调用,处理 skill 动态激活"""
|
||||
result = await handler(request)
|
||||
return self._process_tool_call_result(result, request)
|
||||
|
||||
def wrap_tool_call(
|
||||
self,
|
||||
request: ToolCallRequest,
|
||||
handler: Callable[[ToolCallRequest], Any],
|
||||
):
|
||||
"""同步版本的工具调用包装"""
|
||||
result = handler(request)
|
||||
return self._process_tool_call_result(result, request)
|
||||
|
||||
def _extract_skill_slug_from_skill_md_path(self, file_path: Any) -> str | None:
|
||||
"""从文件路径中提取 skill slug"""
|
||||
if not isinstance(file_path, str):
|
||||
return None
|
||||
raw = file_path.strip()
|
||||
if not raw:
|
||||
return None
|
||||
pure = PurePosixPath(raw if raw.startswith("/") else f"/{raw}")
|
||||
parts = [p for p in pure.parts if p not in ("/", "")]
|
||||
if len(parts) != 3:
|
||||
return None
|
||||
if parts[0] != "skills" or parts[2] != "SKILL.md":
|
||||
return None
|
||||
slug = parts[1]
|
||||
if not is_valid_skill_slug(slug):
|
||||
return None
|
||||
return slug
|
||||
|
||||
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
|
||||
|
||||
def _merge_activated_skill_update(self, result: Any, slug: str):
|
||||
"""合并动态激活的 skill 更新"""
|
||||
from langchain_core.messages import ToolMessage
|
||||
|
||||
if isinstance(result, Command):
|
||||
update = dict(result.update or {})
|
||||
current = update.get("activated_skills") or []
|
||||
update["activated_skills"] = _activated_skills_reducer(current, [slug])
|
||||
return Command(graph=result.graph, update=update, resume=result.resume, goto=result.goto)
|
||||
|
||||
if isinstance(result, ToolMessage):
|
||||
return Command(update={"messages": [result], "activated_skills": [slug]})
|
||||
|
||||
return result
|
||||
|
||||
def _format_skills_locations(self, sources: list[str]) -> str:
|
||||
"""格式化 skills 位置信息"""
|
||||
locations = []
|
||||
for i, source_path in enumerate(sources):
|
||||
name = PurePosixPath(source_path.rstrip("/")).name.capitalize()
|
||||
suffix = " (higher priority)" if i == len(sources) - 1 else ""
|
||||
locations.append(f"**{name} Skills**: `{source_path}`{suffix}")
|
||||
return "\n".join(locations)
|
||||
|
||||
def _format_skills_list(self, skills_meta: list[dict[str, str]]) -> str:
|
||||
"""格式化 skills 列表"""
|
||||
if not skills_meta:
|
||||
return f"(No skills available yet. You can create skills in {' or '.join(self.skills_sources_for_prompt)})"
|
||||
|
||||
lines = []
|
||||
for skill in skills_meta:
|
||||
lines.append(f"- **{skill['name']}**: {skill['description']}")
|
||||
lines.append(f" -> Read `{skill['path']}` for full instructions")
|
||||
return "\n".join(lines)
|
||||
|
||||
def _build_skills_section(self, skills_meta: list[dict[str, str]]) -> str:
|
||||
"""构建 skills 提示段"""
|
||||
skills_locations = self._format_skills_locations(self.skills_sources_for_prompt)
|
||||
skills_list = self._format_skills_list(skills_meta)
|
||||
return SKILLS_SYSTEM_PROMPT.format(
|
||||
skills_locations=skills_locations,
|
||||
skills_list=skills_list,
|
||||
)
|
||||
@ -13,6 +13,7 @@ from src.agents.common import BaseAgent, load_chat_model
|
||||
from src.agents.common.backends import create_agent_composite_backend
|
||||
from src.agents.common.middlewares import RuntimeConfigMiddleware, SummaryOffloadMiddleware, save_attachments_to_fs
|
||||
from src.agents.common.middlewares.knowledge_base_middleware import KnowledgeBaseMiddleware
|
||||
from src.agents.common.middlewares.skills_middleware import SkillsMiddleware
|
||||
from src.agents.common.toolkits.buildin.tools import _create_tavily_search
|
||||
from src.services.mcp_service import get_tools_from_all_servers
|
||||
from src.utils import logger
|
||||
@ -153,6 +154,7 @@ class DeepAgent(BaseAgent):
|
||||
middleware=[
|
||||
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
|
||||
RuntimeConfigMiddleware(extra_tools=all_mcp_tools),
|
||||
SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活)
|
||||
save_attachments_to_fs, # 附件注入提示词
|
||||
TodoListMiddleware(),
|
||||
PatchToolCallsMiddleware(),
|
||||
|
||||
@ -1,226 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TypedDict
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src.repositories.skill_repository import SkillRepository
|
||||
from src.storage.postgres.manager import pg_manager
|
||||
from src.storage.postgres.models_business import Skill
|
||||
from src.utils.logging_config import logger
|
||||
|
||||
|
||||
class SkillPromptMetadata(TypedDict):
|
||||
name: str
|
||||
description: str
|
||||
path: str
|
||||
|
||||
|
||||
class SkillDependencyNode(TypedDict):
|
||||
tools: list[str]
|
||||
mcps: list[str]
|
||||
skills: list[str]
|
||||
|
||||
|
||||
class SkillSessionSnapshot(TypedDict):
|
||||
selected_skills: list[str]
|
||||
visible_skills: list[str]
|
||||
prompt_metadata: dict[str, SkillPromptMetadata]
|
||||
dependency_map: dict[str, SkillDependencyNode]
|
||||
|
||||
|
||||
def normalize_selected_skills(selected_skills: list[str] | None) -> list[str]:
|
||||
return _normalize_string_list(selected_skills)
|
||||
|
||||
|
||||
def is_snapshot_match_selected_skills(
|
||||
snapshot: SkillSessionSnapshot | None,
|
||||
selected_skills: list[str] | None,
|
||||
) -> bool:
|
||||
if not snapshot:
|
||||
return False
|
||||
current = snapshot.get("selected_skills")
|
||||
if not isinstance(current, list):
|
||||
return False
|
||||
return current == normalize_selected_skills(selected_skills)
|
||||
|
||||
|
||||
async def resolve_session_snapshot(
|
||||
selected_skills: list[str] | None,
|
||||
*,
|
||||
db: AsyncSession | None = None,
|
||||
) -> SkillSessionSnapshot:
|
||||
normalized_selected = normalize_selected_skills(selected_skills)
|
||||
skills = await _list_skills_from_db(db)
|
||||
prompt_metadata, dependency_map = _build_maps(skills)
|
||||
visible_skills = expand_skill_closure(normalized_selected, dependency_map)
|
||||
return {
|
||||
"selected_skills": normalized_selected,
|
||||
"visible_skills": visible_skills,
|
||||
"prompt_metadata": prompt_metadata,
|
||||
"dependency_map": dependency_map,
|
||||
}
|
||||
|
||||
|
||||
def collect_prompt_metadata(
|
||||
snapshot: SkillSessionSnapshot | None,
|
||||
slugs: list[str] | None,
|
||||
) -> list[SkillPromptMetadata]:
|
||||
if not snapshot or not slugs:
|
||||
return []
|
||||
|
||||
prompt_metadata = snapshot.get("prompt_metadata") or {}
|
||||
result: list[SkillPromptMetadata] = []
|
||||
seen: set[str] = set()
|
||||
for slug in slugs:
|
||||
if not isinstance(slug, str):
|
||||
continue
|
||||
normalized = slug.strip()
|
||||
if not normalized or normalized in seen:
|
||||
continue
|
||||
seen.add(normalized)
|
||||
|
||||
item = prompt_metadata.get(normalized)
|
||||
if not item:
|
||||
logger.debug(f"Skill slug not found in session snapshot, skip prompt metadata: {normalized}")
|
||||
continue
|
||||
result.append(dict(item))
|
||||
return result
|
||||
|
||||
|
||||
def build_dependency_bundle(
|
||||
snapshot: SkillSessionSnapshot | None,
|
||||
activated_slugs: list[str] | None,
|
||||
) -> dict[str, list[str]]:
|
||||
if not snapshot:
|
||||
return {"tools": [], "mcps": [], "skills": []}
|
||||
|
||||
dependency_map = snapshot.get("dependency_map") or {}
|
||||
closure = expand_skill_closure(activated_slugs or [], dependency_map)
|
||||
tools: list[str] = []
|
||||
mcps: list[str] = []
|
||||
seen_tools: set[str] = set()
|
||||
seen_mcps: set[str] = set()
|
||||
|
||||
for slug in closure:
|
||||
dep = dependency_map.get(slug, {})
|
||||
for tool_name in dep.get("tools", []):
|
||||
if tool_name in seen_tools:
|
||||
continue
|
||||
seen_tools.add(tool_name)
|
||||
tools.append(tool_name)
|
||||
for mcp_name in dep.get("mcps", []):
|
||||
if mcp_name in seen_mcps:
|
||||
continue
|
||||
seen_mcps.add(mcp_name)
|
||||
mcps.append(mcp_name)
|
||||
|
||||
return {"tools": tools, "mcps": mcps, "skills": closure}
|
||||
|
||||
|
||||
def expand_skill_closure(
|
||||
slugs: list[str] | None,
|
||||
dependency_map: dict[str, SkillDependencyNode],
|
||||
) -> list[str]:
|
||||
ordered_roots = _normalize_string_list(slugs)
|
||||
if not ordered_roots:
|
||||
return []
|
||||
|
||||
result: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
def dfs(slug: str, stack: set[str]) -> None:
|
||||
if slug in stack:
|
||||
logger.warning(f"Cycle detected in skill dependencies, skip: {' -> '.join([*stack, slug])}")
|
||||
return
|
||||
if slug in seen:
|
||||
return
|
||||
|
||||
node = dependency_map.get(slug)
|
||||
if not node:
|
||||
logger.warning(f"Skill dependency target not found in DB snapshot, skip: {slug}")
|
||||
return
|
||||
|
||||
seen.add(slug)
|
||||
result.append(slug)
|
||||
next_stack = set(stack)
|
||||
next_stack.add(slug)
|
||||
for dep in node.get("skills", []):
|
||||
dfs(dep, next_stack)
|
||||
|
||||
for root in ordered_roots:
|
||||
dfs(root, set())
|
||||
return result
|
||||
|
||||
|
||||
async def get_skill_options_from_db(
|
||||
*,
|
||||
db: AsyncSession | None = None,
|
||||
) -> list[dict[str, str]]:
|
||||
items = await _list_skills_from_db(db)
|
||||
return [
|
||||
{
|
||||
"id": item.slug,
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
}
|
||||
for item in items
|
||||
]
|
||||
|
||||
|
||||
async def get_skill_slug_set_from_db(
|
||||
*,
|
||||
db: AsyncSession | None = None,
|
||||
) -> set[str]:
|
||||
items = await _list_skills_from_db(db)
|
||||
return {item.slug for item in items}
|
||||
|
||||
|
||||
def _normalize_string_list(values: list[str] | None) -> list[str]:
|
||||
if not values:
|
||||
return []
|
||||
normalized: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for value in values:
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
item = value.strip()
|
||||
if not item or item in seen:
|
||||
continue
|
||||
seen.add(item)
|
||||
normalized.append(item)
|
||||
return normalized
|
||||
|
||||
|
||||
def _build_maps(skills: list[Skill]) -> tuple[dict[str, SkillPromptMetadata], dict[str, SkillDependencyNode]]:
|
||||
prompt_metadata: dict[str, SkillPromptMetadata] = {}
|
||||
dependency_map: dict[str, SkillDependencyNode] = {}
|
||||
for item in skills:
|
||||
prompt_metadata[item.slug] = {
|
||||
"name": item.name,
|
||||
"description": item.description,
|
||||
"path": f"/skills/{item.slug}/SKILL.md",
|
||||
}
|
||||
dependency_map[item.slug] = {
|
||||
"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 prompt_metadata, dependency_map
|
||||
|
||||
|
||||
async def _list_skills_from_db(db: AsyncSession | None) -> list[Skill]:
|
||||
if db is not None:
|
||||
repo = SkillRepository(db)
|
||||
return await repo.list_all()
|
||||
|
||||
try:
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
repo = SkillRepository(session)
|
||||
return await repo.list_all()
|
||||
except RuntimeError:
|
||||
# 在非 FastAPI 生命周期场景(如 worker/脚本)按需初始化
|
||||
pg_manager.initialize()
|
||||
async with pg_manager.get_async_session_context() as session:
|
||||
repo = SkillRepository(session)
|
||||
return await repo.list_all()
|
||||
@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
import shutil
|
||||
import tempfile
|
||||
@ -73,13 +74,6 @@ def is_valid_skill_slug(slug: str) -> bool:
|
||||
return bool(SKILL_SLUG_PATTERN.match(slug.strip()))
|
||||
|
||||
|
||||
def validate_skill_slug(slug: str) -> str:
|
||||
normalized = slug.strip() if isinstance(slug, str) else ""
|
||||
if not is_valid_skill_slug(normalized):
|
||||
raise ValueError("无效 skill slug")
|
||||
return normalized
|
||||
|
||||
|
||||
def get_skills_root_dir() -> Path:
|
||||
root = Path(sys_config.save_dir) / "skills"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
@ -87,19 +81,26 @@ def get_skills_root_dir() -> Path:
|
||||
|
||||
|
||||
async def get_skill_dependency_options(db: AsyncSession) -> dict[str, list[str] | list[dict]]:
|
||||
repo = SkillRepository(db)
|
||||
items = await repo.list_all()
|
||||
|
||||
# 获取所有工具(不仅仅是 buildin 工具),返回 id 和 name
|
||||
# 并行执行三个独立操作
|
||||
from src.services.tool_service import get_tool_metadata
|
||||
|
||||
all_tools = get_tool_metadata()
|
||||
# 返回工具列表,每个包含 id 和 name
|
||||
tool_list = [{"id": tool["id"], "name": tool.get("name", tool["id"])} for tool in all_tools]
|
||||
async def get_skills():
|
||||
repo = SkillRepository(db)
|
||||
return await repo.list_all()
|
||||
|
||||
def get_tools():
|
||||
all_tools = get_tool_metadata()
|
||||
return [{"id": tool["id"], "name": tool.get("name", tool["id"])} for tool in all_tools]
|
||||
|
||||
items, tool_list, mcp_names = await asyncio.gather(
|
||||
get_skills(),
|
||||
asyncio.to_thread(get_tools),
|
||||
asyncio.to_thread(get_mcp_server_names),
|
||||
)
|
||||
|
||||
return {
|
||||
"tools": tool_list,
|
||||
"mcps": get_mcp_server_names(),
|
||||
"mcps": mcp_names,
|
||||
"skills": [item.slug for item in items],
|
||||
}
|
||||
|
||||
@ -380,7 +381,10 @@ async def import_skill_zip(
|
||||
|
||||
|
||||
async def get_skill_or_raise(db: AsyncSession, slug: str) -> Skill:
|
||||
slug = validate_skill_slug(slug)
|
||||
slug = slug.strip() if isinstance(slug, str) else ""
|
||||
if not is_valid_skill_slug(slug):
|
||||
raise ValueError("无效 skill slug")
|
||||
|
||||
repo = SkillRepository(db)
|
||||
item = await repo.get_by_slug(slug)
|
||||
if not item:
|
||||
@ -434,19 +438,12 @@ async def create_skill_node(
|
||||
if not _is_text_path(target):
|
||||
raise ValueError("仅支持创建文本文件")
|
||||
|
||||
parsed_name: str | None = None
|
||||
parsed_desc: str | None = None
|
||||
if target.name == "SKILL.md" and target.parent == skill_dir:
|
||||
parsed_name, parsed_desc, _ = _parse_skill_markdown(content or "")
|
||||
if parsed_name != item.slug:
|
||||
raise ValueError("SKILL.md frontmatter.name 必须与 skill slug 一致")
|
||||
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 先写入文件,再更新元数据
|
||||
target.write_text(content or "", encoding="utf-8")
|
||||
|
||||
if parsed_name is not None and parsed_desc is not None:
|
||||
repo = SkillRepository(db)
|
||||
await repo.update_metadata(item, name=parsed_name, description=parsed_desc, updated_by=updated_by)
|
||||
await _update_skill_metadata_if_skills_md(db, item, content or "", skill_dir, target, updated_by)
|
||||
|
||||
|
||||
async def update_skill_file(
|
||||
@ -465,16 +462,24 @@ async def update_skill_file(
|
||||
if not _is_text_path(target):
|
||||
raise ValueError("仅支持编辑文本文件")
|
||||
|
||||
parsed_name = None
|
||||
parsed_desc = None
|
||||
await _update_skill_metadata_if_skills_md(db, item, content, skill_dir, target, updated_by)
|
||||
|
||||
target.write_text(content, encoding="utf-8")
|
||||
|
||||
|
||||
async def _update_skill_metadata_if_skills_md(
|
||||
db: AsyncSession,
|
||||
item: Skill,
|
||||
content: str,
|
||||
skill_dir: Path,
|
||||
target: Path,
|
||||
updated_by: str | None,
|
||||
) -> None:
|
||||
"""如果目标文件是 SKILL.md,则解析并更新元数据"""
|
||||
if target.name == "SKILL.md" and target.parent == skill_dir:
|
||||
parsed_name, parsed_desc, _ = _parse_skill_markdown(content)
|
||||
if parsed_name != item.slug:
|
||||
raise ValueError("SKILL.md frontmatter.name 必须与 skill slug 一致")
|
||||
|
||||
target.write_text(content, encoding="utf-8")
|
||||
|
||||
if parsed_name is not None and parsed_desc is not None:
|
||||
repo = SkillRepository(db)
|
||||
await repo.update_metadata(item, name=parsed_name, description=parsed_desc, updated_by=updated_by)
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user