CoolFace
Apppublic

dkAmulet/sql-query-optimizer

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
app.py101 linesDownload Raw Back to root
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