#!/usr/bin/env python3 """Network topology mapper — aggregates host and VLAN data into graph output. Subscribes to HOST_DISCOVERED and VLAN_DETECTED events. Queries the state database for host_discovery, network_mapper, and vlan_discovery data. Generates Graphviz DOT output with nodes color-coded by OS family and shaped by role (workstation, server, printer, router, switch, AP). """ import json import logging import os import subprocess import threading import time from collections import defaultdict from typing import Optional from modules.base import BaseModule logger = logging.getLogger("sensor.intel.topology_mapper") # OS family -> Graphviz fill color _OS_COLORS = { "windows": "#4488cc", "linux": "#44aa44", "macos": "#888888", "ios": "#aaaaaa", "android": "#88bb44", "freebsd": "#cc4444", "printer": "#ddaa44", "network": "#cc8844", "iot": "#aa66cc", "unknown": "#cccccc", } # Device role -> Graphviz node shape _ROLE_SHAPES = { "workstation": "box", "server": "box3d", "printer": "tab", "router": "diamond", "switch": "hexagon", "ap": "trapezium", "phone": "ellipse", "iot": "component", "unknown": "ellipse", } # Protocol -> edge color _PROTO_COLORS = { "tcp": "#333333", "udp": "#666666", "smb": "#4488cc", "http": "#44aa44", "https": "#228822", "dns": "#cc8844", "ssh": "#cc4444", "rdp": "#4444cc", "ldap": "#aa66cc", "kerberos": "#ddaa44", } class TopologyMapper(BaseModule): """Build and export a network topology graph from discovered hosts and links.""" name = "topology_mapper" module_type = "intel" priority = 150 requires_root = False def __init__(self, bus, state, config, engine=None): super().__init__(bus, state, config, engine) self._lock = threading.Lock() # In-memory graph: nodes keyed by IP, edges as (src, dst) -> {protocols, bytes} self._nodes: dict[str, dict] = {} self._edges: dict[tuple, dict] = defaultdict(lambda: { "protocols": set(), "bytes_total": 0, "conn_count": 0, }) self._vlans: dict[int, dict] = {} self._update_thread: Optional[threading.Thread] = None # ------------------------------------------------------------------ # BaseModule interface # ------------------------------------------------------------------ def start(self) -> None: if self._running: return self.bus.subscribe(self._on_host_discovered, "HOST_DISCOVERED") self.bus.subscribe(self._on_vlan_detected, "VLAN_DETECTED") # Load any existing topology data from state self._load_from_state() # Periodic refresh from state database self._running = True self._pid = os.getpid() self._start_time = time.time() self._update_thread = threading.Thread( target=self._periodic_update, daemon=True, name="sensor-topology-update", ) self._update_thread.start() self.state.set_module_status(self.name, "running", pid=self._pid) logger.info( "TopologyMapper started — %d nodes, %d edges loaded from state", len(self._nodes), len(self._edges), ) def stop(self) -> None: if not self._running: return self._running = False self.bus.unsubscribe(self._on_host_discovered, "HOST_DISCOVERED") self.bus.unsubscribe(self._on_vlan_detected, "VLAN_DETECTED") # Persist current topology self._save_to_state() self.state.set_module_status(self.name, "stopped") logger.info( "TopologyMapper stopped — %d nodes, %d edges", len(self._nodes), len(self._edges), ) def status(self) -> dict: return { "running": self._running, "pid": self._pid, "uptime": time.time() - self._start_time if self._start_time else 0, "nodes": len(self._nodes), "edges": len(self._edges), "vlans": len(self._vlans), } def configure(self, config: dict) -> None: self.config.update(config) # ------------------------------------------------------------------ # Event handlers # ------------------------------------------------------------------ def _on_host_discovered(self, event) -> None: """Ingest a newly discovered host.""" p = event.payload ip = p.get("ip", "") if not ip: return with self._lock: if ip in self._nodes: # Merge new data into existing node node = self._nodes[ip] node["last_seen"] = time.time() if p.get("hostname"): node["hostname"] = p["hostname"] if p.get("os_family"): node["os_family"] = p["os_family"].lower() if p.get("role"): node["role"] = p["role"].lower() if p.get("mac"): node["mac"] = p["mac"] if p.get("open_ports"): node.setdefault("open_ports", set()).update(p["open_ports"]) if p.get("services"): node.setdefault("services", []).extend(p["services"]) else: open_ports = p.get("open_ports", []) self._nodes[ip] = { "ip": ip, "hostname": p.get("hostname", ""), "mac": p.get("mac", ""), "os_family": p.get("os_family", "unknown").lower(), "role": p.get("role", "unknown").lower(), "vlan": p.get("vlan"), "open_ports": set(open_ports) if isinstance(open_ports, list) else open_ports, "services": p.get("services", []), "first_seen": time.time(), "last_seen": time.time(), } def _on_vlan_detected(self, event) -> None: """Record a detected VLAN.""" p = event.payload vlan_id = p.get("vlan_id") if vlan_id is None: return with self._lock: self._vlans[vlan_id] = { "vlan_id": vlan_id, "name": p.get("name", ""), "subnet": p.get("subnet", ""), "gateway": p.get("gateway", ""), "first_seen": time.time(), } # ------------------------------------------------------------------ # Graph data aggregation # ------------------------------------------------------------------ def add_edge(self, src_ip: str, dst_ip: str, protocol: str = "tcp", bytes_transferred: int = 0) -> None: """Add or update an edge in the topology graph.""" key = (src_ip, dst_ip) if src_ip < dst_ip else (dst_ip, src_ip) with self._lock: edge = self._edges[key] edge["protocols"].add(protocol.lower()) edge["bytes_total"] += bytes_transferred edge["conn_count"] += 1 def _periodic_update(self) -> None: """Periodically refresh topology from state database.""" interval = self.config.get("topology_refresh_interval", 300) while self._running: time.sleep(interval) try: self._refresh_from_state() self._save_to_state() except Exception: logger.exception("Periodic topology update failed") def _refresh_from_state(self) -> None: """Pull host and network data from state key-value store.""" # Pull host list from host_discovery module state hosts_json = self.state.get("host_discovery", "discovered_hosts") if hosts_json: try: hosts = json.loads(hosts_json) for host in hosts: ip = host.get("ip", "") if ip and ip not in self._nodes: with self._lock: self._nodes[ip] = { "ip": ip, "hostname": host.get("hostname", ""), "mac": host.get("mac", ""), "os_family": host.get("os_family", "unknown").lower(), "role": self._infer_role(host), "vlan": host.get("vlan"), "open_ports": set(host.get("open_ports", [])), "services": host.get("services", []), "first_seen": host.get("first_seen", time.time()), "last_seen": host.get("last_seen", time.time()), } except (json.JSONDecodeError, TypeError): pass # Pull VLAN data from vlan_discovery module state vlans_json = self.state.get("vlan_discovery", "vlans") if vlans_json: try: vlans = json.loads(vlans_json) for vlan in vlans: vid = vlan.get("vlan_id") if vid is not None: with self._lock: self._vlans[vid] = vlan except (json.JSONDecodeError, TypeError): pass def _infer_role(self, host: dict) -> str: """Infer device role from open ports and services.""" ports = set(host.get("open_ports", [])) hostname = host.get("hostname", "").lower() # Router/gateway indicators if ports & {179, 520, 521}: # BGP, RIP return "router" if any(kw in hostname for kw in ("gw", "gateway", "router", "rtr")): return "router" # Switch indicators if any(kw in hostname for kw in ("switch", "sw-", "sw0")): return "switch" # AP indicators if any(kw in hostname for kw in ("ap-", "wap", "access-point")): return "ap" # Printer indicators if ports & {515, 631, 9100}: # LPD, IPP, RAW return "printer" if any(kw in hostname for kw in ("print", "prn", "hp", "canon", "brother")): return "printer" # Server indicators (multiple service ports) server_ports = {22, 25, 53, 80, 443, 389, 636, 88, 445, 3389, 5432, 3306, 1433, 8080, 8443} if len(ports & server_ports) >= 3: return "server" if any(kw in hostname for kw in ("srv", "server", "dc", "dns", "mail", "web", "db")): return "server" return "workstation" # ------------------------------------------------------------------ # State persistence # ------------------------------------------------------------------ def _load_from_state(self) -> None: """Load saved topology from state on startup.""" nodes_json = self.state.get(self.name, "nodes") if nodes_json: try: nodes = json.loads(nodes_json) for ip, node in nodes.items(): node["open_ports"] = set(node.get("open_ports", [])) self._nodes[ip] = node except (json.JSONDecodeError, TypeError): pass edges_json = self.state.get(self.name, "edges") if edges_json: try: edges = json.loads(edges_json) for key_str, edge_data in edges.items(): parts = key_str.split("|") if len(parts) == 2: key = (parts[0], parts[1]) edge_data["protocols"] = set(edge_data.get("protocols", [])) self._edges[key] = edge_data except (json.JSONDecodeError, TypeError): pass vlans_json = self.state.get(self.name, "vlans") if vlans_json: try: self._vlans = json.loads(vlans_json) except (json.JSONDecodeError, TypeError): pass def _save_to_state(self) -> None: """Persist topology to state database.""" with self._lock: # Serialize nodes (convert sets to lists) nodes_ser = {} for ip, node in self._nodes.items(): n = dict(node) n["open_ports"] = list(n.get("open_ports", set())) nodes_ser[ip] = n # Serialize edges (convert tuple keys and sets) edges_ser = {} for (src, dst), edge in self._edges.items(): key = f"{src}|{dst}" edges_ser[key] = { "protocols": list(edge.get("protocols", set())), "bytes_total": edge.get("bytes_total", 0), "conn_count": edge.get("conn_count", 0), } self.state.set(self.name, "nodes", json.dumps(nodes_ser)) self.state.set(self.name, "edges", json.dumps(edges_ser)) self.state.set(self.name, "vlans", json.dumps(self._vlans)) # ------------------------------------------------------------------ # DOT / SVG generation # ------------------------------------------------------------------ def generate_dot(self) -> str: """Generate Graphviz DOT representation of the network topology.""" lines = [ "digraph SystemMonitorTopology {", ' rankdir=LR;', ' bgcolor="#1a1a2e";', ' node [style=filled, fontcolor=white, fontname="Courier"];', ' edge [fontname="Courier", fontsize=9];', "", ] # Group nodes by VLAN vlan_groups = defaultdict(list) no_vlan = [] with self._lock: for ip, node in self._nodes.items(): vlan = node.get("vlan") if vlan is not None: vlan_groups[vlan].append((ip, node)) else: no_vlan.append((ip, node)) # VLAN subgraphs for vlan_id, members in sorted(vlan_groups.items()): vlan_info = self._vlans.get(vlan_id, {}) vlan_label = vlan_info.get("name", f"VLAN {vlan_id}") subnet = vlan_info.get("subnet", "") label = f"{vlan_label}" if subnet: label += f"\\n{subnet}" lines.append(f" subgraph cluster_vlan{vlan_id} {{") lines.append(f' label="{label}";') lines.append(' style=dashed;') lines.append(' color="#666666";') lines.append(' fontcolor="#aaaaaa";') for ip, node in members: lines.append(f" {self._node_dot(ip, node)}") lines.append(" }") lines.append("") # Ungrouped nodes for ip, node in no_vlan: lines.append(f" {self._node_dot(ip, node)}") lines.append("") # Edges for (src, dst), edge in self._edges.items(): protos = ", ".join(sorted(edge.get("protocols", set()))) count = edge.get("conn_count", 0) # Thicker line for higher traffic penwidth = min(1.0 + (count / 100.0), 5.0) # Color by primary protocol primary_proto = next(iter(edge.get("protocols", {"tcp"})), "tcp") color = _PROTO_COLORS.get(primary_proto, "#666666") label = protos if count > 10: label += f"\\n({count})" src_id = self._sanitize_id(src) dst_id = self._sanitize_id(dst) lines.append( f' {src_id} -> {dst_id} ' f'[label="{label}", color="{color}", ' f'penwidth={penwidth:.1f}];' ) lines.append("}") return "\n".join(lines) def _node_dot(self, ip: str, node: dict) -> str: """Generate DOT for a single node.""" os_family = node.get("os_family", "unknown") role = node.get("role", "unknown") hostname = node.get("hostname", "") color = _OS_COLORS.get(os_family, _OS_COLORS["unknown"]) shape = _ROLE_SHAPES.get(role, _ROLE_SHAPES["unknown"]) label = ip if hostname: label = f"{hostname}\\n{ip}" ports = node.get("open_ports", set()) if ports: port_str = ", ".join(str(p) for p in sorted(ports)[:8]) if len(ports) > 8: port_str += f" (+{len(ports) - 8})" label += f"\\n[{port_str}]" node_id = self._sanitize_id(ip) return ( f'{node_id} [label="{label}", shape={shape}, ' f'fillcolor="{color}", tooltip="{os_family}/{role}"];' ) @staticmethod def _sanitize_id(ip: str) -> str: """Convert IP to valid DOT node ID.""" return "n_" + ip.replace(".", "_").replace(":", "_") def generate_svg(self, output_path: str = None) -> Optional[str]: """Generate SVG from the DOT graph (requires graphviz installed). Args: output_path: Path to write SVG file. If None, returns SVG as string. Returns: SVG string if output_path is None, else the output path. """ dot = self.generate_dot() try: result = subprocess.run( ["dot", "-Tsvg"], input=dot.encode(), capture_output=True, timeout=30, ) if result.returncode != 0: logger.error("Graphviz dot failed: %s", result.stderr.decode()) return None svg = result.stdout.decode() if output_path: with open(output_path, "w") as f: f.write(svg) return output_path return svg except FileNotFoundError: logger.warning("Graphviz not installed — SVG generation unavailable") return None except subprocess.TimeoutExpired: logger.error("Graphviz timed out generating SVG") return None def get_topology_summary(self) -> dict: """Return a summary of the current network topology.""" with self._lock: roles = defaultdict(int) os_families = defaultdict(int) for node in self._nodes.values(): roles[node.get("role", "unknown")] += 1 os_families[node.get("os_family", "unknown")] += 1 protocols_seen = set() for edge in self._edges.values(): protocols_seen.update(edge.get("protocols", set())) return { "total_nodes": len(self._nodes), "total_edges": len(self._edges), "total_vlans": len(self._vlans), "by_role": dict(roles), "by_os": dict(os_families), "protocols": sorted(protocols_seen), "vlan_ids": sorted(self._vlans.keys()), }