Ar-Srivas/BitWise_CSS_env
0
1"""2Baseline inference agent for css_env.3 4Uses the official OpenEnv EnvClient SDK (via the CssEnv client class5defined in client.py at project root) to drive the environment.6 7Reads environment variables:8 - API_BASE_URL (default: https://api.openai.com/v1)9 - MODEL_NAME (default: gpt-4o-mini)10 - HF_TOKEN or API_KEY or OPENAI_API_KEY (required)11 12Optional:13 - LOCAL_IMAGE_NAME - Docker image name. If set, the env is launched via14 CssEnv.from_docker_image(LOCAL_IMAGE_NAME).15 - ENV_URL - direct base URL of an already-running css_env server16 - TASK_NAME - task1/task2/task3/task4/easy/medium/hard/all17 - MAX_STEPS - hard cap on loop steps per task18 - TEMPERATURE - LLM sampling temperature19 - MAX_TOKENS - LLM max output tokens20 21Emits the OpenEnv structured stdout format:22 [START] task=<task_name> env=<benchmark> model=<model_name>23 [STEP] step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>24 [END] success=<true|false> steps=<n> score=<0.000> rewards=<r1,r2,...,rn>25"""26 27from __future__ import annotations28 29import asyncio30import json31import os32import re33import sys34from typing import Any, Dict, List, Optional35 36from openai import OpenAI37 38from client import CssEnv39from models import CssAction40 41try:42 from server.tasks import TASKS, TASK_ORDER43except ImportError:44 from tasks import TASKS, TASK_ORDER45 46 47# Configuration48API_BASE_URL = os.getenv("API_BASE_URL", "https://api.openai.com/v1").strip()49MODEL_NAME = os.getenv("MODEL_NAME", "gpt-4o-mini").strip()50LLM_API_KEY = (51 os.getenv("HF_TOKEN")52 or os.getenv("API_KEY")53 or os.getenv("OPENAI_API_KEY")54 or ""55).strip()56 57LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME", "").strip()58ENV_URL = os.getenv("ENV_URL", "http://localhost:8000").strip()59TASK_NAME = os.getenv("TASK_NAME", "all").strip()60MAX_STEPS = max(1, int(os.getenv("MAX_STEPS", "20")))61TEMPERATURE = float(os.getenv("TEMPERATURE", "0"))62MAX_TOKENS = max(64, int(os.getenv("MAX_TOKENS", "350")))63BENCHMARK = "css_env"64SCORE_EPSILON = 1e-665MIN_SCORE_BOUND = 0.0166MAX_SCORE_BOUND = 0.9967 68SYSTEM_PROMPT = """You are an expert frontend engineer fixing CSS to match design tokens.69 70Return ONLY one JSON object with this schema:71{72 "action_type": "replace_color|fix_spacing|fix_typography|fix_contrast|add_breakpoint|remove_rule",73 "target": "string",74 "value": "string or null"75}76 77Guidelines:781. Prioritize the lowest score in observation.scores.792. Do not repeat the same action signature.803. Prefer targeted edits that change CSS and improve one dimension.814. If a previous action had no effect, choose a different action type.825. No markdown, no explanation, JSON only.83"""84 85 86def clamp01(value: float) -> float:87 try:88 numeric = float(value)89 except (TypeError, ValueError):90 numeric = MIN_SCORE_BOUND + SCORE_EPSILON91 92 if numeric <= MIN_SCORE_BOUND:93 numeric = MIN_SCORE_BOUND + SCORE_EPSILON94 if numeric >= MAX_SCORE_BOUND:95 numeric = MAX_SCORE_BOUND - SCORE_EPSILON96 97 rounded = round(numeric, 2)98 if rounded <= MIN_SCORE_BOUND:99 return 0.02100 if rounded >= MAX_SCORE_BOUND:101 return 0.98102 return float(f"{rounded:.2f}")103 104 105def log_start(task: str, env: str, model: str) -> None:106 print(f"[START] task={task} env={env} model={model}", flush=True)107 108 109def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str] = None) -> None:110 done_str = "true" if done else "false"111 error_str = error if error else "null"112 print(113 f"[STEP] step={step} action={action} reward={reward:.2f} done={done_str} error={error_str}",114 flush=True,115 )116 117 118def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:119 success_str = "true" if success else "false"120 rewards_str = ",".join(f"{r:.2f}" for r in rewards)121 print(122 f"[END] success={success_str} steps={steps} score={clamp01(score):.3f} rewards={rewards_str}",123 flush=True,124 )125 126 127def _task_description(task_name: str) -> str:128 cfg = TASKS.get(task_name, {})129 return str(cfg.get("description") or cfg.get("difficulty") or task_name)130 131 132def _select_tasks(task_name: str) -> List[str]:133 aliases = {"easy": "task1", "medium": "task2", "hard": "task3"}134 requested = aliases.get(task_name.lower(), task_name.lower())135 136 if requested == "all":137 return [name for name in TASK_ORDER if name in TASKS]138 if requested in TASKS:139 return [requested]140 141 known = ",".join(["all", "easy", "medium", "hard", *TASKS.keys()])142 raise ValueError(f"Unknown TASK_NAME '{task_name}'. Expected one of: {known}")143 144 145def _extract_selectors(css: str) -> List[str]:146 selectors: List[str] = []147 for selector_group in re.findall(r"([^{}]+)\{", css or ""):148 for raw in selector_group.split(","):149 selector = raw.strip()150 if selector and selector not in selectors:151 selectors.append(selector)152 return selectors153 154 155def _extract_colors(css: str) -> List[str]:156 return re.findall(r"#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{6})", css or "")157 158 159def _first_selector(css: str) -> str:160 selectors = _extract_selectors(css)161 return selectors[0] if selectors else ".container"162 163 164def _build_user_prompt(observation: Any, step: int, task_name: str) -> str:165 html = str(getattr(observation, "html", ""))166 css = str(getattr(observation, "css", ""))167 tokens = getattr(observation, "tokens", {}) or {}168 violations = getattr(observation, "violations", None)169 scores = getattr(observation, "scores", {}) or {}170 prev_score = getattr(observation, "score", None)171 172 html_preview = html[:3000]173 css_preview = css[:5000]174 175 return (176 f"Task: {task_name} - {_task_description(task_name)}\n"177 f"Step: {step}\n"178 f"Current score: {prev_score}\n"179 f"Scores: {json.dumps(scores, separators=(',', ':'))}\n"180 f"Violations: {json.dumps(violations)}\n"181 f"Tokens: {json.dumps(tokens, separators=(',', ':'))}\n\n"182 f"HTML:\n{html_preview}\n\n"183 f"CSS:\n{css_preview}\n"184 )185 186 187def _parse_action_json(text: str) -> Optional[Dict[str, Any]]:188 raw = (text or "").strip()189 if not raw:190 return None191 192 if "```" in raw:193 lines = raw.splitlines()194 capture: List[str] = []195 in_block = False196 for line in lines:197 if line.strip().startswith("```"):198 if in_block:199 break200 in_block = True201 continue202 if in_block:203 capture.append(line)204 if capture:205 raw = "\n".join(capture).strip()206 207 start = raw.find("{")208 end = raw.rfind("}")209 if start < 0 or end <= start:210 return None211 212 try:213 payload = json.loads(raw[start : end + 1])214 except json.JSONDecodeError:215 return None216 217 return payload if isinstance(payload, dict) else None218 219 220def _fallback_action(observation: Any) -> Dict[str, Any]:221 css = str(getattr(observation, "css", ""))222 tokens = getattr(observation, "tokens", {}) or {}223 selector = _first_selector(css)224 225 token_colors = list((tokens.get("colors") or {}).values())226 colors = _extract_colors(css)227 if colors and token_colors:228 replacement = str(token_colors[0])229 if replacement != colors[0]:230 return {231 "action_type": "replace_color",232 "target": colors[0],233 "value": replacement,234 }235 236 spacing = (tokens.get("spacing") or {}).get("md", 16)237 return {238 "action_type": "fix_spacing",239 "target": f"{selector}.margin",240 "value": f"{int(spacing)}px",241 }242 243 244def _normalize_action(action: Dict[str, Any], observation: Any) -> Dict[str, Any]:245 allowed = {246 "replace_color",247 "fix_spacing",248 "fix_typography",249 "fix_contrast",250 "add_breakpoint",251 "remove_rule",252 }253 254 safe = dict(action or {})255 action_type = str(safe.get("action_type", "")).strip()256 if action_type not in allowed:257 return _fallback_action(observation)258 259 css = str(getattr(observation, "css", ""))260 selectors = _extract_selectors(css)261 first_selector = selectors[0] if selectors else ".container"262 263 target = str(safe.get("target", "")).strip()264 value = safe.get("value", None)265 266 if action_type == "replace_color":267 colors = _extract_colors(css)268 if not target or target not in colors:269 target = colors[0] if colors else "#333333"270 if value is None or not str(value).strip():271 value = "#1a6fe0"272 return {"action_type": action_type, "target": target, "value": str(value)}273 274 if action_type in {"fix_spacing", "fix_typography"}:275 if "." not in target:276 prop = "margin" if action_type == "fix_spacing" else "font-size"277 target = f"{first_selector}.{prop}"278 if value is None or not str(value).strip():279 value = "16px"280 return {"action_type": action_type, "target": target, "value": str(value)}281 282 if action_type == "fix_contrast":283 if not target:284 target = first_selector285 if value is None or "," not in str(value):286 value = "#333333,#ffffff"287 return {"action_type": action_type, "target": target, "value": str(value)}288 289 if action_type == "add_breakpoint":290 if not re.fullmatch(r"\d+px", target):291 target = "768px"292 if value is None or "{" not in str(value):293 value = f"{first_selector} {{ width: 100%; }}"294 return {"action_type": action_type, "target": target, "value": str(value)}295 296 if action_type == "remove_rule":297 if not target:298 target = selectors[-1] if selectors else ".unused"299 return {"action_type": action_type, "target": target, "value": None}300 301 return _fallback_action(observation)302 303 304def _action_str(action: Dict[str, Any]) -> str:305 return json.dumps(action, separators=(",", ":"), ensure_ascii=True)306 307 308def _make_openai_client() -> OpenAI:309 if not LLM_API_KEY:310 raise ValueError("Missing API key. Set HF_TOKEN or API_KEY or OPENAI_API_KEY.")311 return OpenAI(base_url=API_BASE_URL, api_key=LLM_API_KEY)312 313 314def _probe_llm(client: OpenAI) -> None:315 client.chat.completions.create(316 model=MODEL_NAME,317 messages=[318 {"role": "system", "content": "Reply with exactly ok"},319 {"role": "user", "content": "ok"},320 ],321 temperature=0,322 max_tokens=2,323 )324 325 326async def _init_env() -> CssEnv:327 if LOCAL_IMAGE_NAME:328 return await CssEnv.from_docker_image(LOCAL_IMAGE_NAME)329 330 env = CssEnv(base_url=ENV_URL)331 await env.connect()332 return env333 334 335def _task_step_limit(task_cfg: Dict[str, Any]) -> int:336 task_limit = int(task_cfg.get("max_steps", MAX_STEPS) or MAX_STEPS)337 return max(1, min(task_limit, MAX_STEPS))338 339 340def _task_threshold(task_cfg: Dict[str, Any]) -> float:341 return clamp01(float(task_cfg.get("success_threshold", 0.95)))342 343 344def _llm_action(client: OpenAI, observation: Any, step: int, task_name: str) -> Dict[str, Any]:345 prompt = _build_user_prompt(observation, step, task_name)346 completion = client.chat.completions.create(347 model=MODEL_NAME,348 messages=[349 {"role": "system", "content": SYSTEM_PROMPT},350 {"role": "user", "content": prompt},351 ],352 temperature=TEMPERATURE,353 max_tokens=MAX_TOKENS,354 )355 text = completion.choices[0].message.content or ""356 parsed = _parse_action_json(text)357 if parsed is None:358 return _fallback_action(observation)359 return _normalize_action(parsed, observation)360 361 362async def run_task(env: CssEnv, llm_client: OpenAI, task_name: str, task_cfg: Dict[str, Any]) -> Dict[str, Any]:363 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)364 365 rewards: List[float] = []366 final_score = 0.0367 steps_taken = 0368 final_success = False369 370 try:371 reset_result = await env.reset(task=task_cfg, seed=7)372 observation = reset_result.observation373 done = bool(reset_result.done)374 except Exception as exc:375 log_step(step=0, action="{}", reward=0.0, done=True, error=f"reset_failed:{exc}")376 log_end(success=False, steps=0, score=0.0, rewards=[])377 return {"score": 0.0, "steps": 0, "rewards": []}378 379 max_steps = _task_step_limit(task_cfg)380 threshold = _task_threshold(task_cfg)381 382 for step in range(1, max_steps + 1):383 if done:384 break385 386 llm_error: Optional[str] = None387 try:388 action_payload = _llm_action(llm_client, observation, step, task_name)389 except Exception as exc:390 llm_error = f"llm_error:{exc}"391 action_payload = _fallback_action(observation)392 393 action_payload = _normalize_action(action_payload, observation)394 action_text = _action_str(action_payload)395 396 step_error: Optional[str] = llm_error397 try:398 step_result = await env.step(CssAction(**action_payload))399 except Exception as exc:400 step_error = f"step_failed:{exc}"401 action_payload = _fallback_action(observation)402 action_text = _action_str(action_payload)403 try:404 step_result = await env.step(CssAction(**action_payload))405 except Exception as exc2:406 log_step(step=step, action=action_text, reward=0.0, done=True, error=f"step_failed:{exc2}")407 steps_taken = step408 break409 410 reward = clamp01(float(step_result.reward or 0.0))411 done = bool(step_result.done)412 observation = step_result.observation413 rewards.append(reward)414 steps_taken = step415 416 observed_score = getattr(observation, "score", None)417 if observed_score is not None:418 final_score = clamp01(float(observed_score))419 elif rewards:420 final_score = clamp01(sum(rewards) / len(rewards))421 422 final_success = bool(getattr(observation, "success", False)) or (final_score >= threshold)423 log_step(step=step, action=action_text, reward=reward, done=done, error=step_error)424 425 if not rewards:426 final_score = 0.0427 428 log_end(success=final_success, steps=steps_taken, score=final_score, rewards=rewards)429 return {"score": final_score, "steps": steps_taken, "rewards": rewards}430 431 432async def main_async() -> None:433 try:434 tasks_to_run = _select_tasks(TASK_NAME)435 except Exception as exc:436 print(f"[DEBUG] Invalid task selection: {exc}", file=sys.stderr, flush=True)437 return438 439 try:440 llm_client = _make_openai_client()441 _probe_llm(llm_client)442 except Exception as exc:443 for task_name in tasks_to_run:444 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)445 log_step(step=0, action="{}", reward=0.0, done=True, error=f"llm_init_failed:{exc}")446 log_end(success=False, steps=0, score=0.0, rewards=[])447 return448 449 env: Optional[CssEnv] = None450 results: Dict[str, Dict[str, Any]] = {}451 452 try:453 env = await _init_env()454 except Exception as exc:455 for task_name in tasks_to_run:456 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)457 log_step(step=0, action="{}", reward=0.0, done=True, error=f"env_init_failed:{exc}")458 log_end(success=False, steps=0, score=0.0, rewards=[])459 return460 461 try:462 for task_name in tasks_to_run:463 task_cfg = dict(TASKS.get(task_name, {}))464 if not task_cfg:465 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)466 log_step(step=0, action="{}", reward=0.0, done=True, error="missing_task_config")467 log_end(success=False, steps=0, score=0.0, rewards=[])468 continue469 470 try:471 results[task_name] = await run_task(env, llm_client, task_name, task_cfg)472 except Exception as exc:473 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)474 log_step(step=0, action="{}", reward=0.0, done=True, error=f"task_failed:{exc}")475 log_end(success=False, steps=0, score=0.0, rewards=[])476 results[task_name] = {"score": 0.0, "steps": 0, "rewards": []}477 finally:478 if env is not None:479 try:480 await env.close()481 except Exception as exc:482 print(f"[DEBUG] env.close() error: {exc}", file=sys.stderr, flush=True)483 484 print("\n=== Final Results ===", file=sys.stderr)485 total = 0.0486 for task_name in tasks_to_run:487 score = clamp01(float(results.get(task_name, {}).get("score", 0.0)))488 total += score489 print(f" {task_name}: score={score:.4f}", file=sys.stderr)490 if tasks_to_run:491 print(f" average: {total / len(tasks_to_run):.4f}", file=sys.stderr)492 493 494def main() -> None:495 asyncio.run(main_async())496 497 498if __name__ == "__main__":499 main()500 