Ar-Srivas/BitWise_CSS_env
0
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7# make sure the doc strings are correct, i have done some check the others8 9"""10Css Env Environment Implementation.11"""12 13import re14import sys15import math16from pathlib import Path17from typing import Dict, List, Optional18from uuid import uuid419 20from openenv.core.env_server.interfaces import Environment21from openenv.core.env_server.types import State22 23_ROOT_DIR = Path(__file__).resolve().parents[1]24if str(_ROOT_DIR) not in sys.path:25 sys.path.insert(0, str(_ROOT_DIR))26 27try:28 from .models import CssAction, CssObservation, GradeResult29except ImportError:30 from models import CssAction, CssObservation, GradeResult31 32from reward import compute_reward33from graders import colors, spacing, typography, contrast, layout, cleanliness, design_quality34 35try:36 from .tasks import TASKS37 from .action_engine import apply_action38 from .flaw_injector import inject_flaws39except ImportError:40 from server.tasks import TASKS41 from server.action_engine import apply_action42 from server.flaw_injector import inject_flaws43 44 45class CssEnvironment(Environment):46 """47 CSS Refinement RL Environment48 49 This environment simulates a CSS refinement task where50 an agent iteratively modifies CSS code to improve design quality.51 52 Example:53 >>> env = CssEnvironment()54 >>> obs = env.reset(task, seed=42)55 >>> print(obs.html) # <div class="card">Hello</div>56 >>> print(obs.css) # .card { color: #1a73e7; }57 >>> result = env.step(CssAction(58 ... action_type="replace_color",59 ... target="#1a73e7",60 ... value="#1a6fe0"61 ... ))62 >>> print(result["observation"].css) # .card { color: #1a6fe0; }63 >>> print(result["reward"]) # 0.664 >>> print(result["done"]) # False65 """66 67 # Enable concurrent WebSocket sessions.68 # Set to True if your environment isolates state between instances.69 # When True, multiple WebSocket clients can connect simultaneously, each70 # getting their own environment instance (when using factory mode in app.py).71 SUPPORTS_CONCURRENT_SESSIONS: bool = True72 HIGH_SCORE_GUARD_THRESHOLD: float = 0.9073 HIGH_SCORE_MAX_DROP: float = 0.0374 SCORE_EPSILON: float = 1e-675 MIN_SCORE_BOUND: float = 0.0176 MAX_SCORE_BOUND: float = 0.9977 SUPPORTED_GRADERS = (78 "color",79 "spacing",80 "typography",81 "contrast",82 "layout",83 "cleanliness",84 "design_quality",85 )86 87 def __init__(self):88 """Initialize the css_env environment."""89 self.html = "" # html string in the task90 self.css = "" # css string in the task91 self.tokens = {} # design tokens given in the task as a dict of token_name -> value92 self.config = {} # config given in the task to inject flaws93 self.manifest = [] # list of dicts describing flaws94 self.difficulty = "easy"95 self.success_threshold = 0.9596 self.required_graders = ["color", "spacing", "typography", "contrast", "cleanliness"]97 self.grader_weights = {98 "color": 0.30,99 "spacing": 0.20,100 "typography": 0.20,101 "contrast": 0.20,102 "cleanliness": 0.10,103 }104 self.step_count = 0 # current step count105 self.max_steps = 20 # maximum steps106 self.state_data = {} # tracks progress107 108 @staticmethod109 def _action_signature(action: CssAction) -> str:110 value = "null" if action.value is None else str(action.value)111 return f"{action.action_type}|{action.target}|{value}"112 113 @staticmethod114 def _action_score_key(action_type: str) -> Optional[str]:115 mapping = {116 "replace_color": "color",117 "fix_spacing": "spacing",118 "fix_typography": "typography",119 "fix_contrast": "contrast",120 "add_breakpoint": "layout",121 "remove_rule": "cleanliness",122 }123 return mapping.get(action_type)124 125 def _is_irrelevant_action(self, action: CssAction, scores: Dict[str, float]) -> bool:126 # In easy tasks, breakpoint tuning is usually not part of the objective.127 if self.difficulty == "easy" and action.action_type == "add_breakpoint":128 return True129 130 score_key = self._action_score_key(action.action_type)131 if score_key and scores.get(score_key, 0.0) >= self.success_threshold:132 return True133 134 return False135 136 def reset(self, task: Optional[Dict] = None, seed: int = 0) -> CssObservation:137 """138 task = {139 "html": str,140 "css": str,141 "design_tokens": dict,142 "config": dict,143 "difficulty": "easy" | "medium" | "hard"144 }145 """146 147 self.step_count = 0148 149 if task is None or not isinstance(task, dict) or not task:150 task = TASKS["task1"]151 152 self.html = task["html"]153 self.tokens = task.get("design_tokens", task.get("tokens", {}))154 self.config = task.get("config", task.get("flaw_config", {}))155 self.difficulty = task.get("difficulty", "easy")156 self.max_steps = int(task.get("max_steps", self.max_steps))157 self.success_threshold = float(task.get("success_threshold", 0.95))158 task_weights = self._normalize_grader_weights(task.get("grader_weights", {}))159 task_graders = task.get("required_graders", task.get("graders", self.required_graders))160 required_graders = [161 str(key).strip() for key in (task_graders if isinstance(task_graders, list) else []) if str(key).strip()162 ]163 required_graders = [164 key for key in required_graders if key in self.SUPPORTED_GRADERS165 ]166 if not required_graders:167 required_graders = [168 key for key, weight in task_weights.items() if weight > 0.0 and key in self.SUPPORTED_GRADERS169 ]170 171 if required_graders:172 self.required_graders = required_graders173 else:174 self.required_graders = ["color", "spacing", "typography", "contrast", "cleanliness"]175 176 self.grader_weights = task_weights or self._default_weights_for_graders(self.required_graders)177 178 clean_css = task.get("css", task.get("clean_css", ""))179 180 result = inject_flaws(clean_css, self.tokens, self.config, seed)181 182 self.css = result["css"]183 self.manifest = result["manifest"]184 initial_unused_selectors = self._compute_unused_selectors(self.css, self.html)185 186 if task.get("initial_unused_selectors"):187 initial_unused_selectors = sorted(188 set(initial_unused_selectors) | set(task["initial_unused_selectors"])189 )190 191 self.state_data = {192 "episode_id": str(uuid4()),193 "step_count": 0,194 "initial_manifest": self.manifest,195 "initial_unused_selectors": initial_unused_selectors,196 "last_action_signature": None,197 "same_action_count": 0,198 "reward_drop_streak": 0,199 "last_scores": {},200 "last_reward": None,201 "last_done": False,202 "last_success": False,203 "last_changed": True,204 "last_no_op": False,205 "last_irrelevant": False,206 "action_counts": {},207 "required_graders": list(self.required_graders),208 "grader_weights": dict(self.grader_weights),209 }210 211 initial_grade = self.grade()212 self.state_data["last_scores"] = dict(initial_grade.breakdown)213 self.state_data["last_grade"] = float(initial_grade.score)214 self.state_data["peak_score"] = (215 min(self.state_data["last_scores"].values())216 if self.state_data.get("last_scores")217 else 0.0218 )219 220 violations = None221 if self.difficulty == "easy":222 violations = task.get("violations")223 if violations is None:224 violations = [225 {"type": "hint", "selector": "", "message": hint}226 for hint in result.get("hints", [])227 ]228 229 return CssObservation(230 html=self.html,231 css=self.css,232 tokens=self.tokens,233 violations=violations,234 scores=self.state_data.get("last_scores", {}),235 score=float(self.state_data.get("last_grade", 0.0)),236 success=False,237 changed=True,238 no_op_action=False,239 repeated_action=False,240 terminated_by_max_steps=False,241 )242 243 def step(self, action: CssAction) -> CssObservation:244 """245 Execute a step in the environment by applying the action generated by the agent to the CSS.246 """247 prev_scores = dict(self.state_data.get("last_scores", {}))248 prev_reward = self.state_data.get("last_reward")249 prev_score = float(self.state_data.get("last_grade", min(prev_scores.values()) if prev_scores else 0.0))250 prev_peak_score = float(self.state_data.get("peak_score", prev_score))251 old_css = self.css252 new_css, _ = apply_action(self.css, action)253 css_changed = old_css.strip() != new_css.strip()254 action_sig = self._action_signature(action)255 repeated_action = action_sig == self.state_data.get("last_action_signature")256 same_action_count = int(self.state_data.get("same_action_count", 0)) + 1 if repeated_action else 1257 action_counts = dict(self.state_data.get("action_counts", {}))258 action_counts[action_sig] = int(action_counts.get(action_sig, 0)) + 1259 duplicate_action = action_counts[action_sig] > 1260 no_op_action = not css_changed261 262 self.css = new_css263 grade_result = self.grade()264 scores = dict(grade_result.breakdown)265 success = self._is_done(scores)266 irrelevant_action = self._is_irrelevant_action(action, prev_scores)267 268 # Compute reward with progress and anti-repeat penalties.269 reward = compute_reward(270 scores,271 action_valid=css_changed,272 done=success,273 action_repeated=repeated_action,274 action_duplicate=duplicate_action,275 action_irrelevant=irrelevant_action,276 previous_scores=prev_scores,277 previous_reward=prev_reward,278 )279 280 # FIX 1: block repeated identical actions with explicit penalty.281 if repeated_action:282 reward = max(0.0, reward - 0.05)283 284 # Success must be driven by grader thresholds only.285 # End the episode immediately once all required scores meet the threshold.286 if success:287 self.step_count += 1288 score = float(grade_result.score)289 290 info = {291 "scores": scores,292 "step_count": self.step_count,293 "changed": css_changed,294 "no_op_action": no_op_action,295 "repeated_action": repeated_action,296 "same_action_count": same_action_count,297 "duplicate_action": duplicate_action,298 "irrelevant_action": irrelevant_action,299 "reward_drop_too_much": False,300 "reward_drop_streak": 0,301 "reward_drop_terminated": False,302 "repeated_action_terminated": False,303 "same_action_cap_reached": False,304 "success": True,305 "high_score_degradation": False,306 "terminated_by_max_steps": False,307 "success_threshold": self.success_threshold,308 "score": score,309 }310 311 self.state_data["step_count"] = self.step_count312 self.state_data["last_scores"] = scores313 self.state_data["last_grade"] = score314 self.state_data["last_reward"] = reward315 self.state_data["last_done"] = True316 self.state_data["last_success"] = True317 self.state_data["last_changed"] = css_changed318 self.state_data["last_no_op"] = no_op_action319 self.state_data["last_irrelevant"] = irrelevant_action320 self.state_data["last_action_signature"] = action_sig321 self.state_data["same_action_count"] = same_action_count322 self.state_data["reward_drop_streak"] = 0323 self.state_data["peak_score"] = max(prev_peak_score, score)324 self.state_data["action_counts"] = action_counts325 326 return CssObservation(327 html=self.html,328 css=self.css,329 tokens=self.tokens,330 violations=None,331 scores=scores,332 score=score,333 success=True,334 changed=css_changed,335 no_op_action=no_op_action,336 repeated_action=repeated_action,337 terminated_by_max_steps=False,338 reward=reward,339 done=True,340 metadata=info,341 )342 343 # Track large reward drops, but allow recovery attempts.344 reward_drop_too_much = (345 prev_reward is not None and reward < (float(prev_reward) - 0.2)346 )347 reward_drop_streak = (348 int(self.state_data.get("reward_drop_streak", 0)) + 1349 if reward_drop_too_much350 else 0351 )352 reward_drop_terminated = reward_drop_streak > 2353 354 # FIX 4: cap repeated identical action attempts.355 same_action_cap_reached = same_action_count > 2356 repeated_action_terminated = same_action_cap_reached357 358 self.step_count += 1359 done = (360 success361 or self.step_count >= self.max_steps362 or repeated_action_terminated363 or reward_drop_terminated364 or same_action_cap_reached365 )366 terminated_by_max_steps = done and not success367 score = float(grade_result.score)368 369 high_score_degradation = (370 prev_peak_score >= self.HIGH_SCORE_GUARD_THRESHOLD371 and score < (prev_peak_score - self.HIGH_SCORE_MAX_DROP)372 )373 374 if high_score_degradation:375 # Keep policy from drifting after reaching a strong state.376 self.css = old_css377 scores = prev_scores378 score = prev_score379 success = self._is_done(scores)380 done = True381 reward = max(0.0, (float(prev_reward) - 0.05) if prev_reward is not None else reward)382 383 info = {384 "scores": scores,385 "step_count": self.step_count,386 "changed": css_changed,387 "no_op_action": no_op_action,388 "repeated_action": repeated_action,389 "same_action_count": same_action_count,390 "duplicate_action": duplicate_action,391 "irrelevant_action": irrelevant_action,392 "reward_drop_too_much": reward_drop_too_much,393 "reward_drop_streak": reward_drop_streak,394 "reward_drop_terminated": reward_drop_terminated,395 "repeated_action_terminated": repeated_action_terminated,396 "same_action_cap_reached": same_action_cap_reached,397 "success": success,398 "high_score_degradation": high_score_degradation,399 "terminated_by_max_steps": terminated_by_max_steps,400 "success_threshold": self.success_threshold,401 "score": score,402 }403 404 self.state_data["step_count"] = self.step_count405 self.state_data["last_scores"] = scores406 self.state_data["last_grade"] = score407 self.state_data["last_reward"] = reward408 self.state_data["last_done"] = done409 self.state_data["last_success"] = success410 self.state_data["last_changed"] = css_changed411 self.state_data["last_no_op"] = no_op_action412 self.state_data["last_irrelevant"] = irrelevant_action413 self.state_data["last_action_signature"] = action_sig414 self.state_data["same_action_count"] = same_action_count415 self.state_data["reward_drop_streak"] = reward_drop_streak416 self.state_data["peak_score"] = max(prev_peak_score, score)417 self.state_data["action_counts"] = action_counts418 419 return CssObservation(420 html=self.html,421 css=self.css,422 tokens=self.tokens,423 violations=None,424 scores=scores,425 score=score,426 success=success,427 changed=css_changed,428 no_op_action=no_op_action,429 repeated_action=repeated_action,430 terminated_by_max_steps=terminated_by_max_steps,431 reward=reward,432 done=done,433 metadata=info,434 )435 436 @property437 def state(self) -> State:438 """439 Get the current environment state.440 441 Returns:442 Current State with episode_id and step_count443 """444 return State(445 episode_id=self.state_data.get("episode_id"),446 step_count=self.state_data.get("step_count", 0),447 success=self.state_data.get("last_success", False),448 done=self.state_data.get("last_done", False),449 reward=self.state_data.get("last_reward") if self.state_data.get("last_reward") is not None else 0.0,450 changed=self.state_data.get("last_changed", True),451 no_op_action=self.state_data.get("last_no_op", False),452 scores=self.state_data.get("last_scores", {}),453 )454 455 @staticmethod456 def _normalize_grader_weights(raw_weights: Dict[str, float]) -> Dict[str, float]:457 if not isinstance(raw_weights, dict):458 return {}459 460 normalized: Dict[str, float] = {}461 for grader_key, weight in raw_weights.items():462 key = str(grader_key).strip()463 if not key:464 continue465 try:466 normalized[key] = float(weight)467 except (TypeError, ValueError):468 continue469 470 return normalized471 472 @staticmethod473 def _default_weights_for_graders(graders: List[str]) -> Dict[str, float]:474 normalized = [str(key).strip() for key in graders if str(key).strip()]475 if not normalized:476 return {}477 equal_weight = 1.0 / float(len(normalized))478 return {key: equal_weight for key in normalized}479 480 def _active_grader_weights(self) -> Dict[str, float]:481 candidate = self.state_data.get("grader_weights", self.grader_weights)482 weights = {483 key: value484 for key, value in self._normalize_grader_weights(candidate).items()485 if key in self.SUPPORTED_GRADERS486 }487 if weights:488 return weights489 return self._default_weights_for_graders(list(self.required_graders))490 491 @classmethod492 def _clamp_open_unit_interval(cls, value: float) -> float:493 try:494 numeric = float(value)495 except (TypeError, ValueError):496 numeric = cls.MIN_SCORE_BOUND + cls.SCORE_EPSILON497 if not math.isfinite(numeric):498 numeric = cls.MIN_SCORE_BOUND + cls.SCORE_EPSILON499 if numeric <= cls.MIN_SCORE_BOUND:500 numeric = cls.MIN_SCORE_BOUND + cls.SCORE_EPSILON501 if numeric >= cls.MAX_SCORE_BOUND:502 numeric = cls.MAX_SCORE_BOUND - cls.SCORE_EPSILON503 504 rounded = round(numeric, 2)505 if rounded <= cls.MIN_SCORE_BOUND:506 return 0.02507 if rounded >= cls.MAX_SCORE_BOUND:508 return 0.98509 return float(f"{rounded:.2f}")510 511 def grade(self) -> GradeResult:512 """Grade the current CSS using task-configured grader weights."""513 all_scores = self._run_graders()514 weights = self._active_grader_weights()515 516 breakdown: Dict[str, float] = {}517 details: Dict[str, str] = {}518 weighted_sum = 0.0519 weight_sum = 0.0520 521 for grader_key, weight in weights.items():522 if weight <= 0.0:523 continue524 525 component_score = float(all_scores.get(grader_key, 0.0))526 breakdown[grader_key] = component_score527 weighted_sum += component_score * weight528 weight_sum += weight529 details[grader_key] = (530 f"component={grader_key}, score={component_score:.4f}, weight={weight:.4f}"531 )532 533 if not breakdown:534 required = list(self.required_graders) or ["color", "spacing", "typography", "contrast", "cleanliness"]535 for grader_key in required:536 component_score = float(all_scores.get(grader_key, 0.0))537 breakdown[grader_key] = component_score538 details[grader_key] = (539 f"component={grader_key}, score={component_score:.4f}, weight=auto"540 )541 score = min(breakdown.values()) if breakdown else 0.0542 return GradeResult(score=self._clamp_open_unit_interval(score), breakdown=breakdown, details=details)543 544 score = (weighted_sum / weight_sum) if weight_sum > 0.0 else 0.0545 return GradeResult(score=self._clamp_open_unit_interval(score), breakdown=breakdown, details=details)546 547 def _run_graders(self) -> Dict[str, float]:548 return {549 "color": self._grade_color(),550 "spacing": self._grade_spacing(),551 "typography": self._grade_typography(),552 "contrast": self._grade_contrast(),553 "layout": self._grade_layout(),554 "cleanliness": self._grade_cleanliness(),555 "design_quality": self._grade_design_quality(),556 }557 558 def _grade_color(self) -> float:559 return self._safe_grade(colors.grade)560 561 def _grade_spacing(self) -> float:562 return self._safe_grade(spacing.grade)563 564 def _grade_typography(self) -> float:565 return self._safe_grade(typography.grade)566 567 def _grade_contrast(self) -> float:568 return self._safe_grade(contrast.grade)569 570 def _grade_layout(self) -> float:571 return self._safe_grade(layout.grade)572 573 def _grade_cleanliness(self) -> float:574 return self._safe_grade(cleanliness.grade)575 576 def _grade_design_quality(self) -> float:577 return self._safe_grade(design_quality.grade)578 579 def _is_done(self, scores: Dict[str, float]) -> bool:580 required = list(self.state_data.get("required_graders") or self.required_graders)581 if not required:582 required = ["color", "spacing", "typography", "contrast", "cleanliness"]583 return all(scores.get(key, 0.0) >= self.success_threshold for key in required)584 585 def _safe_grade(self, grader_fn) -> float:586 try:587 score = float(grader_fn(self.html, self.css, self.tokens, self.state_data))588 return self._clamp_open_unit_interval(score)589 except Exception:590 return self.MIN_SCORE_BOUND + self.SCORE_EPSILON591 592 def _compute_unused_selectors(self, css: str, html: str) -> List[str]:593 selector_groups = re.findall(r"([^{}]+)\{", css)594 selectors: List[str] = []595 596 for group in selector_groups:597 for selector in group.split(","):598 clean = selector.strip()599 if clean and clean not in selectors:600 selectors.append(clean)601 602 unused = [s for s in selectors if not self._selector_matches_html(s, html)]603 return sorted(set(unused))604 605 def _selector_matches_html(self, selector: str, html: str) -> bool:606 # Ignore pseudo and combinator suffixes and evaluate only the right-most simple token.607 base = re.split(r"::?|\s+|>|\+|~", selector.strip())[-1]608 if not base or base == "*":609 return True610 611 if base.startswith("."):612 class_name = base[1:]613 return bool(re.search(rf'class\s*=\s*["\'][^"\']*\b{re.escape(class_name)}\b', html))614 615 if base.startswith("#"):616 element_id = base[1:]617 return bool(re.search(rf'id\s*=\s*["\']{re.escape(element_id)}["\']', html))618 619 if re.match(r"^[a-zA-Z][a-zA-Z0-9-]*$", base):620 return bool(re.search(rf"<{re.escape(base)}(?:\s|>)", html))621 622 return False623 