feat: Agent v2 能力定义层 - AgentConfig/ToolRegistry/SkillLoader/MemoryStore

This commit is contained in:
yumoqing 2026-08-10 11:06:37 +08:00
parent 66b8b42064
commit 5fd38f8635
6 changed files with 1161 additions and 0 deletions

22
mysql_v2.ddl.sql Normal file
View File

@ -0,0 +1,22 @@
-- pipeline_user_memory: 跨会话记忆表v2 新增)
-- 对照 Hermes Agent 的 memory 系统
CREATE TABLE IF NOT EXISTS `pipeline_user_memory` (
`id` varchar(32) NOT NULL COMMENT '主键',
`memory_key` varchar(64) NOT NULL COMMENT '记忆唯一标识内容MD5',
`content` text NOT NULL COMMENT '记忆内容',
`category` varchar(32) NOT NULL DEFAULT 'memory' COMMENT '分类: user/memory',
`priority` int NOT NULL DEFAULT 0 COMMENT '优先级(越高越重要)',
`access_count` int NOT NULL DEFAULT 0 COMMENT '被引用次数',
`created_at` datetime NOT NULL DEFAULT CURRENT_TIMESTAMP,
`updated_at` datetime NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP,
PRIMARY KEY (`id`),
UNIQUE KEY `uk_memory` (`memory_key`, `category`),
KEY `idx_category_priority` (`category`, `priority` DESC)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='Agent 跨会话记忆';
-- pipeline_agent_config: 产线级 Agent 配置可选JSON 字段)
ALTER TABLE `pipelines` ADD COLUMN IF NOT EXISTS `agent_config` text COMMENT 'Agent配置JSON';
-- sd_org_settings 增加 agent_config 字段
ALTER TABLE `sd_org_settings` ADD COLUMN IF NOT EXISTS `agent_config` text COMMENT '项目级Agent配置JSON';

View File

@ -12,3 +12,31 @@ from .init import (
publish_pipeline,
load_pipeline_core,
)
# v2: Agent 能力定义层
from .agent_config import (
AgentConfig,
CompressionConfig,
MemoryConfig,
SkillConfig,
ToolDefinition,
SDLC_DEFAULT_CONFIG,
SDLC_DEFAULT_TOOLS,
load_agent_config,
save_agent_config,
)
from .tool_registry import (
ToolRegistry,
get_tool_registry,
register_tool,
)
from .skill_loader import (
Skill,
SkillLoader,
get_skill_loader,
)
from .memory_store import (
MemoryEntry,
MemoryStore,
get_memory_store,
)

View File

