Team Ai
Apppublic

MauryaVivek/sql-data-quality-agent

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