File size: 7,395 Bytes
b43aff5
529657f
b43aff5
c7b11e5
529657f
b43aff5
 
 
 
 
 
529657f
b43aff5
 
 
 
 
529657f
b43aff5
 
 
c7b11e5
b43aff5
 
 
 
c7b11e5
 
 
 
 
b43aff5
 
 
 
 
 
 
c7b11e5
b43aff5
 
 
c7b11e5
 
 
b43aff5
 
 
 
 
 
 
 
 
 
 
 
 
 
c7b11e5
 
 
 
 
b43aff5
c7b11e5
 
 
 
 
 
 
 
b43aff5
c7b11e5
 
 
 
 
 
 
 
 
b43aff5
c7b11e5
 
b43aff5
c7b11e5
 
b43aff5
c7b11e5
b43aff5
 
 
 
 
 
c7b11e5
b43aff5
 
 
 
 
 
 
c7b11e5
 
b43aff5
 
 
529657f
c7b11e5
529657f
 
 
b43aff5
529657f
c7b11e5
 
 
 
 
 
 
 
 
 
 
 
b43aff5
529657f
 
b43aff5
 
c7b11e5
529657f
 
 
c7b11e5
b43aff5
529657f
b43aff5
 
 
 
 
c7b11e5
529657f
c7b11e5
 
b43aff5
c7b11e5
 
 
 
529657f
b43aff5
c7b11e5
 
 
 
 
 
 
 
 
b43aff5
c7b11e5
 
 
 
529657f
 
c7b11e5
 
 
 
 
529657f
c7b11e5
529657f
 
c7b11e5
 
 
 
 
 
 
529657f
 
 
 
c7b11e5
 
529657f
 
 
c7b11e5
 
 
b43aff5
529657f
c7b11e5
 
b43aff5
 
 
529657f
 
b43aff5
 
 
 
 
529657f
c7b11e5
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
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
"""
Inference Script - ConfigDebugEnv
===================================
STDOUT FORMAT - emits [START], [STEP], [END] PER TASK:
    [START] task=<task_id> env=<benchmark> model=<model_name>
    [STEP]  step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>
    [END]   success=<true|false> steps=<n> score=<score> rewards=<r1,r2,...,rn>
"""

import asyncio
import os
import sys
import textwrap
from typing import List, Optional

from openai import OpenAI

MAX_STEPS_PER_TASK = 5
TEMPERATURE = 0.1
MAX_TOKENS = 2000

SYSTEM_PROMPT = textwrap.dedent("""
    You are an expert DevOps engineer specializing in configuration file debugging.
    You will be given a broken configuration file and must fix ALL bugs in it.
    Return ONLY the fixed configuration file content.
    No explanations, no markdown formatting, no code blocks. Just the raw fixed configuration.
""").strip()


def clamp(v: float) -> float:
    return max(0.01, min(0.99, v))


def log_start(task: str, env: str, model: str) -> None:
    print(f"[START] task={task} env={env} model={model}", flush=True)


def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:
    print(f"[STEP] step={step} action={action} reward={clamp(reward):.2f} done={str(done).lower()} error={error or 'null'}", flush=True)


def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:
    s = clamp(score)
    rs = ",".join(f"{clamp(r):.2f}" for r in rewards) if rewards else "0.01"
    print(f"[END] success={str(success).lower()} steps={steps} score={s:.3f} rewards={rs}", flush=True)


def strip_code_blocks(text: str) -> str:
    text = text.strip()
    if text.startswith("```"):
        lines = text.split("\n")
        if lines[-1].strip() == "```":
            lines = lines[1:-1]
        else:
            lines = lines[1:]
        text = "\n".join(lines)
    return text


def get_obs_field(data: dict, field: str, default=None):
    """Get field from response - handles both nested and flat formats."""
    obs = data.get("observation", data)
    return obs.get(field, data.get(field, default))


