Coverage for agentos/background/task_manager.py: 32%

327 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 10:59 +0800

1""" 

2Background Task Manager — v1.11.0 

3 

4Production-grade long-running task execution with: 

5- Submit task → get task_id → poll progress → retrieve result 

6- Persistent task state (SQLite/Postgres) 

7- Progress milestones with phase tracking 

8- Graceful pause/resume/cancel 

9- Timeout and resource budget enforcement 

10- Crash recovery via full checkpoint integration 

11 

12Usage: 

13 mgr = BackgroundTaskManager(loop_factory=my_loop, store=SqliteStore("tasks.db")) 

14 task_id = await mgr.submit("Analyze 10GB dataset", task="...", config=...) 

15 while True: 

16 progress = await mgr.get_progress(task_id) 

17 print(f"{progress.current_phase}: {progress.percent:.0f}%") 

18 if progress.status.is_terminal: 

19 break 

20 result = await mgr.get_result(task_id) 

21""" 

22 

23from __future__ import annotations 

24 

25import asyncio 

26import json 

27import time 

28import uuid 

29from dataclasses import dataclass, field 

30from enum import Enum 

31from typing import Any, Callable 

32 

33 

34# ── Enums ──────────────────────────────────────────────────────── 

35 

36class BackgroundTaskStatus(str, Enum): 

37 """Background task lifecycle states.""" 

38 QUEUED = "queued" # Accepted, waiting to start 

39 RUNNING = "running" # Actively executing 

40 PAUSED = "paused" # Paused by user or system 

41 COMPLETED = "completed" # Finished successfully 

42 FAILED = "failed" # Finished with error 

43 CANCELLED = "cancelled" # Cancelled by user 

44 TIMED_OUT = "timed_out" # Exceeded time budget 

45 

46 @property 

47 def is_terminal(self) -> bool: 

48 return self in (BackgroundTaskStatus.COMPLETED, BackgroundTaskStatus.FAILED, 

49 BackgroundTaskStatus.CANCELLED, BackgroundTaskStatus.TIMED_OUT) 

50 

51 @property 

52 def is_active(self) -> bool: 

53 return self in (BackgroundTaskStatus.QUEUED, BackgroundTaskStatus.RUNNING) 

54 

55 

56# ── Data Models ────────────────────────────────────────────────── 

57 

58@dataclass 

59class ProgressPhase: 

60 """A named phase within task execution.""" 

61 name: str # e.g. "data_loading", "analysis", "reporting" 

62 label: str = "" # Human-readable label 

63 weight: float = 1.0 # Relative weight for percent calculation 

64 started_at: float = 0.0 

65 finished_at: float = 0.0 

66 completed: bool = False 

67 metadata: dict[str, Any] = field(default_factory=dict) 

68 

69 def to_dict(self) -> dict: 

70 return { 

71 "name": self.name, "label": self.label or self.name, 

72 "weight": self.weight, "started_at": self.started_at, 

73 "finished_at": self.finished_at, "completed": self.completed, 

74 "metadata": self.metadata, 

75 } 

76 

77 @classmethod 

78 def from_dict(cls, d: dict) -> "ProgressPhase": 

79 return cls(name=d["name"], label=d.get("label", ""), weight=d.get("weight", 1.0), 

80 started_at=d.get("started_at", 0.0), finished_at=d.get("finished_at", 0.0), 

81 completed=d.get("completed", False), metadata=d.get("metadata", {})) 

82 

83 

84@dataclass 

85class TaskProgress: 

86 """Structured progress report for a background task.""" 

87 task_id: str 

88 status: BackgroundTaskStatus = BackgroundTaskStatus.QUEUED 

89 phases: list[ProgressPhase] = field(default_factory=list) 

90 current_phase: str = "" 

91 current_step: int = 0 

92 total_steps: int = 0 

93 percent: float = 0.0 

94 elapsed_seconds: float = 0.0 

95 estimated_remaining_seconds: float = 0.0 

96 last_update: float = 0.0 

97 message: str = "" 

98 

99 def to_dict(self) -> dict: 

100 return { 

101 "task_id": self.task_id, "status": self.status.value, 

102 "phases": [p.to_dict() for p in self.phases], 

103 "current_phase": self.current_phase, 

104 "current_step": self.current_step, "total_steps": self.total_steps, 

105 "percent": self.percent, "elapsed_seconds": self.elapsed_seconds, 

106 "estimated_remaining_seconds": self.estimated_remaining_seconds, 

107 "last_update": self.last_update, "message": self.message, 

108 } 

109 

110 @classmethod 

