From 393d3b6a32f1e873142448664c9f7c8c74ef8edc Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 5 Mar 2026 01:42:53 +0800 Subject: [PATCH] =?UTF-8?q?refactor(skills):=20=E9=80=9A=E8=BF=87=E5=BC=95?= =?UTF-8?q?=E5=85=A5SkillsMiddleware=E9=87=8D=E6=9E=84=E6=8A=80=E8=83=BD?= =?UTF-8?q?=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 添加了SkillsMiddleware来管理技能提示注入、依赖解析和动态激活。 - 从RuntimeConfigMiddleware中移除了与技能相关的逻辑并将其整合到SkillsMiddleware中。 - 删除了skill_resolver.py,因为其功能已集成到SkillsMiddleware中。 - 更新了各个组件以使用新的SkillsMiddleware进行技能管理。 - 改进了skill_service.py中技能和工具加载的异步处理。 --- src/agents/chatbot/graph.py | 2 + src/agents/common/backends/composite.py | 10 +- src/agents/common/base.py | 4 - .../middlewares/runtime_config_middleware.py | 230 +-------- .../common/middlewares/skills_middleware.py | 475 ++++++++++++++++++ src/agents/deep_agent/graph.py | 2 + src/services/skill_resolver.py | 226 --------- src/services/skill_service.py | 69 +-- 8 files changed, 529 insertions(+), 489 deletions(-) create mode 100644 src/agents/common/middlewares/skills_middleware.py delete mode 100644 src/services/skill_resolver.py diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index c3745874..6f73a9cc 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -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(), diff --git a/src/agents/common/backends/composite.py b/src/agents/common/backends/composite.py index d6119471..d5240261 100644 --- a/src/agents/common/backends/composite.py +++ b/src/agents/common/backends/composite.py @@ -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) diff --git a/src/agents/common/base.py b/src/agents/common/base.py index 1504a2b0..527050f9 100644 --- a/src/agents/common/base.py +++ b/src/agents/common/base.py @@ -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 { diff --git a/src/agents/common/middlewares/runtime_config_middleware.py b/src/agents/common/middlewares/runtime_config_middleware.py index 5bb2926f..73298759 100644 --- a/src/agents/common/middlewares/runtime_config_middleware.py +++ b/src/agents/common/middlewares/runtime_config_middleware.py @@ -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, - ) diff --git a/src/agents/common/middlewares/skills_middleware.py b/src/agents/common/middlewares/skills_middleware.py new file mode 100644 index 00000000..29872699 --- /dev/null +++ b/src/agents/common/middlewares/skills_middleware.py @@ -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, + ) diff --git a/src/agents/deep_agent/graph.py b/src/agents/deep_agent/graph.py index 6c098f5f..b54749dd 100644 --- a/src/agents/deep_agent/graph.py +++ b/src/agents/deep_agent/graph.py @@ -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(), diff --git a/src/services/skill_resolver.py b/src/services/skill_resolver.py deleted file mode 100644 index 45041c4b..00000000 --- a/src/services/skill_resolver.py +++ /dev/null @@ -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() diff --git a/src/services/skill_service.py b/src/services/skill_service.py index a7f671db..a907bb42 100644 --- a/src/services/skill_service.py +++ b/src/services/skill_service.py @@ -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)