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

328 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 21:26 +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 collections.abc import Callable 

30from dataclasses import dataclass, field 

31from enum import StrEnum 

32from typing import Any 

33 

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

35 

36 

37class BackgroundTaskStatus(StrEnum): 

38 """Background task lifecycle states.""" 

39 

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

41 RUNNING = "running" # Actively executing 

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

43 COMPLETED = "completed" # Finished successfully 

44 FAILED = "failed" # Finished with error 

45 CANCELLED = "cancelled" # Cancelled by user 

46 TIMED_OUT = "timed_out" # Exceeded time budget 

47 

48 @property 

49 def is_terminal(self) -> bool: 

50 return self in ( 

51 BackgroundTaskStatus.COMPLETED, 

52 BackgroundTaskStatus.FAILED, 

53 BackgroundTaskStatus.CANCELLED, 

54 BackgroundTaskStatus.TIMED_OUT, 

55 ) 

56 

57 @property 

58 def is_active(self) -> bool: 

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

60 

61 

62# ── Data Models ────────────────────────────────────────────────── 

63 

64 

65@dataclass 

66class ProgressPhase: 

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

68 

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

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

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

72 started_at: float = 0.0 

73 finished_at: float = 0.0 

74 completed: bool = False 

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

76 

77 def to_dict(self) -> dict: 

78 return { 

79 "name": self.name, 

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

81 "weight": self.weight, 

82 "started_at": self.started_at, 

83 "finished_at": self.finished_at, 

84 "completed": self.completed, 

85 "metadata": self.metadata, 

86 } 

87 

88 @classmethod 

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

90 return cls( 

91 name=d["name"], 

92 label=d.get("label", ""), 

93 weight=d.get("weight", 1.0), 

94 started_at=d.get("started_at", 0.0), 

95 finished_at=d.get("finished_at", 0.0), 

96 completed=d.get("completed", False), 

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

98 ) 

99 

100 

101@dataclass 

102class TaskProgress: 

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

104 

105 task_id: str 

106 status: BackgroundTaskStatus = BackgroundTaskStatus.QUEUED 

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

108 current_phase: str = "" 

109 current_step: int = 0 

110 total_steps: int = 0 

111 percent: float = 0.0 

112 elapsed_seconds: float = 0.0 

113 estimated_remaining_seconds: float = 0.0 

114 last_update: float = 0.0 

115 message: str = "" 

116 

117 def to_dict(self) -> dict: 

118 return { 

119 "task_id": self.task_id, 

120 "status": self.status.value, 

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

122 "current_phase": self.current_phase, 

123 "current_step": self.current_step, 

124 "total_steps": self.total_steps, 

125 "percent": self.percent, 

126 "elapsed_seconds": self.elapsed_seconds, 

127 "estimated_remaining_seconds": self.estimated_remaining_seconds, 

128 "last_update": self.last_update, 

129 "message": self.message, 

130 } 

131 

132 @classmethod 

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

134 return cls( 

135 task_id=d["task_id"], 

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

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

138 current_phase=d.get("current_phase", ""), 

139 current_step=d.get("current_step", 0), 

140 total_steps=d.get("total_steps", 0), 

141 percent=d.get("percent", 0.0), 

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

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

144 last_update=d.get("last_update", 0.0), 

145 message=d.get("message", ""), 

146 ) 

147 

148 

149@dataclass 

150class BackgroundTaskConfig: 

151 """Configuration for a background task.""" 

152 

153 max_duration_seconds: float = 3600.0 # 1 hour default 

154 max_cost_usd: float = 10.0 

155 enable_checkpoints: bool = True 

156 checkpoint_interval: int = 20 # iterations between checkpoints 

157 enable_progress: bool = True 

158 progress_report_interval: float = 5.0 # seconds between progress updates 

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

160 max_retries: int = 2 # retries on transient failure 

161 pause_on_cost_warning: bool = True 

162 notify_on_completion: bool = False 

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

164 

165 

166@dataclass 

167class BackgroundTask: 

168 """Complete background task record.""" 

