File size: 3,854 Bytes
a5c9fd4
 
50a0917
a5c9fd4
 
50a0917
a5c9fd4
 
 
 
 
 
 
 
 
 
 
50a0917
 
 
 
 
 
 
 
a5c9fd4
4a63c86
 
 
 
 
 
 
a5c9fd4
 
4a63c86
a5c9fd4
50a0917
4a63c86
 
a5c9fd4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
50a0917
a5c9fd4
 
 
4a63c86
 
a5c9fd4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d5e0274
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a5c9fd4
 
 
 
 
 
 
 
 
9a593a7
 
 
 
 
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
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
from fastapi import FastAPI
from typing import Any, Dict, Tuple
import os

from src.reward import compute_reward
from src.tasks import GRADERS, TASKS

app = FastAPI()


class CodeGuardEnv:
    def __init__(self) -> None:
        self.state: Dict[str, Any] = {}
        self.current_step: int = 0
        self.max_steps: int = 50
        self.threshold: float = 0.95
        self.done: bool = False
        requested_task = os.getenv("TASK", "easy").strip().lower()
        self.task_key: str = requested_task if requested_task in GRADERS else "easy"

    def _get_task_score(self, action: str) -> float:
        grader = GRADERS[self.task_key]
        base_score = float(grader(action))
        # Strict score interval for validator compatibility.
        return max(0.01, min(0.99, base_score))

    def _get_all_task_scores(self, action: str) -> Dict[str, float]:
        scores: Dict[str, float] = {}
        for key, grader in GRADERS.items():
            score = float(grader(action))
            scores[key] = max(0.01, min(0.99, score))
        return scores

    def reset(self) -> Dict[str, Any]:
        self.state = {
            "score": 0.01,
            "history": [],
            "task": TASKS[self.task_key],
            "tasks": list(TASKS.values()),
            "task_scores": {k: 0.01 for k in GRADERS.keys()},
        }
        self.current_step = 0
        self.done = False
        return self.state

    def _is_valid_action(self, action: str) -> bool:
        if not isinstance(action, str):
            return False
        if len(action.strip()) == 0:
            return False
        if len(action) > 1000:
            return False
        return True

    def step(self, action: str) -> Tuple[Dict[str, Any], float, bool, Dict[str, Any]]:
        if self.done:
            return self.state, 0.0, True, {"error": "episode_done"}

        self.current_step += 1

        info: Dict[str, Any] = {"error": None}

        if not self._is_valid_action(action):
            self.done = True
            return self.state, -1.0, True, {"error": "invalid_action"}

        base_score: float = self._get_task_score(action)

        reward: float = compute_reward(self.state, action, base_score)

        self.state["score"] = max(0.01, min(0.99, base_score))
        self.state["task_scores"] = self._get_all_task_scores(action)
        self.state["history"].append(
            {
                "step": self.current_step,
                "action": action,
                "reward": reward,
            }
        )

        if reward <= -2.0:
            self.done = True
        elif reward >= self.threshold:
            self.done = True
        elif self.current_step >= self.max_steps:
            self.done = True

        return self.state, reward, self.done, info


# --- FastAPI Routes ---

env_instance = CodeGuardEnv()


@app.get("/")
def health_check() -> Dict[str, str]:
    return {"status": "ok"}


@app.post("/reset")
def reset() -> Dict[str, Any]:
    return env_instance.reset()


@app.get("/tasks")
def tasks() -> Dict[str, Any]:
    return {
        "count": len(TASKS),
        "tasks": list(TASKS.values()),
        "graders": sorted(GRADERS.keys()),
    }


@app.post("/grade")
def grade(task_id: str, candidate_code: str) -> Dict[str, Any]:
    key = task_id.strip().lower()
    if key not in GRADERS:
        return {"error": "unknown_task", "task_id": task_id}

    score = float(GRADERS[key](candidate_code))
    score = max(0.01, min(0.99, score))
    return {"task_id": key, "score": score}


@app.post("/step")
def step(action: str) -> Dict[str, Any]:
    state, reward, done, info = env_instance.step(action)
    return {
        "state": state,
        "reward": reward,
        "done": done,
        "info": info,
    }


@app.get("/state")
def state() -> Dict[str, Any]:
    return env_instance.state