| import numpy as np |
| import gymnasium as gym |
| from pettingzoo import ParallelEnv |
| from config import Config |
|
|
| class CyberPettingZooEnv(ParallelEnv): |
| metadata = {'render_modes': ['human'], "name": "cyber_network_defense_v2_gnn"} |
|
|
| def __init__(self): |
| super().__init__() |
| |
| self.topology = Config.NETWORK_TOPOLOGY |
| self.total_nodes = Config.TOTAL_HOSTS() |
| self.features_per_node = 2 |
| |
| self.possible_agents = ["red", "blue"] |
| self.agents = self.possible_agents[:] |
| |
| self.observation_spaces = { |
| agent: gym.spaces.Dict({ |
| 'node_features': gym.spaces.Box(low=0.0, high=1.0, shape=(self.total_nodes, self.features_per_node), dtype=np.float32), |
| 'adjacency_matrix': gym.spaces.Box(low=0.0, high=1.0, shape=(self.total_nodes, self.total_nodes), dtype=np.float32) |
| }) for agent in self.possible_agents |
| } |
| |
| self.action_spaces = { |
| 'red': gym.spaces.Discrete(self.total_nodes * 4), |
| 'blue': gym.spaces.Discrete(self.total_nodes * 2) |
| } |
| |
| self.node_id_to_specs = {} |
| self._build_network_topology_map() |
| self._generate_static_adjacency_matrix() |
|
|
| def _build_network_topology_map(self): |
| node_id = 0 |
| for subnet_name, specs in self.topology["subnets"].items(): |
| for _ in range(specs["num_hosts"]): |
| self.node_id_to_specs[node_id] = { |
| "subnet": subnet_name, |
| "base_vulnerability": specs["base_vulnerability"] |
| } |
| node_id += 1 |
|
|
| def _generate_static_adjacency_matrix(self): |
| """Programmatically wires lateral communication pathways between subnets.""" |
| A = np.zeros((self.total_nodes, self.total_nodes), dtype=np.float32) |
| |
| |
| allowed_connections = { |
| "public_dmz": ["public_dmz", "dns_services"], |
| "dns_services": ["public_dmz", "dns_services", "corporate", "active_directory"], |
| "corporate": ["dns_services", "corporate", "active_directory"], |
| "active_directory": ["dns_services", "corporate", "active_directory", "secure_core"], |
| "secure_core": ["active_directory", "secure_core"] |
| } |
| |
| |
| for i in range(self.total_nodes): |
| for j in range(self.total_nodes): |
| sub_i = self.node_id_to_specs[i]["subnet"] |
| sub_j = self.node_id_to_specs[j]["subnet"] |
| if sub_j in allowed_connections[sub_i]: |
| A[i, j] = 1.0 |
| |
| |
| A_tilde = A + np.eye(self.total_nodes, dtype=np.float32) |
| |
| |
| row_sum = np.sum(A_tilde, axis=1) |
| d_inv_sqrt = np.power(row_sum, -0.5, where=row_sum > 0) |
| d_inv_sqrt[row_sum == 0] = 0.0 |
| D_inv_sqrt = np.diag(d_inv_sqrt) |
| |
| self.normalized_adj = np.matmul(np.matmul(D_inv_sqrt, A_tilde), D_inv_sqrt).astype(np.float32) |
|
|
| def reset(self, seed=None, options=None): |
| self.agents = self.possible_agents[:] |
| self.state_matrix = np.zeros((self.total_nodes, self.features_per_node), dtype=np.float32) |
| |
| obs_dict = { |
| 'node_features': self.state_matrix.copy(), |
| 'adjacency_matrix': self.normalized_adj.copy() |
| } |
| observations = {'red': obs_dict, 'blue': obs_dict} |
| infos = {'score_difference': 0.0, 'team_progress': 0.0} |
| return observations, infos |
|
|
| def _is_subnet_compromised(self, subnet_name): |
| for node_id, specs in self.node_id_to_specs.items(): |
| if specs["subnet"] == subnet_name and self.state_matrix[node_id, 0] == 1.0: |
| return True |
| return False |
|
|
| def step(self, actions): |
| if not actions: |
| self.agents = [] |
| return {}, {}, {}, {}, {} |
|
|
| red_action = actions.get('red', 0) |
| blue_action = actions.get('blue', 0) |
|
|
| red_target_node = red_action // 4 |
| red_action_type = red_action % 4 |
| blue_target_node = blue_action // 2 |
| blue_action_type = blue_action % 2 |
|
|
| red_target_node = min(max(0, red_target_node), self.total_nodes - 1) |
| blue_target_node = min(max(0, blue_target_node), self.total_nodes - 1) |
|
|
| red_target_specs = self.node_id_to_specs[red_target_node] |
| red_target_subnet = red_target_specs["subnet"] |
|
|
| |
| if blue_action_type == 1: |
| self.state_matrix[blue_target_node, 1] = 1.0 |
|
|
| |
| if red_action_type == 1: |
| if self.state_matrix[red_target_node, 1] == 0.0: |
| exploit_probability = red_target_specs["base_vulnerability"] |
| |
| if red_target_subnet == "secure_core": |
| if not self._is_subnet_compromised("active_directory"): |
| exploit_probability = 0.0 |
| elif red_target_subnet == "active_directory": |
| if self._is_subnet_compromised("dns_services"): |
| exploit_probability = min(1.0, exploit_probability * 2.0) |
| |
| if np.random.rand() < exploit_probability: |
| self.state_matrix[red_target_node, 0] = 1.0 |
|
|
| |
| total_compromised = np.sum(self.state_matrix[:, 0]) |
| total_patched = np.sum(self.state_matrix[:, 1]) |
| score_diff = float(total_compromised - total_patched) |
| progress = float(total_compromised / self.total_nodes) if self.total_nodes > 0 else 0.0 |
|
|
| obs_dict = { |
| 'node_features': self.state_matrix.copy(), |
| 'adjacency_matrix': self.normalized_adj.copy() |
| } |
| observations = {'red': obs_dict, 'blue': obs_dict} |
| rewards = {'red': 0.0, 'blue': 0.0} |
| terminations = {'red': False, 'blue': False} |
| truncations = {'red': False, 'blue': False} |
| infos = {'score_difference': score_diff, 'team_progress': progress} |
|
|
| return observations, rewards, terminations, truncations, infos |