from __future__ import annotations import os import json import contextlib 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 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, "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(): artifacts.append( { "path": str(path), "name": path.name, "size": path.stat().st_size, } ) return artifacts[:200] 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