@ -0,0 +1,428 @@
"""
pipeline-core: Agent 能力定义层
对照 Hermes Agent config.yaml + system prompt builder + tool registry
在此定义每个产线的 Agent 配置模型工具技能记忆上下文压缩
每个产线pipeline可以有不同的 AgentConfig通过 sd_org_settings
pipelines 表的 agent_config 字段存储pipeline-service 执行时读取此配置
"""
import json
import os
from dataclasses import dataclass, field
from typing import Dict, List, Optional
# ═══════════════════════════════════════════════════════════
# 数据模型
# ═══════════════════════════════════════════════════════════
@dataclass
class ToolDefinition:
"""工具定义 — 等价于 HA 的 tool schema"""
name: str
description: str
parameters: dict = field(default_factory=dict) # JSON Schema for params
enabled: bool = True
category: str = "general" # project / task / repo / shell / agent
requires_confirmation: bool = False # 是否需要用户确认
@dataclass
class CompressionConfig:
"""上下文压缩配置"""
enabled: bool = True
threshold: float = 0.50 # 达到上下文窗口的 50% 时触发压缩
target_ratio: float = 0.20 # 压缩到 20%
keep_recent: int = 8 # 保留最近 N 轮
@dataclass
class MemoryConfig:
"""记忆系统配置"""
enabled: bool = True
user_profile_enabled: bool = True # 用户偏好
cross_session_enabled: bool = True # 跨会话事实
max_entries: int = 50 # 最多保留条数
@dataclass
class SkillConfig:
"""技能配置"""
enabled: bool = True
dirs: List[str] = field(default_factory=lambda: ["skills/common", "skills/sdlc"])
max_skills_per_turn: int = 5
@dataclass
class AgentConfig:
"""Agent 完整配置 — 一个产线一个配置"""
# ── 模型 ──
model_name: str = ""
temperature: float = 0.4
max_turns: int = 30 # 最大 tool-calling 轮次(替代硬编码 10
# ── 系统提示词 ──
system_prompt: str = "" # 产线专属系统提示词
personality: str = "" # 人设描述
# ── 上下文 ──
compression: CompressionConfig = field(default_factory=CompressionConfig)
context_limit: int = 64000 # token 上限估计
# ── 记忆 ──
memory: MemoryConfig = field(default_factory=MemoryConfig)
# ── 技能 ──
skills: SkillConfig = field(default_factory=SkillConfig)
# ── 工具 ──
tools: List[ToolDefinition] = field(default_factory=list)
# ── 会话 ──
session_isolation: str = "project" # "project" | "user" | "none"
history_limit: int = 50 # 加载历史消息数
# ── 安全 ──
require_approval: bool = False # 危险命令需确认
allowed_workdirs: List[str] = field(default_factory=list)
def to_dict(self) -> dict:
"""序列化为 JSON"""
return {
"model_name": self.model_name,
"temperature": self.temperature,
"max_turns": self.max_turns,
"system_prompt": self.system_prompt,
"personality": self.personality,
"compression": {
"enabled": self.compression.enabled,
"threshold": self.compression.threshold,
"target_ratio": self.compression.target_ratio,
"keep_recent": self.compression.keep_recent,
},
"memory": {
"enabled": self.memory.enabled,
"user_profile_enabled": self.memory.user_profile_enabled,
"cross_session_enabled": self.memory.cross_session_enabled,
"max_entries": self.memory.max_entries,
},
"skills": {
"enabled": self.skills.enabled,
"dirs": self.skills.dirs,
"max_skills_per_turn": self.skills.max_skills_per_turn,
},
"tools": [
{
"name": t.name,
"description": t.description,
"parameters": t.parameters,
"enabled": t.enabled,
"category": t.category,
}
for t in self.tools
],
"session_isolation": self.session_isolation,
"history_limit": self.history_limit,
"context_limit": self.context_limit,
}
@classmethod
def from_dict(cls, data: dict) -> "AgentConfig":
"""从 JSON 反序列化"""
comp = data.get("compression", {})
mem = data.get("memory", {})
sk = data.get("skills", {})
return cls(
model_name=data.get("model_name", ""),
temperature=data.get("temperature", 0.4),
max_turns=data.get("max_turns", 30),
system_prompt=data.get("system_prompt", ""),
personality=data.get("personality", ""),
compression=CompressionConfig(
enabled=comp.get("enabled", True),
threshold=comp.get("threshold", 0.50),
target_ratio=comp.get("target_ratio", 0.20),
keep_recent=comp.get("keep_recent", 8),
),
memory=MemoryConfig(
enabled=mem.get("enabled", True),
user_profile_enabled=mem.get("user_profile_enabled", True),
cross_session_enabled=mem.get("cross_session_enabled", True),
max_entries=mem.get("max_entries", 50),
),
skills=SkillConfig(
enabled=sk.get("enabled", True),
dirs=sk.get("dirs", ["skills/common", "skills/sdlc"]),
max_skills_per_turn=sk.get("max_skills_per_turn", 5),
),
tools=[ToolDefinition(**t) if isinstance(t, dict) else t for t in data.get("tools", [])],
session_isolation=data.get("session_isolation", "project"),
history_limit=data.get("history_limit", 50),
context_limit=data.get("context_limit", 64000),
)
# ═══════════════════════════════════════════════════════════
# SDLC 默认配置
# ═══════════════════════════════════════════════════════════
SDLC_DEFAULT_TOOLS = [
ToolDefinition(
name="switch_project",
description="切换到指定项目。用户说「切换到XXX」时调用。",
parameters={"project_name": "项目名称"},
category="project",
),
ToolDefinition(
name="create_project",
description="创建新的软件项目",
parameters={"name": "项目名称", "description": "项目描述"},
category="project",
),
ToolDefinition(
name="create_task",
description="创建开发任务分配给角色agent执行",
parameters={
"title": "任务标题",
"role": "目标角色(requirement/design/develop/test/deploy)",
"description": "任务详细描述",
},
category="task",
),
ToolDefinition(
name="list_tasks",
description="查看当前项目的任务列表",
parameters={"role": "按角色筛选(可选)", "state": "按状态筛选(可选)"},
category="task",
),
ToolDefinition(
name="task_detail",
description="查看任务详情",
parameters={"task_id": "任务ID"},
category="task",
),
ToolDefinition(
name="start_agents",
description="启动角色agent执行已提交的任务",
parameters={},
category="agent",
),
ToolDefinition(
name="diagnose_project",
description="诊断项目:汇总卡点/失败/问题/运行中任务",
parameters={},
category="agent",
),
ToolDefinition(
name="list_deliverables",
description="查看交付件列表",
parameters={"role": "按角色筛选(可选)"},
category="task",
),
ToolDefinition(
name="view_deliverable",
description="查看交付件详细内容",
parameters={"deliverable_id": "交付件ID"},
category="task",
),
ToolDefinition(
name="list_questions",
description="查看待回答的问题",
parameters={},
category="agent",
),
ToolDefinition(
name="answer_question",
description="回答agent提出的问题",
parameters={"question_id": "问题ID", "answer": "回答内容"},
category="agent",
),
ToolDefinition(
name="add_repo",
description="添加项目关联的Git仓库",
parameters={"repo_url": "仓库URL", "repo_name": "仓库名称", "default_branch": "默认分支(默认main)"},
category="repo",
),
ToolDefinition(
name="list_repos",
description="查看项目关联的仓库列表",
parameters={},
category="repo",
),
ToolDefinition(
name="run_command",
description="在工作空间中执行shell命令",
parameters={"command": "命令"},
category="shell",
requires_confirmation=True,
),
ToolDefinition(
name="check_progress",
description="查看项目整体进展",
parameters={},
category="agent",
),
ToolDefinition(
name="add_bug",
description="提交Bug",
parameters={"title": "Bug标题", "description": "描述", "severity": "严重程度"},
category="task",
),
# ── v2 新增 ──
ToolDefinition(
name="ask_user",
description="向用户提问澄清意图(不确定时使用)",
parameters={"question": "问题"},
category="agent",
),
ToolDefinition(
name="delegate_subtask",
description="派生子agent调查子任务并行执行",
parameters={"goal": "子任务目标", "context": "背景信息"},
category="agent",
),
]
SDLC_DEFAULT_CONFIG = AgentConfig(
model_name="deepseek-v4-pro",
temperature=0.4,
max_turns=30,
system_prompt="""你是 SDLC 开发产线的 AI 驾驶舱助理。
## 核心原则
- 说做就做承诺调用工具时立刻调不要只描述计划
- 每次只输出一个 JSON
- 工具调用后等待结果再决定下一步
- 不确定用户意图时使用 ask_user 工具澄清
## 项目上下文
当前项目{current_project}
{project_context}
## 可用工具
{tools_description}
## 输出格式
每次只输出一个 JSON 对象
- 回复用户{"action":"reply","message":"回复内容"}
- 调用工具{"action":"tool_call","tool":"工具名","params":{}}
- 提问澄清{"action":"tool_call","tool":"ask_user","params":{"question":"问题"}}
""",
tools=SDLC_DEFAULT_TOOLS,
session_isolation="project",
history_limit=50,
compression=CompressionConfig(enabled=True, threshold=0.50, target_ratio=0.20, keep_recent=6),
memory=MemoryConfig(enabled=True),
skills=SkillConfig(enabled=True, dirs=["skills/common", "skills/sdlc"]),
)
# ═══════════════════════════════════════════════════════════
# 配置加载器
# ═══════════════════════════════════════════════════════════
async def load_agent_config(pipeline_id: str = None, project_id: str = None) -> AgentConfig:
"""加载产线的 Agent 配置。
优先级
1. 项目级 sd_org_settings.agent_config
2. 产线级 pipelines.agent_config
3. SDLC_DEFAULT_CONFIG
"""
try:
from sqlor.dbpools import DBPools
db = DBPools()
async with db.sqlorContext("pipeline") as sor:
# 项目级配置
if project_id:
recs = await sor.sqlExe(
"SELECT agent_config FROM sd_org_settings WHERE project_id=${pid}$",
{"pid": project_id},
)
if recs:
raw = getattr(recs[0], "agent_config", "")
if raw:
try:
data = json.loads(raw) if isinstance(raw, str) else raw
data["tools"] = _merge_tools(data.get("tools", []), SDLC_DEFAULT_TOOLS)
return AgentConfig.from_dict(data)
except (json.JSONDecodeError, TypeError):
pass
# 产线级配置
if pipeline_id:
recs = await sor.sqlExe(
"SELECT agent_config FROM pipelines WHERE id=${pid}$",
{"pid": pipeline_id},
)
if recs:
raw = getattr(recs[0], "agent_config", "")
if raw:
try:
data = json.loads(raw) if isinstance(raw, str) else raw
data["tools"] = _merge_tools(data.get("tools", []), SDLC_DEFAULT_TOOLS)
return AgentConfig.from_dict(data)
except (json.JSONDecodeError, TypeError):
pass
except Exception:
pass
return SDLC_DEFAULT_CONFIG
def _merge_tools(custom_tools: list, default_tools: list) -> list:
"""合并自定义工具和默认工具。自定义覆盖同名工具。"""
merged = {t["name"] if isinstance(t, dict) else t.name: t for t in default_tools}
for t in custom_tools:
name = t["name"] if isinstance(t, dict) else t.name
merged[name] = t
return list(merged.values())
async def save_agent_config(project_id: str, config: AgentConfig):
"""保存项目级 Agent 配置到 sd_org_settings"""
from sqlor.dbpools import DBPools
db = DBPools()
async with db.sqlorContext("pipeline") as sor:
config_json = json.dumps(config.to_dict(), ensure_ascii=False)
# UPDATE-first 防止竞态
await sor.sqlExe(
"UPDATE sd_org_settings SET agent_config=${cfg}$ WHERE project_id=${pid}$",
{"cfg": config_json, "pid": project_id},
)
existing = await sor.sqlExe(
"SELECT id FROM sd_org_settings WHERE project_id=${pid}$",
{"pid": project_id},
)
if not existing:
from appPublic.uniqueID import getID
await sor.C("sd_org_settings", {
"id": getID(),
"project_id": project_id,
"agent_config": config_json,
})
# ═══════════════════════════════════════════════════════════
# 模块加载
# ═══════════════════════════════════════════════════════════
def load_pipeline_core():
"""注册 Agent 配置管理函数到 ServerEnv"""
from ahserver.serverenv import ServerEnv
env = ServerEnv()
env.load_agent_config = load_agent_config
env.save_agent_config = save_agent_config
env.AgentConfig = AgentConfig
env.SDLC_DEFAULT_CONFIG = SDLC_DEFAULT_CONFIG
MODULE_NAME = "pipeline_core"