111 def from_dict(cls, d: dict) -> "TaskProgress": 

112 return cls( 

113 task_id=d["task_id"], status=BackgroundTaskStatus(d.get("status", "queued")), 

114 phases=[ProgressPhase.from_dict(p) for p in d.get("phases", [])], 

115 current_phase=d.get("current_phase", ""), current_step=d.get("current_step", 0), 

116 total_steps=d.get("total_steps", 0), percent=d.get("percent", 0.0), 

117 elapsed_seconds=d.get("elapsed_seconds", 0.0), 

118 estimated_remaining_seconds=d.get("estimated_remaining_seconds", 0.0), 

119 last_update=d.get("last_update", 0.0), message=d.get("message", ""), 

120 ) 

121 

122 

123@dataclass 

124class BackgroundTaskConfig: 

125 """Configuration for a background task.""" 

126 max_duration_seconds: float = 3600.0 # 1 hour default 

127 max_cost_usd: float = 10.0 

128 enable_checkpoints: bool = True 

129 checkpoint_interval: int = 20 # iterations between checkpoints 

130 enable_progress: bool = True 

131 progress_report_interval: float = 5.0 # seconds between progress updates 

132 auto_resume: bool = True # auto-resume from checkpoint on restart 

133 max_retries: int = 2 # retries on transient failure 

134 pause_on_cost_warning: bool = True 

135 notify_on_completion: bool = False 

136 metadata: dict[str, Any] = field(default_factory=dict) 

137 

138 

139@dataclass 

140class BackgroundTask: 

141 """Complete background task record.""" 

142 id: str = field(default_factory=lambda: uuid.uuid4().hex[:16]) 

143 name: str = "" 

144 task_description: str = "" 

145 status: BackgroundTaskStatus = BackgroundTaskStatus.QUEUED 

146 config: BackgroundTaskConfig = field(default_factory=BackgroundTaskConfig) 

147 progress: TaskProgress = field(default_factory=lambda: TaskProgress(task_id="")) 

148 result: Any = None 

149 error: str = "" 

150 created_at: float = field(default_factory=time.time) 

151 started_at: float = 0.0 

152 finished_at: float = 0.0 

153 cost_usd: float = 0.0 

154 tokens_used: int = 0 

155 checkpoint_id: str = "" 

156 metadata: dict[str, Any] = field(default_factory=dict) 

157 

158 def __post_init__(self): 

159 if not self.progress.task_id: 

160 self.progress.task_id = self.id 

161 

162 @property 

163 def duration_seconds(self) -> float: 

164 end = self.finished_at or time.time() 

165 start = self.started_at or self.created_at 

166 return end - start 

167 

168 def to_dict(self) -> dict: 

169 return { 

170 "id": self.id, "name": self.name, "task_description": self.task_description, 

171 "status": self.status.value, "config": { 

172 "max_duration_seconds": self.config.max_duration_seconds, 

173 "max_cost_usd": self.config.max_cost_usd, 

174 "enable_checkpoints": self.config.enable_checkpoints, 

175 "checkpoint_interval": self.config.checkpoint_interval, 

176 "auto_resume": self.config.auto_resume, 

177 "max_retries": self.config.max_retries, 

178 }, 

179 "progress": self.progress.to_dict(), 

180 "result": self.result, "error": self.error, 

181 "created_at": self.created_at, "started_at": self.started_at, 

182 "finished_at": self.finished_at, "cost_usd": self.cost_usd, 

183 "tokens_used": self.tokens_used, "checkpoint_id": self.checkpoint_id, 

184 "metadata": self.metadata, 

185 } 

186 

187 @classmethod 

188 def from_dict(cls, d: dict) -> "BackgroundTask": 

189 cfg_d = d.get("config", {}) 

190 return cls( 

191 id=d["id"], name=d.get("name", ""), 

192 task_description=d.get("task_description", ""), 

193 status=BackgroundTaskStatus(d.get("status", "queued")), 

194 config=BackgroundTaskConfig( 

195 max_duration_seconds=cfg_d.get("max_duration_seconds", 3600.0), 

196 max_cost_usd=cfg_d.get("max_cost_usd", 10.0), 

197 enable_checkpoints=cfg_d.get("enable_checkpoints", True), 

198 checkpoint_interval=cfg_d.get("checkpoint_interval", 20), 

199 auto_resume=cfg_d.get("auto_resume", True), 

200 max_retries=cfg_d.get("max_retries", 2), 

201 ), 

202 progress=TaskProgress.from_dict(d.get("progress", {"task_id": d["id"]})), 

203 result=d.get("result"), error=d.get("error", ""), 

204 created_at=d.get("created_at", 0.0), started_at=d.get("started_at", 0.0), 

205 finished_at=d.get("finished_at", 0.0), cost_usd=d.get("cost_usd", 0.0), 

206 tokens_used=d.get("tokens_used", 0), checkpoint_id=d.get("checkpoint_id", ""), 

207 metadata=d.get("metadata", {}), 

208 ) 

