Coverage for src\genesis\orchestrator\task_queue.py: 100%

98 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-26 14:20 +0800

1"""任务级持久化(T16,OV7)。 

2 

3背景:api-design §5 定义 TaskQueue 抽象(v1 仅 InMemoryQueue),但崩溃恢复 

4只到会话级(agent-runtime §3.6),任务层数据丢失。OV7 裁定任务级持久化。 

5 

6实现:PersistentTaskQueue —— TaskQueue 抽象 + SQLite 落盘。 

7 - 任务状态 / payload / result 全部写入 SQLite(零外部依赖,标准库 sqlite3) 

8 - 幂等去重(§5.3):同 (session_id, step, chapter_id) 已完成 → 返回缓存结果 

9 - recover():重启后 running → failed(中断标记),pending 保留待执行 

10""" 

11 

12from __future__ import annotations 

13 

14import json 

15import sqlite3 

16from abc import ABC, abstractmethod 

17from dataclasses import dataclass, field 

18from datetime import datetime, timezone 

19from enum import Enum 

20from pathlib import Path 

21from typing import Any 

22 

23from pydantic import BaseModel, Field 

24 

25 

26class TaskStatus(str, Enum): 

27 PENDING = "pending" 

28 RUNNING = "running" 

29 COMPLETED = "completed" 

30 FAILED = "failed" 

31 CANCELLED = "cancelled" 

32 

33 

34# 终态:不可再转移(cancel 仅对非终态生效) 

35_TERMINAL = frozenset({TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED}) 

36 

37 

38class TaskSpec(BaseModel): 

39 """任务投递规格(api-design §5.1)。""" 

40 

41 task_id: str 

42 session_id: str 

43 step: str 

44 chapter_id: str | None = None 

45 payload: dict[str, Any] = Field(default_factory=dict) 

46 idempotency_key: str 

47 

48 

49@dataclass 

50class TaskHandle: 

51 """任务句柄(含状态与结果)。""" 

52 

53 task_id: str 

54 session_id: str 

55 step: str 

56 chapter_id: str | None = None 

57 payload: dict[str, Any] = field(default_factory=dict) 

58 idempotency_key: str = "" 

59 status: TaskStatus = TaskStatus.PENDING 

60 result: Any | None = None 

61 retry_count: int = 0 

62 created_at: str = "" 

63 updated_at: str = "" 

64 

65 

66class TaskQueue(ABC): 

67 """统一任务队列抽象(v1 仅 PersistentTaskQueue;Redis/Valkey 为 v2 预留)。""" 

68 

69 @abstractmethod 

70 def enqueue(self, task: TaskSpec) -> TaskHandle: ... 

71 

72 @abstractmethod 

73 def poll(self, session_id: str) -> list[TaskHandle]: ... 

74 

75 @abstractmethod 

76 def update_status(self, task_id: str, status: TaskStatus, result: Any = None) -> None: ... 

77 

78 @abstractmethod 

79 def get(self, task_id: str) -> TaskHandle | None: ... 

80 

81 @abstractmethod 

82 def cancel(self, task_id: str) -> bool: ... 

83 

84 @abstractmethod 

85 def recover(self) -> list[TaskHandle]: ... 

86 

87 @abstractmethod 

88 def close(self) -> None: ... 

89 

90 

91def _now() -> str: 

92 return datetime.now(timezone.utc).isoformat() 

93 

94 

95class PersistentTaskQueue(TaskQueue): 

96 """SQLite 持久化任务队列(T16)。 

97 

98 表结构 tasks: 

99 task_id PK | session_id | step | chapter_id | idempotency_key 

100 payload JSON | status | result JSON | created_at | updated_at 

101 """ 

102 

103 def __init__(self, db_path: Path | str) -> None: 

104 self._db_path = str(db_path) 

105 self._conn = sqlite3.connect(self._db_path) 

106 self._conn.row_factory = sqlite3.Row 

107 self._init_schema() 

108 

109 # ---------- 生命周期 ---------- 

110 

111 def _init_schema(self) -> None: 

112 self._conn.execute( 

113 """ 

114 CREATE TABLE IF NOT EXISTS tasks ( 

115 task_id TEXT PRIMARY KEY, 

116 session_id TEXT NOT NULL, 

117 step TEXT NOT NULL, 

118 chapter_id TEXT, 

119 idempotency_key TEXT NOT NULL, 

120 payload TEXT NOT NULL, 

121 status TEXT NOT NULL, 

122 result TEXT, 

123 created_at TEXT NOT NULL, 

124 updated_at TEXT NOT NULL 

125 ) 

126 """ 

127 ) 

128 self._conn.execute( 

129 "CREATE INDEX IF NOT EXISTS idx_tasks_session ON tasks(session_id)" 

130 ) 

131 self._conn.commit() 

132 

133 def close(self) -> None: 

134 self._conn.close() 

135 

136 # ---------- TaskQueue 接口 ---------- 

137 

