File size: 2,969 Bytes
022e674
 
1dd8866
022e674
1dd8866
022e674
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
from fastapi import FastAPI
from env.environment import DataCleaningEnv

app = FastAPI()

TASKS = [
    {
        "id": "task1",
        "name": "Basic Data Cleaning",
        "difficulty": "easy",
        "objective": "Handle missing values, remove duplicates, fix simple type issues, and preserve dataset structure.",
        "grader": "env.graders.task1_grader.grade_task1",
        "grader_path": "env/graders/task1_grader.py",
        "grader_function": "grade_task1",
        "task_number": 1
    },
    {
        "id": "task2",
        "name": "Intermediate Data Cleaning",
        "difficulty": "medium",
        "objective": "Normalize numeric columns, reduce correlation issues, and preserve useful records.",
        "grader": "env.graders.task2_grader.grade_task2",
        "grader_path": "env/graders/task2_grader.py",
        "grader_function": "grade_task2",
        "task_number": 2
    },
    {
        "id": "task3",
        "name": "Full Data Cleaning Pipeline",
        "difficulty": "hard",
        "objective": "Perform end-to-end cleaning with null handling, duplicate removal, structure preservation, and minimal data loss.",
        "grader": "env.graders.task3_grader.grade_task3",
        "grader_path": "env/graders/task3_grader.py",
        "grader_function": "grade_task3",
        "task_number": 3
    }
]

TASK_BY_ID = {task["id"]: task for task in TASKS}

env = DataCleaningEnv(task=1)

def parse_task(payload):
    payload = payload or {}
    task_value = payload.get("task", payload.get("task_id", "task1"))

    if isinstance(task_value, int):
        if task_value in [1, 2, 3]:
            return task_value
        return 1

    task_value = str(task_value).lower()

    if task_value in TASK_BY_ID:
        return TASK_BY_ID[task_value]["task_number"]

    if task_value in ["1", "2", "3"]:
        return int(task_value)

    return 1

@app.get("/")
def root():
    return {
        "name": "data-cleaning-env-final",
        "tasks": TASKS,
        "message": "OpenEnv data cleaning environment with 3 tasks and graders."
    }

@app.get("/tasks")
def tasks():
    return {
        "count": len(TASKS),
        "tasks": TASKS
    }

@app.post("/reset")
def reset(payload: dict = None):
    global env

    task_number = parse_task(payload)
    env = DataCleaningEnv(task=task_number)
    obs = env.reset()

    return {
        "task": f"task{task_number}",
        "task_number": task_number,
        "observation": obs,
        "available_tasks": TASKS
    }

@app.post("/step")
def step(action: dict):
    obs, reward, done, info = env.step(action)

    return {
        "obs": obs,
        "reward": reward,
        "done": done,
        "info": info
    }

@app.get("/state")
def state():
    current_state = env.state()
    current_state["task"] = f"task{env.task}"
    current_state["available_tasks"] = TASKS
    return current_state

@app.post("/submit")
def submit():
    return env.submit_cleaned_data(env.dirty_df)