File size: 2,032 Bytes
9682da5
ebedc9d
 
79f731d
9682da5
 
 
 
 
 
 
ebedc9d
 
 
 
 
9682da5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
07c8a8a
 
 
9682da5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79f731d
 
ebedc9d
 
 
 
91496f8
ebedc9d
 
 
 
 
 
 
 
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
# server/app.py β€” FastAPI server with all required endpoints
import os
import uvicorn

from fastapi import FastAPI
from fastapi.responses import JSONResponse
from pydantic import BaseModel
from typing import Optional
from server.pipeline_environment import PipelineEnvironment
from models import PipelineAction, RepairAction

# ── Only build the Gradio UI if running on HF Spaces ───
if os.getenv("SPACE_ID"):
    import gradio as gr
    from ui import demo as gradio_demo

app = FastAPI(title="PipelineEnv", version="1.0.0")

# Single global environment instance
env = PipelineEnvironment()


# ─── REQUEST MODELS ──────────────────────────────────
class ResetRequest(BaseModel):
    task_id: str = "easy"


class StepRequest(BaseModel):
    action: str
    target: Optional[str] = None
    value:  Optional[str] = None


# ─── ENDPOINTS ───────────────────────────────────────
@app.get("/health")
def health():
    return {"status": "healthy"}


@app.post("/reset")
def reset(req: Optional[ResetRequest] = None):
    task_id = req.task_id if req else "easy"
    obs = env.reset(task_id=task_id)
    return obs.model_dump()


@app.post("/step")
def step(req: StepRequest):
    try:
        action = PipelineAction(
            action=RepairAction(req.action),
            target=req.target,
            value=req.value,
        )
    except Exception as e:
        return JSONResponse(status_code=400, content={"error": str(e)})

    result = env.step(action)
    return result


@app.get("/state")
def state():
    return env.state.model_dump()


# ─── GRADIO MOUNT (only on HF Spaces) ──────────────────
if os.getenv("SPACE_ID"):
    from ui import demo as _demo
    app = gr.mount_gradio_app(app, _demo, path="/")


def main():
    uvicorn.run(app, host="0.0.0.0", port=7860)


if __name__ == "__main__":
    main()