View File

@ -0,0 +1,271 @@
"""
pipeline-core: Memory Store 跨会话记忆系统
对照 Hermes Agent memory 工具 + MEMORY.md / USER.md
- 用户偏好存储user profile
- 跨会话事实存储agent memory
- 记忆注入到 system prompt 每轮对话
- 自动去重和淘汰
存储方式MySQL pipeline_user_memory + 内存缓存
"""
import json
import logging
import time
from dataclasses import dataclass, field
from typing import Dict, List, Optional
logger = logging.getLogger("pipeline.memory_store")
@dataclass
class MemoryEntry:
"""一条记忆"""
key: str # 唯一标识
content: str # 记忆内容
category: str = "memory" # "user" | "memory"
priority: int = 0 # 优先级(越高越重要)
created_at: float = 0.0
updated_at: float = 0.0
access_count: int = 0 # 被引用次数
def to_dict(self) -> dict:
return {
"key": self.key,
"content": self.content,
"category": self.category,
"priority": self.priority,
"created_at": self.created_at,
"updated_at": self.updated_at,
"access_count": self.access_count,
}
@classmethod
def from_dict(cls, data: dict) -> "MemoryEntry":
return cls(
key=data.get("key", ""),
content=data.get("content", ""),
category=data.get("category", "memory"),
priority=data.get("priority", 0),
created_at=data.get("created_at", time.time()),
updated_at=data.get("updated_at", time.time()),
access_count=data.get("access_count", 0),
)
class MemoryStore:
"""跨会话记忆存储。
两层架构
1. MySQL pipeline_user_memory持久化
2. 内存缓存加速读取
"""
# 优先级分组
PRIORITY_HIGH = 10 # 用户偏好、环境配置
PRIORITY_MEDIUM = 5 # 项目约定、工作流
PRIORITY_LOW = 1 # 临时备注
def __init__(self):
self._cache: Dict[str, Dict[str, MemoryEntry]] = {} # {category: {key: entry}}
self._cache_loaded = False
async def _ensure_cache(self):
"""从 DB 加载缓存"""
if self._cache_loaded:
return
try:
from sqlor.dbpools import DBPools
db = DBPools()
async with db.sqlorContext("pipeline") as sor:
recs = await sor.sqlExe(
"SELECT memory_key, content, category, priority, created_at, updated_at, access_count "
"FROM pipeline_user_memory ORDER BY priority DESC, updated_at DESC LIMIT 200",
{},
)
for r in (recs or []):
entry = MemoryEntry(
key=getattr(r, "memory_key", ""),
content=getattr(r, "content", ""),
category=getattr(r, "category", "memory"),
priority=getattr(r, "priority", 0),
created_at=_ts(getattr(r, "created_at", None)),
updated_at=_ts(getattr(r, "updated_at", None)),
access_count=getattr(r, "access_count", 0),
)
if entry.category not in self._cache:
self._cache[entry.category] = {}
self._cache[entry.category][entry.key] = entry
except Exception as e:
logger.warning(f"MemoryStore cache load failed: {e}")
self._cache_loaded = True
async def add(self, content: str, category: str = "memory",
priority: int = PRIORITY_MEDIUM, key: str = None):
"""添加一条记忆。自动生成 key取内容前80字符的哈希"""
await self._ensure_cache()
if not key:
key = _make_key(content)
now = time.time()
entry = MemoryEntry(
key=key,
content=content,
category=category,
priority=priority,
created_at=now,
updated_at=now,
)
# 更新缓存
if category not in self._cache:
self._cache[category] = {}
self._cache[category][key] = entry
# 持久化
try:
await self._db_upsert(entry)
except Exception as e:
logger.error(f"MemoryStore add failed: {e}")
# 淘汰低优先级旧条目
await self._evict(max_entries=100)
async def get(self, category: str = None, key: str = None) -> List[MemoryEntry]:
"""获取记忆。可按分类和 key 筛选。"""
await self._ensure_cache()
results = []
cats = [category] if category else list(self._cache.keys())
for cat in cats:
if cat in self._cache:
for k, entry in self._cache[cat].items():
if key and k != key:
continue
results.append(entry)
entry.access_count += 1
results.sort(key=lambda e: e.priority, reverse=True)
return results
async def remove(self, key: str, category: str = "memory"):
"""删除一条记忆"""
await self._ensure_cache()
if category in self._cache:
self._cache[category].pop(key, None)
try:
from sqlor.dbpools import DBPools
db = DBPools()
async with db.sqlorContext("pipeline") as sor:
await sor.sqlExe(
"DELETE FROM pipeline_user_memory WHERE memory_key=${k}$ AND category=${c}$",
{"k": key, "c": category},
)
except Exception as e:
logger.error(f"MemoryStore remove failed: {e}")
async def build_prompt_block(self, category: str = None, max_entries: int = 20) -> str:
"""构建注入 system prompt 的记忆段落。
高优先级记忆注入完整内容低优先级只注入摘要
"""
entries = await self.get(category)
if not entries:
return ""
blocks = []
high_priority = [e for e in entries if e.priority >= self.PRIORITY_MEDIUM]
low_priority = [e for e in entries if e.priority < self.PRIORITY_MEDIUM]
for e in high_priority[:max_entries]:
tag = "用户偏好" if e.category == "user" else "已知信息"
blocks.append(f"[{tag}] {e.content}")
if low_priority:
blocks.append(f"[上下文] 相关背景: {'; '.join(e.content for e in low_priority[:3])}")
return "\n".join(blocks) if blocks else ""
async def _db_upsert(self, entry: MemoryEntry):
"""写入/更新 DB"""
from sqlor.dbpools import DBPools
db = DBPools()
async with db.sqlorContext("pipeline") as sor:
existing = await sor.sqlExe(
"SELECT id FROM pipeline_user_memory WHERE memory_key=${k}$ AND category=${c}$",
{"k": entry.key, "c": entry.category},
)
if existing:
rid = existing[0].id
await sor.U("pipeline_user_memory", {
"id": rid,
"content": entry.content,
"priority": entry.priority,
"updated_at": "NOW()",
"access_count": entry.access_count,
})
else:
from appPublic.uniqueID import getID
await sor.C("pipeline_user_memory", {
"id": getID(),
"memory_key": entry.key,
"content": entry.content,
"category": entry.category,
"priority": entry.priority,
})
async def _evict(self, max_entries: int = 100):
"""淘汰低优先级的旧记忆LRU"""
total = sum(len(v) for v in self._cache.values())
if total <= max_entries:
return
# 找出可淘汰的条目(低优先级 + 低访问次数)
all_entries = []
for cat_entries in self._cache.values():
all_entries.extend(cat_entries.values())
all_entries.sort(key=lambda e: (e.priority, e.access_count, e.updated_at))
to_remove = all_entries[: (total - max_entries)]
for entry in to_remove:
await self.remove(entry.key, entry.category)
# ── 辅助 ──
def _make_key(content: str) -> str:
"""从内容生成稳定 key"""
import hashlib
return hashlib.md5(content.encode("utf-8")).hexdigest()[:16]
def _ts(val) -> float:
"""转换时间戳"""
if val is None:
return time.time()
if isinstance(val, (int, float)):
return float(val)
return time.time()
# ── 全局单例 ──
_store: Optional[MemoryStore] = None
def get_memory_store() -> MemoryStore:
global _store
if _store is None:
_store = MemoryStore()
return _store

