File size: 5,882 Bytes
66ee87e
 
 
 
 
 
43c3b3c
66ee87e
 
43c3b3c
 
 
 
 
 
 
 
 
 
66ee87e
 
 
 
 
 
 
 
21c5089
66ee87e
 
 
 
 
 
 
43c3b3c
66ee87e
 
 
 
43c3b3c
66ee87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
43c3b3c
66ee87e
 
 
 
 
 
 
 
 
 
43c3b3c
 
 
 
 
21c5089
 
66ee87e
43c3b3c
 
 
 
 
 
 
 
 
 
 
 
66ee87e
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
"""DecisionLab on a Hugging Face Space (Gradio SDK, ZeroGPU hardware).

The Space runs this file. It serves the SAME FastAPI app as the container, so the lab page at "/" (HTML, CSS and
JavaScript) and /api/* are the container's. A Gradio app is mounted beside it at /gradio: a small form and a "decide"
API endpoint for gradio_client, both running the same decision code (app.main.run_decision).

ZeroGPU (Hugging Face docs and the `spaces` package source, checked 2026-09-30):
  - `import spaces` comes first: it patches torch so CUDA looks available everywhere; outside @spaces.GPU a CUDA
    emulation mode lets models be moved to "cuda" without a real GPU. A real GPU exists only inside @spaces.GPU calls.
  - ZeroGPU's startup step (spaces/zero/__init__.py: `gradio.one_launch(torch.pack)`) runs inside Gradio's own
    Blocks.launch(). So this file starts the server with demo.launch(), NOT uvicorn + mount_gradio_app: without
    launch() the step never runs and the Space fails with "No @spaces.GPU function detected during startup"
    (2.1.0 and 2.1.1, 2026-09-30).
  - Models are loaded BEFORE launch(), so the startup step packs them (the documented ZeroGPU pattern).
  - The lab's routes (the page at "/", /static, /api/*) are put in front of Gradio's inside Gradio's server, so "/"
    is the container's page. Gradio's own routes stay for gradio_client (api_name "/decide").
  - Only the model-running step is decorated; the request checks and the concurrency limit stay in the main process.
    Each model times itself inside the GPU call, so model timings stay comparable.
  - Gradio's server-side rendering is off: its Node proxy would take port 7860.
@spaces.GPU does nothing on non-ZeroGPU hardware, so this file also runs on a CPU Space.

No token auth anywhere (operator ruling 2026-09-28): on a public Space, anyone can use the page and the API.
"""
import spaces  # noqa: I001  -- must be the first import (patches torch for ZeroGPU)

import json
import os
import threading
from pathlib import Path

HERE = Path(__file__).resolve().parent
os.environ.setdefault("MODELS_DIR", str(HERE / "models"))      # the Space's own models/ folder
os.environ.setdefault("DLAB_WARMUP", "0")                       # no real GPU outside @spaces.GPU: no warm-up run

import gradio as gr  # noqa: E402
from starlette.middleware import Middleware  # noqa: E402

from app import main as lab  # noqa: E402
from app.main import VERSION, Busy, app, run_decision  # noqa: E402
from app.models import load_all  # noqa: E402
from app.security import SecurityMiddleware  # noqa: E402

GPU_SECONDS = int(os.getenv("GPU_SECONDS", "60"))   # longest a single decision may hold the GPU


@spaces.GPU(duration=GPU_SECONDS)
def run_models_on_gpu(state, questions, keys):
    return lab.run_models(state, questions, keys)


lab.RUN_MODELS = run_models_on_gpu


EXAMPLE_STATE = "I was charged twice for March. Refund the duplicate or I cancel."
EXAMPLE_QUESTIONS = json.dumps({
    "dept": {"type": "choice", "instructions": "Which team?",
             "criteria": {"billing": "Payments and refunds", "tech": "Bugs and errors"}},
    "churn": {"type": "noul", "instructions": "The customer threatens to leave"},
}, indent=2)


def decide(state: str, questions_json: str, models: str = "") -> dict:
    """Run one decision. state: text, or a JSON object. questions_json: the questions as JSON.
    models: comma-separated model keys (see /api/status); empty runs every model."""
    try:
        questions = json.loads(questions_json)
    except ValueError as exc:
        raise gr.Error(f"Questions must be valid JSON: {exc}") from None
    if not isinstance(questions, dict):
        raise gr.Error("Questions must be a JSON object: {name: {type, instructions, criteria}}.")
    text = (state or "").strip()
    try:
        parsed = json.loads(text) if text.startswith("{") else text
    except ValueError:
        parsed = text
    keys = [k.strip() for k in (models or "").split(",") if k.strip()] or None
    try:
        return run_decision(parsed, questions, keys)
    except (ValueError, Busy) as exc:
        raise gr.Error(str(exc)) from None


def build_blocks() -> gr.Blocks:
    with gr.Blocks(title="DecisionLab API") as demo:
        gr.Markdown(f"# DecisionLab {VERSION}: Gradio API\n"
                    "The full lab is the main page. `gradio_client` (api_name \"/decide\") runs the same decision code.")
        with gr.Row():
            state = gr.Textbox(label="State (text or JSON)", lines=8, value=EXAMPLE_STATE)
            questions = gr.Code(label="Questions (JSON)", language="json", value=EXAMPLE_QUESTIONS)
        models = gr.Textbox(label="Models (comma-separated keys from /api/status; empty = every model)", value="")
        run = gr.Button("Run", variant="primary")
        out = gr.JSON(label="Result")
        run.click(decide, [state, questions, models], out, api_name="decide")
    return demo


def lab_first(server_app) -> None:
    """Put the lab's routes (page, /static, /api/*) in front of Gradio's, so "/" is DecisionLab's page."""
    lab_routes = list(app.routes)
    rest = [r for r in server_app.router.routes if r not in lab_routes]
    server_app.router.routes[:] = lab_routes + rest


def main() -> None:
    load_all()                              # before launch(): ZeroGPU's startup step packs the loaded models
    demo = build_blocks()
    demo.launch(
        server_name=os.getenv("GRADIO_SERVER_NAME") or "0.0.0.0",
        server_port=int(os.getenv("GRADIO_SERVER_PORT") or os.getenv("PORT") or 7860),
        ssr_mode=False,
        prevent_thread_lock=True,
        app_kwargs={"docs_url": None, "redoc_url": None, "openapi_url": None,
                    "middleware": [Middleware(SecurityMiddleware, max_body=lab.MAX_BODY_BYTES)]},
    )
    lab_first(demo.app)
    demo.block_thread()


if __name__ == "__main__":
    main()