protocols.py 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263
  1. """
  2. Trace Storage Protocol - Trace 存储接口定义
  3. 使用 Protocol 定义接口,允许不同的存储实现(内存、PostgreSQL、Neo4j 等)
  4. """
  5. from typing import Protocol, List, Optional, Dict, Any, runtime_checkable
  6. from .models import Trace, Message
  7. from .goal_models import GoalTree, Goal
  8. from .attachments import AttachmentRef
  9. @runtime_checkable
  10. class TraceStore(Protocol):
  11. """Trace + Message + GoalTree 存储接口"""
  12. # ===== Trace 操作 =====
  13. async def create_trace(self, trace: Trace) -> str:
  14. """
  15. 创建新的 Trace
  16. Args:
  17. trace: Trace 对象
  18. Returns:
  19. trace_id
  20. """
  21. ...
  22. async def get_trace(self, trace_id: str) -> Optional[Trace]:
  23. """获取 Trace"""
  24. ...
  25. async def update_trace(self, trace_id: str, **updates) -> None:
  26. """
  27. 更新 Trace
  28. Args:
  29. trace_id: Trace ID
  30. **updates: 要更新的字段
  31. """
  32. ...
  33. async def list_traces(
  34. self,
  35. mode: Optional[str] = None,
  36. agent_type: Optional[str] = None,
  37. uid: Optional[str] = None,
  38. status: Optional[str] = None,
  39. agent_role: Optional[str] = None,
  40. parent_trace_id: Optional[str] = None,
  41. root_trace_id: Optional[str] = None,
  42. task_id: Optional[str] = None,
  43. attempt_id: Optional[str] = None,
  44. validation_id: Optional[str] = None,
  45. operation_id: Optional[str] = None,
  46. limit: int = 50,
  47. ) -> List[Trace]:
  48. """列出 Traces"""
  49. ...
  50. # ===== GoalTree 操作 =====
  51. async def get_goal_tree(self, trace_id: str) -> Optional[GoalTree]:
  52. """
  53. 获取 GoalTree
  54. Args:
  55. trace_id: Trace ID
  56. Returns:
  57. GoalTree 对象,如果不存在返回 None
  58. """
  59. ...
  60. async def update_goal_tree(self, trace_id: str, tree: GoalTree) -> None:
  61. """
  62. 更新完整 GoalTree
  63. Args:
  64. trace_id: Trace ID
  65. tree: GoalTree 对象
  66. """
  67. ...
  68. async def add_goal(self, trace_id: str, goal: Goal) -> None:
  69. """
  70. 添加 Goal 到 GoalTree
  71. Args:
  72. trace_id: Trace ID
  73. goal: Goal 对象
  74. """
  75. ...
  76. async def update_goal(
  77. self,
  78. trace_id: str,
  79. goal_id: str,
  80. cascade_completion: bool = True,
  81. **updates,
  82. ) -> None:
  83. """
  84. 更新 Goal 字段
  85. Args:
  86. trace_id: Trace ID
  87. goal_id: Goal ID
  88. cascade_completion: 是否自动级联完成父 Goal(legacy 默认开启)
  89. **updates: 要更新的字段(如 status, summary, self_stats, cumulative_stats)
  90. """
  91. ...
  92. # ===== Message 操作 =====
  93. async def add_message(self, message: Message) -> str:
  94. """
  95. 添加 Message
  96. 自动更新关联 Goal 的 stats(self_stats 和祖先的 cumulative_stats)
  97. Args:
  98. message: Message 对象
  99. Returns:
  100. message_id
  101. """
  102. ...
  103. async def get_message(self, message_id: str) -> Optional[Message]:
  104. """获取 Message"""
  105. ...
  106. async def get_trace_messages(
  107. self,
  108. trace_id: str,
  109. ) -> List[Message]:
  110. """
  111. 获取 Trace 的所有 Messages(按 sequence 排序)
  112. 返回该 Trace 下所有消息(包含所有分支)。
  113. 如需获取特定主路径的消息,使用 get_main_path_messages()。
  114. Args:
  115. trace_id: Trace ID
  116. Returns:
  117. Message 列表
  118. """
  119. ...
  120. async def get_model_usage(self, trace_id: str) -> Dict[str, Any]:
  121. """Return the durable provider usage summary for one Trace."""
  122. ...
  123. async def get_main_path_messages(
  124. self, trace_id: str, head_sequence: int
  125. ) -> List[Message]:
  126. """
  127. 获取主路径上的消息(从 head_sequence 沿 parent_sequence 链回溯到 root)
  128. Args:
  129. trace_id: Trace ID
  130. head_sequence: 主路径头节点的 sequence
  131. Returns:
  132. 按 sequence 正序排列的主路径 Message 列表
  133. """
  134. ...
  135. async def get_messages_by_goal(self, trace_id: str, goal_id: str) -> List[Message]:
  136. """
  137. 获取指定 Goal 关联的所有 Messages
  138. Args:
  139. trace_id: Trace ID
  140. goal_id: Goal ID
  141. Returns:
  142. Message 列表
  143. """
  144. ...
  145. async def update_message(self, message_id: str, **updates) -> None:
  146. """
  147. 更新 Message 字段(用于状态变更、错误记录等)
  148. Args:
  149. message_id: Message ID
  150. **updates: 要更新的字段
  151. """
  152. ...
  153. async def abandon_messages_after(
  154. self, trace_id: str, cutoff_sequence: int
  155. ) -> List[str]:
  156. """
  157. 将 cutoff_sequence 之后的所有 active 消息标记为 abandoned(回溯专用)
  158. Args:
  159. trace_id: Trace ID
  160. cutoff_sequence: 截断点(该 sequence 及之前的消息保留)
  161. Returns:
  162. 被标记为 abandoned 的 message_id 列表
  163. """
  164. ...
  165. # ===== 事件流操作(用于 WebSocket 断线续传)=====
  166. async def get_events(
  167. self, trace_id: str, since_event_id: int = 0
  168. ) -> List[Dict[str, Any]]:
  169. """
  170. 获取事件流(用于 WS 断线续传)
  171. Args:
  172. trace_id: Trace ID
  173. since_event_id: 从哪个事件 ID 开始(0 表示全部)
  174. Returns:
  175. 事件列表(按 event_id 排序)
  176. """
  177. ...
  178. async def append_event(
  179. self, trace_id: str, event_type: str, payload: Dict[str, Any]
  180. ) -> int:
  181. """
  182. 追加事件,返回 event_id
  183. Args:
  184. trace_id: Trace ID
  185. event_type: 事件类型
  186. payload: 事件数据
  187. Returns:
  188. event_id: 新事件的 ID
  189. """
  190. ...
  191. @runtime_checkable
  192. class TraceAttachmentStore(Protocol):
  193. """Optional port for content-addressed message attachments."""
  194. async def store_message_attachment(
  195. self,
  196. *,
  197. trace_id: str,
  198. message_id: str,
  199. media_type: str,
  200. content: bytes,
  201. sha256: str,
  202. ) -> AttachmentRef:
  203. """Persist bytes after verifying their caller-supplied digest."""
  204. ...
  205. async def read_message_attachment(self, ref: AttachmentRef) -> bytes:
  206. """Read bytes and verify the durable reference."""
  207. ...