View File

@ -0,0 +1,246 @@
"""
pipeline-core: Skill Loader 技能加载与注入
对照 Hermes Agent skills 系统
- 从目录扫描 SKILL.md 文件
- 解析 YAML frontmatter
- 按产线/角色注入到 system prompt
- 支持热更新agent 执行时可重新扫描
目录结构 HA 一致
skills/
common/ 所有角色共用
code-style/SKILL.md
sdlc/ SDLC 专用
git-workflow/SKILL.md
requirement/ 角色专用
design/
develop/
test/
deploy/
"""
import json
import logging
import os
import re
from dataclasses import dataclass, field
from typing import Dict, List, Optional
logger = logging.getLogger("pipeline.skill_loader")
@dataclass
class Skill:
"""单个技能"""
name: str
path: str # SKILL.md 绝对路径
description: str = ""
content: str = "" # SKILL.md 完整内容
category: str = "common" # common / sdlc / requirement / design / ...
trigger_keywords: List[str] = field(default_factory=list)
version: str = "1.0.0"
@classmethod
def from_file(cls, filepath: str, category: str = "common") -> "Skill":
"""从 SKILL.md 文件加载技能"""
name = os.path.basename(os.path.dirname(filepath))
content = ""
description = ""
trigger_keywords = []
version = "1.0.0"
try:
with open(filepath, "r", encoding="utf-8") as f:
content = f.read()
except Exception as e:
logger.warning(f"Skill read failed: {filepath} err={e}")
return cls(name=name, path=filepath, category=category)
# 解析 YAML frontmatter如果存在
fm = _parse_frontmatter(content)
if fm:
description = fm.get("description", description)
trigger_keywords = fm.get("trigger_keywords", trigger_keywords)
version = fm.get("version", version)
# 无 frontmatter 时从内容首行提取描述
if not description:
for line in content.split("\n"):
line = line.strip()
if line and not line.startswith("#") and not line.startswith("---"):
description = line[:200]
break
return cls(
name=name,
path=filepath,
description=description,
content=content,
category=category,
trigger_keywords=trigger_keywords,
version=version,
)
def to_prompt_block(self) -> str:
"""生成注入到 system prompt 的技能文本块"""
return f"""## 技能: {self.name}
{self.description}
{self.content}
"""
def _parse_frontmatter(content: str) -> Optional[dict]:
"""解析 YAML frontmatter--- 之间的内容)"""
if not content.startswith("---"):
return None
end = content.find("---", 3)
if end == -1:
return None
fm_text = content[3:end].strip()
result = {}
# 简单的 YAML 解析(不依赖 PyYAML
for line in fm_text.split("\n"):
line = line.strip()
if ":" in line:
key, _, val = line.partition(":")
key = key.strip()
val = val.strip().strip('"').strip("'")
# 处理列表
if val.startswith("[") and val.endswith("]"):
val = [v.strip().strip('"').strip("'") for v in val[1:-1].split(",")]
result[key] = val
return result
class SkillLoader:
"""技能加载器 — 管理技能目录并注入到 prompt"""
def __init__(self):
self._skills: Dict[str, Dict[str, Skill]] = {} # {category: {name: Skill}}
self._base_dirs: List[str] = []
def add_skill_dir(self, base_dir: str):
"""添加技能目录。目录下每个子目录是一个技能。"""
if base_dir not in self._base_dirs:
self._base_dirs.append(base_dir)
self._scan_dir(base_dir)
def _scan_dir(self, base_dir: str):
"""扫描目录下的所有 SKILL.md"""
if not os.path.isdir(base_dir):
logger.debug(f"Skill dir not found: {base_dir}")
return
# base_dir 下的每个子目录是一个技能分类
for category in os.listdir(base_dir):
cat_path = os.path.join(base_dir, category)
if not os.path.isdir(cat_path):
continue
# 直接是 SKILL.md
skill_file = os.path.join(cat_path, "SKILL.md")
if os.path.isfile(skill_file):
skill = Skill.from_file(skill_file, category)
self._register(skill)
continue
# 子目录下可能有多个技能
for sub in os.listdir(cat_path):
sub_path = os.path.join(cat_path, sub)
if os.path.isdir(sub_path):
skill_file = os.path.join(sub_path, "SKILL.md")
if os.path.isfile(skill_file):
skill = Skill.from_file(skill_file, category)
self._register(skill)
def _register(self, skill: Skill):
"""注册技能到内存"""
if skill.category not in self._skills:
self._skills[skill.category] = {}
self._skills[skill.category][skill.name] = skill
logger.debug(f"SkillLoader: loaded {skill.category}/{skill.name}")
def get_all(self, categories: List[str] = None) -> List[Skill]:
"""获取指定分类的所有技能"""
result = []
cats = categories or list(self._skills.keys())
for cat in cats:
if cat in self._skills:
result.extend(self._skills[cat].values())
return result
def get_by_trigger(self, user_input: str, max_skills: int = 5) -> List[Skill]:
"""根据用户输入匹配相关技能(基于 trigger_keywords"""
scored = []
user_lower = user_input.lower()
for cat_skills in self._skills.values():
for skill in cat_skills.values():
score = 0
for kw in skill.trigger_keywords:
if kw.lower() in user_lower:
score += 1
# 标题匹配加分
if skill.name.lower() in user_lower:
score += 2
if score > 0:
scored.append((score, skill))
scored.sort(key=lambda x: x[0], reverse=True)
return [s[1] for s in scored[:max_skills]]
def get_skill_count(self) -> int:
"""已加载技能总数"""
return sum(len(skills) for skills in self._skills.values())
def reload(self):
"""重新扫描所有技能目录(热更新)"""
self._skills = {}
for d in self._base_dirs:
self._scan_dir(d)
def build_prompt_block(self, categories: List[str] = None, user_input: str = None,
max_skills: int = 5) -> str:
"""构建注入 system prompt 的技能段落。
如果有 user_input优先匹配相关技能
"""
if user_input:
skills = self.get_by_trigger(user_input, max_skills)
else:
skills = self.get_all(categories)[:max_skills]
if not skills:
return ""
blocks = ["## 可用技能\n"]
for s in skills:
# 只注入技能摘要,不注入完整内容(节省 token
blocks.append(f"- **{s.category}/{s.name}**: {s.description[:200]}")
blocks.append("")
return "\n".join(blocks)
def build_full_prompt_block(self, skill_names: List[str]) -> str:
"""注入指定技能的完整内容到 system prompt"""
blocks = []
for cat_skills in self._skills.values():
for name, skill in cat_skills.items():
if name in skill_names:
blocks.append(skill.to_prompt_block())
return "\n".join(blocks)
# ── 全局单例 ──
_loader: Optional[SkillLoader] = None
def get_skill_loader() -> SkillLoader:
global _loader
if _loader is None:
_loader = SkillLoader()
return _loader

