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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 14:20 +0800
1"""任务级持久化(T16,OV7)。
3背景:api-design §5 定义 TaskQueue 抽象(v1 仅 InMemoryQueue),但崩溃恢复
4只到会话级(agent-runtime §3.6),任务层数据丢失。OV7 裁定任务级持久化。
6实现:PersistentTaskQueue —— TaskQueue 抽象 + SQLite 落盘。
7 - 任务状态 / payload / result 全部写入 SQLite(零外部依赖,标准库 sqlite3)
8 - 幂等去重(§5.3):同 (session_id, step, chapter_id) 已完成 → 返回缓存结果
9 - recover():重启后 running → failed(中断标记),pending 保留待执行
10"""
12from __future__ import annotations
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
23from pydantic import BaseModel, Field
26class TaskStatus(str, Enum):
27 PENDING = "pending"
28 RUNNING = "running"
29 COMPLETED = "completed"
30 FAILED = "failed"
31 CANCELLED = "cancelled"
34# 终态:不可再转移(cancel 仅对非终态生效)
35_TERMINAL = frozenset({TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED})
38class TaskSpec(BaseModel):
39 """任务投递规格(api-design §5.1)。"""
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
49@dataclass
50class TaskHandle:
51 """任务句柄(含状态与结果)。"""
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 = ""
66class TaskQueue(ABC):
67 """统一任务队列抽象(v1 仅 PersistentTaskQueue;Redis/Valkey 为 v2 预留)。"""
69 @abstractmethod
70 def enqueue(self, task: TaskSpec) -> TaskHandle: ...
72 @abstractmethod
73 def poll(self, session_id: str) -> list[TaskHandle]: ...
75 @abstractmethod
76 def update_status(self, task_id: str, status: TaskStatus, result: Any = None) -> None: ...
78 @abstractmethod
79 def get(self, task_id: str) -> TaskHandle | None: ...
81 @abstractmethod
82 def cancel(self, task_id: str) -> bool: ...
84 @abstractmethod
85 def recover(self) -> list[TaskHandle]: ...
87 @abstractmethod
88 def close(self) -> None: ...
91def _now() -> str:
92 return datetime.now(timezone.utc).isoformat()
95class PersistentTaskQueue(TaskQueue):
96 """SQLite 持久化任务队列(T16)。
98 表结构 tasks:
99 task_id PK | session_id | step | chapter_id | idempotency_key
100 payload JSON | status | result JSON | created_at | updated_at
101 """
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()
109 # ---------- 生命周期 ----------
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()
133 def close(self) -> None:
134 self._conn.close()
136 # ---------- TaskQueue 接口 ----------
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
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)
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]
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})")
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()
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
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
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]
218 # ---------- 内部 ----------
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
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 )