"""Data layer for the embedding-fleet control plane. One poll tick = bucket reads (run manifest + worker heartbeats) + Jobs API reads (stage, durations, one metrics sample per running job). All aggregation to the run level happens here; app.py only renders. Cost figures are client-side estimates (flavor unit price x running time), NOT billing — always presented as "~$". """ from __future__ import annotations import json import logging import tempfile import time from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from pathlib import Path from huggingface_hub import ( download_bucket_files, fetch_job_metrics, list_bucket_tree, list_jobs, list_jobs_hardware, ) TERMINAL_OK = {"COMPLETED"} TERMINAL_BAD = {"ERROR", "CANCELED", "DELETED"} def _stage_name(job) -> str: stage = job.status.stage if job.status else None return getattr(stage, "value", None) or str(stage or "UNKNOWN") @dataclass class WorkerRow: rank: int job_id: str | None = None stage: str = "UNKNOWN" rows_done: int = 0 rows_total: int | None = None rows_per_sec: float = 0.0 tokens_done_est: int | None = None gpu_util: float | None = None cost_usd: float | None = None state: str | None = None # worker-reported: running/done/error @dataclass class RunView: run_id: str manifest: dict workers: list[WorkerRow] = field(default_factory=list) # run-level aggregates rows_done: int = 0 rows_total: int = 0 tokens_done_est: int = 0 cost_usd: float = 0.0 cost_ceiling_usd: float | None = None eta_secs: float | None = None gpu_util: float | None = None healthy: int = 0 errored: int = 0 done: int = 0 num_shards: int = 0 _PRICING: dict | None = None def pricing() -> dict: global _PRICING if _PRICING is None: _PRICING = {hw.name: hw for hw in list_jobs_hardware()} return _PRICING def parse_timeout_secs(timeout) -> float | None: """'20m' / '1h' / '90s' / plain seconds -> seconds.""" if timeout is None: return None s = str(timeout).strip().lower() try: mult = {"s": 1, "m": 60, "h": 3600, "d": 86400}.get(s[-1]) return float(s[:-1]) * mult if mult else float(s) except (ValueError, IndexError): return None def list_runs(bucket: str) -> list[str]: """Run ids under runs/, newest first (ids are timestamp-prefixed).""" try: ids = [Path(e.path.rstrip("/")).name for e in list_bucket_tree(bucket, prefix="runs/", recursive=False) if e.__class__.__name__ == "BucketFolder"] # Timestamp-prefixed ids first (newest first), ad-hoc ids after. return sorted(set(ids), key=lambda r: (r[:8].isdigit(), r), reverse=True) if ids else [] except Exception: return [] def _read_bucket_json(bucket: str, paths: list[str]) -> dict[str, dict]: """Fetch small JSON files from the bucket; missing files are skipped.""" out: dict[str, dict] = {} if not paths: return out with tempfile.TemporaryDirectory() as td: pairs = [(p, Path(td) / p.replace("/", "__")) for p in paths] try: download_bucket_files(bucket, pairs, raise_on_missing_files=False) except Exception: return out for src, dst in pairs: if dst.exists(): try: out[src] = json.loads(dst.read_text()) except (json.JSONDecodeError, OSError): pass return out def _sample_gpu_util(job_id: str, timeout: float = 3.0) -> float | None: """One metrics sample -> mean GPU utilization. Never blocks past `timeout`.""" def _one(): gen = iter(fetch_job_metrics(job_id=job_id)) try: raw = next(gen) finally: getattr(gen, "close", lambda: None)() gpus = raw.get("gpus") or {} utils = [float(g.get("utilization") or 0) for g in gpus.values()] return sum(utils) / len(utils) if utils else None with ThreadPoolExecutor(max_workers=1) as pool: fut = pool.submit(_one) try: return fut.result(timeout=timeout) except Exception: return None def _accrued_cost(job, hw_pricing: dict) -> float | None: flavor = getattr(job.flavor, "value", None) or (str(job.flavor) if job.flavor else None) hw = hw_pricing.get(flavor) if not hw: return None secs = job.durations.running_secs if job.durations else None if not secs and job.started_at and _stage_name(job) == "RUNNING": secs = time.time() - job.started_at.timestamp() if not secs: return None return secs / 60.0 * hw.unit_cost_usd def load_run(bucket: str, run_id: str, namespace: str | None = None) -> RunView | None: """One full poll tick: manifest + heartbeats + job stages + metrics samples -> RunView.""" manifest = _read_bucket_json(bucket, [f"runs/{run_id}/run.json"]).get(f"runs/{run_id}/run.json") if not manifest: return None n = manifest["num_shards"] view = RunView(run_id=run_id, manifest=manifest, num_shards=n, rows_total=manifest.get("rows_total") or 0) status_paths = [f"runs/{run_id}/status/{i:05d}.json" for i in range(n)] statuses = _read_bucket_json(bucket, status_paths) # Jobs by label (server-side filter); fall back to manifest job_ids via list comprehension. jobs_by_id = {} try: for j in list_jobs(labels={"embedding-fleet-run": run_id}, namespace=namespace): jobs_by_id[j.id] = j except Exception as e: logging.getLogger("control-plane").warning(f"list_jobs failed: {e!r}") manifest_job_ids = manifest.get("job_ids") or [] hw_pricing = pricing() workers: list[WorkerRow] = [] running_job_ids: list[str] = [] for rank in range(n): row = WorkerRow(rank=rank) st = statuses.get(f"runs/{run_id}/status/{rank:05d}.json") if st: row.state = st.get("state") row.rows_done = st.get("rows_done") or 0 row.rows_total = st.get("rows_total") row.rows_per_sec = st.get("rows_per_sec") or 0.0 row.tokens_done_est = st.get("tokens_done_est") row.job_id = st.get("job_id") if row.job_id is None and rank < len(manifest_job_ids): row.job_id = manifest_job_ids[rank] job = jobs_by_id.get(row.job_id) if job is None and str(rank) in {j.labels.get("rank") for j in jobs_by_id.values() if j.labels}: job = next(j for j in jobs_by_id.values() if (j.labels or {}).get("rank") == str(rank)) if job is not None: row.stage = _stage_name(job) row.cost_usd = _accrued_cost(job, hw_pricing) if row.stage == "RUNNING": running_job_ids.append(row.job_id) workers.append(row) # One GPU sample per running job, in parallel, bounded. if running_job_ids: with ThreadPoolExecutor(max_workers=min(8, len(running_job_ids))) as pool: samples = dict(zip(running_job_ids, pool.map(_sample_gpu_util, running_job_ids))) for row in workers: if row.job_id in samples: row.gpu_util = samples[row.job_id] # Consolidator cost (labeled role=consolidate) counts toward the run. consolidator_cost = sum( _accrued_cost(j, hw_pricing) or 0.0 for j in jobs_by_id.values() if (j.labels or {}).get("role") == "consolidate" ) # ---- aggregate ---- view.workers = workers view.rows_done = sum(w.rows_done for w in workers) view.tokens_done_est = sum(w.tokens_done_est or 0 for w in workers) view.cost_usd = sum(w.cost_usd or 0.0 for w in workers) + consolidator_cost view.done = sum(1 for w in workers if w.state == "done" or w.stage in TERMINAL_OK) view.errored = sum(1 for w in workers if w.state == "error" or w.stage in TERMINAL_BAD) view.healthy = n - view.errored gpu_vals = [w.gpu_util for w in workers if w.gpu_util is not None] view.gpu_util = sum(gpu_vals) / len(gpu_vals) if gpu_vals else None timeout_secs = parse_timeout_secs(manifest.get("timeout")) hw = hw_pricing.get(manifest.get("flavor")) if timeout_secs and hw: view.cost_ceiling_usd = n * timeout_secs / 60.0 * hw.unit_cost_usd active_rps = sum(w.rows_per_sec for w in workers if w.state == "running" and w.stage not in TERMINAL_BAD) remaining = max((view.rows_total or 0) - view.rows_done, 0) if active_rps > 0 and remaining > 0: view.eta_secs = remaining / active_rps elif remaining == 0 and view.rows_total: view.eta_secs = 0.0 return view