138 def enqueue(self, task: TaskSpec) -> TaskHandle: 

139 # 幂等去重(§5.3):同幂等键已存在 → 返回已有句柄(completed 带缓存结果) 

140 existing = self._find_by_idem(task.session_id, task.step, task.chapter_id, task.idempotency_key) 

141 if existing is not None: 

142 return existing 

143 

144 now = _now() 

145 self._conn.execute( 

146 """ 

147 INSERT INTO tasks (task_id, session_id, step, chapter_id, idempotency_key, 

148 payload, status, result, created_at, updated_at) 

149 VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) 

150 """, 

151 ( 

152 task.task_id, 

153 task.session_id, 

154 task.step, 

155 task.chapter_id, 

156 task.idempotency_key, 

157 json.dumps(task.payload, ensure_ascii=False), 

158 TaskStatus.PENDING.value, 

159 None, 

160 now, 

161 now, 

162 ), 

163 ) 

164 self._conn.commit() 

165 return self._row_to_handle(task.task_id) 

166 

167 def poll(self, session_id: str) -> list[TaskHandle]: 

168 rows = self._conn.execute( 

169 "SELECT * FROM tasks WHERE session_id = ? ORDER BY created_at", 

170 (session_id,), 

171 ).fetchall() 

172 return [self._row_to_handle(row["task_id"], row=row) for row in rows] 

173 

174 def update_status(self, task_id: str, status: TaskStatus, result: Any = None) -> None: 

175 current = self.get(task_id) 

176 if current is None: 

177 raise KeyError(f"任务不存在: {task_id}") 

178 if current.status in _TERMINAL: 

179 raise ValueError(f"终态任务不可再转移: {task_id} ({current.status})") 

180 

181 self._conn.execute( 

182 "UPDATE tasks SET status = ?, result = ?, updated_at = ? WHERE task_id = ?", 

183 ( 

184 status.value, 

185 json.dumps(result, ensure_ascii=False) if result is not None else None, 

186 _now(), 

187 task_id, 

188 ), 

189 ) 

190 self._conn.commit() 

191 

192 def get(self, task_id: str) -> TaskHandle | None: 

193 row = self._conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone() 

194 return self._row_to_handle(task_id, row=row) if row else None 

195 

196 def cancel(self, task_id: str) -> bool: 

197 current = self.get(task_id) 

198 if current is None or current.status in _TERMINAL: 

199 return False 

200 self.update_status(task_id, TaskStatus.CANCELLED) 

201 return True 

202 

203 def recover(self) -> list[TaskHandle]: 

204 """崩溃恢复:running → failed(中断标记);pending 保留;返回全部未完成。""" 

205 rows = self._conn.execute("SELECT * FROM tasks WHERE status = ?", (TaskStatus.RUNNING.value,)).fetchall() 

206 for row in rows: 

207 self._conn.execute( 

208 "UPDATE tasks SET status = ?, updated_at = ? WHERE task_id = ?", 

209 (TaskStatus.FAILED.value, _now(), row["task_id"]), 

210 ) 

211 self._conn.commit() 

212 incomplete = self._conn.execute( 

213 "SELECT * FROM tasks WHERE status IN (?, ?) ORDER BY created_at", 

214 (TaskStatus.PENDING.value, TaskStatus.FAILED.value), 

215 ).fetchall() 

216 return [self._row_to_handle(row["task_id"], row=row) for row in incomplete] 

217 

218 # ---------- 内部 ---------- 

219 

220 def _find_by_idem(self, session_id: str, step: str, chapter_id: str | None, idem: str) -> TaskHandle | None: 

221 row = self._conn.execute( 

222 "SELECT * FROM tasks WHERE session_id = ? AND step = ? AND idempotency_key = ? AND " 

223 "chapter_id IS ?", 

224 (session_id, step, idem, chapter_id), 

225 ).fetchone() 

226 if row is None: 

227 # chapter_id 可为 NULL(SQL 的 IS 处理);此处统一按精确匹配 

228 row = self._conn.execute( 

229 "SELECT * FROM tasks WHERE session_id = ? AND step = ? AND idempotency_key = ?", 

230 (session_id, step, idem), 

231 ).fetchone() 

232 return self._row_to_handle(row["task_id"], row=row) if row else None 

233 

234 def _row_to_handle(self, task_id: str, row: sqlite3.Row | None = None) -> TaskHandle: 

235 if row is None: 

236 row = self._conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone() 

237 if row is None: 

238 raise KeyError(f"任务不存在: {task_id}") 

239 return TaskHandle( 

240 task_id=row["task_id"], 

241 session_id=row["session_id"], 

242 step=row["step"], 

243 chapter_id=row["chapter_id"], 

244 payload=json.loads(row["payload"]), 

245 idempotency_key=row["idempotency_key"], 

246 status=TaskStatus(row["status"]), 

247 result=json.loads(row["result"]) if row["result"] else None, 

248 created_at=row["created_at"], 

249 updated_at=row["updated_at"], 

250 )