From 85b28f3753c9d14747d10ba4569ab23d388c3120 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=82=96=E6=B3=BD=E6=B6=9B?= Date: Sun, 22 Feb 2026 17:41:35 +0800 Subject: [PATCH] fix: correct activated_skills state reducer annotation --- .../middlewares/runtime_config_middleware.py | 158 ++++++++++++++++-- 1 file changed, 143 insertions(+), 15 deletions(-) diff --git a/src/agents/common/middlewares/runtime_config_middleware.py b/src/agents/common/middlewares/runtime_config_middleware.py index 8f4b4ece..be283b1d 100644 --- a/src/agents/common/middlewares/runtime_config_middleware.py +++ b/src/agents/common/middlewares/runtime_config_middleware.py @@ -2,20 +2,42 @@ from __future__ import annotations from collections.abc import Callable from pathlib import PurePosixPath -from typing import Any +from typing import Annotated, Any, NotRequired from deepagents.middleware.skills import SKILLS_SYSTEM_PROMPT +from langchain.agents import AgentState from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse -from langchain_core.messages import SystemMessage +from langchain.tools.tool_node import ToolCallRequest +from langchain_core.messages import SystemMessage, ToolMessage +from langgraph.types import Command from src.agents.common import load_chat_model from src.agents.common.tools import get_buildin_tools, get_kb_based_tools from src.services.mcp_service import get_enabled_mcp_tools -from src.services.skill_service import get_skill_prompt_metadata_by_slugs +from src.services.skill_service import get_dependency_bundle_for_activated_skills, get_skill_prompt_metadata_by_slugs 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]] + + class RuntimeConfigMiddleware(AgentMiddleware): """运行时配置中间件 - 应用模型/工具/知识库/MCP/提示词配置 @@ -24,6 +46,7 @@ class RuntimeConfigMiddleware(AgentMiddleware): 支持自定义上下文字段名称,以便在不同场景(如主智能体/子智能体)使用不同的配置字段 """ + state_schema = RuntimeConfigState def __init__( self, @@ -104,13 +127,23 @@ class RuntimeConfigMiddleware(AgentMiddleware): # 2. 工具覆盖(可选) if self.enable_tools_override: - enabled_tools = await self.get_tools_from_context(runtime_context) + activated_skills = request.state.get("activated_skills", []) if request.state else [] + if not isinstance(activated_skills, list): + activated_skills = [] + deps_bundle = get_dependency_bundle_for_activated_skills(activated_skills) + enabled_tools = await self.get_tools_from_context( + runtime_context, + extra_tool_names=deps_bundle["tools"], + extra_mcps=deps_bundle["mcps"], + ) existing_tools = list(request.tools or []) + enabled_tool_names = {t.name for t in enabled_tools} + managed_tool_names = {t.name for t in self.tools} merged_tools = [] for t_bind in existing_tools: # (1) 已启用的工具保留 # (2) 非本中间件管理的工具保留 - if t_bind.name in [t.name for t in enabled_tools] or t_bind.name not in [t.name for t in self.tools]: + if t_bind.name in enabled_tool_names or t_bind.name not in managed_tool_names: merged_tools.append(t_bind) overrides["tools"] = merged_tools logger.debug(f"RuntimeConfigMiddleware selected tools: {[t.name for t in merged_tools]}") @@ -142,17 +175,36 @@ class RuntimeConfigMiddleware(AgentMiddleware): return await handler(request) - async def get_tools_from_context(self, context) -> list: + async def get_tools_from_context( + self, + context, + *, + extra_tool_names: list[str] | None = None, + extra_mcps: list[str] | None = None, + ) -> list: """从上下文配置中获取工具列表""" selected_tools = [] + selected_tool_names: set[str] = set() # 1. 基础工具 (从 context.tools 中筛选) - tools = getattr(context, self.tools_context_name, None) - if tools: - tools_map = {t.name: t for t in self.tools} - for tool_name in tools: - if tool_name in tools_map: - selected_tools.append(tools_map[tool_name]) + tools = getattr(context, self.tools_context_name, None) or [] + all_tool_names = [] + for tool_name in tools: + if isinstance(tool_name, str): + all_tool_names.append(tool_name) + 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: + if tool_name in selected_tool_names: + continue + if tool_name in tools_map: + selected_tools.append(tools_map[tool_name]) + selected_tool_names.add(tool_name) + continue + logger.warning(f"RuntimeConfigMiddleware: tool dependency not found, skip: {tool_name}") # 2. 知识库工具 knowledges = getattr(context, self.knowledges_context_name, None) @@ -161,14 +213,90 @@ class RuntimeConfigMiddleware(AgentMiddleware): selected_tools.extend(kb_tools) # 3. MCP 工具(使用统一入口,自动过滤 disabled_tools) - mcps = getattr(context, self.mcps_context_name, None) - if mcps: - for server_name in mcps: + mcps = getattr(context, self.mcps_context_name, None) or [] + all_mcp_names: list[str] = [] + for server_name in mcps: + if isinstance(server_name, str): + all_mcp_names.append(server_name) + 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: + if server_name in selected_mcp_servers: + continue + selected_mcp_servers.add(server_name) + try: mcp_tools = await get_enabled_mcp_tools(server_name) + if not mcp_tools: + logger.warning(f"RuntimeConfigMiddleware: mcp dependency unavailable, skip: {server_name}") selected_tools.extend(mcp_tools) + except Exception as e: + logger.warning(f"RuntimeConfigMiddleware: failed to load mcp dependency '{server_name}': {e}") return selected_tools + 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 + 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 + 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 + return parts[1] + + 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 _supports_skill_prompt(self, request: ModelRequest) -> bool: """仅当请求工具中包含 read_file 时,才注入 skills 指引。""" for tool in request.tools or []: