| from fastapi import FastAPI, HTTPException |
| from fastapi.responses import HTMLResponse |
| from fastapi.staticfiles import StaticFiles |
| from pydantic import BaseModel |
| from typing import Dict, Any |
| import sys |
| import os |
|
|
| sys.path.append(os.path.dirname(os.path.abspath(__file__))) |
|
|
| from models import Action, Observation, State |
| from server.llm_env import LLMEnv |
|
|
| app = FastAPI(title="LLM Control OpenEnv") |
|
|
| |
| envs: Dict[str, LLMEnv] = {} |
| completed_episodes: Dict[str, Dict[str, Any]] = {} |
|
|
| |
| default_env = LLMEnv() |
| default_env.reset() |
|
|
| class ResetRequest(BaseModel): |
| task: str = "easy" |
|
|
| class StepRequest(BaseModel): |
| action: Action |
| episode_id: str | None = None |
|
|
| class GraderRequest(BaseModel): |
| episode_id: str |
|
|
| |
| default_env = LLMEnv() |
|
|
| @app.get("/", response_class=HTMLResponse) |
| async def serve_gui(): |
| path = os.path.join(os.path.dirname(__file__), "index.html") |
| try: |
| with open(path, "r") as f: |
| return f.read() |
| except FileNotFoundError: |
| return "GUI index.html not found. Check the root directory." |
|
|
| @app.post("/reset") |
| async def reset(req: ResetRequest = ResetRequest()): |
| if req.task not in ["easy", "medium", "hard"]: |
| raise HTTPException(status_code=400, detail="Invalid task") |
| |
| env = LLMEnv(task=req.task) |
| obs = env.reset() |
| state = env.state |
| envs[state.episode_id] = env |
| |
| |
| global default_env |
| default_env = env |
| |
| return { |
| "observation": obs.model_dump(), |
| "state": state.model_dump() |
| } |
|
|
| @app.post("/step") |
| async def step(req: StepRequest): |
| |
| env = default_env |
| if req.episode_id and req.episode_id in envs: |
| env = envs[req.episode_id] |
| |
| obs, reward, done, info = env.step(req.action) |
| |
| if done: |
| |
| completed_episodes[env.state.episode_id] = { |
| "reward": env.state.cumulative_reward, |
| "bounds": env._reward_bounds() |
| } |
| |
| return { |
| "observation": obs.model_dump(), |
| "reward": reward, |
| "done": done, |
| "info": info |
| } |
|
|
| @app.get("/state", response_model=State) |
| async def get_state(episode_id: str | None = None): |
| env = default_env |
| if episode_id and episode_id in envs: |
| env = envs[episode_id] |
| return env.state |
|
|
| @app.post("/baseline") |
| async def run_baseline(): |
| import subprocess |
| try: |
| |
| baseline_path = os.path.join(os.path.dirname(__file__), "baseline.py") |
| result = subprocess.run([sys.executable, baseline_path], capture_output=True, text=True, check=True) |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| import json |
| out = result.stdout.strip().splitlines()[-1] |
| |
| scores = json.loads(out) |
| return scores |
| except subprocess.CalledProcessError as e: |
| raise HTTPException(status_code=500, detail=f"Baseline failed: {e.stderr}") |
| except Exception as e: |
| raise HTTPException(status_code=500, detail=f"Baseline error: {str(e)}") |
|
|
| @app.post("/grader") |
| async def grader(req: GraderRequest): |
| if req.episode_id not in completed_episodes: |
| |
| if req.episode_id in envs: |
| env = envs[req.episode_id] |
| r = env.state.cumulative_reward |
| b_min, b_max = env._reward_bounds() |
| score = (norm * 0.998) + 0.001 |
| return {"score": score} |
| |
| raise HTTPException(status_code=404, detail="Episode not found or not finished") |
| |
| data = completed_episodes[req.episode_id] |
| r = data["reward"] |
| b_min, b_max = data["bounds"] |
| norm = (r - b_min) / (b_max - b_min) |
| |
| |
| score = (norm * 0.998) + 0.001 |
| return {"score": score} |
|
|
| @app.get("/tasks") |
| async def get_tasks(): |
| return { |
| "tasks": ["easy", "medium", "hard"], |
| "action_schema": Action.model_json_schema() |
| } |
|
|
| if __name__ == "__main__": |
| import uvicorn |
| uvicorn.run(app, host="0.0.0.0", port=7860) |
|
|