169 

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

171 name: str = "" 

172 task_description: str = "" 

173 status: BackgroundTaskStatus = BackgroundTaskStatus.QUEUED 

174 config: BackgroundTaskConfig = field(default_factory=BackgroundTaskConfig) 

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

176 result: Any = None 

177 error: str = "" 

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

179 started_at: float = 0.0 

180 finished_at: float = 0.0 

181 cost_usd: float = 0.0 

182 tokens_used: int = 0 

183 checkpoint_id: str = "" 

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

185 

186 def __post_init__(self): 

187 if not self.progress.task_id: 

188 self.progress.task_id = self.id 

189 

190 @property 

191 def duration_seconds(self) -> float: 

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

193 start = self.started_at or self.created_at 

194 return end - start 

195 

196 def to_dict(self) -> dict: 

197 return { 

198 "id": self.id, 

199 "name": self.name, 

200 "task_description": self.task_description, 

201 "status": self.status.value, 

202 "config": { 

203 "max_duration_seconds": self.config.max_duration_seconds, 

204 "max_cost_usd": self.config.max_cost_usd, 

205 "enable_checkpoints": self.config.enable_checkpoints, 

206 "checkpoint_interval": self.config.checkpoint_interval, 

207 "auto_resume": self.config.auto_resume, 

208 "max_retries": self.config.max_retries, 

209 }, 

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

211 "result": self.result, 

212 "error": self.error, 

213 "created_at": self.created_at, 

214 "started_at": self.started_at, 

215 "finished_at": self.finished_at, 

216 "cost_usd": self.cost_usd, 

217 "tokens_used": self.tokens_used, 

218 "checkpoint_id": self.checkpoint_id, 

219 "metadata": self.metadata, 

220 } 

221 

222 @classmethod 

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

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

225 return cls( 

226 id=d["id"], 

227 name=d.get("name", ""), 

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

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

230 config=BackgroundTaskConfig( 

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

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

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

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

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

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

237 ), 

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

239 result=d.get("result"), 

240 error=d.get("error", ""), 

241 created_at=d.get("created_at", 0.0), 

242 started_at=d.get("started_at", 0.0), 

243 finished_at=d.get("finished_at", 0.0), 

244 cost_usd=d.get("cost_usd", 0.0), 

245 tokens_used=d.get("tokens_used", 0), 

246 checkpoint_id=d.get("checkpoint_id", ""), 

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

248 ) 

249 

250 

251# ── Callback types ─────────────────────────────────────────────── 

252 

253ProgressCallback = Callable[[TaskProgress], None] 

254CompletionCallback = Callable[[BackgroundTask], None] 

255 

256 

257# ── Background Task Manager ────────────────────────────────────── 

258 

259 

260class BackgroundTaskManager: 

261 """ 

262 Manages long-running background agent tasks. 

263 

264 Features: 

265 - Async task submission with configurable budgets 

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

267 - Progress tracking with named phases 

268 - Pause/resume/cancel by task ID 

269 - Crash recovery with checkpoint replay 

270 - Concurrent task execution with configurable max workers 

271 """ 

272 

273 def __init__( 

274 self, 

275 max_concurrent: int = 5, 

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

277 ): 

278 self.max_concurrent = max_concurrent 

279 self._store = store 

280 self._semaphore = asyncio.Semaphore(max_concurrent) 

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

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

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

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

285 

286 # ── Public API ─────────────────────────────────────────────── 

287 

288 async def submit( 

289 self, 

290 name: str, 

291 task: str, 

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

293 agent_loop: Any = None, 

294 config: BackgroundTaskConfig | None = None, 

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

296 ) -> str: 

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

298 bt = BackgroundTask( 

299 name=name, 

300 task_description=task, 

301 config=config or BackgroundTaskConfig(), 

302 ) 

303 if phases: 

304 bt.progress.phases = phases 

305 bt.progress.total_steps = len(phases) 

306 

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

308 self._tasks[bt.id] = bt 

309 

310 if self._store: 

311 await self._persist(bt) 

312 

