refactor(skills): 通过引入SkillsMiddleware重构技能处理

- 添加了SkillsMiddleware来管理技能提示注入、依赖解析和动态激活。
- 从RuntimeConfigMiddleware中移除了与技能相关的逻辑并将其整合到SkillsMiddleware中。
- 删除了skill_resolver.py,因为其功能已集成到SkillsMiddleware中。
- 更新了各个组件以使用新的SkillsMiddleware进行技能管理。
- 改进了skill_service.py中技能和工具加载的异步处理。
This commit is contained in:
Wenjie Zhang 2026-03-05 01:42:53 +08:00
parent 24cbb41f4f
commit 393d3b6a32
8 changed files with 529 additions and 489 deletions

View File

@ -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(),

View File

@ -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)

View File

@ -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 {

View File

@ -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,
)

View 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,
)

View File

@ -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(),

View File

@ -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()

View File

@ -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)