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
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +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 dataclasses import dataclass, field
30from enum import Enum
31from typing import Any, Callable
34# ── Enums ────────────────────────────────────────────────────────
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
46 @property
47 def is_terminal(self) -> bool:
48 return self in (BackgroundTaskStatus.COMPLETED, BackgroundTaskStatus.FAILED,
49 BackgroundTaskStatus.CANCELLED, BackgroundTaskStatus.TIMED_OUT)
51 @property
52 def is_active(self) -> bool:
53 return self in (BackgroundTaskStatus.QUEUED, BackgroundTaskStatus.RUNNING)
56# ── Data Models ──────────────────────────────────────────────────
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)
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 }
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", {}))
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 = ""
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 }
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 )
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)
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)
158 def __post_init__(self):
159 if not self.progress.task_id:
160 self.progress.task_id = self.id
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
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 }
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 )
211# ── Callback types ───────────────────────────────────────────────
213ProgressCallback = Callable[[TaskProgress], None]
214CompletionCallback = Callable[[BackgroundTask], None]
217# ── Background Task Manager ──────────────────────────────────────
219class BackgroundTaskManager:
220 """
221 Manages long-running background agent tasks.
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 """
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]] = {}
245 # ── Public API ───────────────────────────────────────────────
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)
266 bt.progress.last_update = time.time()
267 self._tasks[bt.id] = bt
269 if self._store:
270 await self._persist(bt)
272 # Start in background
273 coro = self._run_task(bt, loop_factory, agent_loop)
274 self._running[bt.id] = asyncio.create_task(coro)
276 return bt.id
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
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
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
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
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
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
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]
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)
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)
368 # ── Progress Reporting ───────────────────────────────────────
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
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)
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()
397 prog.current_phase = phase_name
398 if step:
399 prog.current_step = step
400 if message:
401 prog.message = message
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
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
422 await self._update_progress(task_id)
424 # ── Internal ─────────────────────────────────────────────────
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)
441 try:
442 # Timeout enforcement
443 timeout = bt.config.max_duration_seconds
444 start = time.time()
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")
453 # Inject progress callback into the loop
454 original_on_iteration = getattr(loop, 'on_iteration', None)
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)
473 loop.on_iteration = progress_on_iteration
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
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)
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
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
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
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