""" HyperTower distributed job server. Manages a SQLite job queue and a registry of connected clients. Clients poll for work, push status updates, and report completion. Usage: python -m v3.distributed.server --port 8765 --token Environment: HT_TOKEN — fallback if --token is not passed """ from __future__ import annotations import argparse import json import os import sqlite3 import threading import time import uuid from contextlib import contextmanager from datetime import datetime, timezone from pathlib import Path from typing import Optional from fastapi import Depends, FastAPI, Header, HTTPException import uvicorn from .protocol import ( ClientInfo, JobResult, JobSpec, JobSubmit, PollResponse, RegisterRequest, RegisterResponse, StatusPush, ) # ────────────────────────────────────────────────────────────── # Global state # ────────────────────────────────────────────────────────────── _TOKEN: str = "" _DB_PATH: Path = Path("v3/distributed/jobs.db") _CLIENT_TTL: int = 120 # seconds before a client is considered gone _MAX_ATTEMPTS: int = 3 # max times a job is retried before being left as failed _clients: dict[str, ClientInfo] = {} _clients_lock = threading.Lock() def _reap_stale_clients(): """Background thread: remove silent clients and re-queue their running jobs.""" while True: time.sleep(30) cutoff = datetime.now(timezone.utc).timestamp() - _CLIENT_TTL # Step 1: evict timed-out clients from registry with _clients_lock: stale = [ cid for cid, c in _clients.items() if datetime.fromisoformat(c.last_seen).timestamp() < cutoff ] for cid in stale: print(f"[server] reaped stale client {cid} ({_clients[cid].hostname})", flush=True) del _clients[cid] known_ids = set(_clients.keys()) # Step 2: reset any running job whose assigned client is no longer known with _db() as conn: rows = conn.execute( "SELECT job_id, assigned_to FROM jobs WHERE state='running'" ).fetchall() for row in rows: if row["assigned_to"] not in known_ids: conn.execute( "UPDATE jobs SET state='pending', assigned_to=NULL, started_at=NULL " "WHERE job_id=?", (row["job_id"],) ) print(f"[server] re-queued job {row['job_id']} " f"(client {row['assigned_to']} unknown)", flush=True) # ────────────────────────────────────────────────────────────── # Database helpers # ────────────────────────────────────────────────────────────── @contextmanager def _db(): conn = sqlite3.connect(str(_DB_PATH)) conn.row_factory = sqlite3.Row try: yield conn conn.commit() finally: conn.close() def _init_db(): _DB_PATH.parent.mkdir(parents=True, exist_ok=True) with _db() as conn: conn.execute(""" CREATE TABLE IF NOT EXISTS jobs ( job_id TEXT PRIMARY KEY, run_name TEXT NOT NULL, module TEXT NOT NULL, args TEXT NOT NULL, -- JSON list output_dir TEXT NOT NULL DEFAULT 'v3/results', state TEXT NOT NULL DEFAULT 'pending', priority INTEGER NOT NULL DEFAULT 0, assigned_to TEXT, created_at TEXT NOT NULL, started_at TEXT, completed_at TEXT, error_msg TEXT, attempts INTEGER NOT NULL DEFAULT 0 ) """) # Add attempts column to existing DBs that predate this field try: conn.execute("ALTER TABLE jobs ADD COLUMN attempts INTEGER NOT NULL DEFAULT 0") except Exception: pass # column already exists # Note: running jobs are NOT reset on startup — active clients will re-register # via poll/status and the reaper will clean up any that don't reconnect within TTL. def _now() -> str: return datetime.now(timezone.utc).isoformat() def _ensure_client(client_id: str, hostname: str = "", gpu_info: str = "") -> bool: """Re-register a client that survived a server restart. Returns True if the client was unknown (placeholder created) so the caller can ask the client to re-register with full info.""" if client_id not in _clients: _clients[client_id] = ClientInfo( client_id=client_id, hostname=hostname or client_id, gpu_info=gpu_info or "unknown", status=StatusPush(state="idle"), last_seen=_now(), ) print(f"[server] re-registered {client_id} (survived restart)", flush=True) return True return False # ────────────────────────────────────────────────────────────── # FastAPI app # ────────────────────────────────────────────────────────────── app = FastAPI(title="HyperTower Job Server") def _check_token(x_token: str = Header(...)): if x_token != _TOKEN: raise HTTPException(status_code=403, detail="Invalid token") # ── Registration ────────────────────────────────────────────── @app.post("/register", response_model=RegisterResponse, dependencies=[Depends(_check_token)]) def register(req: RegisterRequest, reuse_id: Optional[str] = None): with _clients_lock: client_id = reuse_id if (reuse_id and reuse_id in _clients) else str(uuid.uuid4())[:8] existing_status = _clients[client_id].status if client_id in _clients else StatusPush(state="idle") _clients[client_id] = ClientInfo( client_id=client_id, hostname=req.hostname, gpu_info=req.gpu_info, status=existing_status, last_seen=_now(), ) action = "re-registered" if reuse_id else "registered" print(f"[server] {action} {client_id} ({req.hostname} | {req.gpu_info})", flush=True) return RegisterResponse(client_id=client_id) # ── Job polling ─────────────────────────────────────────────── @app.post("/poll", response_model=PollResponse, dependencies=[Depends(_check_token)]) def poll(client_id: str): with _clients_lock: needs_reregister = _ensure_client(client_id) _clients[client_id].last_seen = _now() with _db() as conn: row = conn.execute( "SELECT * FROM jobs WHERE state='pending' " "ORDER BY priority DESC, created_at ASC LIMIT 1" ).fetchone() if row is None: return PollResponse(job=None, please_reregister=needs_reregister) job_id = row["job_id"] # Reset any other running jobs for this client — client can only work on one at a time. # This cleans up orphans left over from server restarts. cur = conn.execute( "UPDATE jobs SET state='pending', assigned_to=NULL, started_at=NULL " "WHERE assigned_to=? AND state='running' AND job_id!=?", (client_id, job_id), ) if cur.rowcount: print(f"[server] reset {cur.rowcount} orphaned running job(s) for {client_id}", flush=True) conn.execute( "UPDATE jobs SET state='running', assigned_to=?, started_at=? WHERE job_id=?", (client_id, _now(), job_id), ) job = JobSpec( job_id=job_id, run_name=row["run_name"], module=row["module"], args=json.loads(row["args"]), output_dir=row["output_dir"], ) with _clients_lock: _clients[client_id].status = StatusPush( state="syncing", job_id=job_id, run_name=row["run_name"] ) print(f"[server] dispatched {job_id} ({row['run_name']}) → {client_id}", flush=True) return PollResponse(job=job, please_reregister=needs_reregister) # ── Status ──────────────────────────────────────────────────── @app.post("/status/{client_id}", dependencies=[Depends(_check_token)]) def push_status(client_id: str, status: StatusPush): with _clients_lock: needs_reregister = _ensure_client(client_id) _clients[client_id].status = status _clients[client_id].last_seen = _now() return {"ok": True, "please_reregister": needs_reregister} @app.get("/clients", dependencies=[Depends(_check_token)]) def list_clients(): with _clients_lock: return list(_clients.values()) @app.get("/clients/{client_id}", dependencies=[Depends(_check_token)]) def get_client(client_id: str): with _clients_lock: if client_id not in _clients: raise HTTPException(status_code=404, detail="Unknown client") return _clients[client_id] # ── Job completion ──────────────────────────────────────────── @app.post("/complete", dependencies=[Depends(_check_token)]) def complete(result: JobResult): with _db() as conn: if result.success: conn.execute( "UPDATE jobs SET state='done', completed_at=?, error_msg=NULL WHERE job_id=?", (_now(), result.job_id), ) print(f"[server] job {result.job_id} → done", flush=True) # Auto-clear run if all jobs for this run_name are now done run_row = conn.execute( "SELECT run_name FROM jobs WHERE job_id=?", (result.job_id,) ).fetchone() if run_row: run_name = run_row["run_name"] remaining = conn.execute( "SELECT COUNT(*) FROM jobs WHERE run_name=? AND state != 'done'", (run_name,) ).fetchone()[0] if remaining == 0: total = conn.execute( "SELECT COUNT(*) FROM jobs WHERE run_name=?", (run_name,) ).fetchone()[0] conn.execute("DELETE FROM jobs WHERE run_name=?", (run_name,)) print(f"[server] run '{run_name}' complete ({total} jobs) — cleared", flush=True) else: row = conn.execute( "SELECT attempts FROM jobs WHERE job_id=?", (result.job_id,) ).fetchone() attempts = (row["attempts"] if row else 0) + 1 if attempts < _MAX_ATTEMPTS: conn.execute( "UPDATE jobs SET state='pending', assigned_to=NULL, started_at=NULL, " "attempts=?, error_msg=? WHERE job_id=?", (attempts, result.error_msg, result.job_id), ) print(f"[server] job {result.job_id} failed (attempt {attempts}/{_MAX_ATTEMPTS}), " f"re-queuing", flush=True) else: conn.execute( "UPDATE jobs SET state='failed', completed_at=?, attempts=?, error_msg=? " "WHERE job_id=?", (_now(), attempts, result.error_msg, result.job_id), ) print(f"[server] job {result.job_id} failed permanently after " f"{attempts} attempts", flush=True) return {"ok": True} # ── Job queue management ────────────────────────────────────── @app.post("/jobs", dependencies=[Depends(_check_token)]) def submit_job(job: JobSubmit): job_id = str(uuid.uuid4())[:12] with _db() as conn: conn.execute( "INSERT INTO jobs " "(job_id, run_name, module, args, output_dir, priority, created_at) " "VALUES (?,?,?,?,?,?,?)", (job_id, job.run_name, job.module, json.dumps(job.args), job.output_dir, job.priority, _now()), ) print(f"[server] queued {job_id} ({job.run_name})", flush=True) return {"job_id": job_id} @app.get("/jobs", dependencies=[Depends(_check_token)]) def list_jobs(state: Optional[str] = None): with _db() as conn: if state: rows = conn.execute( "SELECT * FROM jobs WHERE state=? ORDER BY created_at DESC", (state,) ).fetchall() else: rows = conn.execute( "SELECT * FROM jobs ORDER BY created_at DESC" ).fetchall() return [dict(r) for r in rows] @app.post("/jobs/clear", dependencies=[Depends(_check_token)]) def clear_jobs(body: dict): with _db() as conn: if body.get("all"): cur = conn.execute("DELETE FROM jobs") elif body.get("run_name"): cur = conn.execute("DELETE FROM jobs WHERE run_name=?", (body["run_name"],)) else: states = body.get("states", ["done", "failed", "cancelled"]) placeholders = ",".join("?" * len(states)) cur = conn.execute(f"DELETE FROM jobs WHERE state IN ({placeholders})", states) print(f"[server] cleared {cur.rowcount} jobs", flush=True) return {"cleared": cur.rowcount} @app.delete("/jobs/{job_id}", dependencies=[Depends(_check_token)]) def cancel_job(job_id: str): with _db() as conn: conn.execute( "UPDATE jobs SET state='cancelled' WHERE job_id=? AND state='pending'", (job_id,), ) return {"ok": True} # ────────────────────────────────────────────────────────────── # Entry point # ────────────────────────────────────────────────────────────── def main(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--port", type=int, default=8765) ap.add_argument("--host", default="0.0.0.0") ap.add_argument("--token", default=os.environ.get("HT_TOKEN", ""), help="Shared secret (or set HT_TOKEN env var)") ap.add_argument("--db", default="v3/distributed/jobs.db", help="Path to SQLite job database") ap.add_argument("--client-ttl", type=int, default=120, help="Seconds of silence before a client is reaped (default: 120)") ap.add_argument("--max-attempts", type=int, default=3, help="Max times a failed job is retried before being left as failed (default: 3)") args = ap.parse_args() if not args.token: ap.error("--token is required (or set HT_TOKEN)") global _TOKEN, _DB_PATH, _CLIENT_TTL, _MAX_ATTEMPTS _TOKEN = args.token _DB_PATH = Path(args.db) _CLIENT_TTL = args.client_ttl _MAX_ATTEMPTS = args.max_attempts _init_db() reaper = threading.Thread(target=_reap_stale_clients, daemon=True) reaper.start() print(f"[server] listening on {args.host}:{args.port} client_ttl={_CLIENT_TTL}s", flush=True) uvicorn.run(app, host=args.host, port=args.port, log_level="warning") if __name__ == "__main__": main()