MauryaVivek/sql-data-quality-agent
0
1"""2app.py3======4FastAPI server exposing the SQL Data Quality Agent as an OpenEnv HTTP API.5 6Endpoints:7 GET /health -> health check (required for HF Space ping)8 GET /tasks -> list all available tasks9 POST /reset -> start a new episode10 POST /step -> take one action11 GET /state -> current environment state12 GET /docs -> auto-generated Swagger UI (FastAPI built-in)13"""14 15import os16from typing import Any, Dict, Optional17 18from fastapi import FastAPI, HTTPException, Request19from fastapi.middleware.cors import CORSMiddleware20from fastapi.responses import JSONResponse21from pydantic import BaseModel, field_validator22 23from environment import DataQualityEnv, DataQualityAction, DataQualityObservation, DataQualityState24from tasks import list_tasks, clamp_score, clamp_ratio25 26# ---------------------------------------------------------------------------27# App setup28# ---------------------------------------------------------------------------29 30app = FastAPI(31 title="SQL Data Quality Agent — OpenEnv",32 description=(33 "An OpenEnv environment where an AI agent acts as a Data Quality Engineer, "34 "fixing dirty SQL databases through iterative SQL statements. "35 "Supports 3 tasks: easy -> medium -> hard."36 ),37 version="1.0.0",38 docs_url="/docs",39 redoc_url="/redoc",40)41 42app.add_middleware(43 CORSMiddleware,44 allow_origins=["*"],45 allow_credentials=True,46 allow_methods=["*"],47 allow_headers=["*"],48)49 50# Single global environment instance (stateful per-container)51env = DataQualityEnv()52 53 54# ---------------------------------------------------------------------------55# Helper: recursively clamp all float values in a dict that look like scores56# ---------------------------------------------------------------------------57 58def _clamp_response_scores(obj: Any) -> Any:59 """Recursively walk a dict/list and clamp any float values that are60 exactly 0.0 or 1.0, or any float keys named *score*, *ratio*, *reward*."""61 if isinstance(obj, dict):62 result = {}63 for k, v in obj.items():64 if isinstance(v, float):65 # Clamp any float field whose name suggests it's a score/ratio/reward66 key_lower = str(k).lower()67 if any(word in key_lower for word in ["score", "ratio", "reward"]):68 result[k] = clamp_score(v) if "score" in key_lower or "reward" in key_lower else clamp_ratio(v)69 elif v <= 0.0 or v >= 1.0:70 # For other floats that happen to be 0.0 or 1.0, leave them71 # (they might be legitimate values like amounts, counts, etc.)72 result[k] = v73 else:74 result[k] = v75 else:76 result[k] = _clamp_response_scores(v)77 return result78 elif isinstance(obj, list):79 return [_clamp_response_scores(item) for item in obj]80 else:81 return obj82 83 84# ---------------------------------------------------------------------------85# Request / Response models86# ---------------------------------------------------------------------------87 88class ResetRequest(BaseModel):89 task_id: str = "null_patrol"90 seed: int = 4291 92 class Config:93 extra = "ignore"94 95 96class StepRequest(BaseModel):97 sql: str98 rationale: str = ""99 100 101class StepResponse(BaseModel):102 observation: DataQualityObservation103 reward: float104 done: bool105 info: Dict[str, Any]106 107 @field_validator("reward", mode="before")108 @classmethod109 def _clamp_reward(cls, v: float) -> float:110 """Ensure reward is strictly in (0, 1) before serialization."""111 return clamp_score(v)112 113 114# ---------------------------------------------------------------------------115# Endpoints116# ---------------------------------------------------------------------------117 118@app.get("/", tags=["meta"])119def root():120 """Root endpoint — required by HF Spaces health probe."""121 return {"status": "healthy", "environment": "sql-data-quality-agent", "version": "1.0.0"}122 123 124@app.get("/health", tags=["meta"])125def health_check():126 """Health check — openenv validate expects {"status": "healthy"}."""127 return {"status": "healthy", "environment": "sql-data-quality-agent", "version": "1.0.0"}128 129 130@app.get("/metadata", tags=["meta"])131def metadata():132 """Metadata endpoint — openenv validate expects name + description."""133 return {134 "name": "sql-data-quality-agent",135 "description": (136 "An OpenEnv environment where an AI agent acts as a Data Quality Engineer. "137 "Given a dirty SQLite database (NULL values, duplicate rows, type errors, "138 "foreign-key violations), the agent issues SQL statements to bring the "139 "dataset to a target quality threshold."140 ),141 "version": "1.0.0",142 "author": "Vivek Kumar Maurya",143 }144 145 146@app.get("/schema", tags=["meta"])147def schema():148 """Schema endpoint — openenv validate expects action, observation, state JSON schemas."""149 return {150 "action": DataQualityAction.model_json_schema(),151 "observation": DataQualityObservation.model_json_schema(),152 "state": DataQualityState.model_json_schema(),153 }154 155 156@app.get("/tasks", tags=["meta"])157def get_tasks():158 """List all available tasks with metadata."""159 return {"tasks": list_tasks()}160 161 162@app.post("/reset", tags=["environment"])163async def reset(request: Request):164 """165 Start a new episode for the given task.166 Returns the initial observation (table schema, sample rows, quality report).167 Accepts: full JSON body, partial body, empty body, or NO body at all (uses defaults).168 """169 task_id = "null_patrol"170 seed = 42171 try:172 body = await request.body()173 if body:174 data = await request.json()175 task_id = data.get("task_id", task_id)176 seed = int(data.get("seed", seed))177 except Exception:178 pass # No/invalid body — use defaults179 try:180 obs = env.reset(task_id=task_id, seed=seed)181 # Serialize and clamp all score/ratio/reward fields182 response_data = obs.model_dump()183 response_data = _clamp_response_scores(response_data)184 return JSONResponse(content=response_data)185 except ValueError as e:186 raise HTTPException(status_code=400, detail=str(e))187 except Exception as e:188 raise HTTPException(status_code=500, detail=f"Internal error: {str(e)}")189 190 191@app.post("/step", tags=["environment"])192def step(request: StepRequest):193 """194 Execute one SQL action.195 Returns the new observation, reward, done flag, and episode info.196 """197 try:198 action = DataQualityAction(sql=request.sql, rationale=request.rationale)199 obs, reward, done, info = env.step(action)200 # Final safety: clamp reward at the API boundary201 reward = clamp_score(reward)202 # Also clamp cumulative_reward in info203 if "cumulative_reward" in info:204 info["cumulative_reward"] = clamp_score(info["cumulative_reward"])205 # Serialize and clamp all score/ratio/reward fields206 response_data = {207 "observation": obs.model_dump(),208 "reward": clamp_score(reward),209 "done": done,210 "info": info,211 }212 response_data = _clamp_response_scores(response_data)213 return JSONResponse(content=response_data)214 except RuntimeError as e:215 raise HTTPException(status_code=400, detail=str(e))216 except Exception as e:217 raise HTTPException(status_code=500, detail=f"Internal error: {str(e)}")218 219 220@app.get("/state", tags=["environment"])221def state():222 """Return current episode metadata and state."""223 try:224 st = env.state()225 response_data = st.model_dump()226 response_data = _clamp_response_scores(response_data)227 return JSONResponse(content=response_data)228 except RuntimeError as e:229 raise HTTPException(status_code=400, detail=str(e))230 except Exception as e:231 raise HTTPException(status_code=500, detail=f"Internal error: {str(e)}")232 233 234# ---------------------------------------------------------------------------235# Entry point236# ---------------------------------------------------------------------------237 238if __name__ == "__main__":239 import uvicorn240 port = int(os.environ.get("PORT", 7860))241 print(f"Starting SQL Data Quality Agent on port {port} ...")242 uvicorn.run("app:app", host="0.0.0.0", port=port, reload=False)243 