Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
transforms.py95 linesDownload Raw Back to server
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