Team Ai
Apppublic

vinayaknandi05/sql-optimization-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
app.py113 linesDownload Raw Back to server
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