209 

210 

211# ── Callback types ─────────────────────────────────────────────── 

212 

213ProgressCallback = Callable[[TaskProgress], None] 

214CompletionCallback = Callable[[BackgroundTask], None] 

215 

216 

217# ── Background Task Manager ────────────────────────────────────── 

218 

219class BackgroundTaskManager: 

220 """ 

221 Manages long-running background agent tasks. 

222 

223 Features: 

224 - Async task submission with configurable budgets 

225 - Persistent task state (in-memory + optional DB store) 

226 - Progress tracking with named phases 

227 - Pause/resume/cancel by task ID 

228 - Crash recovery with checkpoint replay 

229 - Concurrent task execution with configurable max workers 

230 """ 

231 

232 def __init__( 

233 self, 

234 max_concurrent: int = 5, 

235 store: Any = None, # Optional CheckpointStore-like persistence 

236 ): 

237 self.max_concurrent = max_concurrent 

238 self._store = store 

239 self._semaphore = asyncio.Semaphore(max_concurrent) 

240 self._tasks: dict[str, BackgroundTask] = {} 

241 self._running: dict[str, asyncio.Task] = {} 

242 self._progress_callbacks: dict[str, list[ProgressCallback]] = {} 

243 self._completion_callbacks: dict[str, list[CompletionCallback]] = {} 

244 

245 # ── Public API ─────────────────────────────────────────────── 

246 

247 async def submit( 

248 self, 

249 name: str, 

250 task: str, 

251 loop_factory: Callable[[], Any] | None = None, 

252 agent_loop: Any = None, 

253 config: BackgroundTaskConfig | None = None, 

254 phases: list[ProgressPhase] | None = None, 

255 ) -> str: 

256 """Submit a task for background execution. Returns task_id.""" 

257 bt = BackgroundTask( 

258 name=name, 

259 task_description=task, 

260 config=config or BackgroundTaskConfig(), 

261 ) 

262 if phases: 

263 bt.progress.phases = phases 

264 bt.progress.total_steps = len(phases) 

265 

266 bt.progress.last_update = time.time() 

267 self._tasks[bt.id] = bt 

268 

269 if self._store: 

270 await self._persist(bt) 

271 

272 # Start in background 

273 coro = self._run_task(bt, loop_factory, agent_loop) 

274 self._running[bt.id] = asyncio.create_task(coro) 

275 

276 return bt.id 

277 

278 async def get_task(self, task_id: str) -> BackgroundTask | None: 

279 """Get full task record.""" 

280 if task_id in self._tasks: 

281 return self._tasks[task_id] 

282 if self._store: 

283 return await self._load(task_id) 

284 return None 

285 

286 async def get_progress(self, task_id: str) -> TaskProgress | None: 

287 """Get current progress for a task.""" 

288 t = await self.get_task(task_id) 

289 return t.progress if t else None 

290 

291 async def get_result(self, task_id: str) -> Any: 

292 """Get task result (blocks if still running).""" 

293 t = await self.get_task(task_id) 

294 if not t: 

295 raise KeyError(f"Task {task_id} not found") 

296 if t.status.is_active: 

297 # Wait for completion 

298 running_task = self._running.get(task_id) 

299 if running_task and not running_task.done(): 

300 await running_task 

301 t = self._tasks.get(task_id) 

302 if not t: 

303 raise KeyError(f"Task {task_id} vanished") 

304 if t.status == BackgroundTaskStatus.FAILED: 

305 raise RuntimeError(f"Task {task_id} failed: {t.error}") 

306 return t.result 

307 

308 async def pause(self, task_id: str) -> bool: 

309 """Pause a running task.""" 

310 t = self._tasks.get(task_id) 

311 if not t or not t.status.is_active: 

312 return False 

313 t.status = BackgroundTaskStatus.PAUSED 

314 t.progress.status = BackgroundTaskStatus.PAUSED 

315 await self._update_progress(task_id) 

316 return True 

317 

318 async def resume(self, task_id: str) -> bool: 

319 """Resume a paused task.""" 

320 t = self._tasks.get(task_id) 

