Data-cleaning / main.py
Mihir Mungara
graders result clampped between 0 and 1
747aba3
Raw
History Blame Contribute Delete
11.7 kB
import sys
import os
sys.path.insert(0, os.path.dirname(__file__))
from fastapi import FastAPI, HTTPException
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
import os
from fastapi.middleware.cors import CORSMiddleware
from typing import Dict, Any, Optional
from pydantic import BaseModel, Field
from models import Action, StepResult, Reward
from environment import DataCleaningEnv
# ─────────────────────────────────────────
# App Setup
# ─────────────────────────────────────────
app = FastAPI(
title="Data Cleaning OpenEnv",
description=(
"An OpenEnv-compliant environment where AI agents "
"learn to clean messy real-world datasets step by step."
),
version="1.0.0"
)
# Serve UI
os.makedirs("static", exist_ok=True)
app.mount("/static", StaticFiles(directory="static"), name="static")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
# ─────────────────────────────────────────
# One environment instance per task
# ─────────────────────────────────────────
VALID_TASKS = [
"easy_dedup_rename",
"medium_missing_dtype",
"hard_full_pipeline",
"expert_sales_pipeline"
]
envs: Dict[str, DataCleaningEnv] = {
task_id: DataCleaningEnv(task_id=task_id)
for task_id in VALID_TASKS
}
# Tracks active task for generic OpenEnv endpoints.
current_task_id = VALID_TASKS[0]
class ResetRequest(BaseModel):
task_id: str = Field(default=VALID_TASKS[0])
class GenericStepRequest(BaseModel):
task_id: Optional[str] = Field(default=None)
operation: str
parameters: Dict[str, Any] = Field(default_factory=dict)
def get_env(task_id: str) -> DataCleaningEnv:
if task_id not in envs:
raise HTTPException(
status_code=404,
detail=(
f"Task '{task_id}' not found. "
f"Valid tasks: {VALID_TASKS}"
)
)
return envs[task_id]
# ─────────────────────────────────────────
# ROUTES
# ─────────────────────────────────────────
@app.get("/ui")
def ui():
return FileResponse("static/index.html")
@app.get("/")
def root():
return {
"name": "Data Cleaning OpenEnv",
"version": "1.0.0",
"status": "running",
"tasks": VALID_TASKS,
"endpoints": {
"reset": ["POST /reset", "POST /reset/{task_id}"],
"step": ["POST /step", "POST /step/{task_id}"],
"state": ["GET /state", "GET /state/{task_id}"],
"tasks": "GET /tasks",
"health": "GET /health",
"docs": "GET /docs"
}
}
@app.get("/health")
def health():
return {
"status": "ok",
"tasks_loaded": len(envs)
}
@app.get("/tasks")
def list_tasks():
return {
"tasks": [
{
"task_id": "easy_dedup_rename",
"difficulty": "easy",
"description": (
"Remove duplicate rows and rename columns "
"to snake_case in an employee dataset."
),
"max_steps": 10,
"operations": ["remove_duplicates", "rename_columns", "finish"]
},
{
"task_id": "medium_missing_dtype",
"difficulty": "medium",
"description": (
"Fill missing values using correct strategies "
"and fix wrong data types in a customer dataset."
),
"max_steps": 15,
"operations": ["fill_missing", "fix_dtype", "finish"]
},
{
"task_id": "hard_full_pipeline",
"difficulty": "hard",
"description": (
"Run a full cleaning pipeline: remove duplicates, "
"fill missing values, fix dtypes, remove outliers, "
"and validate schema on an orders dataset."
),
"max_steps": 20,
"operations": [
"remove_duplicates", "fill_missing", "fix_dtype",
"remove_outliers", "validate_schema", "finish"
]
},
{
"task_id": "expert_sales_pipeline",
"difficulty": "expert",
"description": (
"Expert level: Full sales data cleaning pipeline "
"requiring correct order of operations including "
"case standardization, outlier removal, and schema validation."
),
"max_steps": 25,
"operations": [
"remove_duplicates", "fill_missing", "fix_dtype",
"remove_outliers", "rename_columns",
"validate_schema", "finish"
]
}
]
}
@app.post("/reset/{task_id}")
def reset(task_id: str):
"""Reset environment and start fresh episode."""
global current_task_id
env = get_env(task_id)
try:
current_task_id = task_id
result = env.reset()
return result.dict()
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"Reset failed: {str(e)}"
)
@app.post("/reset")
def reset_generic(payload: Optional[ResetRequest] = None):
"""OpenEnv-compatible reset endpoint."""
task_id = payload.task_id if payload else VALID_TASKS[0]
return reset(task_id)
@app.post("/step/{task_id}")
def step(task_id: str, action: Action):
"""Take one action in the environment."""
env = get_env(task_id)
if env.current_df is None:
raise HTTPException(
status_code=400,
detail="Environment not initialized. Call /reset/{task_id} first."
)
try:
result = env.step(action)
return result.dict()
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"Step failed: {str(e)}"
)
@app.post("/step")
def step_generic(payload: Dict[str, Any]):
"""OpenEnv-compatible step endpoint."""
global current_task_id
task_id = payload.get("task_id") or current_task_id
action_payload: Dict[str, Any]
# Support both body formats:
# 1) {"task_id": "...", "operation": "...", "parameters": {...}}
# 2) {"task_id": "...", "action": {"operation": "...", "parameters": {...}}}
if isinstance(payload.get("action"), dict):
action_payload = payload["action"]
else:
action_payload = payload
operation = action_payload.get("operation")
parameters = action_payload.get("parameters", {})
if not operation:
raise HTTPException(status_code=400, detail="Missing 'operation' in request body")
current_task_id = task_id
action = Action(operation=operation, parameters=parameters)
return step(task_id, action)
@app.get("/state/{task_id}")
def state(task_id: str):
"""Get current environment state."""
env = get_env(task_id)
try:
return env.state()
except Exception as e:
raise HTTPException(
status_code=500,
detail=f"State failed: {str(e)}"
)
@app.get("/state")
def state_generic(task_id: Optional[str] = None):
"""OpenEnv-compatible state endpoint."""
active_task = task_id or current_task_id
return state(active_task)
@app.get("/validate")
def validate():
"""OpenEnv spec validation endpoint."""
results = {}
for task_id in VALID_TASKS:
try:
env = DataCleaningEnv(task_id=task_id)
# Test reset
reset_result = env.reset()
assert reset_result.observation is not None
assert reset_result.reward is not None
assert reset_result.done == False
# Test step
from models import Action
action = Action(
operation="remove_duplicates",
parameters={}
)
step_result = env.step(action)
assert step_result.observation is not None
assert 0.0 < step_result.reward.total < 1.0 # strict bounds required by grader
# Test state
state_result = env.state()
assert "task_id" in state_result
results[task_id] = {
"status": "passed",
"reset": "ok",
"step": "ok",
"state": "ok",
"reward_range": f"{step_result.reward.total}"
}
except Exception as e:
results[task_id] = {
"status": "failed",
"error": str(e)
}
all_passed = all(r["status"] == "passed" for r in results.values())
return {
"openenv_valid": all_passed,
"tasks": results
}
# In memory leaderboard
leaderboard_data = []
@app.post("/leaderboard/submit")
def submit_score(entry: Dict[str, Any]):
"""Submit a score to the leaderboard."""
required = ["model_name", "task_id", "score"]
for field in required:
if field not in entry:
raise HTTPException(
status_code=400,
detail=f"Missing field: {field}"
)
if not 0.0 < float(entry["score"]) < 1.0:
raise HTTPException(
status_code=400,
detail="Score must be strictly between 0.0 and 1.0 (not 0.0 or 1.0)"
)
leaderboard_data.append({
"model_name": entry["model_name"],
"task_id": entry["task_id"],
"score": round(max(0.0001, min(0.9999, float(entry["score"]))), 4),
"steps": entry.get("steps", 0),
"timestamp": __import__("datetime").datetime.utcnow().isoformat()
})
return {"status": "submitted", "entry": leaderboard_data[-1]}
@app.get("/leaderboard")
def get_leaderboard():
"""Get current leaderboard rankings."""
if not leaderboard_data:
# Return baseline scores
return {
"leaderboard": [
{
"rank": 1,
"model_name": "gpt-4o-mini (baseline)",
"easy_score": 0.9999,
"medium_score": 0.6643,
"hard_score": 0.8386,
"avg_score": 0.8343
}
],
"total_submissions": 1
}
# Group by model
from collections import defaultdict
model_scores = defaultdict(dict)
for entry in leaderboard_data:
model_scores[entry["model_name"]][entry["task_id"]] = entry["score"]
ranked = []
for model, scores in model_scores.items():
avg = sum(scores.values()) / len(scores) if scores else 0
ranked.append({
"model_name": model,
"scores": scores,
"avg_score": round(max(0.0001, min(0.9999, avg)), 4)
})
ranked.sort(key=lambda x: x["avg_score"], reverse=True)
for i, r in enumerate(ranked):
r["rank"] = i + 1
return {
"leaderboard": ranked,
"total_submissions": len(leaderboard_data)
}