Coverage for agentos/background/task_manager.py: 32%
328 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 01:44 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 01:44 +0800
1"""
2Background Task Manager — v1.11.0
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
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"""
23from __future__ import annotations
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
34# ── Enums ────────────────────────────────────────────────────────
37class BackgroundTaskStatus(StrEnum):
38 """Background task lifecycle states."""
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
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 )
57 @property
58 def is_active(self) -> bool:
59 return self in (BackgroundTaskStatus.QUEUED, BackgroundTaskStatus.RUNNING)
62# ── Data Models ──────────────────────────────────────────────────
65@dataclass
66class ProgressPhase:
67 """A named phase within task execution."""
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)
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 }
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 )
101@dataclass
102class TaskProgress:
103 """Structured progress report for a background task."""
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 = ""
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 }
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 )
149@dataclass
150class BackgroundTaskConfig:
151 """Configuration for a background task."""
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)
166@dataclass
167class BackgroundTask:
168 """Complete background task record."""
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)
186 def __post_init__(self):
187 if not self.progress.task_id:
188 self.progress.task_id = self.id
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
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 }
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 )
251# ── Callback types ───────────────────────────────────────────────
253ProgressCallback = Callable[[TaskProgress], None]
254CompletionCallback = Callable[[BackgroundTask], None]
257# ── Background Task Manager ──────────────────────────────────────
260class BackgroundTaskManager:
261 """
262 Manages long-running background agent tasks.
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 """
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]] = {}
286 # ── Public API ───────────────────────────────────────────────
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)
307 bt.progress.last_update = time.time()
308 self._tasks[bt.id] = bt
310 if self._store:
311 await self._persist(bt)
313 # Start in background
314 coro = self._run_task(bt, loop_factory, agent_loop)
315 self._running[bt.id] = asyncio.create_task(coro)
317 return bt.id
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
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
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
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
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
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
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]
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)
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)
409 # ── Progress Reporting ───────────────────────────────────────
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
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)
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()
442 prog.current_phase = phase_name
443 if step:
444 prog.current_step = step
445 if message:
446 prog.message = message
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
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
473 await self._update_progress(task_id)
475 # ── Internal ─────────────────────────────────────────────────
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)
492 try:
493 # Timeout enforcement
494 timeout = bt.config.max_duration_seconds
495 start = time.time()
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")
504 # Inject progress callback into the loop
505 original_on_iteration = getattr(loop, "on_iteration", None)
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)
528 loop.on_iteration = progress_on_iteration
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
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)
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
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
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
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