ForcePilot/backend/package/yuxi/channel/runtime/boot.py

62 lines
2.0 KiB
Python
Raw Normal View History

from __future__ import annotations
import asyncio
import logging
from typing import Any, Awaitable, Callable
from .startup import StartupTaskRunner
logger = logging.getLogger(__name__)
StartupTask = Callable[[], Awaitable[Any]]
class BootManager:
"""启动管理器 — 协调系统启动阶段的任务执行顺序。"""
def __init__(self):
self._phases: dict[str, list[tuple[int, StartupTask]]] = {
"early": [],
"config": [],
"plugins": [],
"channels": [],
"services": [],
"late": [],
}
self._runner = StartupTaskRunner()
def register(self, phase: str, task: StartupTask, *, priority: int = 0) -> None:
"""注册一个启动任务到指定阶段。
priority 值越大优先级越高同阶段内优先执行
"""
if phase not in self._phases:
raise ValueError(f"Unknown boot phase: {phase}")
self._phases[phase].append((priority, task))
async def run_startup_tasks(self) -> None:
"""按阶段顺序执行所有启动任务,阶段内按 priority 降序执行。"""
for phase, tasks in self._phases.items():
if not tasks:
continue
ordered = sorted(tasks, key=lambda x: x[0], reverse=True)
logger.info("Boot phase: %s (%d tasks)", phase, len(ordered))
for _, task in ordered:
try:
await self._runner.run(task)
except Exception:
logger.exception("Boot task failed in phase '%s'", phase)
raise
def get_phases(self) -> dict[str, list[StartupTask]]:
return {k: [t for _, t in v] for k, v in self._phases.items()}
def to_dict(self) -> dict[str, Any]:
return {
"phases": {
phase: [getattr(t, "__name__", str(t)) for t in tasks]
for phase, tasks in self.get_phases().items()
},
"total_tasks": sum(len(v) for v in self._phases.values()),
}