vinayaknandi05/sql-optimization-openenv
0
1"""2FastAPI application exposing the SQLOptimizationEnv via HTTP,3compatible with the OpenEnv spec used by the competition judge.4"""5 6import json7import uuid8from typing import Optional9from fastapi import FastAPI, HTTPException10from fastapi.responses import JSONResponse11from pydantic import BaseModel12 13import sys, os14sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))15from env.environment import SQLOptimizationEnv, SQLAction16 17app = FastAPI(18 title="SQL Optimization OpenEnv",19 description="An OpenEnv environment for SQL query optimization tasks.",20 version="1.0.0",21)22 23# Simple in-memory session store (fine for single-container HF Space)24_sessions: dict[str, SQLOptimizationEnv] = {}25 26 27def _get_or_create(session_id: str, task_id: Optional[str] = None) -> SQLOptimizationEnv:28 if session_id not in _sessions:29 _sessions[session_id] = SQLOptimizationEnv(task_id=task_id)30 return _sessions[session_id]31 32 33# ── Health / ping ────────────────────────────────────────────────────────────34 35@app.get("/")36def root():37 return {"status": "ok", "env": "sql-optimization-openenv"}38 39 40@app.get("/health")41def health():42 return {"status": "ok"}43 44 45# ── OpenEnv endpoints ────────────────────────────────────────────────────────46 47class ResetRequest(BaseModel):48 session_id: Optional[str] = None49 task_id: Optional[str] = None # "task_easy" | "task_medium" | "task_hard"50 51 class Config:52 # Allow both empty and populated bodies53 extra = "ignore"54 55 56@app.post("/reset")57def reset(req: Optional[ResetRequest] = None):58 if req is None:59 req = ResetRequest()60 sid = req.session_id or str(uuid.uuid4())61 env = _get_or_create(sid, task_id=req.task_id)62 obs = env.reset(task_id=req.task_id)63 return {"session_id": sid, "observation": obs.model_dump()}64 65 66class StepRequest(BaseModel):67 session_id: str68 action: dict # {"query": "...", "message": "..."}69 70 71@app.post("/step")72def step(req: StepRequest):73 if req.session_id not in _sessions:74 raise HTTPException(status_code=404, detail="Session not found. Call /reset first.")75 env = _sessions[req.session_id]76 try:77 action = SQLAction(**req.action)78 except Exception as e:79 raise HTTPException(status_code=422, detail=f"Invalid action: {e}")80 81 obs, reward, done, info = env.step(action)82 return {83 "session_id": req.session_id,84 "observation": obs.model_dump(),85 "reward": reward,86 "done": done,87 "info": info,88 }89 90 91@app.get("/state")92def state(session_id: str):93 if session_id not in _sessions:94 raise HTTPException(status_code=404, detail="Session not found.")95 return _sessions[session_id].state()96 97 98@app.get("/tasks")99def list_tasks():100 """List all available tasks."""101 from env.environment import TASKS102 return {"tasks": [{"id": t["id"], "name": t["name"], "difficulty": t["difficulty"]} for t in TASKS]}103 104 105def main():106 """Entry point for server."""107 import uvicorn108 uvicorn.run(app, host="0.0.0.0", port=7860)109 110 111if __name__ == "__main__":112 main()113 