fix: correct activated_skills state reducer annotation

This commit is contained in:
肖泽涛 2026-02-22 17:41:35 +08:00
parent 46b676db07
commit 85b28f3753

View File

@ -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 []: