fix(skills): 修复 skills 加载以及按需暴露问题,并新增 reporter 技能,已可以替代数据库报表智能体
This commit is contained in:
parent
96ce4dbe8a
commit
186012e5e8
@ -93,7 +93,7 @@ class BaseContext:
|
|||||||
metadata={
|
metadata={
|
||||||
"name": "Skills",
|
"name": "Skills",
|
||||||
"options": [],
|
"options": [],
|
||||||
"description": "可选技能列表(由超级管理员维护)。运行时仅挂载并只读暴露选中的 skills。",
|
"description": "可选技能列表(由超级管理员维护)。运行时仅挂载并只读暴露选中的 skills。技能依赖的工具和 MCP 服务器也会被自动挂载。",
|
||||||
"type": "list",
|
"type": "list",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
@ -13,6 +13,7 @@ from langchain.tools.tool_node import ToolCallRequest
|
|||||||
from langgraph.types import Command
|
from langgraph.types import Command
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.agents.common.toolkits import get_all_tool_instances
|
||||||
from src.repositories.skill_repository import SkillRepository
|
from src.repositories.skill_repository import SkillRepository
|
||||||
from src.services.mcp_service import get_enabled_mcp_tools
|
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.services.skill_service import _normalize_string_list, is_valid_skill_slug
|
||||||
@ -171,6 +172,21 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
self.skills_context_name = skills_context_name
|
self.skills_context_name = skills_context_name
|
||||||
self.enable_skills_prompt = enable_skills_prompt
|
self.enable_skills_prompt = enable_skills_prompt
|
||||||
self.skills_sources_for_prompt = skills_sources_for_prompt or ["/skills/"]
|
self.skills_sources_for_prompt = skills_sources_for_prompt or ["/skills/"]
|
||||||
|
# 实例级缓存:避免每次模型调用都查数据库
|
||||||
|
self._dependency_map_cache: dict[str, SkillDependencyNode] | None = None
|
||||||
|
self._prompt_metadata_cache: dict[str, SkillPromptMetadata] | None = None
|
||||||
|
|
||||||
|
async def _get_dependency_map_cached(self) -> dict[str, SkillDependencyNode]:
|
||||||
|
"""获取依赖映射(带缓存)"""
|
||||||
|
if self._dependency_map_cache is None:
|
||||||
|
self._dependency_map_cache = await get_dependency_map()
|
||||||
|
return self._dependency_map_cache
|
||||||
|
|
||||||
|
async def _get_prompt_metadata_cached(self) -> dict[str, SkillPromptMetadata]:
|
||||||
|
"""获取提示词元数据(带缓存)"""
|
||||||
|
if self._prompt_metadata_cache is None:
|
||||||
|
self._prompt_metadata_cache = await get_prompt_metadata()
|
||||||
|
return self._prompt_metadata_cache
|
||||||
|
|
||||||
async def abefore_agent(self, state: SkillsState, runtime) -> dict[str, Any] | None:
|
async def abefore_agent(self, state: SkillsState, runtime) -> dict[str, Any] | None:
|
||||||
"""在 agent 执行前注入 skills 提示词"""
|
"""在 agent 执行前注入 skills 提示词"""
|
||||||
@ -182,8 +198,8 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
if getattr(runtime_context, "_skills_prompt_injected", False):
|
if getattr(runtime_context, "_skills_prompt_injected", False):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# 从数据库加载 skills 数据
|
# 从数据库加载 skills 数据(使用缓存)
|
||||||
dependency_map = await get_dependency_map()
|
dependency_map = await self._get_dependency_map_cached()
|
||||||
|
|
||||||
# 获取配置的 skills
|
# 获取配置的 skills
|
||||||
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
||||||
@ -219,8 +235,8 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
"""包装模型调用,处理动态激活和依赖展开"""
|
"""包装模型调用,处理动态激活和依赖展开"""
|
||||||
runtime_context = request.runtime.context
|
runtime_context = request.runtime.context
|
||||||
|
|
||||||
# 从数据库加载 skills 数据
|
# 从缓存加载 skills 数据
|
||||||
dependency_map = await get_dependency_map()
|
dependency_map = await self._get_dependency_map_cached()
|
||||||
|
|
||||||
# 1. 获取配置的 skills
|
# 1. 获取配置的 skills
|
||||||
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
configured_skills = getattr(runtime_context, self.skills_context_name, None) or []
|
||||||
@ -239,40 +255,47 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
# 4. 更新 runtime_context 中的 visible_skills
|
# 4. 更新 runtime_context 中的 visible_skills
|
||||||
setattr(runtime_context, "_visible_skills", visible_skills)
|
setattr(runtime_context, "_visible_skills", visible_skills)
|
||||||
|
|
||||||
# 5. 构建依赖包
|
# 5. 构建依赖包(只从直接激活的 skills 获取依赖,不包含闭包展开的依赖)
|
||||||
deps_bundle = await self._build_dependency_bundle(visible_skills)
|
deps_bundle = await self._build_dependency_bundle(activated)
|
||||||
|
|
||||||
# 6. 加载依赖的工具
|
# 6. 加载依赖的工具(普通工具 + MCP 工具)
|
||||||
if deps_bundle["tools"] or deps_bundle["mcps"]:
|
enabled_tools = []
|
||||||
enabled_tools = await self._get_tools_from_context(
|
|
||||||
|
# 6.1 从 toolkits 获取普通工具
|
||||||
|
if deps_bundle["tools"]:
|
||||||
|
all_tools = get_all_tool_instances()
|
||||||
|
required_tool_names = set(deps_bundle["tools"])
|
||||||
|
enabled_tools = [t for t in all_tools if t.name in required_tool_names]
|
||||||
|
|
||||||
|
# 6.2 获取 MCP 工具
|
||||||
|
if deps_bundle["mcps"]:
|
||||||
|
mcp_tools = await self._get_mcp_tools_from_context(
|
||||||
runtime_context,
|
runtime_context,
|
||||||
extra_tool_names=deps_bundle["tools"],
|
|
||||||
extra_mcps=deps_bundle["mcps"],
|
extra_mcps=deps_bundle["mcps"],
|
||||||
)
|
)
|
||||||
|
enabled_tools.extend(mcp_tools)
|
||||||
|
|
||||||
# 合并工具
|
# 合并工具:保留原有工具 + 追加依赖的新工具
|
||||||
if enabled_tools:
|
if enabled_tools:
|
||||||
existing_tools = list(request.tools or [])
|
existing_tool_names = {t.name for t in request.tools or []}
|
||||||
enabled_tool_names = {t.name for t in enabled_tools}
|
merged_tools = list(request.tools or [])
|
||||||
merged_tools = []
|
for t in enabled_tools:
|
||||||
for t_bind in existing_tools:
|
if t.name not in existing_tool_names:
|
||||||
if t_bind.name in enabled_tool_names:
|
merged_tools.append(t)
|
||||||
merged_tools.append(t_bind)
|
request = request.override(tools=merged_tools)
|
||||||
if merged_tools:
|
|
||||||
request = request.override(tools=merged_tools)
|
|
||||||
|
|
||||||
return await handler(request)
|
return await handler(request)
|
||||||
|
|
||||||
async def _build_dependency_bundle(self, visible_skills: list[str]) -> dict[str, list[str]]:
|
async def _build_dependency_bundle(self, activated_skills: list[str]) -> dict[str, list[str]]:
|
||||||
"""根据 visible_skills 构建依赖包"""
|
"""根据直接激活的 skills 构建依赖包(不包含闭包展开的依赖)"""
|
||||||
dependency_map = await get_dependency_map()
|
dependency_map = await self._get_dependency_map_cached()
|
||||||
|
|
||||||
tools: list[str] = []
|
tools: list[str] = []
|
||||||
mcps: list[str] = []
|
mcps: list[str] = []
|
||||||
seen_tools: set[str] = set()
|
seen_tools: set[str] = set()
|
||||||
seen_mcps: set[str] = set()
|
seen_mcps: set[str] = set()
|
||||||
|
|
||||||
for slug in visible_skills:
|
for slug in activated_skills:
|
||||||
dep = dependency_map.get(slug, {})
|
dep = dependency_map.get(slug, {})
|
||||||
for tool_name in dep.get("tools", []):
|
for tool_name in dep.get("tools", []):
|
||||||
if tool_name in seen_tools:
|
if tool_name in seen_tools:
|
||||||
@ -285,11 +308,11 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
seen_mcps.add(mcp_name)
|
seen_mcps.add(mcp_name)
|
||||||
mcps.append(mcp_name)
|
mcps.append(mcp_name)
|
||||||
|
|
||||||
return {"tools": tools, "mcps": mcps, "skills": visible_skills}
|
return {"tools": tools, "mcps": mcps, "skills": activated_skills}
|
||||||
|
|
||||||
async def _collect_prompt_metadata(self, slugs: list[str]) -> list[SkillPromptMetadata]:
|
async def _collect_prompt_metadata(self, slugs: list[str]) -> list[SkillPromptMetadata]:
|
||||||
"""收集指定 slugs 的提示词元数据"""
|
"""收集指定 slugs 的提示词元数据"""
|
||||||
prompt_metadata = await get_prompt_metadata()
|
prompt_metadata = await self._get_prompt_metadata_cached()
|
||||||
|
|
||||||
result: list[SkillPromptMetadata] = []
|
result: list[SkillPromptMetadata] = []
|
||||||
seen: set[str] = set()
|
seen: set[str] = set()
|
||||||
@ -310,28 +333,16 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
async def _get_tools_from_context(
|
async def _get_mcp_tools_from_context(
|
||||||
self,
|
self,
|
||||||
context,
|
context,
|
||||||
*,
|
*,
|
||||||
extra_tool_names: list[str] | None = None,
|
|
||||||
extra_mcps: list[str] | None = None,
|
extra_mcps: list[str] | None = None,
|
||||||
) -> list:
|
) -> list:
|
||||||
"""从上下文配置中获取工具列表"""
|
"""从上下文配置中获取 MCP 工具列表"""
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
selected_tools = []
|
# MCP 工具(并行加载)
|
||||||
|
|
||||||
# 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 []
|
mcps = getattr(context, "mcps", None) or []
|
||||||
all_mcp_names: list[str] = []
|
all_mcp_names: list[str] = []
|
||||||
for server_name in mcps:
|
for server_name in mcps:
|
||||||
@ -357,6 +368,7 @@ class SkillsMiddleware(AgentMiddleware):
|
|||||||
|
|
||||||
# 并行加载所有 MCP 工具
|
# 并行加载所有 MCP 工具
|
||||||
results = await asyncio.gather(*[load_mcp_tools(name) for name in unique_mcp_names])
|
results = await asyncio.gather(*[load_mcp_tools(name) for name in unique_mcp_names])
|
||||||
|
selected_tools = []
|
||||||
for tools in results:
|
for tools in results:
|
||||||
selected_tools.extend(tools)
|
selected_tools.extend(tools)
|
||||||
|
|
||||||
|
|||||||
@ -46,18 +46,12 @@ def get_connection_manager() -> MySQLConnectionManager:
|
|||||||
return _connection_manager
|
return _connection_manager
|
||||||
|
|
||||||
|
|
||||||
class TableListModel(BaseModel):
|
|
||||||
"""获取表名列表的参数模型"""
|
|
||||||
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
@tool(
|
@tool(
|
||||||
category="mysql",
|
category="mysql",
|
||||||
tags=["数据库", "查询"],
|
tags=["数据库", "查询"],
|
||||||
display_name="列出MySQL表",
|
display_name="列出MySQL表",
|
||||||
name_or_callable="mysql_list_tables",
|
name_or_callable="mysql_list_tables",
|
||||||
args_schema=TableListModel,
|
|
||||||
)
|
)
|
||||||
def mysql_list_tables() -> str:
|
def mysql_list_tables() -> str:
|
||||||
"""【查询表名及说明】获取数据库中的所有表名
|
"""【查询表名及说明】获取数据库中的所有表名
|
||||||
|
|||||||
28
src/agents/skills/reporter/SKILLS.md
Normal file
28
src/agents/skills/reporter/SKILLS.md
Normal file
@ -0,0 +1,28 @@
|
|||||||
|
---
|
||||||
|
name: sql-reporter
|
||||||
|
description: "生成 SQL 查询报表并生成可视化图表。当用户需要查询数据库并以报表形式展示结果时使用此技能,包括:统计销售数据、分析用户行为、生成业务报表、查询业务指标等。"
|
||||||
|
---
|
||||||
|
|
||||||
|
# SQL 报表技能
|
||||||
|
|
||||||
|
根据用户的指令,使用数据库工具和图表绘制工具,构建 SQL 查询报告。
|
||||||
|
|
||||||
|
## 操作流程
|
||||||
|
|
||||||
|
1. 理解用户的指令,明确报表的需求和目标
|
||||||
|
2. 使用 MySQL 工具生成正确的 SQL 查询
|
||||||
|
3. 执行查询并获取结果
|
||||||
|
4. 使用 Charts MCP 生成图表
|
||||||
|
5. 将图表以 markdown 图片格式嵌入报表
|
||||||
|
|
||||||
|
## 关键约束
|
||||||
|
|
||||||
|
- 生成的 SQL 查询必须正确且高效,避免全表扫描
|
||||||
|
- 图表生成工具的返回结果不会默认渲染,必须在最终报表中以 `` 格式嵌入
|
||||||
|
- 只返回报表相关的结论,不要返回原始 SQL 查询语句
|
||||||
|
|
||||||
|
## 允许的工具
|
||||||
|
|
||||||
|
- MySQL 工具:执行 SQL 查询
|
||||||
|
- Charts MCP:生成可视化图表
|
||||||
|
- 网络检索工具:必要时补充背景信息
|
||||||
Loading…
Reference in New Issue
Block a user