fix: correct activated_skills state reducer annotation
This commit is contained in:
parent
46b676db07
commit
85b28f3753
@ -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 []:
|
||||
|
||||
Loading…
Reference in New Issue
Block a user