321 if not t or t.status != BackgroundTaskStatus.PAUSED: 

322 return False 

323 t.status = BackgroundTaskStatus.RUNNING 

324 t.progress.status = BackgroundTaskStatus.RUNNING 

325 await self._update_progress(task_id) 

326 return True 

327 

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

329 """Cancel a task.""" 

330 t = self._tasks.get(task_id) 

331 if not t: 

332 return False 

333 t.status = BackgroundTaskStatus.CANCELLED 

334 t.progress.status = BackgroundTaskStatus.CANCELLED 

335 t.finished_at = time.time() 

336 running = self._running.pop(task_id, None) 

337 if running and not running.done(): 

338 running.cancel() 

339 await self._update_progress(task_id) 

340 await self._notify_completion(task_id) 

341 if self._store: 

342 await self._persist(t) 

343 return True 

344 

345 async def list_tasks( 

346 self, 

347 status: BackgroundTaskStatus | None = None, 

348 limit: int = 50, 

349 ) -> list[BackgroundTask]: 

350 """List tasks, optionally filtered by status.""" 

351 tasks = list(self._tasks.values()) 

352 if status: 

353 tasks = [t for t in tasks if t.status == status] 

354 return sorted(tasks, key=lambda t: t.created_at, reverse=True)[:limit] 

355 

356 def on_progress(self, task_id: str, callback: ProgressCallback): 

357 """Register a progress callback for a task.""" 

358 if task_id not in self._progress_callbacks: 

359 self._progress_callbacks[task_id] = [] 

360 self._progress_callbacks[task_id].append(callback) 

361 

362 def on_completion(self, task_id: str, callback: CompletionCallback): 

363 """Register a completion callback for a task.""" 

364 if task_id not in self._completion_callbacks: 

365 self._completion_callbacks[task_id] = [] 

366 self._completion_callbacks[task_id].append(callback) 

367 

368 # ── Progress Reporting ─────────────────────────────────────── 

369 

370 async def update_phase( 

371 self, task_id: str, phase_name: str, 

372 completed: bool = False, step: int = 0, message: str = "", 

373 ): 

374 """Update a named phase in the task progress.""" 

375 t = self._tasks.get(task_id) 

376 if not t or not t.config.enable_progress: 

377 return 

378 

379 prog = t.progress 

380 # Find or create phase 

381 phase = None 

382 for p in prog.phases: 

383 if p.name == phase_name: 

384 phase = p 

385 break 

386 if not phase: 

387 phase = ProgressPhase(name=phase_name, label=phase_name) 

388 prog.phases.append(phase) 

389 prog.total_steps = len(prog.phases) 

390 

391 if completed: 

392 phase.completed = True 

393 phase.finished_at = time.time() 

394 elif not phase.started_at: 

395 phase.started_at = time.time() 

396 

397 prog.current_phase = phase_name 

398 if step: 

399 prog.current_step = step 

400 if message: 

401 prog.message = message 

402 

403 # Calculate percent from phase weights 

404 total_weight = sum(p.weight for p in prog.phases) 

405 completed_weight = sum(p.weight for p in prog.phases if p.completed) 

406 if prog.current_phase and total_weight > 0: 

407 current_phase_obj = phase 

408 if current_phase_obj and not current_phase_obj.completed and current_phase_obj.weight > 0: 

409 # Partial credit for current phase 

410 partial = current_phase_obj.weight * min(step / max(t.config.checkpoint_interval, 1), 1.0) 

411 completed_weight += partial 

412 prog.percent = min(completed_weight / total_weight * 100, 99.9) 

413 elif completed_weight >= total_weight: 

414 prog.percent = 100.0 

415 

416 prog.last_update = time.time() 

417 elapsed = prog.last_update - (t.started_at or t.created_at) 

418 prog.elapsed_seconds = elapsed 

419 if prog.percent > 0: 

420 prog.estimated_remaining_seconds = elapsed / (prog.percent / 100) - elapsed 

421 

422 await self._update_progress(task_id) 

423 

424 # ── Internal ───────────────────────────────────────────────── 

425 

426 async def _run_task( 

427 self, 

428 bt: BackgroundTask, 

429 loop_factory: Callable[[], Any] | None, 

430 agent_loop: Any, 

431 ): 

432 """Execute a background task with full lifecycle management.""" 

433 async with self._semaphore: 

434 bt.status = BackgroundTaskStatus.RUNNING 

435 bt.progress.status = BackgroundTaskStatus.RUNNING 

436 bt.started_at = time.time() 

437 await self._update_progress(bt.id) 

438 if self._store: 

