TestCaseGenerator_env / inference.py
Vidhikoul's picture
Upload folder using huggingface_hub
7769e75 verified
Raw
History Blame Contribute Delete
5.93 kB
import asyncio
import os
import textwrap
from typing import List, Optional
from openai import OpenAI
# from .UnitTestCaseGenerator_environment import (
# UnittestcasegeneratorEnvironment,
# UnittestcasegeneratorAction,
# )
from client import UnittestcasegeneratorEnv, UnittestcasegeneratorAction
# ── CONFIG ─────────────────────────────────────────────
API_KEY = os.getenv("HF_TOKEN") or os.getenv("API_KEY")
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 = "unit_test_env"
MAX_STEPS = 1
SUCCESS_SCORE_THRESHOLD = 0.5
DIFFICULTIES = ["easy", "medium", "hard"]
# ── PROMPT ─────────────────────────────────────────────
SYSTEM_PROMPT = textwrap.dedent(
"""
You are an expert Java developer. You write JUnit 5 unit tests.
Rules:
- ALWAYS read the source code carefully before writing tests
- ALWAYS use the exact class name specified in the task hint
- NEVER write tests for a different class than what is given
- Use @Test annotation on every test method
- Always import: import org.junit.jupiter.api.Test;
- Always import: import static org.junit.jupiter.api.Assertions.*;
- Use assertEquals, assertTrue, assertFalse, assertThrows
- Reply with ONLY Java code, no explanation
"""
).strip()
# ── LOGGING (STRICT FORMAT) ─────────────────────────────
def log_start(task: str, env: str, model: str):
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]):
error_val = error if error else "null"
done_str = "true" if done else "false"
print(
f"[STEP] step={step} action={action[:100]} reward={reward} done={done_str} error={error_val}",
flush=True,
)
def log_end(success: bool, steps: int, score: float, rewards: List[float]):
# Update 0.0 and 1.0 to be 0 and 1 without decimal places
for i, r in enumerate(rewards, 1):
if abs(r - 0.0) < 1e-6:
rewards[i - 1] = 0
elif abs(r - 1.0) < 1e-6:
rewards[i - 1] = 1
rewards_str = ",".join(f"{r}" for r in rewards)
success_str = "true" if success else "false"
print(
f"[END] success={success_str} steps={steps} score={score} rewards={rewards_str}",
flush=True,
)
# ── MODEL CALL ─────────────────────────────────────────────
def get_tests_from_model(
client: OpenAI,
source_code: str,
task_hint: str,
feedback: Optional[str],
) -> str:
prompt = f"""
{task_hint}
Source code:
{source_code}
Feedback: {feedback or "None"}
Write JUnit 5 tests:
"""
try:
completion = client.chat.completions.create(
model=MODEL_NAME,
messages=[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": prompt},
],
temperature=0.3,
max_tokens=800,
)
text = (completion.choices[0].message.content or "").strip()
return text.replace("```java", "").replace("```", "").strip()
except Exception:
return "public class PlaceholderTest {}"
# ── MAIN ─────────────────────────────────────────────
async def main():
client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
from client import UnittestcasegeneratorEnv
all_rewards = []
success = False
score = 0
for difficulty in DIFFICULTIES:
# env = UnittestcasegeneratorEnv(base_url="http://localhost:8000")
with UnittestcasegeneratorEnv(base_url="http://localhost:8000").sync() as env:
rewards: List[float] = []
steps_taken = 0
feedback = None
log_start(task=difficulty, env=BENCHMARK, model=MODEL_NAME)
try:
result = env.reset(difficulty=difficulty)
source_code = result.observation.source_code
task_hint = result.observation.task_hint
for step in range(1, MAX_STEPS + 1):
if result.observation.done:
break
action = get_tests_from_model(client, source_code, task_hint, feedback)
result = env.step(UnittestcasegeneratorAction(test_code=action))
reward = result.observation.reward or 0.0
done = result.observation.done
error = getattr(result, "error", None)
feedback = f"Passed {result.observation.passed}/{result.observation.total}"
rewards.append(reward)
steps_taken = step
log_step(step, action, reward, done, error)
if done:
break
score = max(rewards) if rewards else 0.0
_EPS = 0.001
score = min(max(score, _EPS), 1.0 - _EPS)
success = score >= SUCCESS_SCORE_THRESHOLD
finally:
try:
if hasattr(env, "close"):
env.close()
except Exception:
pass
log_end(success, steps_taken, score, rewards)
all_rewards.append(score)
# ── RUN ─────────────────────────────────────────────
if __name__ == "__main__":
asyncio.run(main())