File size: 4,551 Bytes
a5c9fd4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9a593a7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a5c9fd4
 
 
 
9a593a7
a5c9fd4
 
 
9a593a7
 
 
a5c9fd4
 
 
9a593a7
a5c9fd4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9a593a7
 
 
a5c9fd4
 
 
 
 
 
9a593a7
 
 
 
a5c9fd4
 
 
 
 
9a593a7
a5c9fd4
 
 
 
 
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
141
142
143
144
145
146
147
148
149
150
151
152
import os
from typing import Any, Dict, List, Optional

from openai import OpenAI

from src.env import CodeGuardEnv


def _get_env_var(name: str, default: Optional[str] = None, required: bool = False) -> str:
    value = os.getenv(name, default)
    if required and not value:
        raise ValueError(f"Missing required environment variable: {name}")
    return value or ""


def _format_bool(value: bool) -> str:
    return "true" if value else "false"


def _safe_str(value: Optional[str]) -> str:
    if value is None:
        return "null"
    return str(value)


def _get_api_key() -> str:
    # Submission-required variable first, with compatibility fallbacks.
    for name in ("HF_TOKEN", "API_KEY", "GEMINI_API_KEY", "OPENAI_API_KEY"):
        value = os.getenv(name)
        if value:
            return value
    raise ValueError(
        "Missing API key. Set one of: HF_TOKEN, API_KEY, GEMINI_API_KEY, OPENAI_API_KEY"
    )


def _resolve_base_url(api_key: str) -> str:
    explicit_base_url = os.getenv("API_BASE_URL") or os.getenv("BASE_URL")
    if explicit_base_url:
        return explicit_base_url

    # Default should reflect active inference setup for submission.
    # Keep provider-aware fallback behavior.
    if api_key.startswith("hf_"):
        return "https://router.huggingface.co/v1"
    if api_key.startswith("AIza"):
        return "https://generativelanguage.googleapis.com/v1beta/openai"
    return "https://router.huggingface.co/v1"


def _resolve_model(base_url: str) -> str:
    explicit_model = os.getenv("MODEL_NAME") or os.getenv("MODEL")
    if explicit_model:
        return explicit_model

    if "router.huggingface.co" in base_url:
        return "Qwen/Qwen2.5-72B-Instruct"
    if "generativelanguage.googleapis.com" in base_url:
        return "gemini-1.5-flash"
    return "gpt-4.1-mini"


def _clamp_score(value: float) -> float:
    return max(0.0, min(1.0, value))


def main() -> None:
    success: bool = False
    rewards: List[float] = []
    step_count: int = 0
    score: float = 0.0

    try:
        # --- ENV VARS ---
        api_key: str = _get_api_key()
        base_url: str = _resolve_base_url(api_key)
        model: str = _resolve_model(base_url)
        task_name: str = _get_env_var("TASK", "easy")

        # --- CLIENT INIT ---
        client = OpenAI(base_url=base_url, api_key=api_key)

        # --- ENV INIT ---
        env = CodeGuardEnv()
        state: Dict[str, Any] = env.reset()

        # --- START LOG ---
        print(f"[START] task={task_name} env=codeguard model={model}")

        done: bool = False

        while not done:
            step_count += 1

            # --- LLM CALL ---
            try:
                response = client.chat.completions.create(
                    model=model,
                    messages=[
                        {"role": "system", "content": "You are a code-fixing agent."},
                        {"role": "user", "content": str(state)},
                    ],
                    temperature=0.0,
                )
                action: str = response.choices[0].message.content.strip()
                error_msg: Optional[str] = None
            except Exception as e:
                action = ""
                error_msg = str(e)

            # --- ENV STEP ---
            next_state, reward, done, info = env.step(action)

            rewards.append(reward)

            # Merge error sources
            error_output: Optional[str] = error_msg or info.get("error")

            # --- STEP LOG ---
            print(
                f"[STEP] step={step_count} action={action} "
                f"reward={reward:.2f} done={_format_bool(done)} "
                f"error={_safe_str(error_output)}"
            )

            state = next_state

            # Keep score in [0, 1] for validator compatibility.
            score = _clamp_score(float(state.get("score", 0.0)))

            # Success condition
            if done and reward >= env.threshold:
                success = True

        # --- END LOG ---
        reward_str = ",".join(f"{r:.2f}" for r in rewards)
        print(
            f"[END] success={_format_bool(success)} steps={step_count} "
            f"score={score:.2f} rewards={reward_str}"
        )

    except Exception as e:
        # Ensure END always prints
        _ = str(e)
        reward_str = ",".join(f"{r:.2f}" for r in rewards) if rewards else "0.00"
        print(f"[END] success=false steps={step_count} score=0.00 rewards={reward_str}")


if __name__ == "__main__":
    main()