439 await self._persist(bt) 

440 

441 try: 

442 # Timeout enforcement 

443 timeout = bt.config.max_duration_seconds 

444 start = time.time() 

445 

446 if loop_factory: 

447 loop = loop_factory() 

448 elif agent_loop: 

449 loop = agent_loop 

450 else: 

451 raise ValueError("Must provide loop_factory or agent_loop") 

452 

453 # Inject progress callback into the loop 

454 original_on_iteration = getattr(loop, 'on_iteration', None) 

455 

456 async def progress_on_iteration(iteration: int, tool_results: list): 

457 elapsed = time.time() - start 

458 if elapsed > timeout: 

459 raise asyncio.TimeoutError("Task exceeded max duration") 

460 if bt.status == BackgroundTaskStatus.PAUSED: 

461 # Spin-wait for resume (or timeout) 

462 while bt.status == BackgroundTaskStatus.PAUSED: 

463 await asyncio.sleep(0.5) 

464 if time.time() - start > timeout: 

465 raise asyncio.TimeoutError("Task timed out while paused") 

466 if bt.config.enable_checkpoints and iteration % bt.config.checkpoint_interval == 0: 

467 bt.progress.current_step = iteration 

468 await self.update_phase(bt.id, "execution", step=iteration, 

469 message=f"Step {iteration}") 

470 if original_on_iteration: 

471 original_on_iteration(iteration, tool_results) 

472 

473 loop.on_iteration = progress_on_iteration 

474 

475 # Run 

476 result = await loop.run(bt.task_description, session_id=bt.id) 

477 bt.result = result.output if hasattr(result, 'output') else result 

478 bt.cost_usd = getattr(result, 'cost_usd', 0.0) 

479 bt.tokens_used = sum(getattr(result, 'tokens_used', {}).values()) 

480 bt.status = BackgroundTaskStatus.COMPLETED 

481 bt.progress.status = BackgroundTaskStatus.COMPLETED 

482 bt.progress.percent = 100.0 

483 

484 except asyncio.TimeoutError: 

485 bt.status = BackgroundTaskStatus.TIMED_OUT 

486 bt.progress.status = BackgroundTaskStatus.TIMED_OUT 

487 bt.error = f"Exceeded max duration of {bt.config.max_duration_seconds}s" 

488 except asyncio.CancelledError: 

489 bt.status = BackgroundTaskStatus.CANCELLED 

490 bt.progress.status = BackgroundTaskStatus.CANCELLED 

491 except Exception as e: 

492 bt.status = BackgroundTaskStatus.FAILED 

493 bt.progress.status = BackgroundTaskStatus.FAILED 

494 bt.error = str(e) 

495 finally: 

496 bt.finished_at = time.time() 

497 bt.progress.last_update = time.time() 

498 bt.progress.elapsed_seconds = bt.finished_at - bt.started_at 

499 await self._update_progress(bt.id) 

500 await self._notify_completion(bt.id) 

501 if self._store: 

502 await self._persist(bt) 

503 self._running.pop(bt.id, None) 

504 

505 async def _update_progress(self, task_id: str): 

506 """Fire progress callbacks.""" 

507 callbacks = self._progress_callbacks.get(task_id, []) 

508 if not callbacks: 

509 return 

510 t = self._tasks.get(task_id) 

511 if t: 

512 for cb in callbacks: 

513 try: 

514 cb(t.progress) 

515 except Exception: 

516 pass 

517 

518 async def _notify_completion(self, task_id: str): 

519 """Fire completion callbacks.""" 

520 callbacks = self._completion_callbacks.get(task_id, []) 

521 if not callbacks: 

522 return 

523 t = self._tasks.get(task_id) 

524 if t: 

525 for cb in callbacks: 

526 try: 

527 cb(t) 

528 except Exception: 

529 pass 

530 

531 async def _persist(self, bt: BackgroundTask): 

532 """Persist task to store.""" 

533 if not self._store: 

534 return 

535 try: 

536 data = json.dumps(bt.to_dict()) 

537 if hasattr(self._store, 'save'): 

538 await self._store.save(f"bg_task:{bt.id}", {"data": data}) 

539 except Exception: 

540 pass 

541 

542 async def _load(self, task_id: str) -> BackgroundTask | None: 

543 """Load task from store.""" 

544 if not self._store: 

545 return None 

546 try: 

547 snap = await self._store.load(f"bg_task:{task_id}") 

548 if snap and "data" in snap: 

549 return BackgroundTask.from_dict(json.loads(snap["data"])) 

550 except Exception: 

551 pass 

552 return None