store.py 43 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098109911001101110211031104110511061107110811091110111111121113111411151116111711181119112011211122112311241125112611271128112911301131113211331134113511361137113811391140114111421143114411451146114711481149115011511152115311541155115611571158115911601161116211631164116511661167116811691170117111721173117411751176117711781179118011811182118311841185118611871188118911901191119211931194119511961197119811991200120112021203120412051206
  1. """
  2. FileSystem Trace Store - 文件系统存储实现
  3. 用于跨进程数据共享,数据持久化到 .trace/ 目录
  4. 目录结构:
  5. .trace/{trace_id}/
  6. ├── meta.json # Trace 元数据
  7. ├── goal.json # GoalTree(扁平 JSON,通过 parent_id 构建层级)
  8. ├── messages/ # Messages(每条独立文件)
  9. │ ├── {message_id}.json
  10. │ └── ...
  11. └── events.jsonl # 事件流(WebSocket 续传)
  12. Sub-Trace 是完全独立的 Trace,有自己的目录:
  13. .trace/{parent_id}@{mode}-{timestamp}-{seq}/
  14. ├── meta.json # parent_trace_id 指向父 Trace
  15. ├── goal.json
  16. ├── messages/
  17. └── events.jsonl
  18. """
  19. import hashlib
  20. import json
  21. import logging
  22. import os
  23. import uuid
  24. from datetime import datetime
  25. from pathlib import Path
  26. from typing import Dict, List, Optional, Any
  27. from .attachments import AttachmentRef
  28. from .models import Trace, Message
  29. from .goal_models import GoalTree, Goal, GoalStats
  30. logger = logging.getLogger(__name__)
  31. class FileSystemTraceStore:
  32. """文件系统 Trace 存储"""
  33. def __init__(self, base_path: str = ".trace"):
  34. self.base_path = Path(base_path)
  35. self.base_path.mkdir(exist_ok=True)
  36. def _get_trace_dir(self, trace_id: str) -> Path:
  37. """获取 trace 目录"""
  38. return self.base_path / trace_id
  39. def _get_meta_file(self, trace_id: str) -> Path:
  40. """获取 meta.json 文件路径"""
  41. return self._get_trace_dir(trace_id) / "meta.json"
  42. def _get_goal_file(self, trace_id: str) -> Path:
  43. """获取 goal.json 文件路径"""
  44. return self._get_trace_dir(trace_id) / "goal.json"
  45. def _get_messages_dir(self, trace_id: str) -> Path:
  46. """获取 messages 目录"""
  47. return self._get_trace_dir(trace_id) / "messages"
  48. def _get_message_file(self, trace_id: str, message_id: str) -> Path:
  49. """获取 message 文件路径"""
  50. return self._get_messages_dir(trace_id) / f"{message_id}.json"
  51. def _get_events_file(self, trace_id: str) -> Path:
  52. """获取 events.jsonl 文件路径"""
  53. return self._get_trace_dir(trace_id) / "events.jsonl"
  54. def _get_model_usage_file(self, trace_id: str) -> Path:
  55. """获取 model_usage.json 文件路径"""
  56. return self._get_trace_dir(trace_id) / "model_usage.json"
  57. def _get_attachments_dir(self, trace_id: str) -> Path:
  58. """Return the durable attachment directory for one trace."""
  59. return self._get_trace_dir(trace_id) / "attachments"
  60. # ===== Trace 操作 =====
  61. async def create_trace(self, trace: Trace) -> str:
  62. """创建新的 Trace"""
  63. trace_dir = self._get_trace_dir(trace.trace_id)
  64. trace_dir.mkdir(exist_ok=True)
  65. # 创建 messages 目录
  66. messages_dir = self._get_messages_dir(trace.trace_id)
  67. messages_dir.mkdir(exist_ok=True)
  68. self._get_attachments_dir(trace.trace_id).mkdir(exist_ok=True)
  69. # 写入 meta.json
  70. meta_file = self._get_meta_file(trace.trace_id)
  71. meta_file.write_text(
  72. json.dumps(trace.to_dict(), indent=2, ensure_ascii=False), encoding="utf-8"
  73. )
  74. # 创建空的 events.jsonl
  75. events_file = self._get_events_file(trace.trace_id)
  76. events_file.touch()
  77. return trace.trace_id
  78. async def store_message_attachment(
  79. self,
  80. *,
  81. trace_id: str,
  82. message_id: str,
  83. media_type: str,
  84. content: bytes,
  85. sha256: str,
  86. ) -> AttachmentRef:
  87. """Persist verified bytes without exposing filesystem paths."""
  88. _validate_path_segment(trace_id, "trace_id")
  89. _validate_path_segment(message_id, "message_id")
  90. if not isinstance(content, bytes):
  91. raise TypeError("content must be bytes")
  92. actual = f"sha256:{hashlib.sha256(content).hexdigest()}"
  93. if sha256 != actual:
  94. raise ValueError("attachment digest mismatch")
  95. if not self._get_meta_file(trace_id).exists():
  96. raise ValueError(f"Trace not found: {trace_id}")
  97. message = await self.get_message(message_id)
  98. if message is None or message.trace_id != trace_id:
  99. raise ValueError("attachment message does not belong to the Trace")
  100. reference = AttachmentRef(
  101. trace_id=trace_id,
  102. message_id=message_id,
  103. media_type=media_type,
  104. sha256=actual,
  105. size_bytes=len(content),
  106. )
  107. digest_hex = actual.removeprefix("sha256:")
  108. attachment_dir = self._get_attachments_dir(trace_id) / message_id
  109. attachment_dir.mkdir(parents=True, exist_ok=True)
  110. target = attachment_dir / f"{digest_hex}.bin"
  111. metadata = attachment_dir / f"{digest_hex}.json"
  112. if target.exists():
  113. if target.read_bytes() != content:
  114. raise ValueError("attachment content conflict")
  115. else:
  116. temporary = attachment_dir / f".{digest_hex}.{uuid.uuid4().hex}.tmp"
  117. try:
  118. with temporary.open("xb") as output:
  119. output.write(content)
  120. output.flush()
  121. os.fsync(output.fileno())
  122. os.replace(temporary, target)
  123. _fsync_directory(attachment_dir)
  124. finally:
  125. temporary.unlink(missing_ok=True)
  126. if metadata.exists():
  127. try:
  128. stored_reference = AttachmentRef.from_dict(
  129. json.loads(metadata.read_text(encoding="utf-8"))
  130. )
  131. except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
  132. raise ValueError("attachment metadata is invalid") from exc
  133. if stored_reference != reference:
  134. raise ValueError("attachment reference conflict")
  135. else:
  136. _write_atomic_json(metadata, reference.to_dict())
  137. _fsync_directory(attachment_dir)
  138. return reference
  139. async def read_message_attachment(self, ref: AttachmentRef) -> bytes:
  140. """Read and verify one content-addressed attachment."""
  141. if not isinstance(ref, AttachmentRef):
  142. raise TypeError("ref must be an AttachmentRef")
  143. _validate_path_segment(ref.trace_id, "trace_id")
  144. _validate_path_segment(ref.message_id, "message_id")
  145. digest_hex = ref.sha256.removeprefix("sha256:")
  146. target = (
  147. self._get_attachments_dir(ref.trace_id)
  148. / ref.message_id
  149. / f"{digest_hex}.bin"
  150. )
  151. metadata = target.with_suffix(".json")
  152. if not target.is_file():
  153. raise FileNotFoundError("trace attachment not found")
  154. if not metadata.is_file():
  155. raise ValueError("trace attachment metadata not found")
  156. try:
  157. stored_reference = AttachmentRef.from_dict(
  158. json.loads(metadata.read_text(encoding="utf-8"))
  159. )
  160. except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc:
  161. raise ValueError("trace attachment metadata is invalid") from exc
  162. if stored_reference != ref:
  163. raise ValueError("trace attachment reference does not match metadata")
  164. content = target.read_bytes()
  165. actual = f"sha256:{hashlib.sha256(content).hexdigest()}"
  166. if actual != ref.sha256 or len(content) != ref.size_bytes:
  167. raise ValueError("trace attachment integrity check failed")
  168. return content
  169. async def get_trace(self, trace_id: str) -> Optional[Trace]:
  170. """获取 Trace"""
  171. meta_file = self._get_meta_file(trace_id)
  172. if not meta_file.exists():
  173. return None
  174. data = json.loads(meta_file.read_text(encoding="utf-8"))
  175. # 解析 datetime 字段
  176. if data.get("created_at"):
  177. data["created_at"] = datetime.fromisoformat(data["created_at"])
  178. if data.get("completed_at"):
  179. data["completed_at"] = datetime.fromisoformat(data["completed_at"])
  180. return Trace.from_dict(data)
  181. async def update_trace(self, trace_id: str, **updates) -> None:
  182. """更新 Trace"""
  183. trace = await self.get_trace(trace_id)
  184. if not trace:
  185. return
  186. # 更新字段
  187. for key, value in updates.items():
  188. if hasattr(trace, key):
  189. setattr(trace, key, value)
  190. # 写回文件
  191. meta_file = self._get_meta_file(trace_id)
  192. meta_file.write_text(
  193. json.dumps(trace.to_dict(), indent=2, ensure_ascii=False), encoding="utf-8"
  194. )
  195. async def list_traces(
  196. self,
  197. mode: Optional[str] = None,
  198. agent_type: Optional[str] = None,
  199. uid: Optional[str] = None,
  200. status: Optional[str] = None,
  201. agent_role: Optional[str] = None,
  202. parent_trace_id: Optional[str] = None,
  203. root_trace_id: Optional[str] = None,
  204. task_id: Optional[str] = None,
  205. attempt_id: Optional[str] = None,
  206. validation_id: Optional[str] = None,
  207. operation_id: Optional[str] = None,
  208. limit: int = 50,
  209. ) -> List[Trace]:
  210. """列出 Traces"""
  211. traces = []
  212. if not self.base_path.exists():
  213. return []
  214. for trace_dir in self.base_path.iterdir():
  215. if not trace_dir.is_dir():
  216. continue
  217. meta_file = trace_dir / "meta.json"
  218. if not meta_file.exists():
  219. continue
  220. try:
  221. data = json.loads(meta_file.read_text(encoding="utf-8"))
  222. # 过滤
  223. if mode and data.get("mode") != mode:
  224. continue
  225. if agent_type and data.get("agent_type") != agent_type:
  226. continue
  227. if uid and data.get("uid") != uid:
  228. continue
  229. if status and data.get("status") != status:
  230. continue
  231. if agent_role and data.get("agent_role", "legacy") != agent_role:
  232. continue
  233. if parent_trace_id and data.get("parent_trace_id") != parent_trace_id:
  234. continue
  235. context = data.get("context") or {}
  236. if not isinstance(context, dict):
  237. continue
  238. if root_trace_id and not (
  239. data.get("trace_id") == root_trace_id
  240. or context.get("root_trace_id") == root_trace_id
  241. ):
  242. continue
  243. if task_id and context.get("task_id") != task_id:
  244. continue
  245. if attempt_id and context.get("attempt_id") != attempt_id:
  246. continue
  247. if validation_id and context.get("validation_id") != validation_id:
  248. continue
  249. if operation_id and context.get("operation_id") != operation_id:
  250. continue
  251. # 解析 datetime
  252. if data.get("created_at"):
  253. data["created_at"] = datetime.fromisoformat(data["created_at"])
  254. if data.get("completed_at"):
  255. data["completed_at"] = datetime.fromisoformat(data["completed_at"])
  256. traces.append(Trace.from_dict(data))
  257. except Exception:
  258. continue
  259. # 排序(最新的在前)
  260. traces.sort(key=lambda t: t.created_at, reverse=True)
  261. return traces[:limit]
  262. # ===== GoalTree 操作 =====
  263. async def get_goal_tree(self, trace_id: str) -> Optional[GoalTree]:
  264. """获取 GoalTree"""
  265. goal_file = self._get_goal_file(trace_id)
  266. if not goal_file.exists():
  267. return None
  268. try:
  269. data = json.loads(goal_file.read_text(encoding="utf-8"))
  270. return GoalTree.from_dict(data)
  271. except Exception:
  272. return None
  273. async def update_goal_tree(self, trace_id: str, tree: GoalTree) -> None:
  274. """更新完整 GoalTree"""
  275. goal_file = self._get_goal_file(trace_id)
  276. goal_file.write_text(
  277. json.dumps(tree.to_dict(), indent=2, ensure_ascii=False), encoding="utf-8"
  278. )
  279. async def add_goal(self, trace_id: str, goal: Goal) -> None:
  280. """添加 Goal 到 GoalTree"""
  281. tree = await self.get_goal_tree(trace_id)
  282. if not tree:
  283. return
  284. tree.goals.append(goal)
  285. await self.update_goal_tree(trace_id, tree)
  286. # 推送 goal_added 事件
  287. event_data = {"goal": goal.to_dict(), "parent_id": goal.parent_id}
  288. await self.append_event(trace_id, "goal_added", event_data)
  289. # 打印详细的 goal 信息
  290. desc_preview = (
  291. goal.description[:80] + "..."
  292. if len(goal.description) > 80
  293. else goal.description
  294. )
  295. print(f"[Goal Added] ID={goal.id}, Parent={goal.parent_id or 'root'}")
  296. print(f" 📝 {desc_preview}")
  297. if goal.reason:
  298. reason_preview = (
  299. goal.reason[:60] + "..." if len(goal.reason) > 60 else goal.reason
  300. )
  301. print(f" 💡 {reason_preview}")
  302. async def update_goal(
  303. self,
  304. trace_id: str,
  305. goal_id: str,
  306. cascade_completion: bool = True,
  307. **updates,
  308. ) -> None:
  309. """更新 Goal 字段"""
  310. tree = await self.get_goal_tree(trace_id)
  311. if not tree:
  312. return
  313. goal = tree.find(goal_id)
  314. if not goal:
  315. return
  316. # 更新字段
  317. for key, value in updates.items():
  318. if hasattr(goal, key):
  319. # 特殊处理 stats 字段(可能是 dict)
  320. if key in ["self_stats", "cumulative_stats"] and isinstance(
  321. value, dict
  322. ):
  323. value = GoalStats.from_dict(value)
  324. setattr(goal, key, value)
  325. await self.update_goal_tree(trace_id, tree)
  326. # 推送 goal_updated 事件
  327. # 如果状态变为 completed,检查是否需要级联完成父 Goal
  328. affected_goals = [{"goal_id": goal_id, "updates": updates}]
  329. if cascade_completion and updates.get("status") == "completed":
  330. # 检查级联完成:如果所有兄弟 Goal 都完成,父 Goal 也完成
  331. cascade_completed = await self._check_cascade_completion(trace_id, goal)
  332. affected_goals.extend(cascade_completed)
  333. await self.append_event(
  334. trace_id,
  335. "goal_updated",
  336. {"goal_id": goal_id, "updates": updates, "affected_goals": affected_goals},
  337. )
  338. print(
  339. f"[DEBUG] Pushed goal_updated event: goal_id={goal_id}, updates={updates}, affected={len(affected_goals)}"
  340. )
  341. # Goal 完成时触发知识评估
  342. if updates.get("status") in ["completed", "abandoned"]:
  343. pending = await self.get_pending_knowledge_entries(trace_id)
  344. if pending:
  345. # 在trace.context中设置标志,由runner主循环检查
  346. trace = await self.get_trace(trace_id)
  347. if trace:
  348. if not trace.context:
  349. trace.context = {}
  350. trace.context["pending_knowledge_eval"] = True
  351. trace.context["knowledge_eval_trigger"] = "goal_completion"
  352. await self.update_trace(trace_id, context=trace.context)
  353. logger.info(
  354. f"[Knowledge Eval] Goal {goal_id} 完成,设置评估标志,待评估知识: {len(pending)} 条"
  355. )
  356. async def _check_cascade_completion(
  357. self, trace_id: str, completed_goal: Goal
  358. ) -> List[Dict[str, Any]]:
  359. """
  360. 检查级联完成:如果一个 Goal 的所有子 Goal 都完成,则自动完成父 Goal
  361. Args:
  362. trace_id: Trace ID
  363. completed_goal: 刚完成的 Goal
  364. Returns:
  365. 受影响的父 Goals 列表(自动完成的)
  366. """
  367. if not completed_goal.parent_id:
  368. return []
  369. tree = await self.get_goal_tree(trace_id)
  370. if not tree:
  371. return []
  372. affected = []
  373. parent = tree.find(completed_goal.parent_id)
  374. if not parent:
  375. return []
  376. # 获取父 Goal 的所有子 Goal
  377. children = tree.get_children(parent.id)
  378. # 检查是否所有子 Goal 都已完成(排除 abandoned)
  379. all_completed = all(
  380. child.status in ["completed", "abandoned"] for child in children
  381. )
  382. if all_completed and parent.status != "completed":
  383. # 自动完成父 Goal
  384. parent.status = "completed"
  385. if not parent.summary:
  386. # 生成自动摘要
  387. completed_count = sum(1 for c in children if c.status == "completed")
  388. parent.summary = f"所有子目标已完成 ({completed_count}/{len(children)})"
  389. await self.update_goal_tree(trace_id, tree)
  390. affected.append(
  391. {
  392. "goal_id": parent.id,
  393. "status": "completed",
  394. "summary": parent.summary,
  395. "cumulative_stats": parent.cumulative_stats.to_dict(),
  396. }
  397. )
  398. # 递归检查祖父 Goal
  399. grandparent_affected = await self._check_cascade_completion(
  400. trace_id, parent
  401. )
  402. affected.extend(grandparent_affected)
  403. return affected
  404. # ===== Message 操作 =====
  405. async def add_message(self, message: Message) -> str:
  406. """
  407. 添加 Message
  408. 自动更新关联 Goal 的 stats(self_stats 和祖先的 cumulative_stats)
  409. """
  410. trace_id = message.trace_id
  411. # 1. 写入 message 文件
  412. messages_dir = self._get_messages_dir(trace_id)
  413. message_file = messages_dir / f"{message.message_id}.json"
  414. message_file.write_text(
  415. json.dumps(message.to_dict(), indent=2, ensure_ascii=False),
  416. encoding="utf-8",
  417. )
  418. # 2. 更新 trace 统计
  419. trace = await self.get_trace(trace_id)
  420. if trace:
  421. trace.total_messages += 1
  422. trace.last_sequence = max(trace.last_sequence, message.sequence)
  423. # 累计 tokens(完整版)
  424. if message.prompt_tokens:
  425. trace.total_prompt_tokens += message.prompt_tokens
  426. if message.completion_tokens:
  427. trace.total_completion_tokens += message.completion_tokens
  428. if message.reasoning_tokens:
  429. trace.total_reasoning_tokens += message.reasoning_tokens
  430. if message.cache_creation_tokens:
  431. trace.total_cache_creation_tokens += message.cache_creation_tokens
  432. if message.cache_read_tokens:
  433. trace.total_cache_read_tokens += message.cache_read_tokens
  434. # 向后兼容:也更新 total_tokens
  435. if message.tokens:
  436. trace.total_tokens += message.tokens
  437. elif message.prompt_tokens or message.completion_tokens:
  438. trace.total_tokens += (message.prompt_tokens or 0) + (
  439. message.completion_tokens or 0
  440. )
  441. if message.cost:
  442. trace.total_cost += message.cost
  443. if message.duration_ms:
  444. trace.total_duration_ms += message.duration_ms
  445. # 更新 Trace
  446. await self.update_trace(
  447. trace_id,
  448. total_messages=trace.total_messages,
  449. last_sequence=trace.last_sequence,
  450. total_tokens=trace.total_tokens,
  451. total_prompt_tokens=trace.total_prompt_tokens,
  452. total_completion_tokens=trace.total_completion_tokens,
  453. total_reasoning_tokens=trace.total_reasoning_tokens,
  454. total_cache_creation_tokens=trace.total_cache_creation_tokens,
  455. total_cache_read_tokens=trace.total_cache_read_tokens,
  456. total_cost=trace.total_cost,
  457. total_duration_ms=trace.total_duration_ms,
  458. )
  459. # 3. 更新 Goal stats
  460. await self._update_goal_stats(trace_id, message)
  461. # 4. 追加 message_added 事件
  462. affected_goals = await self._get_affected_goals(trace_id, message)
  463. event_id = await self.append_event(
  464. trace_id,
  465. "message_added",
  466. {"message": message.to_dict(), "affected_goals": affected_goals},
  467. )
  468. if event_id:
  469. try:
  470. from . import websocket as trace_ws
  471. await trace_ws.broadcast_message_added(
  472. trace_id=trace_id,
  473. event_id=event_id,
  474. message_dict=message.to_dict(),
  475. affected_goals=affected_goals,
  476. )
  477. except Exception:
  478. logger.exception(
  479. "Failed to broadcast message_added (trace_id=%s, event_id=%s)",
  480. trace_id,
  481. event_id,
  482. )
  483. return message.message_id
  484. async def _update_goal_stats(self, trace_id: str, message: Message) -> None:
  485. """更新 Goal 的 self_stats 和祖先的 cumulative_stats"""
  486. tree = await self.get_goal_tree(trace_id)
  487. if not tree:
  488. return
  489. # 找到关联的 Goal
  490. goal = tree.find(message.goal_id)
  491. if not goal:
  492. return
  493. # 更新自身 self_stats
  494. goal.self_stats.message_count += 1
  495. if message.tokens:
  496. goal.self_stats.total_tokens += message.tokens
  497. if message.cost:
  498. goal.self_stats.total_cost += message.cost
  499. # TODO: 更新 preview(工具调用摘要)
  500. # 更新自身 cumulative_stats
  501. goal.cumulative_stats.message_count += 1
  502. if message.tokens:
  503. goal.cumulative_stats.total_tokens += message.tokens
  504. if message.cost:
  505. goal.cumulative_stats.total_cost += message.cost
  506. # 沿祖先链向上更新 cumulative_stats
  507. current_goal = goal
  508. while current_goal.parent_id:
  509. parent = tree.find(current_goal.parent_id)
  510. if not parent:
  511. break
  512. parent.cumulative_stats.message_count += 1
  513. if message.tokens:
  514. parent.cumulative_stats.total_tokens += message.tokens
  515. if message.cost:
  516. parent.cumulative_stats.total_cost += message.cost
  517. current_goal = parent
  518. # 保存更新后的 tree
  519. await self.update_goal_tree(trace_id, tree)
  520. async def _get_affected_goals(
  521. self, trace_id: str, message: Message
  522. ) -> List[Dict[str, Any]]:
  523. """获取受影响的 Goals(自身 + 所有祖先)"""
  524. tree = await self.get_goal_tree(trace_id)
  525. if not tree:
  526. return []
  527. goal = tree.find(message.goal_id)
  528. if not goal:
  529. return []
  530. affected = []
  531. # 添加自身(包含 self_stats 和 cumulative_stats)
  532. affected.append(
  533. {
  534. "goal_id": goal.id,
  535. "self_stats": goal.self_stats.to_dict(),
  536. "cumulative_stats": goal.cumulative_stats.to_dict(),
  537. }
  538. )
  539. # 添加所有祖先(仅 cumulative_stats)
  540. current_goal = goal
  541. while current_goal.parent_id:
  542. parent = tree.find(current_goal.parent_id)
  543. if not parent:
  544. break
  545. affected.append(
  546. {
  547. "goal_id": parent.id,
  548. "cumulative_stats": parent.cumulative_stats.to_dict(),
  549. }
  550. )
  551. current_goal = parent
  552. return affected
  553. async def get_message(self, message_id: str) -> Optional[Message]:
  554. """获取 Message(扫描所有 trace)"""
  555. for trace_dir in self.base_path.iterdir():
  556. if not trace_dir.is_dir():
  557. continue
  558. # 检查 messages 目录
  559. message_file = trace_dir / "messages" / f"{message_id}.json"
  560. if message_file.exists():
  561. try:
  562. data = json.loads(message_file.read_text(encoding="utf-8"))
  563. return Message.from_dict(data)
  564. except Exception:
  565. pass
  566. return None
  567. async def get_trace_messages(
  568. self,
  569. trace_id: str,
  570. ) -> List[Message]:
  571. """获取 Trace 的所有 Messages(包含所有分支,按 sequence 排序)"""
  572. messages_dir = self._get_messages_dir(trace_id)
  573. if not messages_dir.exists():
  574. return []
  575. messages = []
  576. for message_file in messages_dir.glob("*.json"):
  577. try:
  578. data = json.loads(message_file.read_text(encoding="utf-8"))
  579. msg = Message.from_dict(data)
  580. messages.append(msg)
  581. except Exception:
  582. continue
  583. # 按 sequence 排序
  584. messages.sort(key=lambda m: m.sequence)
  585. return messages
  586. async def get_main_path_messages(
  587. self, trace_id: str, head_sequence: int
  588. ) -> List[Message]:
  589. """
  590. 获取从 head_sequence 沿 parent_sequence 链回溯到 root 的完整路径
  591. 此函数是通用的路径追溯函数,返回从指定 head 到 root 的完整消息链。
  592. 只要 trace.head_sequence 管理正确(指向主路径),此函数自然返回主路径消息。
  593. 侧分支消息通过 parent_sequence 链自然被跳过(因为主路径的 parent 不指向侧分支)。
  594. Returns:
  595. 按 sequence 正序排列的路径 Message 列表
  596. """
  597. # 加载所有消息,建立 sequence -> Message 索引
  598. all_messages = await self.get_trace_messages(trace_id)
  599. messages_by_seq = {m.sequence: m for m in all_messages}
  600. # 从 head 沿 parent chain 回溯
  601. path = []
  602. seq = head_sequence
  603. while seq is not None:
  604. msg = messages_by_seq.get(seq)
  605. if not msg:
  606. break
  607. path.append(msg)
  608. seq = msg.parent_sequence
  609. # 反转为正序(root → head)
  610. path.reverse()
  611. return path
  612. async def get_messages_by_goal(self, trace_id: str, goal_id: str) -> List[Message]:
  613. """获取指定 Goal 关联的所有 Messages"""
  614. all_messages = await self.get_trace_messages(trace_id)
  615. return [m for m in all_messages if m.goal_id == goal_id]
  616. async def update_message(self, message_id: str, **updates) -> None:
  617. """更新 Message 字段"""
  618. message = await self.get_message(message_id)
  619. if not message:
  620. return
  621. # 更新字段
  622. for key, value in updates.items():
  623. if hasattr(message, key):
  624. setattr(message, key, value)
  625. # 确定文件路径
  626. messages_dir = self._get_messages_dir(message.trace_id)
  627. message_file = messages_dir / f"{message_id}.json"
  628. message_file.write_text(
  629. json.dumps(message.to_dict(), indent=2, ensure_ascii=False),
  630. encoding="utf-8",
  631. )
  632. async def abandon_messages_after(
  633. self, trace_id: str, cutoff_sequence: int
  634. ) -> List[str]:
  635. """
  636. 将 sequence > cutoff_sequence 的 active messages 标记为 abandoned。
  637. 返回被 abandon 的 message_id 列表。
  638. """
  639. all_messages = await self.get_trace_messages(trace_id)
  640. abandoned_ids = []
  641. now = datetime.now()
  642. for msg in all_messages:
  643. if msg.sequence > cutoff_sequence and msg.status == "active":
  644. msg.status = "abandoned"
  645. msg.abandoned_at = now
  646. # 直接写回文件
  647. message_file = (
  648. self._get_messages_dir(trace_id) / f"{msg.message_id}.json"
  649. )
  650. message_file.write_text(
  651. json.dumps(msg.to_dict(), indent=2, ensure_ascii=False),
  652. encoding="utf-8",
  653. )
  654. abandoned_ids.append(msg.message_id)
  655. return abandoned_ids
  656. # ===== 模型使用追踪 =====
  657. async def record_model_usage(
  658. self,
  659. trace_id: str,
  660. sequence: int,
  661. role: str,
  662. model: str,
  663. prompt_tokens: int,
  664. completion_tokens: int,
  665. cache_read_tokens: int = 0,
  666. tool_name: Optional[str] = None,
  667. ) -> None:
  668. """
  669. 记录模型使用情况到 model_usage.json
  670. Args:
  671. trace_id: Trace ID
  672. sequence: 消息序号
  673. role: 角色(assistant/tool)
  674. model: 模型名称
  675. prompt_tokens: 输入tokens
  676. completion_tokens: 输出tokens
  677. cache_read_tokens: 缓存读取tokens
  678. tool_name: 工具名称(role=tool时)
  679. """
  680. usage_file = self._get_model_usage_file(trace_id)
  681. # 读取现有数据
  682. if usage_file.exists():
  683. data = json.loads(usage_file.read_text(encoding="utf-8"))
  684. else:
  685. data = {
  686. "summary": {
  687. "total_models": 0,
  688. "total_tokens": 0,
  689. "total_cache_read_tokens": 0,
  690. "agent_tokens": 0,
  691. "tool_tokens": 0,
  692. },
  693. "models": [],
  694. "timeline": [],
  695. }
  696. # 更新summary
  697. total_tokens = prompt_tokens + completion_tokens
  698. data["summary"]["total_tokens"] += total_tokens
  699. data["summary"]["total_cache_read_tokens"] += cache_read_tokens
  700. if role == "assistant":
  701. data["summary"]["agent_tokens"] += total_tokens
  702. source = "agent"
  703. else:
  704. data["summary"]["tool_tokens"] += total_tokens
  705. source = f"tool:{tool_name}" if tool_name else "tool"
  706. # 更新models列表
  707. model_entry = None
  708. for m in data["models"]:
  709. if m["model"] == model and m["source"] == source:
  710. model_entry = m
  711. break
  712. if model_entry:
  713. model_entry["prompt_tokens"] += prompt_tokens
  714. model_entry["completion_tokens"] += completion_tokens
  715. model_entry["total_tokens"] += total_tokens
  716. model_entry["cache_read_tokens"] += cache_read_tokens
  717. model_entry["call_count"] += 1
  718. else:
  719. data["models"].append(
  720. {
  721. "model": model,
  722. "source": source,
  723. "prompt_tokens": prompt_tokens,
  724. "completion_tokens": completion_tokens,
  725. "total_tokens": total_tokens,
  726. "cache_read_tokens": cache_read_tokens,
  727. "call_count": 1,
  728. }
  729. )
  730. data["summary"]["total_models"] = len(data["models"])
  731. # 添加到timeline
  732. timeline_entry = {
  733. "sequence": sequence,
  734. "role": role,
  735. "model": model,
  736. "prompt_tokens": prompt_tokens,
  737. "completion_tokens": completion_tokens,
  738. }
  739. if cache_read_tokens > 0:
  740. timeline_entry["cache_read_tokens"] = cache_read_tokens
  741. if tool_name:
  742. timeline_entry["tool_name"] = tool_name
  743. data["timeline"].append(timeline_entry)
  744. # 写回文件
  745. usage_file.write_text(
  746. json.dumps(data, indent=2, ensure_ascii=False), encoding="utf-8"
  747. )
  748. # ===== 事件流操作(用于 WebSocket 断线续传)=====
  749. async def get_events(
  750. self, trace_id: str, since_event_id: int = 0
  751. ) -> List[Dict[str, Any]]:
  752. """获取事件流"""
  753. events_file = self._get_events_file(trace_id)
  754. if not events_file.exists():
  755. return []
  756. events = []
  757. with events_file.open("r", encoding="utf-8") as f:
  758. for line in f:
  759. try:
  760. event = json.loads(line.strip())
  761. if event.get("event_id", 0) > since_event_id:
  762. events.append(event)
  763. except Exception:
  764. continue
  765. return events
  766. async def append_event(
  767. self, trace_id: str, event_type: str, payload: Dict[str, Any]
  768. ) -> int:
  769. """追加事件,返回 event_id"""
  770. # 获取 trace 并递增 event_id
  771. trace = await self.get_trace(trace_id)
  772. if not trace:
  773. return 0
  774. trace.last_event_id += 1
  775. event_id = trace.last_event_id
  776. # 更新 trace 的 last_event_id
  777. await self.update_trace(trace_id, last_event_id=event_id)
  778. # 创建事件
  779. event = {
  780. "event_id": event_id,
  781. "event": event_type,
  782. "ts": datetime.now().isoformat(),
  783. **payload,
  784. }
  785. # 追加到 events.jsonl
  786. events_file = self._get_events_file(trace_id)
  787. with events_file.open("a", encoding="utf-8") as f:
  788. f.write(json.dumps(event, ensure_ascii=False) + "\n")
  789. return event_id
  790. # ===== Cognition Log 管理 =====
  791. def _get_cognition_log_file(self, trace_id: str) -> Path:
  792. """获取 cognition_log.json 文件路径"""
  793. return self._get_trace_dir(trace_id) / "cognition_log.json"
  794. def _get_knowledge_log_file(self, trace_id: str) -> Path:
  795. """兼容旧接口:优先使用 cognition_log,回退到 knowledge_log"""
  796. cognition_file = self._get_cognition_log_file(trace_id)
  797. if cognition_file.exists():
  798. return cognition_file
  799. legacy_file = self._get_trace_dir(trace_id) / "knowledge_log.json"
  800. if legacy_file.exists():
  801. return legacy_file
  802. return cognition_file # 新建时用 cognition_log
  803. async def get_cognition_log(self, trace_id: str) -> Dict[str, Any]:
  804. """读取认知日志"""
  805. log_file = self._get_cognition_log_file(trace_id)
  806. if log_file.exists():
  807. return json.loads(log_file.read_text(encoding="utf-8"))
  808. # 兼容旧格式:如果只有 knowledge_log.json,读取并转换
  809. legacy_file = self._get_trace_dir(trace_id) / "knowledge_log.json"
  810. if legacy_file.exists():
  811. return json.loads(legacy_file.read_text(encoding="utf-8"))
  812. return {"trace_id": trace_id, "events": []}
  813. async def get_knowledge_log(self, trace_id: str) -> Dict[str, Any]:
  814. """兼容旧接口"""
  815. log = await self.get_cognition_log(trace_id)
  816. # 旧格式用 entries,新格式用 events
  817. if "entries" not in log and "events" in log:
  818. log["entries"] = log["events"]
  819. return log
  820. async def append_cognition_event(
  821. self,
  822. trace_id: str,
  823. event: Dict[str, Any],
  824. ) -> None:
  825. """追加认知事件到 cognition_log.json。
  826. 所有事件共有字段:
  827. type: str 事件类型(见下表)
  828. timestamp: str ISO 格式时间戳(框架自动写入)
  829. 已定义的事件类型及典型字段:
  830. type="query" — 知识注入查询(goal focus 时触发)
  831. sequence, goal_id, query, response, source_ids, sources
  832. type="evaluation" — 知识评估(Goal 完成/压缩前/任务结束触发)
  833. knowledge_id, eval_result{relevance, utility, notes}, trigger_event
  834. type="extraction_pending" — 反思侧分支暂存的待审核提取(Phase 1.2+)
  835. extraction_id, sequence, goal_id, branch_id, payload
  836. (payload 字段与 knowledge_save 参数一一对应)
  837. type="extraction_reviewed" — 人工审核决策(CLI / HTTP API 写入)
  838. extraction_id, decision("approve"/"edit"/"discard"), edited_payload?
  839. type="extraction_committed" — 已上传到 KnowHub
  840. extraction_id, knowledge_id
  841. type="reflection" — Dream 的 per-trace 反思摘要(Phase 2.4 / 3.1)
  842. sequence_range: [start, end] 本次反思覆盖的消息区间
  843. summary: str LLM 生成的反思摘要
  844. consumed_at: 可选, ISO 时间戳 当跨 trace 整合已消化此反思时写入
  845. 其他字段可按需附加,不做强校验(演进友好)。
  846. """
  847. log = await self.get_cognition_log(trace_id)
  848. if "events" not in log:
  849. log["events"] = log.pop("entries", [])
  850. event["timestamp"] = datetime.now().isoformat()
  851. log["events"].append(event)
  852. log_file = self._get_cognition_log_file(trace_id)
  853. log_file.write_text(
  854. json.dumps(log, indent=2, ensure_ascii=False), encoding="utf-8"
  855. )
  856. async def append_knowledge_entry(
  857. self,
  858. trace_id: str,
  859. knowledge_id: str,
  860. goal_id: str,
  861. injected_at_sequence: int,
  862. task: str,
  863. content: str,
  864. ) -> None:
  865. """兼容旧接口:追加知识注入记录(转换为 query 事件)"""
  866. await self.append_cognition_event(
  867. trace_id=trace_id,
  868. event={
  869. "type": "query",
  870. "sequence": injected_at_sequence,
  871. "goal_id": goal_id,
  872. "query": task,
  873. "response": "",
  874. "source_ids": [knowledge_id],
  875. "sources": [
  876. {"id": knowledge_id, "task": task, "content": content[:500]}
  877. ],
  878. },
  879. )
  880. async def update_knowledge_evaluation(
  881. self,
  882. trace_id: str,
  883. knowledge_id: str,
  884. eval_result: Dict[str, Any],
  885. trigger_event: str,
  886. ) -> None:
  887. """更新知识评估结果(兼容旧格式 + 新 cognition_log 格式)
  888. 旧格式:更新 entries[] 中匹配 knowledge_id 的条目的 eval_result
  889. 新格式:追加 evaluation 事件到 events[]
  890. """
  891. log = await self.get_cognition_log(trace_id)
  892. events = log.get("events", log.get("entries", []))
  893. # 旧格式兼容:直接更新 entries 中的 eval_result 字段
  894. if "entries" in log:
  895. matching = [
  896. (i, e)
  897. for i, e in enumerate(log["entries"])
  898. if e.get("knowledge_id") == knowledge_id
  899. and e.get("eval_result") is None
  900. ]
  901. if matching:
  902. matching.sort(
  903. key=lambda x: x[1].get("injected_at_sequence", 0), reverse=True
  904. )
  905. _, entry = matching[0]
  906. entry["eval_result"] = eval_result
  907. entry["evaluated_at"] = datetime.now().isoformat()
  908. entry["evaluated_at_trigger"] = trigger_event
  909. log_file = self._get_knowledge_log_file(trace_id)
  910. log_file.write_text(
  911. json.dumps(log, indent=2, ensure_ascii=False), encoding="utf-8"
  912. )
  913. return
  914. # 新格式:追加 evaluation 事件
  915. # 找到包含该 knowledge_id 的最近 query 事件
  916. query_events = [
  917. e
  918. for e in events
  919. if e.get("type") == "query" and knowledge_id in e.get("source_ids", [])
  920. ]
  921. query_sequence = query_events[-1]["sequence"] if query_events else None
  922. await self.append_cognition_event(
  923. trace_id=trace_id,
  924. event={
  925. "type": "evaluation",
  926. "sequence": max((e.get("sequence", 0) for e in events), default=0) + 1,
  927. "query_sequence": query_sequence,
  928. "trigger": trigger_event,
  929. "assessments": [
  930. {
  931. "source_id": knowledge_id,
  932. "status": eval_result.get("eval_status", ""),
  933. "reason": eval_result.get("reason", ""),
  934. }
  935. ],
  936. },
  937. )
  938. async def get_pending_knowledge_entries(
  939. self, trace_id: str
  940. ) -> List[Dict[str, Any]]:
  941. """获取所有待评估的知识条目(兼容旧格式 + 新格式)"""
  942. log = await self.get_cognition_log(trace_id)
  943. # 旧格式
  944. if "entries" in log:
  945. return [e for e in log["entries"] if e.get("eval_result") is None]
  946. # 新格式:找没有对应 evaluation 事件的 query 事件
  947. events = log.get("events", [])
  948. query_events = [e for e in events if e.get("type") == "query"]
  949. eval_events = [e for e in events if e.get("type") == "evaluation"]
  950. # 已评估的 query sequences
  951. evaluated_sequences = {e.get("query_sequence") for e in eval_events}
  952. pending = []
  953. for qe in query_events:
  954. if qe.get("sequence") not in evaluated_sequences:
  955. # 转为旧格式兼容(runner 中的评估逻辑期望此格式)
  956. for source in qe.get("sources", []):
  957. pending.append(
  958. {
  959. "knowledge_id": source.get("id", ""),
  960. "goal_id": qe.get("goal_id", ""),
  961. "injected_at_sequence": qe.get("sequence", 0),
  962. "task": source.get("task", ""),
  963. "content": source.get("content", ""),
  964. "query_sequence": qe.get("sequence"),
  965. }
  966. )
  967. return pending
  968. async def update_user_feedback(
  969. self, trace_id: str, knowledge_id: str, user_feedback: Dict[str, Any]
  970. ) -> None:
  971. """记录用户对知识的反馈(confirm/override)"""
  972. log = await self.get_cognition_log(trace_id)
  973. # 旧格式
  974. if "entries" in log:
  975. matching = [
  976. (i, e)
  977. for i, e in enumerate(log["entries"])
  978. if e.get("knowledge_id") == knowledge_id
  979. ]
  980. if matching:
  981. matching.sort(
  982. key=lambda x: x[1].get("injected_at_sequence", 0), reverse=True
  983. )
  984. _, entry = matching[0]
  985. entry["user_feedback"] = user_feedback
  986. log_file = self._get_knowledge_log_file(trace_id)
  987. log_file.write_text(
  988. json.dumps(log, indent=2, ensure_ascii=False), encoding="utf-8"
  989. )
  990. return
  991. # 新格式:追加 user_feedback 事件(或直接记录在 evaluation 上)
  992. await self.append_cognition_event(
  993. trace_id=trace_id,
  994. event={
  995. "type": "user_feedback",
  996. "knowledge_id": knowledge_id,
  997. "feedback": user_feedback,
  998. },
  999. )
  1000. def _validate_path_segment(value: str, field_name: str) -> None:
  1001. if (
  1002. not isinstance(value, str)
  1003. or not value
  1004. or value in {".", ".."}
  1005. or Path(value).name != value
  1006. or "/" in value
  1007. or "\\" in value
  1008. ):
  1009. raise ValueError(f"{field_name} must be a safe path segment")
  1010. def _fsync_directory(path: Path) -> None:
  1011. descriptor = os.open(path, os.O_RDONLY)
  1012. try:
  1013. os.fsync(descriptor)
  1014. finally:
  1015. os.close(descriptor)
  1016. def _write_atomic_json(path: Path, value: Dict[str, Any]) -> None:
  1017. temporary = path.parent / f".{path.name}.{uuid.uuid4().hex}.tmp"
  1018. try:
  1019. with temporary.open("x", encoding="utf-8") as output:
  1020. json.dump(value, output, ensure_ascii=False, sort_keys=True)
  1021. output.flush()
  1022. os.fsync(output.fileno())
  1023. os.replace(temporary, path)
  1024. finally:
  1025. temporary.unlink(missing_ok=True)