diff --git a/mysql_v2.ddl.sql b/mysql_v2.ddl.sql new file mode 100644 index 0000000..84f340e --- /dev/null +++ b/mysql_v2.ddl.sql @@ -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'; diff --git a/pipeline_core/__init__.py b/pipeline_core/__init__.py index 34a9d89..da825ef 100644 --- a/pipeline_core/__init__.py +++ b/pipeline_core/__init__.py @@ -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, +) diff --git a/pipeline_core/agent_config.py b/pipeline_core/agent_config.py new file mode 100644 index 0000000..efefedc --- /dev/null +++ b/pipeline_core/agent_config.py @@ -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" diff --git a/pipeline_core/memory_store.py b/pipeline_core/memory_store.py new file mode 100644 index 0000000..e92f72a --- /dev/null +++ b/pipeline_core/memory_store.py @@ -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 diff --git a/pipeline_core/skill_loader.py b/pipeline_core/skill_loader.py new file mode 100644 index 0000000..1ede6f1 --- /dev/null +++ b/pipeline_core/skill_loader.py @@ -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 diff --git a/pipeline_core/tool_registry.py b/pipeline_core/tool_registry.py new file mode 100644 index 0000000..6d49626 --- /dev/null +++ b/pipeline_core/tool_registry.py @@ -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