Team Ai
Apppublic

khushmagrawal/devsecops_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
devsecops_env_environment.py342 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"""8DevSecOps Environment Implementation.9 10Core state machine for the PR gatekeeper RL environment.11Manages episode state, tool execution, reward calculation, and episode transitions.12"""13 14from uuid import uuid415from typing import Dict, Any, Optional16import copy17 18from openenv.core.env_server.interfaces import Environment19from openenv.core.env_server.types import State20 21import sys22from pathlib import Path23 24try:25    from devsecops_env.models import (26        DevsecopsAction,27        DevsecopsObservation,28        PullRequest,29        RepositoryContext,30        Budget,31        ToolUseRecord,32    )33except ImportError:34    # Fallback for direct import35    sys.path.insert(0, str(Path(__file__).parent.parent))36    from models import (37        DevsecopsAction,38        DevsecopsObservation,39        PullRequest,40        RepositoryContext,41        Budget,42        ToolUseRecord,43    )44from .scenarios import load_scenario, get_reward_calculator45from .mock_tools import dispatch_tool46from .graders import compute_reward47 48 49class DevSecOpsEnvironment(Environment):50    """51    DevSecOps Gatekeeper RL Environment.52    53    Trains agents to make security-aware decisions on incoming Pull Requests54    by using simulated tools: code inspection, CI execution, vulnerability scanning.55    56    The environment is fully stateful per-episode:57    - Each reset() creates a new PR scenario (task1, task2, or task3)58    - Each step() executes a tool and updates the observable state59    - State tracking enables task2 code patching and task3 malware detection60    61    The environment supports concurrent sessions (SUPPORTS_CONCURRENT_SESSIONS=True)62    because each instance gets its own reset() call and maintains isolated state.63    64    Example:65        >>> env = DevSecOpsEnvironment()66        >>> obs = env.reset()67        >>> print(obs.task_id)  # "task1", "task2", or "task3"68        >>> 69        >>> # Run CI to check the PR70        >>> action = DevsecopsAction(tool_name="run_ci", scope="unit_only")71        >>> obs = env.step(action)72        >>> print(obs.reward)73        >>>74        >>> # Make final decision75        >>> action = DevsecopsAction(76        ...     tool_name="make_decision",77        ...     verdict="MERGE",78        ...     justification="Documentation changes only"79        ... )80        >>> obs = env.step(action)81        >>> print(obs.done)  # True82        >>> print(obs.episode_reward)  # Total accumulated reward83    """84    85    # Enable concurrent WebSocket sessions86    # Each client connection gets its own environment instance87    SUPPORTS_CONCURRENT_SESSIONS: bool = True88    89    # Environment metadata for Gymnasium/OpenEnv compliance90    render_mode: str = "text"91    spec: Optional[Dict[str, Any]] = None  # Can be populated with EnvSpec if needed92    93    def __init__(self):94        """Initialize the environment."""95        self._state = State(episode_id=str(uuid4()), step_count=0)96        97        # Per-episode state - initialized in reset()98        self._scenario: Optional[Dict[str, Any]] = None99        self._observation_state = {100            "pr": None,101            "repo_context": None,102            "budget": None,103            "pipeline_history": [],104            "done": False,105            "reward": 0.0,106            "episode_reward": 0.0,107            "step_count": 0,108            "last_tool_output": None,109            "task_id": None,110        }111        112        # Mutable state tracking across tool calls in same episode113        self._internal_state: Dict[str, Any] = {}114        115        # For reward calculation at episode end116        self._verdict: Optional[str] = None117        self._ci_runs_used: int = 0118        self._code_patched: bool = False119        self._final_justification: str = ""120    121    # ========================================================================122    # RESET: Initialize new episode123    # ========================================================================124    125    def reset(self, options: Optional[Dict[str, Any]] = None) -> DevsecopsObservation:126        """127        Reset environment for a new episode.128        129        Args:130            options: Optional dict with:131                - "task": Specific task to load ("task1", "task2", "task3")132                         If not provided, randomly selects one133                         134        Returns:135            DevsecopsObservation: Initial state for the episode136        """137        138        # Reset episode tracking139        self._state = State(episode_id=str(uuid4()), step_count=0)140        141        # Load scenario (random or specified)142        task_id = None143        if options and "task" in options:144            task_id = options["task"]145        self._scenario = load_scenario(task_id)146        147        # Reset internal state tracking148        self._internal_state = {}149        self._verdict = None150        self._ci_runs_used = 0151        self._code_patched = False152        self._final_justification = ""153        154        # Initialize observation state from scenario155        pr = self._scenario["pr"]156        repo = self._scenario["repo_context"]157        budget = self._scenario["budget"]158        159        # Make deep copies to avoid mutation of scenario160        self._observation_state = {161            "pr": copy.deepcopy(pr),162            "repo_context": copy.deepcopy(repo),163            "budget": copy.deepcopy(budget),164            "pipeline_history": [],165            "done": False,166            "reward": 0.0,167            "episode_reward": 0.0,168            "step_count": 0,169            "last_tool_output": None,170            "task_id": self._scenario["task_id"],171            "internal_state": {},172        }173        174        return self._make_observation(reward=0.0)175    176    # ========================================================================177    # STEP: Execute action (tool call)178    # ========================================================================179    180    def step(self, action: DevsecopsAction) -> DevsecopsObservation:181        """182        Execute a tool action.183        184        Args:185            action: DevsecopsAction specifying tool and parameters186            187        Returns:188            DevsecopsObservation: Updated state after tool execution189        """190        191        if self._scenario is None:192            raise RuntimeError("Must call reset() before step()")193        194        # Increment step count195        self._state.step_count += 1196        self._observation_state["step_count"] = self._state.step_count197        198        # ====================================================================199        # STEP 1: Dispatch tool and get result200        # ====================================================================201        202        tool_output, tool_metadata = dispatch_tool(203            action=action,204            scenario=self._scenario,205            environment_state=self._internal_state,206        )207        208        # Update internal state tracking209        # (e.g., task2_code_patched flag set by patch_code tool)210        if action.tool_name == "patch_code" and tool_metadata.get("success"):211            self._code_patched = True212            self._internal_state["task2_code_patched"] = True213        214        # Track CI usage215        if action.tool_name == "run_ci":216            self._ci_runs_used += 1217            self._observation_state["budget"].use_ci()218        219        # ====================================================================220        # STEP 2: Record tool call in pipeline history221        # ====================================================================222        223        tool_record = ToolUseRecord(224            step=self._state.step_count,225            tool_name=action.tool_name,226            arguments={227                k: v for k, v in action.dict().items()228                if v is not None and k != "tool_name"229            },230            result=tool_output[:500],  # Truncate for storage231        )232        self._observation_state["pipeline_history"].append(tool_record)233        234        # ====================================================================235        # STEP 3: Check if episode is done (make_decision called)236        # ====================================================================237        238        is_done = False239        step_reward = 0.0240        241        if action.tool_name == "make_decision":242            is_done = True243            self._verdict = action.verdict244            self._final_justification = action.justification or ""245            246            # Calculate step reward [0, 1] based on verdict + work done247            step_reward = compute_reward(248                task_id=self._scenario["task_id"],249                verdict=self._verdict,250                ci_runs_used=self._ci_runs_used,251                code_patched=self._code_patched,252                justification=self._final_justification,253            )254        else:255            # Intermediate step - small penalty for each step/tool256            # (encourages efficiency)257            step_reward = 0.0  # Or -0.1 per step if you want to encourage finishing258        259        # ====================================================================260        # STEP 4: Update observation state261        # ====================================================================262        263        self._observation_state["done"] = is_done264        self._observation_state["reward"] = step_reward265        self._observation_state["episode_reward"] += step_reward266        self._observation_state["last_tool_output"] = tool_output267        self._observation_state["internal_state"] = copy.deepcopy(self._internal_state)268        269        # Check budget exceeded270        if self._observation_state["step_count"] >= self._observation_state["budget"].step_limit:271            self._observation_state["done"] = True272        273        return self._make_observation(reward=step_reward)274    275    # ========================================================================276    # HELPERS: Observation construction277    # ========================================================================278    279    def _make_observation(self, reward: float) -> DevsecopsObservation:280        """281        Construct a DevsecopsObservation from current state.282        283        Args:284            reward: Reward from the last step285            286        Returns:287            DevsecopsObservation instance288        """289        290        obs = DevsecopsObservation(291            task_id=self._observation_state["task_id"],292            pr=self._observation_state["pr"],293            repo_context=self._observation_state["repo_context"],294            pipeline_history=self._observation_state["pipeline_history"],295            budget=self._observation_state["budget"],296            last_tool_output=self._observation_state["last_tool_output"],297            done=self._observation_state["done"],298            reward=reward,299            episode_reward=self._observation_state["episode_reward"],300            step_count=self._observation_state["step_count"],301            internal_state=self._observation_state["internal_state"],302            metadata={303                "episode_id": self._state.episode_id,304                "verdict": self._verdict,305                "ci_runs_used": self._ci_runs_used,306                "code_patched": self._code_patched,307            },308        )309        310        return obs311    312    # ========================================================================313    # STATE ACCESS (for debugging)314    # ========================================================================315    316    @property317    def state(self) -> State:318        """Get current environment state (episode_id, step_count)."""319        return self._state320    321    def get_episode_summary(self) -> Dict[str, Any]:322        """323        Get a summary of the current episode.324        325        Useful for debugging and analysis.326        """327        328        return {329            "episode_id": self._state.episode_id,330            "task_id": self._observation_state["task_id"],331            "step_count": self._state.step_count,332            "done": self._observation_state["done"],333            "episode_reward": self._observation_state["episode_reward"],334            "ci_runs_used": self._ci_runs_used,335            "code_patched": self._code_patched,336            "verdict": self._verdict,337            "tool_calls": [338                {"step": t.step, "tool": t.tool_name}339                for t in self._observation_state["pipeline_history"]340            ],341        }342