313 # Start in background 

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

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

316 

317 return bt.id 

318 

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

320 """Get full task record.""" 

321 if task_id in self._tasks: 

322 return self._tasks[task_id] 

323 if self._store: 

324 return await self._load(task_id) 

325 return None 

326 

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

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

329 t = await self.get_task(task_id) 

330 return t.progress if t else None 

331 

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

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

334 t = await self.get_task(task_id) 

335 if not t: 

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

337 if t.status.is_active: 

338 # Wait for completion 

339 running_task = self._running.get(task_id) 

340 if running_task and not running_task.done(): 

341 await running_task 

342 t = self._tasks.get(task_id) 

343 if not t: 

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

345 if t.status == BackgroundTaskStatus.FAILED: 

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

347 return t.result 

348 

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

350 """Pause a running task.""" 

351 t = self._tasks.get(task_id) 

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

353 return False 

354 t.status = BackgroundTaskStatus.PAUSED 

355 t.progress.status = BackgroundTaskStatus.PAUSED 

356 await self._update_progress(task_id) 

357 return True 

358 

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

360 """Resume a paused task.""" 

361 t = self._tasks.get(task_id) 

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

363 return False 

364 t.status = BackgroundTaskStatus.RUNNING 

365 t.progress.status = BackgroundTaskStatus.RUNNING 

366 await self._update_progress(task_id) 

367 return True 

368 

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

370 """Cancel a task.""" 

371 t = self._tasks.get(task_id) 

372 if not t: 

373 return False 

374 t.status = BackgroundTaskStatus.CANCELLED 

375 t.progress.status = BackgroundTaskStatus.CANCELLED 

376 t.finished_at = time.time() 

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

378 if running and not running.done(): 

379 running.cancel() 

380 await self._update_progress(task_id) 

381 await self._notify_completion(task_id) 

382 if self._store: 

383 await self._persist(t) 

384 return True 

385 

386 async def list_tasks( 

387 self, 

388 status: BackgroundTaskStatus | None = None, 

389 limit: int = 50, 

390 ) -> list[BackgroundTask]: 

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

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

393 if status: 

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

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

396 

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

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

399 if task_id not in self._progress_callbacks: 

400 self._progress_callbacks[task_id] = [] 

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

402 

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

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

405 if task_id not in self._completion_callbacks: 

406 self._completion_callbacks[task_id] = [] 

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

408 

409 # ── Progress Reporting ─────────────────────────────────────── 

410 

411 async def update_phase( 

412 self, 

413 task_id: str, 

414 phase_name: str, 

415 completed: bool = False, 

416 step: int = 0, 

417 message: str = "", 

418 ): 

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

420 t = self._tasks.get(task_id) 

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

422 return 

423 

424 prog = t.progress 

425 # Find or create phase 

426 phase = None 

427 for p in prog.phases: 

428 if p.name == phase_name: 

429 phase = p 

430 break 

431 if not phase: 

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

433 prog.phases.append(phase) 

434 prog.total_steps = len(prog.phases) 

435 

436 if completed: 

437 phase.completed = True 

438 phase.finished_at = time.time() 

439 elif not phase.started_at: 

440 phase.started_at = time.time() 

441 

442 prog.current_phase = phase_name 

443 if step: 

444 prog.current_step = step 

445 if message: 

446 prog.message = message 

447 

448 # Calculate percent from phase weights 

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

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

451 if prog.current_phase and total_weight > 0: 

452 current_phase_obj = phase 

453 if ( 

454 current_phase_obj 

455 and not current_phase_obj.completed 

456 and current_phase_obj.weight > 0 

457 ): 

458 # Partial credit for current phase 

459 partial = current_phase_obj.weight * min( 

460 step / max(t.config.checkpoint_interval, 1), 1.0 

461 ) 

462 completed_weight += partial 

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

464 elif completed_weight >= total_weight: 

465 prog.percent = 100.0 

466 

467 prog.last_update = time.time() 

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

469 prog.elapsed_seconds = elapsed 

470 if prog.percent > 0: 

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

472 

