rtrm's picture
rtrm HF Staff
feat: HF OAuth sign-in, probe endpoints restricted to Hugging Face org members
88e4a42 unverified
Raw
History Blame Contribute Delete
5.41 kB
import logging
import os
import re
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from contextlib import asynccontextmanager
import uvicorn
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import FileResponse
from pydantic import BaseModel
from . import auth, config, history, prober, state
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s %(message)s")
log = logging.getLogger("gpu_capacity")
INDEX_HTML = os.path.join(os.path.dirname(__file__), "static", "index.html")
def _sweep(instance_type: str, regions: list[str]) -> dict:
with ThreadPoolExecutor(max_workers=config.MAX_PROBE_WORKERS) as pool:
results = dict(
zip(regions, pool.map(lambda r: prober.probe_region(r, instance_type), regions))
)
for region, result in results.items():
state.update(instance_type, region, result)
try:
history.record(instance_type, results)
except Exception:
log.exception("history record failed")
return results
def _auto_refresh_loop():
# If restored state is still fresh (e.g. restart right after a sweep),
# don't re-probe everything — wait out the remainder of the period.
# Types added to the config since the state was written are probed anyway.
newest = state.newest_checked_at()
if newest is not None:
wait = config.AUTO_REFRESH_MINUTES * 60 - (time.time() - newest)
if wait > 0:
for instance_type in config.INSTANCE_TYPES:
if instance_type not in state.get():
_sweep(instance_type, config.REGIONS)
log.info("restored state is fresh, first full sweep in %.0fs", wait)
time.sleep(wait)
while True:
log.info("auto-refresh: probing all configured types")
started = time.time()
for instance_type in config.INSTANCE_TYPES:
_sweep(instance_type, config.REGIONS)
log.info("auto-refresh: sweep done in %.0fs", time.time() - started)
time.sleep(max(60, config.AUTO_REFRESH_MINUTES * 60 - (time.time() - started)))
@asynccontextmanager
async def lifespan(app: FastAPI):
state.load()
history.load()
history.start_writer()
threading.Thread(target=prober.cleanup_leaked, daemon=True).start()
if config.AUTO_REFRESH_MINUTES > 0:
log.info("auto-refresh enabled: every %d min", config.AUTO_REFRESH_MINUTES)
threading.Thread(target=_auto_refresh_loop, daemon=True).start()
yield
app = FastAPI(title="gpu-capacity", lifespan=lifespan)
app.include_router(auth.router)
class ProbeRequest(BaseModel):
instance_type: str
region: str | None = None
TYPE_RE = re.compile(r"[a-z0-9\-]+\.[a-z0-9\-]+")
def _validate(instance_type: str, region: str | None):
if not TYPE_RE.fullmatch(instance_type):
raise HTTPException(400, f"invalid instance type {instance_type!r}")
if region is not None and region not in config.REGIONS:
raise HTTPException(400, f"unknown region {region!r}")
@app.get("/")
def index():
return FileResponse(INDEX_HTML)
@app.get("/api/state")
def get_state():
current = state.get()
extras = sorted(t for t in current if t not in config.INSTANCE_TYPES)
types = config.INSTANCE_TYPES + extras
return {
"instance_types": types,
"regions": config.REGIONS,
"stale_after_seconds": config.STALE_AFTER_SECONDS,
"gpu_info": prober.gpu_info(types),
"state": current,
}
@app.get("/api/me")
def me(request: Request):
user = auth.current_user(request)
return {
"authenticated": user is not None,
"username": user.get("u") if user else None,
"can_probe": (not auth.ENABLED) or bool(user and user.get("hf")),
"auth_enabled": auth.ENABLED,
}
@app.get("/api/availability/{instance_type}")
def availability(request: Request, instance_type: str,
refresh: bool = False, region: str | None = None):
"""Cached availability for one instance type; refresh=true probes first."""
_validate(instance_type, region)
regions = [region] if region else config.REGIONS
if refresh:
auth.require_probe_rights(request)
_sweep(instance_type, regions)
data = state.get().get(instance_type, {})
data = {r: data[r] for r in regions if r in data}
return {
"instance_type": instance_type,
"available_regions": sorted(r for r, v in data.items() if v["status"] == "available"),
"regions": data,
}
@app.get("/api/history/{instance_type}")
def history_series(instance_type: str, region: str, hours: float = 48):
_validate(instance_type, region)
return {
"instance_type": instance_type,
"region": region,
"series": history.series(instance_type, region, hours),
}
@app.post("/api/probe")
def probe(req: ProbeRequest, request: Request):
auth.require_probe_rights(request)
_validate(req.instance_type, req.region)
regions = [req.region] if req.region else config.REGIONS
log.info("probing %s in %s", req.instance_type, ",".join(regions))
return {"instance_type": req.instance_type, "results": _sweep(req.instance_type, regions)}
def run():
uvicorn.run(app, host=os.environ.get("HOST", "127.0.0.1"), port=int(os.environ.get("PORT", "8300")))
if __name__ == "__main__":
run()