khushmagrawal/devsecops_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"""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 