View File

@ -0,0 +1,166 @@
"""
pipeline-core: Tool Registry 工具注册与发现
对照 Hermes Agent tools/registry.py
- 集中注册所有可用工具
- 生成 LLM function-calling schema
- 按产线启用/禁用工具
- 工具分类管理
"""
import json
import logging
from typing import Callable, Dict, List, Optional
from .agent_config import ToolDefinition
logger = logging.getLogger("pipeline.tool_registry")
class ToolRegistry:
"""全局工具注册表。每个工具包含定义schema和执行函数。"""
def __init__(self):
self._tools: Dict[str, ToolDefinition] = {}
self._handlers: Dict[str, Callable] = {}
def register(
self,
tool: ToolDefinition,
handler: Callable = None,
):
"""注册一个工具。
Args:
tool: 工具定义名称描述参数 schema
handler: 执行函数 async def handler(params, context) -> str
"""
self._tools[tool.name] = tool
if handler:
self._handlers[tool.name] = handler
logger.debug(f"ToolRegistry: registered {tool.name}")
def unregister(self, name: str):
"""移除工具"""
self._tools.pop(name, None)
self._handlers.pop(name, None)
def enable(self, name: str):
"""启用工具"""
if name in self._tools:
self._tools[name].enabled = True
def disable(self, name: str):
"""禁用工具"""
if name in self._tools:
self._tools[name].enabled = False
def get(self, name: str) -> Optional[ToolDefinition]:
"""获取工具定义"""
return self._tools.get(name)
def get_handler(self, name: str) -> Optional[Callable]:
"""获取工具执行函数"""
return self._handlers.get(name)
def list_enabled(self, category: str = None) -> List[ToolDefinition]:
"""列出启用的工具,可按分类筛选"""
tools = [t for t in self._tools.values() if t.enabled]
if category:
tools = [t for t in tools if t.category == category]
return tools
def to_openai_schema(self, category: str = None) -> List[dict]:
"""生成 OpenAI function-calling 格式的 tool 定义。
等价于 HA system prompt 中注入的 tool descriptions
"""
tools = self.list_enabled(category)
return [
{
"type": "function",
"function": {
"name": t.name,
"description": t.description,
"parameters": {
"type": "object",
"properties": {
k: {"type": "string", "description": v}
for k, v in (t.parameters or {}).items()
},
"required": list(t.parameters.keys()) if t.parameters else [],
},
},
}
for t in tools
]
def to_text_description(self, category: str = None) -> str:
"""生成纯文本工具描述(用于不支持 native function calling 的模型)。
等价于 HA tool description 注入到 system prompt
"""
tools = self.list_enabled(category)
lines = []
for t in tools:
params = ", ".join(
f"{k}: {v}" for k, v in (t.parameters or {}).items()
)
lines.append(f"- {t.name}({params}): {t.description}")
return "\n".join(lines)
def get_tool_names(self) -> List[str]:
"""获取所有已注册工具名称"""
return list(self._tools.keys())
def get_by_category(self) -> Dict[str, List[str]]:
"""按分类分组工具名"""
groups: Dict[str, List[str]] = {}
for t in self._tools.values():
groups.setdefault(t.category, []).append(t.name)
return groups
# ── 全局单例 ──
_registry: Optional[ToolRegistry] = None
def get_tool_registry() -> ToolRegistry:
"""获取全局工具注册表单例"""
global _registry
if _registry is None:
_registry = ToolRegistry()
return _registry
# ── 便捷函数 ──
def register_tool(
name: str,
description: str,
parameters: dict = None,
category: str = "general",
requires_confirmation: bool = False,
):
"""装饰器:注册工具。
用法:
@register_tool("search_web", "搜索网页", {"query": "搜索关键词"})
async def search_web(params, context):
...
"""
def decorator(handler: Callable):
tool = ToolDefinition(
name=name,
description=description,
parameters=parameters or {},
category=category,
requires_confirmation=requires_confirmation,
)
registry = get_tool_registry()
registry.register(tool, handler)
return handler
return decorator