openenv/coding_env
21
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"""Transforms specific to coding environments."""8 9import ast10import re11 12from openenv.core.env_server.base_transforms import CompositeTransform13from openenv.core.env_server.interfaces import Transform14from openenv.core.env_server.types import Observation15 16from ..models import CodeObservation17 18 19class CodeSafetyTransform(Transform):20 """Evaluates code safety and assigns penalties for dangerous patterns."""21 22 def __init__(self, penalty: float = -1.0):23 self.penalty = penalty24 self.dangerous_patterns = [25 r"import\s+os",26 r"import\s+subprocess",27 r"eval\(",28 r"exec\(",29 r"__import__",30 r"open\(",31 ]32 33 def __call__(self, observation: Observation) -> Observation:34 if not isinstance(observation, CodeObservation):35 return observation36 37 if "last_code" in observation.metadata:38 code = observation.metadata["last_code"]39 for pattern in self.dangerous_patterns:40 if re.search(pattern, code):41 observation.reward = self.penalty42 observation.metadata["safety_violation"] = pattern43 break44 else:45 if observation.reward is None:46 observation.reward = 0.047 48 return observation49 50 51class CodeQualityTransform(Transform):52 """Evaluates and rewards code quality metrics."""53 54 def __init__(55 self,56 concise_bonus: float = 0.1,57 max_length_threshold: int = 100,58 syntax_penalty: float = -0.2,59 ):60 self.concise_bonus = concise_bonus61 self.max_length_threshold = max_length_threshold62 self.syntax_penalty = syntax_penalty63 64 def __call__(self, observation: Observation) -> Observation:65 if not isinstance(observation, CodeObservation):66 return observation67 68 quality_score = 0.069 70 if "last_code" in observation.metadata:71 code = observation.metadata["last_code"]72 73 # Reward concise code74 if len(code.strip()) <= self.max_length_threshold:75 quality_score += self.concise_bonus76 77 # Check syntax (redundant but useful for quality assessment)78 try:79 ast.parse(code)80 except SyntaxError:81 quality_score += self.syntax_penalty82 83 # Add to existing reward84 if observation.reward is None:85 observation.reward = quality_score86 else:87 observation.reward += quality_score88 89 return observation90 91 92def create_safe_coding_transform() -> CompositeTransform:93 """Create a transform focused on safe coding practices and quality."""94 return CompositeTransform([CodeSafetyTransform(), CodeQualityTransform()])95 