dkAmulet/sql-query-optimizer
0
1"""2FastAPI application — HTTP wrapper around SQLQueryOptimizerEnv.3"""4from __future__ import annotations5import os6from pathlib import Path7from typing import Optional8 9import uvicorn10from fastapi import FastAPI, HTTPException, Query11from fastapi.middleware.cors import CORSMiddleware12from fastapi.responses import FileResponse, JSONResponse13from fastapi.staticfiles import StaticFiles14 15from env import SQLQueryOptimizerEnv16from models import SQLAction, SQLObservation, StepResult, EnvironmentState, SQLReward, RewardBreakdown17 18app = FastAPI(title="SQL Query Optimizer Environment", version="1.0.0",19 docs_url="/docs", redoc_url="/redoc")20 21app.add_middleware(CORSMiddleware, allow_origins=["*"],22 allow_methods=["*"], allow_headers=["*"])23 24_static_dir = Path(__file__).parent / "static"25if _static_dir.exists():26 app.mount("/static", StaticFiles(directory=str(_static_dir)), name="static")27 28_env = SQLQueryOptimizerEnv()29 30 31def _clamp(v: float) -> float:32 """Force score strictly into (0.001, 0.999) — validator requires exclusive (0,1)."""33 return max(0.001, min(0.999, v))34 35 36def _clamp_result(result: StepResult) -> StepResult:37 """Clamp every score field in a StepResult before returning over HTTP."""38 bd = result.reward.breakdown39 clamped_bd = RewardBreakdown(40 validity=_clamp(bd.validity),41 correctness=_clamp(bd.correctness),42 performance=_clamp(bd.performance),43 style=_clamp(bd.style),44 )45 clamped_reward = SQLReward(46 value=_clamp(result.reward.value),47 breakdown=clamped_bd,48 feedback=result.reward.feedback,49 )50 return StepResult(51 observation=result.observation,52 reward=clamped_reward,53 done=result.done,54 info=result.info,55 )56 57 58@app.get("/", include_in_schema=False)59def root():60 index = _static_dir / "index.html"61 if index.exists():62 return FileResponse(str(index), media_type="text/html")63 return JSONResponse({"name": "SQL Query Optimizer", "version": "1.0.0", "docs": "/docs"})64 65 66@app.get("/health", tags=["meta"])67def health():68 return {"status": "ok", "environment": "sql-query-optimizer", "version": "1.0.0"}69 70 71@app.get("/tasks", tags=["meta"])72def list_tasks():73 return _env.list_tasks()74 75 76@app.post("/reset", response_model=SQLObservation, tags=["openenv"])77def reset(task_id: Optional[str] = Query(default=None)):78 try:79 return _env.reset(task_id)80 except ValueError as exc:81 raise HTTPException(status_code=400, detail=str(exc))82 83 84@app.post("/step", response_model=StepResult, tags=["openenv"])85def step(action: SQLAction):86 try:87 result = _env.step(action)88 return _clamp_result(result) # ← clamp here before sending over HTTP89 except RuntimeError as exc:90 raise HTTPException(status_code=400, detail=str(exc))91 92 93@app.get("/state", response_model=EnvironmentState, tags=["openenv"])94def state():95 return _env.state()96 97 98if __name__ == "__main__":99 port = int(os.environ.get("PORT", 7860))100 uvicorn.run(app, host="0.0.0.0", port=port, log_level="info")101 