473 await self._update_progress(task_id) 

474 

475 # ── Internal ───────────────────────────────────────────────── 

476 

477 async def _run_task( 

478 self, 

479 bt: BackgroundTask, 

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

481 agent_loop: Any, 

482 ): 

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

484 async with self._semaphore: 

485 bt.status = BackgroundTaskStatus.RUNNING 

486 bt.progress.status = BackgroundTaskStatus.RUNNING 

487 bt.started_at = time.time() 

488 await self._update_progress(bt.id) 

489 if self._store: 

490 await self._persist(bt) 

491 

492 try: 

493 # Timeout enforcement 

494 timeout = bt.config.max_duration_seconds 

495 start = time.time() 

496 

497 if loop_factory: 

498 loop = loop_factory() 

499 elif agent_loop: 

500 loop = agent_loop 

501 else: 

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

503 

504 # Inject progress callback into the loop 

505 original_on_iteration = getattr(loop, "on_iteration", None) 

506 

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

508 elapsed = time.time() - start 

509 if elapsed > timeout: 

510 raise TimeoutError("Task exceeded max duration") 

511 if bt.status == BackgroundTaskStatus.PAUSED: 

512 # Spin-wait for resume (or timeout) 

513 while bt.status == BackgroundTaskStatus.PAUSED: 

514 await asyncio.sleep(0.5) 

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

516 raise TimeoutError("Task timed out while paused") 

517 if ( 

518 bt.config.enable_checkpoints 

519 and iteration % bt.config.checkpoint_interval == 0 

520 ): 

521 bt.progress.current_step = iteration 

522 await self.update_phase( 

523 bt.id, "execution", step=iteration, message=f"Step {iteration}" 

524 ) 

525 if original_on_iteration: 

526 original_on_iteration(iteration, tool_results) 

527 

528 loop.on_iteration = progress_on_iteration 

529 

530 # Run 

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

532 bt.result = result.output if hasattr(result, "output") else result 

533 bt.cost_usd = getattr(result, "cost_usd", 0.0) 

534 bt.tokens_used = sum(getattr(result, "tokens_used", {}).values()) 

535 bt.status = BackgroundTaskStatus.COMPLETED 

536 bt.progress.status = BackgroundTaskStatus.COMPLETED 

537 bt.progress.percent = 100.0 

538 

539 except TimeoutError: 

540 bt.status = BackgroundTaskStatus.TIMED_OUT 

541 bt.progress.status = BackgroundTaskStatus.TIMED_OUT 

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

543 except asyncio.CancelledError: 

544 bt.status = BackgroundTaskStatus.CANCELLED 

545 bt.progress.status = BackgroundTaskStatus.CANCELLED 

546 except Exception as e: 

547 bt.status = BackgroundTaskStatus.FAILED 

548 bt.progress.status = BackgroundTaskStatus.FAILED 

549 bt.error = str(e) 

550 finally: 

551 bt.finished_at = time.time() 

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

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

554 await self._update_progress(bt.id) 

555 await self._notify_completion(bt.id) 

556 if self._store: 

557 await self._persist(bt) 

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

559 

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

561 """Fire progress callbacks.""" 

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

563 if not callbacks: 

564 return 

565 t = self._tasks.get(task_id) 

566 if t: 

567 for cb in callbacks: 

568 try: 

569 cb(t.progress) 

570 except Exception: 

571 pass 

572 

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

574 """Fire completion callbacks.""" 

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

576 if not callbacks: 

577 return 

578 t = self._tasks.get(task_id) 

579 if t: 

580 for cb in callbacks: 

581 try: 

582 cb(t) 

583 except Exception: 

584 pass 

585 

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

587 """Persist task to store.""" 

588 if not self._store: 

589 return 

590 try: 

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

592 if hasattr(self._store, "save"): 

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

594 except Exception: 

595 pass 

596 

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

598 """Load task from store.""" 

599 if not self._store: 

600 return None 

601 try: 

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

603 if snap and "data" in snap: 

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

605 except Exception: 

606 pass 

607 return None