git-conflict-resolver / baseline.py
Bharath Vikas Tadepalli
Fix: update baseline to follow OpenEnv checklist requirements
63803bf
Raw
History Blame Contribute Delete
4.31 kB
"""
Baseline inference script for Git Conflict Resolution Environment.
Follows OpenEnv submission requirements exactly:
- Uses API_BASE_URL, MODEL_NAME, HF_TOKEN environment variables
- Uses OpenAI client configured via these variables
- Stdout logs follow START/STEP/END structured format
"""
import os
import json
import sys
from environment import make_env, Action
from openai import OpenAI
# Environment variables — API_BASE_URL and MODEL_NAME have defaults, HF_TOKEN does not
API_BASE_URL = os.getenv("API_BASE_URL", "https://api.groq.com/openai/v1")
MODEL_NAME = os.getenv("MODEL_NAME", "llama-3.3-70b-versatile")
HF_TOKEN = os.getenv("HF_TOKEN")
# Optional — used if from_docker_image() is called
LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
# Configure OpenAI client via API_BASE_URL and HF_TOKEN
client = OpenAI(
base_url=API_BASE_URL,
api_key=HF_TOKEN or os.getenv("GROQ_API_KEY", ""),
)
SYSTEM_PROMPT = """\
You are an expert software engineer specializing in resolving Git merge conflicts.
You will be given files containing Git conflict markers (<<<<<<<, =======, >>>>>>>).
Your job is to resolve each conflict by producing clean, correct code.
Rules:
1. Remove ALL conflict markers from your output
2. Merge the changes intelligently — preserve all intended features from both branches
3. The resolved code must be syntactically valid Python
4. Return ONLY a JSON object with this structure:
{
"resolved_files": {
"filename.py": "...full resolved content..."
}
}
No explanation, no markdown, no backticks. Pure JSON only.
"""
def format_prompt(obs: dict) -> str:
lines = [
f"Task: {obs['task_description']}",
f"Hint: {obs['test_cases_hint']}",
"",
"Conflicted files to resolve:",
]
for fname, content in obs["conflicted_files"].items():
lines.append(f"\n--- {fname} ---\n{content}")
lines.append("\nReturn the resolved files as JSON.")
return "\n".join(lines)
def run_task(task_id: str) -> float:
env = make_env(task_id)
obs = env.reset()
print(f"START task_id={task_id}")
best_score = 0.0
for attempt in range(3):
prompt = format_prompt(obs.model_dump())
try:
response = client.chat.completions.create(
model=MODEL_NAME,
messages=[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
temperature=0.2,
max_tokens=2000,
)
raw = response.choices[0].message.content.strip()
# Strip markdown code fences if present
if raw.startswith("```"):
raw = raw.split("```")[1]
if raw.startswith("json"):
raw = raw[4:]
raw = raw.strip()
data = json.loads(raw)
resolved_files = data.get("resolved_files", {})
except json.JSONDecodeError as e:
print(f"STEP task_id={task_id} attempt={attempt+1} score=0.0 error=json_parse")
continue
except Exception as e:
print(f"STEP task_id={task_id} attempt={attempt+1} score=0.0 error=api_error")
break
action = Action(resolved_files=resolved_files)
obs, reward, done, info = env.step(action)
score = info["score"]
best_score = max(best_score, score)
print(f"STEP task_id={task_id} attempt={attempt+1} score={score:.2f} feedback={info['feedback']}")
if done:
break
print(f"END task_id={task_id} score={best_score:.2f}")
return best_score
def main():
tasks = ["easy", "medium", "hard"]
scores = {}
for task_id in tasks:
score = run_task(task_id)
scores[task_id] = score
avg = sum(scores.values()) / len(scores)
print(f"START summary")
for task_id, score in scores.items():
print(f"STEP task_id={task_id} final_score={score:.2f}")
print(f"END summary average_score={avg:.2f}")
with open("baseline_scores.json", "w") as f:
json.dump({
"scores": scores,
"average": avg,
"model": MODEL_NAME,
"api_base_url": API_BASE_URL,
}, f, indent=2)
if __name__ == "__main__":
main()