from __future__ import annotations import os import json import contextlib import hashlib import signal import subprocess import time from dataclasses import dataclass, field from pathlib import Path from typing import Any TERMINAL_STATUSES = {"completed", "failed", "stopped"} @dataclass class ManagedProcess: id: str name: str command: list[str] work_dir: str log_path: Path output_dir: str gpus: list[int] process: subprocess.Popen[Any] | None created_at: float pid: int | None = None status: str = "running" progress: int = 5 artifacts: list[dict[str, Any]] = field(default_factory=list) class ProcessManager: def __init__(self, log_root: str) -> None: self.log_root = Path(log_root) self.log_root.mkdir(parents=True, exist_ok=True) self.registry_path = self.log_root / "compute-jobs.json" self.jobs: dict[str, ManagedProcess] = {} self._load_registry() def create_job(self, payload: dict[str, Any], command: list[str], work_dir: str) -> dict[str, Any]: job_id = str(payload.get("id") or f"job_{int(time.time() * 1000)}") if job_id in self.jobs and self.jobs[job_id].status not in TERMINAL_STATUSES: raise ValueError(f"job {job_id} is already running") output_dir = str(payload.get("output_dir") or f"/data/yg-ft/outputs/{payload.get('name', job_id)}") Path(output_dir).mkdir(parents=True, exist_ok=True) log_path = self.log_root / f"{job_id}.log" env = os.environ.copy() gpus = [int(item) for item in payload.get("gpus") or []] locked = self.locked_gpus() conflict = sorted(set(gpus).intersection(locked)) if conflict: raise ValueError(f"gpu already locked: {conflict}") if gpus: env["CUDA_VISIBLE_DEVICES"] = ",".join(str(item) for item in gpus) env.update({str(k): str(v) for k, v in payload.get("env", {}).items()}) cwd = work_dir if Path(work_dir).exists() else None with log_path.open("ab") as log_file: log_file.write(f"[INFO] starting job_id={job_id} command={' '.join(command)}\n".encode("utf-8")) process = subprocess.Popen( command, cwd=cwd, env=env, stdout=log_file, stderr=subprocess.STDOUT, ) managed = ManagedProcess( id=job_id, name=str(payload.get("name") or job_id), command=command, work_dir=work_dir, log_path=log_path, output_dir=output_dir, gpus=gpus, process=process, created_at=time.time(), pid=process.pid, progress=10, ) self.jobs[job_id] = managed data = self.serialize(managed) self._save_registry() return data def get_job(self, job_id: str) -> dict[str, Any] | None: job = self.jobs.get(job_id) if not job: return None return self.serialize(job) def list_jobs(self) -> list[dict[str, Any]]: return [self.serialize(job) for job in self.jobs.values()] def stop_job(self, job_id: str) -> dict[str, Any] | None: job = self.jobs.get(job_id) if not job: return None if job.status not in TERMINAL_STATUSES: try: if job.process is not None and os.name == "nt": job.process.terminate() elif job.pid is not None: os.kill(job.pid, signal.SIGTERM) if job.process is not None: job.process.wait(timeout=10) except Exception: if job.process is not None: job.process.kill() elif job.pid is not None: with contextlib.suppress(Exception): os.kill(job.pid, signal.SIGKILL) job.status = "stopped" job.progress = min(job.progress, 99) data = self.serialize(job) self._save_registry() return data def logs(self, job_id: str) -> str: job = self.jobs.get(job_id) if not job or not job.log_path.exists(): return "" return job.log_path.read_text(encoding="utf-8", errors="replace") def serialize(self, job: ManagedProcess) -> dict[str, Any]: code = job.process.poll() if job.process is not None else None checkpoints = self._collect_checkpoints(job.output_dir) if job.status not in TERMINAL_STATUSES: if job.process is None and job.pid is not None and not self._pid_alive(job.pid): job.status = "failed" job.progress = min(job.progress, 99) code = -1 elif code is None: job.status = "running" elapsed = max(0, int(time.time() - job.created_at)) job.progress = min(95, max(job.progress, 10 + elapsed // 6)) elif code == 0: job.status = "completed" job.progress = 100 job.artifacts = self._collect_artifacts(job.output_dir) else: job.status = "failed" job.progress = min(job.progress, 99) self._save_registry() return { "id": job.id, "name": job.name, "status": job.status, "progress": job.progress, "pid": job.pid, "gpus": job.gpus, "created_at": job.created_at, "command": job.command, "work_dir": job.work_dir, "output_dir": job.output_dir, "log_file": str(job.log_path), "artifacts": job.artifacts, "checkpoints": checkpoints, "return_code": code, } def locked_gpus(self) -> set[int]: locked: set[int] = set() for job in self.jobs.values(): status = self.serialize(job)["status"] if status in {"queued", "running"}: locked.update(job.gpus) return locked def _collect_artifacts(self, output_dir: str) -> list[dict[str, Any]]: root = Path(output_dir) if not root.exists(): return [] artifacts: list[dict[str, Any]] = [] for path in root.rglob("*"): if path.is_file(): digest = hashlib.sha256() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(1024 * 1024), b""): digest.update(chunk) size = path.stat().st_size artifacts.append( { "path": str(path), "name": path.name, "size": size, "size_bytes": size, "checksum_sha256": digest.hexdigest(), } ) return artifacts[:200] def _collect_checkpoints(self, output_dir: str) -> list[dict[str, Any]]: root = Path(output_dir) if not root.exists(): return [] checkpoints: list[dict[str, Any]] = [] for path in root.glob("checkpoint-*"): if not path.is_dir(): continue step = 0 try: step = int(path.name.rsplit("-", 1)[-1]) except ValueError: step = 0 size_bytes = sum(item.stat().st_size for item in path.rglob("*") if item.is_file()) checkpoints.append( { "step": step, "name": path.name, "path": str(path), "size_bytes": size_bytes, "create_time": path.stat().st_mtime, } ) return sorted(checkpoints, key=lambda item: (int(item.get("step") or 0), str(item.get("name") or ""))) def _save_registry(self) -> None: items = [] for job in self.jobs.values(): items.append( { "id": job.id, "name": job.name, "command": job.command, "work_dir": job.work_dir, "log_path": str(job.log_path), "output_dir": job.output_dir, "gpus": job.gpus, "pid": job.pid, "created_at": job.created_at, "status": job.status, "progress": job.progress, "artifacts": job.artifacts, } ) self.registry_path.write_text(json.dumps(items, ensure_ascii=False, indent=2), encoding="utf-8") def _load_registry(self) -> None: if not self.registry_path.exists(): return try: items = json.loads(self.registry_path.read_text(encoding="utf-8")) except json.JSONDecodeError: return for item in items if isinstance(items, list) else []: if not isinstance(item, dict): continue pid = item.get("pid") status = item.get("status", "failed") if status not in TERMINAL_STATUSES and pid and not self._pid_alive(int(pid)): status = "failed" job = ManagedProcess( id=str(item["id"]), name=str(item.get("name") or item["id"]), command=[str(part) for part in item.get("command") or []], work_dir=str(item.get("work_dir") or ""), log_path=Path(item.get("log_path") or self.log_root / f"{item['id']}.log"), output_dir=str(item.get("output_dir") or ""), gpus=[int(gpu) for gpu in item.get("gpus") or []], process=None, pid=int(pid) if pid else None, created_at=float(item.get("created_at") or time.time()), status=status, progress=int(item.get("progress") or 0), artifacts=item.get("artifacts") or [], ) self.jobs[job.id] = job def _pid_alive(self, pid: int) -> bool: if pid <= 0: return False try: os.kill(pid, 0) return True except OSError: return False