90shikhar08's picture
fix: inherit from openenv.core.Environment base class
0903e23
Raw
History Blame Contribute Delete
9.47 kB
"""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': <log_index> 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"]}