def get_model_fix(client, obs_data, step, history, model_name):
    file_type = get_obs_field(obs_data, "file_type", "config")
    desc = get_obs_field(obs_data, "task_description", "")
    difficulty = get_obs_field(obs_data, "difficulty", "")
    num_bugs = get_obs_field(obs_data, "num_bugs", 0)
    bugs_found = get_obs_field(obs_data, "bugs_found_so_far", 0)
    error_msg = get_obs_field(obs_data, "error_message", "")
    broken = get_obs_field(obs_data, "broken_config", "")

    history_block = "\n".join(history[-4:]) if history else "None"
    prompt = f"""Fix the following broken {file_type} configuration file.

Task: {desc}
Difficulty: {difficulty}
Number of bugs to find: {num_bugs}
Bugs fixed so far: {bugs_found}
Error message: {error_msg}
Step: {step}

Previous attempts:
{history_block}

Broken configuration:
{broken}

Return ONLY the fixed configuration file content."""

    try:
        completion = client.chat.completions.create(
            model=model_name,
            messages=[
                {"role": "system", "content": SYSTEM_PROMPT},
                {"role": "user", "content": prompt},
            ],
            temperature=TEMPERATURE,
            max_tokens=MAX_TOKENS,
            stream=False,
        )
        text = (completion.choices[0].message.content or "").strip()
        return strip_code_blocks(text) if text else ""
    except Exception as e:
        print(f"[DEBUG] LLM error: {e}", flush=True)
        return ""


class HTTPEnvClient:
    def __init__(self, base_url):
        import httpx
        self.base_url = base_url.rstrip("/")
        self.http = httpx.AsyncClient(timeout=60.0)

    async def reset(self):
        r = await self.http.post(f"{self.base_url}/reset")
        r.raise_for_status()
        return r.json()

    async def step(self, fixed_config: str):
        """Send step with correct OpenEnv format: {"action": {"fixed_config": "..."}}"""
        r = await self.http.post(
            f"{self.base_url}/step",
            json={"action": {"fixed_config": fixed_config}},
        )
        r.raise_for_status()
        return r.json()

    async def close(self):
        await self.http.aclose()


async def main():
    api_key = os.getenv("HF_TOKEN") or os.getenv("API_KEY") or ""
    api_base_url = os.getenv("API_BASE_URL") or "https://router.huggingface.co/v1"
    model_name = os.getenv("MODEL_NAME") or "Qwen/Qwen2.5-72B-Instruct"
    benchmark = "config_debug_env"

    client = OpenAI(base_url=api_base_url, api_key=api_key)
    env_url = sys.argv[1] if len(sys.argv) > 1 else "http://localhost:7860"
    env = HTTPEnvClient(env_url)

    try:
        result = await env.reset()
        current_task = get_obs_field(result, "task_id", "unknown")
        task_step = 0
        task_rewards = []
        task_history = []

        log_start(task=current_task, env=benchmark, model=model_name)

        for global_step in range(1, 50):
            fixed_config = get_model_fix(client, result, task_step + 1, task_history, model_name)
            task_step += 1

            try:
                step_result = await env.step(fixed_config)
            except Exception as step_err:
                print(f"[DEBUG] Step failed: {step_err}", flush=True)
                log_step(task_step, f"fix({current_task})", 0.01, False, str(step_err))
                task_rewards.append(0.01)
                # End this task on step failure
                log_end(False, task_step, 0.01, task_rewards)
                break

            reward = get_obs_field(step_result, "reward", 0.01)
            if reward is None:
                reward = 0.01
            reward = clamp(float(reward))
            task_rewards.append(reward)

            new_task = get_obs_field(step_result, "task_id", current_task)
            is_done = get_obs_field(step_result, "done", False)
            error = get_obs_field(step_result, "error_message", None)
            if error == "All checks passed!":
                error = None

            log_step(task_step, f"fix({current_task})", reward, bool(is_done), error)
            task_history.append(f"Step {task_step}: reward {reward:.2f}")

            # Detect task transition
            task_changed = (new_task != current_task) and (new_task != "unknown")

            if task_changed or is_done:
                # End current task
                task_score = clamp(sum(task_rewards) / len(task_rewards)) if task_rewards else 0.01
                log_end(task_score >= 0.5, task_step, task_score, task_rewards)

                if is_done:
                    break

                # Start next task
                current_task = new_task
                task_step = 0
                task_rewards = []
                task_history = []
                log_start(task=current_task, env=benchmark, model=model_name)

            result = step_result

    except Exception as e:
        print(f"[DEBUG] Fatal error: {e}", flush=True)
        log_end(False, 0, 0.01, [0.01])
    finally:
        try:
            await env.close()
        except Exception:
            pass


if __name__ == "__main__":
    try:
        asyncio.run(main())
    except Exception:
        print("[END] success=false steps=0 score=0.01 rewards=0.01", flush=True)