"""Core environment logic for Incident-Response-Detective.""" import os import sys import uuid from typing import Any, Optional # Ensure project root (/app) is on sys.path so models and task_definitions are importable # regardless of how this module is loaded (as server.environment or directly). _project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) if _project_root not in sys.path: sys.path.insert(0, _project_root) from openenv.core import Environment from models import IncidentAction, IncidentObservation, IncidentState from task_definitions import TASKS, ACTIONS, ADVERSARIAL_OVERLAYS, compute_reward class IncidentResponseEnvironment(Environment): """ OpenEnv environment for incident triage. Implements reset(), step(), state following the Gymnasium-style API. """ SUPPORTS_CONCURRENT_SESSIONS = True def __init__(self): super().__init__() self._episodes: dict[str, dict] = {} def get_tasks(self) -> list[dict]: return [ { "id": t["id"], "name": t["name"], "difficulty": t["difficulty"], "description": t["description"], "max_steps": t["max_steps"], } for t in TASKS.values() ] def reset( self, seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs: Any, ) -> tuple[str, dict]: """Start a new episode. Returns (episode_id, observation_dict). Kwargs: task_id (str): Task to run. Defaults to "task_easy". adversarial (bool): Use adversarial chat overlay. Defaults to False. """ task_id: str = kwargs.get("task_id", "task_easy") adversarial: bool = kwargs.get("adversarial", False) if task_id not in TASKS: raise ValueError(f"Unknown task: {task_id}. Choose from: {list(TASKS.keys())}") new_episode_id = str(uuid.uuid4()) task = TASKS[task_id] chat_history = ( ADVERSARIAL_OVERLAYS[task_id] if adversarial and task_id in ADVERSARIAL_OVERLAYS else task["observation"]["chat_history"] ) self._episodes[new_episode_id] = { "task_id": task_id, "step_count": 0, "done": False, "resolved": False, "actions_taken": [], "rewards": [], "cumulative_reward": 0.0, "adversarial": adversarial, } obs = { "task_id": task_id, "task_name": task["name"], "task_description": task["description"], "logs": task["observation"]["logs"], "chat_history": chat_history, "runbook": task["observation"]["runbook"], "available_actions": ACTIONS, "step": 0, "max_steps": task["max_steps"], "done": False, "score": 0.0, "last_reward": 0.0, "reward_breakdown": {}, "feedback": "Episode started. Analyze the observation and choose a remediation action.", "last_action_error": None, } return new_episode_id, obs def step( self, action: Any, timeout_s: Optional[float] = None, **kwargs: Any, ) -> dict: """Execute an action. Returns observation dict. Args: action: Action dict with keys 'action' (str) and optionally 'evidence' (int). Kwargs: episode_id (str): The episode to step. """ action_dict: dict = action if isinstance(action, dict) else {} episode_id: str = kwargs.get("episode_id", "") if episode_id not in self._episodes: raise ValueError(f"Unknown episode_id: {episode_id}") ep = self._episodes[episode_id] if ep["done"]: raise ValueError("Episode already finished.") action_str = action_dict.get("action", "") if action_str not in ACTIONS: # Return error observation without consuming a step task = TASKS[ep["task_id"]] return { "task_id": ep["task_id"], "task_name": task["name"], "task_description": task["description"], "logs": task["observation"]["logs"], "chat_history": task["observation"]["chat_history"], "runbook": task["observation"]["runbook"], "available_actions": ACTIONS, "step": ep["step_count"], "max_steps": task["max_steps"], "done": False, "score": ep["cumulative_reward"], "last_reward": 0.0, "reward_breakdown": {}, "feedback": f"Invalid action: {action_str}", "last_action_error": f"Invalid action '{action_str}'. Choose from: {ACTIONS}", } ep["step_count"] += 1 ep["actions_taken"].append(action_str) reward_info = compute_reward(ep["task_id"], action_str, ep["step_count"]) task = TASKS[ep["task_id"]] # Evidence validation evidence_penalty = 0.0 evidence_warning = None evidence = action_dict.get("evidence") if evidence is None: evidence_penalty = 0.1 evidence_warning = "No evidence provided. Supply 'evidence': to justify your action." else: try: idx = int(evidence) log_count = len(task["observation"]["logs"]) if idx < 0 or idx >= log_count: evidence_penalty = 0.1 evidence_warning = f"Evidence index {idx} out of bounds (valid: 0–{log_count - 1})." except (TypeError, ValueError): evidence_penalty = 0.1 evidence_warning = f"Evidence must be an integer log index, got {evidence!r}." adjusted_reward = round(max(0.0, reward_info["reward"] - evidence_penalty), 3) ep["rewards"].append(adjusted_reward) ep["cumulative_reward"] = round(sum(ep["rewards"]), 3) if reward_info["done"]: ep["done"] = True ep["resolved"] = reward_info["resolved"] feedback_parts = [ reward_info["safety"]["reason"], reward_info["efficiency"]["reason"], ] if evidence_warning: feedback_parts.append(evidence_warning) if ep["done"]: if ep["resolved"]: feedback_parts.append("INCIDENT RESOLVED.") else: feedback_parts.append("INCIDENT NOT RESOLVED. Episode ended.") return { "task_id": ep["task_id"], "task_name": task["name"], "task_description": task["description"], "logs": task["observation"]["logs"], "chat_history": task["observation"]["chat_history"], "runbook": task["observation"]["runbook"], "available_actions": ACTIONS, "step": ep["step_count"], "max_steps": task["max_steps"], "done": ep["done"], "score": ep["cumulative_reward"], "last_reward": adjusted_reward, "reward_breakdown": { "safety": reward_info["safety"], "efficiency": reward_info["efficiency"], }, "feedback": " | ".join(feedback_parts), "last_action_error": None, } @property def state(self) -> dict: """Returns a summary of all active episode states.""" return { eid: { "task_id": ep["task_id"], "step_count": ep["step_count"], "done": ep["done"], "resolved": ep["resolved"], "cumulative_reward": ep["cumulative_reward"], } for eid, ep in self._episodes.items() } def get_state(self, episode_id: str) -> dict: """Per-episode state lookup used by the server /state endpoint.""" if episode_id not in self._episodes: raise ValueError(f"Unknown episode_id: {episode_id}") ep = self._episodes[episode_id] return { "episode_id": episode_id, "task_id": ep["task_id"], "step_count": ep["step_count"], "done": ep["done"], "resolved": ep["resolved"], "actions_taken": ep["actions_taken"], "rewards": ep["rewards"], "cumulative_reward": ep["cumulative_reward"], } def grade(self, episode_id: str) -> dict: """Grade an episode. Returns score in 0.0-1.0.""" if episode_id not in self._episodes: raise ValueError(f"Unknown episode_id: {episode_id}") ep = self._episodes[episode_id] if ep["resolved"]: base = 0.999 if ep["step_count"] == 1 else max(0.5, 0.999 - 0.15 * (ep["step_count"] - 1)) return {"score": round(base, 3), "resolved": True, "steps": ep["step_count"]} else: task = TASKS[ep["task_id"]] dangerous_taken = [a for a in ep["actions_taken"] if a in task["dangerous_actions"]] if dangerous_taken: return {"score": 0.001, "resolved": False, "steps": ep["step_count"]} else: return {"score": 0.15, "resolved": False, "steps": ep["step_count"]}