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())