"""黑板(Blackboard)/ 工作记忆:专家系统风格的共享工作区(零依赖)。 - TaskNode:子任务节点(DAG 顶点),由 Planner 创建、Router 按拓扑序执行 - TaskGraph:子任务 DAG,提供拓扑排序与状态查询 - WorkingMemory:黑板,各知识源(执行器/规则)写入部分解,最后合并为最终答案 对齐《可行性调研与落地实现路线报告》第八章: "黑板协作:多知识源(领域专家/执行器)通过共享黑板协作,而不是一个模型全包"。 """ from __future__ import annotations from dataclasses import dataclass, field from typing import Any, Dict, List, Optional @dataclass class TaskNode: """一个子任务节点。""" id: str kind: str # analyze | design | implement | solve | diagnose | fix # | retrieve | conclude | advise | explain | disclaimer | verify domain: str query: str # 子任务输入(通常为原始查询) status: str = "pending" # pending | running | done | failed | skipped output: Optional[str] = None rule_trace: List[str] = field(default_factory=list) deps: List[str] = field(default_factory=list) desc: str = "" error: Optional[str] = None class TaskGraph: """子任务 DAG:节点 + 依赖边。""" def __init__(self): self._nodes: Dict[str, TaskNode] = {} def add_node(self, node: TaskNode) -> None: if node.id in self._nodes: raise ValueError(f"节点 id 重复: {node.id}") self._nodes[node.id] = node def get(self, node_id: str) -> Optional[TaskNode]: return self._nodes.get(node_id) def nodes(self) -> List[TaskNode]: return list(self._nodes.values()) def topo_order(self) -> List[TaskNode]: """Kahn 拓扑排序:依赖在前。循环依赖时按插入序兜底(不崩溃)。""" indeg: Dict[str, int] = {} for n in self._nodes.values(): indeg[n.id] = 0 for n in self._nodes.values(): for d in n.deps: if d in indeg: indeg[n.id] += 1 ready = [n for n in self._nodes.values() if indeg[n.id] == 0] ready.sort(key=lambda n: list(self._nodes.keys()).index(n.id)) order: List[TaskNode] = [] while ready: n = ready.pop(0) order.append(n) for m in self._nodes.values(): if n.id in m.deps: indeg[m.id] -= 1 if indeg[m.id] == 0 and m not in order: ready.append(m) if len(order) < len(self._nodes): # 循环依赖兜底:剩余节点按插入序追加 for n in self._nodes.values(): if n not in order: order.append(n) return order def all_done(self) -> bool: return all(n.status == "done" for n in self._nodes.values()) def failed(self) -> List[TaskNode]: return [n for n in self._nodes.values() if n.status == "failed"] def __len__(self) -> int: return len(self._nodes) class WorkingMemory: """黑板:facts(槽位事实)+ sections(章节部分解)+ trace(推理轨迹)。""" def __init__(self): self.facts: Dict[str, Any] = {} self.sections: Dict[str, str] = {} self.trace: List[str] = [] # ---- 事实 ---- def write_fact(self, key: str, value: Any, rule_id: Optional[str] = None) -> None: if key in self.facts: self.trace.append(f"overwrite:{key}@{rule_id or '?'}") self.facts[key] = value if rule_id: self.trace.append(f"fact:{key}={str(value)[:40]}@rule:{rule_id}") def get_fact(self, key: str, default: Any = None) -> Any: return self.facts.get(key, default) # ---- 章节 ---- def write_section(self, sid: str, text: str) -> None: """写入章节;同 id 覆盖(记录 trace)。""" if sid in self.sections: self.trace.append(f"overwrite_section:{sid}") self.sections[sid] = text def section(self, sid: str) -> Optional[str]: return self.sections.get(sid) def merge(self, order: Optional[List[str]] = None) -> str: """按 order(章节顺序)合并为最终答案;order 为空则按写入顺序。""" if order: parts = [self.sections[s] for s in order if s in self.sections] if parts: return "\n\n".join(parts) return "\n\n".join(self.sections.values()) # ---- 轨迹 ---- def add_trace(self, item: str) -> None: self.trace.append(item) def explain(self) -> List[str]: return list(self.trace)