Team Ai
Apppublic

lablab-ai-amd-developer-hackathon/gpu-goblin

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
server.py329 linesDownload Raw Back to agent
1"""FastAPI server for GPU Goblin.2 3One audit endpoint plus a health probe. Streams the agent loop's `SSEEvent`s4to the UI via Server-Sent Events. CORS is wide open because Streamlit runs on5a different port — fine for a hackathon.6 7The agent runs on Qwen via Hugging Face Inference Providers. HF_TOKEN is8read at startup; if it's missing the server still starts (so the offline-9replay UI lane keeps working) but `/audit` yields a single error event.10We never crash on missing keys.11"""12 13from __future__ import annotations14 15import asyncio16import json17import os18import subprocess19import sys20import tempfile21from collections.abc import AsyncIterator22from pathlib import Path23from typing import Any24 25from fastapi import FastAPI, File, HTTPException, UploadFile26from fastapi.middleware.cors import CORSMiddleware27from pydantic import BaseModel, Field28from sse_starlette.sse import EventSourceResponse29 30from agent.backends import active_backend_name31from agent.loop import run_audit32from agent.schemas import SSEEvent33from agent.tools import ALL_TOOLS34 35_REPO_ROOT = Path(__file__).resolve().parent.parent36_AUTO_TUNE_SCRIPT = _REPO_ROOT / "scripts" / "auto_tune.py"37 38app = FastAPI(title="GPU Goblin Agent", version="0.1.0")39 40app.add_middleware(41    CORSMiddleware,42    allow_origins=["*"],43    allow_credentials=False,44    allow_methods=["*"],45    allow_headers=["*"],46)47 48 49def _has_hf_token() -> bool:50    return bool(os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACEHUB_API_TOKEN"))51 52 53@app.get("/healthz")54async def healthz() -> dict:55    """Liveness + tool inventory + active backend. UI uses this to confirm56    the agent is reachable and configured."""57    name = active_backend_name()58    base = {59        "ok": True,60        "tools": [t.name for t in ALL_TOOLS],61        "backend": name,62    }63    if name == "qwen-vllm":64        base.update(65            {66                "model": os.environ.get(67                    "GOBLIN_QWEN_VLLM_MODEL", "Qwen/Qwen2.5-7B-Instruct"68                ),69                "vllm_url": os.environ.get(70                    "GOBLIN_QWEN_VLLM_URL", "http://localhost:8000/v1"71                ),72                "has_api_key": True,  # vLLM doesn't require one by default73            }74        )75    else:76        base.update(77            {78                "model": os.environ.get(79                    "GOBLIN_QWEN_MODEL", "Qwen/Qwen2.5-7B-Instruct"80                ),81                "provider": os.environ.get("GOBLIN_QWEN_PROVIDER", "auto"),82                "has_api_key": _has_hf_token(),83            }84        )85    return base86 87 88async def _stream_audit(file_path: str) -> AsyncIterator[dict]:89    """Bridge `run_audit`'s SSEEvent generator into the dict shape that90    sse-starlette expects. Each yielded dict becomes one `data: ...\\n\\n`91    SSE message.92    """93    if not _has_hf_token():94        # Surface a clean error instead of letting the loop crash on missing key.95        yield {96            "data": SSEEvent(97                type="error",98                data={99                    "message": (100                        "HF_TOKEN not set on the server — Qwen agent loop is "101                        "unavailable. Set HF_TOKEN (or HUGGINGFACEHUB_API_TOKEN) "102                        "or use the offline-replay UI lane."103                    )104                },105            ).model_dump_json()106        }107        return108 109    try:110        async for event in run_audit(file_path):111            yield {"data": event.model_dump_json()}112    except Exception as exc:  # defence in depth — run_audit already wraps itself113        yield {114            "data": SSEEvent(115                type="error", data={"message": f"server: {exc}"}116            ).model_dump_json()117        }118 119 120@app.post("/audit")121async def audit(file: UploadFile = File(...)) -> EventSourceResponse:122    """Accept a multipart file upload and stream the agent's audit events.123 124    The uploaded file is saved to a tempfile (preserving the extension so125    `parse_config`'s extension-dispatched parser picks the right path) and126    handed to `run_audit`. We don't delete the temp file here — the audit127    might still be reading it; the OS reaps it eventually and `bench_cache/`128    is gitignored.129    """130    suffix = Path(file.filename or "").suffix or ".bin"131    fd, tmp_path = tempfile.mkstemp(prefix="goblin_upload_", suffix=suffix)132    try:133        with os.fdopen(fd, "wb") as f:134            f.write(await file.read())135    except Exception:136        # If we couldn't even land the upload, surface that immediately.137        async def _err() -> AsyncIterator[dict]:138            yield {139                "data": SSEEvent(140                    type="error",141                    data={"message": "Failed to save uploaded file."},142                ).model_dump_json()143            }144 145        return EventSourceResponse(_err())146 147    return EventSourceResponse(_stream_audit(tmp_path))148 149 150# ---------------------------------------------------------------------------151# Auto-tune endpoint — lets a UI on a CPU-only host (e.g. an HF Space) drive152# scripts/auto_tune.py running on a remote MI300X server. The endpoint153# spawns the CLI, tails its --events NDJSON stream, and re-emits each line154# as an SSE message. Subprocess output is discarded; everything the UI155# needs is in the structured events.156# ---------------------------------------------------------------------------157 158 159class AutoTuneRequest(BaseModel):160    """JSON shape the /auto-tune endpoint accepts. Mirrors the auto_tune.py161    CLI surface so the UI just sends what the user picked in the form."""162 163    model: str | None = Field(164        default=None,165        description="HuggingFace model id (e.g. Qwen/Qwen2.5-7B-Instruct). "166        "Mutually exclusive with `workload`.",167    )168    workload: str | None = Field(169        default=None,170        description="Path to a workload script ON THE SERVER's filesystem. "171        "Mutually exclusive with `model`.",172    )173    mode: str = Field(default="hardcoded", pattern="^(hardcoded|llm|llm-explore)$")174    candidates_per_iteration: int = Field(default=3, ge=2, le=10)175    steps: int = Field(default=20, ge=1, le=500)176    max_iterations: int = Field(default=10, ge=1, le=50)177    early_stop_after: int = Field(default=3, ge=1, le=20)178    max_crashes: int = Field(default=4, ge=1, le=20)179    improvement_threshold: float = Field(default=0.0, ge=0.0, le=20.0)180 181 182def _build_auto_tune_cmd(req: AutoTuneRequest, events_file: Path) -> list[str]:183    cmd: list[str] = [sys.executable, "-u", str(_AUTO_TUNE_SCRIPT)]184    if req.model:185        cmd.extend(["--model", req.model])186    elif req.workload:187        cmd.append(req.workload)188    cmd.extend([189        "--mode", req.mode,190        "--steps", str(req.steps),191        "--max-iterations", str(req.max_iterations),192        "--early-stop-after", str(req.early_stop_after),193        "--max-crashes", str(req.max_crashes),194        "--improvement-threshold", str(req.improvement_threshold),195        "--events", str(events_file),196    ])197    if req.mode == "llm-explore":198        cmd.extend(["--candidates-per-iteration", str(req.candidates_per_iteration)])199    return cmd200 201 202async def _stream_auto_tune(req: AutoTuneRequest) -> AsyncIterator[dict]:203    """Spawn auto_tune.py and forward its NDJSON --events stream as SSE.204 205    Each event is forwarded verbatim — the UI gets the same structured206    payload it would see when running auto_tune.py locally. We discard207    the subprocess's stdout/stderr; any errors are surfaced via the208    `summary` event's absence at process exit.209    """210    events_file = Path(tempfile.mktemp(prefix="auto_tune_events_", suffix=".ndjson"))211    events_file.write_text("")212 213    cmd = _build_auto_tune_cmd(req, events_file)214 215    # Validate at least one of model/workload was provided. (Pydantic216    # can't express "exactly one of A or B" cleanly, so we check here.)217    if not req.model and not req.workload:218        yield {"data": json.dumps({219            "type": "error",220            "message": "Pass either `model` or `workload`, not neither."221        })}222        return223    if req.model and req.workload:224        yield {"data": json.dumps({225            "type": "error",226            "message": "Pass either `model` or `workload`, not both."227        })}228        return229 230    proc = subprocess.Popen(231        cmd,232        cwd=str(_REPO_ROOT),233        stdout=subprocess.DEVNULL,234        stderr=subprocess.DEVNULL,235        env={**os.environ},236    )237 238    seen_bytes = 0239    try:240        while True:241            # Poll the events file for new lines242            try:243                with events_file.open("r") as f:244                    f.seek(seen_bytes)245                    chunk = f.read()246                    new_seen = f.tell()247            except OSError:248                chunk = ""249                new_seen = seen_bytes250 251            if chunk:252                # Drop a trailing partial line — re-read it next tick once253                # the writer has flushed the rest.254                lines = chunk.splitlines(keepends=True)255                if lines and not lines[-1].endswith("\n"):256                    partial = lines.pop()257                    new_seen -= len(partial.encode("utf-8"))258                for line in lines:259                    line = line.strip()260                    if line:261                        yield {"data": line}262            seen_bytes = new_seen263 264            if proc.poll() is not None:265                # Subprocess exited. Drain whatever's left on disk.266                try:267                    with events_file.open("r") as f:268                        f.seek(seen_bytes)269                        tail = f.read()270                except OSError:271                    tail = ""272                for line in tail.splitlines():273                    line = line.strip()274                    if line:275                        yield {"data": line}276                if proc.returncode != 0:277                    yield {"data": json.dumps({278                        "type": "process_exit",279                        "returncode": proc.returncode,280                        "message": (281                            f"auto_tune.py exited with code {proc.returncode}. "282                            "Check the server's stderr or check `last_runner_failure_*` "283                            "in `bench_cache/` for goblin_runner.sh failure logs."284                        ),285                    })}286                break287 288            await asyncio.sleep(0.5)289    finally:290        if proc.poll() is None:291            proc.terminate()292            try:293                proc.wait(timeout=3)294            except subprocess.TimeoutExpired:295                proc.kill()296        try:297            events_file.unlink()298        except OSError:299            pass300 301 302@app.post("/auto-tune")303async def auto_tune_endpoint(req: AutoTuneRequest) -> EventSourceResponse:304    """Stream auto_tune.py events back to the caller as SSE.305 306    Run a UI on any host (HF Spaces, local laptop), point it at this307    endpoint, and the actual GPU work happens on the server hosting the308    FastAPI app. Subprocess output is discarded — only the --events309    NDJSON stream crosses the wire, one structured event per SSE message.310    """311    if not _AUTO_TUNE_SCRIPT.exists():312        raise HTTPException(313            status_code=500,314            detail=f"auto_tune.py not found at {_AUTO_TUNE_SCRIPT}",315        )316    return EventSourceResponse(_stream_auto_tune(req))317 318 319# Convenience: support `python -m uvicorn agent.server:app --reload`.320__all__ = ["app"]321 322 323def _decode_event(raw: str) -> dict:324    """Helper for the CLI driver — parse a serialized SSEEvent JSON payload.325 326    Lives here so __main__.py and tests can share one parser.327    """328    return json.loads(raw)329