diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index f3ebff8..d576e7c 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -1,36 +1,55 @@ -import socket -import struct -import threading import asyncio -import asyncpg -import os import logging -from typing import Dict, List +import threading +import os +from typing import List, Dict + +import asyncpg +from scapy.all import ( + Ether, + ARP, + IP, + IPv6, + TCP, + UDP, + ICMP, + ICMPv6Unknown, + Dot1Q, + Raw, + sniff, + conf, +) logging.basicConfig(level=logging.INFO) -logger = logging.getLogger("afpacket_sniffer") +logger = logging.getLogger("af_packet_sniffer") DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" -# Active sniffer threads and stop flags sniffer_threads: Dict[str, threading.Thread] = {} thread_stop_flags: Dict[str, threading.Event] = {} - -# Cache for bridge -> ports bridge_ports_cache: Dict[str, List[str]] = {} # ------------------------------------------------------------------- -# ASYNC LOOP FOR DB INSERTS +# Async loop for DB inserts # ------------------------------------------------------------------- async_loop = asyncio.new_event_loop() -threading.Thread(target=lambda: async_loop.run_forever(), daemon=True).start() + + +def start_async_loop(loop): + asyncio.set_event_loop(loop) + loop.run_forever() + + +threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start() + # ------------------------------------------------------------------- -# BRIDGE PORT HANDLING +# Interface / bridge checks # ------------------------------------------------------------------- def check_interface_exists(iface: str) -> bool: return os.path.isdir(f"/sys/class/net/{iface}") + def check_interface_up(iface: str) -> bool: try: with open(f"/sys/class/net/{iface}/operstate", "r") as f: @@ -38,6 +57,7 @@ def check_interface_up(iface: str) -> bool: except FileNotFoundError: return False + def get_bridge_ports(bridge: str) -> List[str]: if bridge in bridge_ports_cache: return bridge_ports_cache[bridge] @@ -47,59 +67,57 @@ def get_bridge_ports(bridge: str) -> List[str]: logger.error(f"Bridge '{bridge}' does not exist") return [] - ports = [] - try: - for p in os.listdir(base): - if check_interface_exists(p): - ports.append(p) - else: - logger.warning(f"Port '{p}' listed in bridge but does not exist") - except PermissionError: - logger.error(f"No permission to read bridge '{bridge}' ports") - + ports = [p for p in os.listdir(base) if check_interface_exists(p)] bridge_ports_cache[bridge] = ports logger.info(f"Bridge {bridge} ports: {ports}") return ports + def determine_direction(pkt_iface: str, bridge: str): ports = get_bridge_ports(bridge) ingress = pkt_iface egress = [p for p in ports if p != pkt_iface] return ingress, egress + # ------------------------------------------------------------------- -# DATABASE INSERTION +# Database insertion # ------------------------------------------------------------------- async def db_insert_packet(pkt_info: dict): conn = None try: conn = await asyncpg.connect(DB_DSN) - await conn.execute(""" + await conn.execute( + """ INSERT INTO packets( iface, direction, src_mac, dst_mac, eth_type, + vlan_id, src_ip, dst_ip, - protocol, + ip_proto, + src_port, + dst_port, length, - ebpf_verdict, raw - ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) - """, - pkt_info["iface"], - "unknown", - pkt_info["src_mac"], - pkt_info["dst_mac"], - pkt_info["eth_type"], - pkt_info["src_ip"], - pkt_info["dst_ip"], - pkt_info["protocol"], - pkt_info["length"], - str(pkt_info["egress"]), - pkt_info["raw"] + ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13) + """, + pkt_info["ingress"], + "unknown", + pkt_info["src_mac"], + pkt_info["dst_mac"], + pkt_info["eth_type"], + pkt_info.get("vlan_id"), + pkt_info.get("src_ip"), + pkt_info.get("dst_ip"), + pkt_info.get("protocol_name"), + pkt_info.get("src_port"), + pkt_info.get("dst_port"), + pkt_info["length"], + pkt_info["raw"], ) except Exception as e: logger.exception(f"DB insert failed: {e}") @@ -107,78 +125,123 @@ async def db_insert_packet(pkt_info: dict): if conn: await conn.close() + # ------------------------------------------------------------------- -# PACKET HANDLER +# Packet parsing: AF_PACKET / full stack # ------------------------------------------------------------------- -def handle_packet(pkt_bytes: bytes, iface: str, bridge: str): - # Ethernet header - if len(pkt_bytes) < 14: +def parse_packet(pkt, bridge: str): + pkt_iface = getattr(pkt, "sniffed_on", None) + if not pkt_iface: return - eth_header = pkt_bytes[:14] - dst_mac, src_mac, eth_type = struct.unpack("!6s6sH", eth_header) - dst_mac = ':'.join('%02x' % b for b in dst_mac) - src_mac = ':'.join('%02x' % b for b in src_mac) - eth_type = socket.ntohs(eth_type) - # IP header - src_ip = dst_ip = None - protocol = None - if eth_type == 0x0800 and len(pkt_bytes) >= 34: - ip_header = pkt_bytes[14:34] - iph = struct.unpack('!BBHHHBBH4s4s', ip_header) - src_ip = socket.inet_ntoa(iph[8]) - dst_ip = socket.inet_ntoa(iph[9]) - protocol = iph[6] - - ingress, egress = determine_direction(iface, bridge) + ingress, egress = determine_direction(pkt_iface, bridge) pkt_info = { - "iface": iface, + "ingress": ingress, "egress": egress, - "src_mac": src_mac, - "dst_mac": dst_mac, - "eth_type": hex(eth_type), - "src_ip": src_ip, - "dst_ip": dst_ip, - "protocol": protocol, - "length": len(pkt_bytes), - "raw": pkt_bytes + "length": len(pkt), + "raw": bytes(pkt), + "src_mac": None, + "dst_mac": None, + "eth_type": None, + "vlan_id": None, + "src_ip": None, + "dst_ip": None, + "protocol_name": None, + "src_port": None, + "dst_port": None, } + # Ethernet + if Ether in pkt: + pkt_info["src_mac"] = pkt[Ether].src + pkt_info["dst_mac"] = pkt[Ether].dst + pkt_info["eth_type"] = hex(pkt[Ether].type) + + # VLAN + if Dot1Q in pkt: + pkt_info["vlan_id"] = pkt[Dot1Q].vlan + + # ARP + if ARP in pkt: + pkt_info["protocol_name"] = "ARP" + pkt_info["src_ip"] = pkt[ARP].psrc + pkt_info["dst_ip"] = pkt[ARP].pdst + pkt_info["src_port"] = None + pkt_info["dst_port"] = None + + # IPv4 + if IP in pkt: + pkt_info["src_ip"] = pkt[IP].src + pkt_info["dst_ip"] = pkt[IP].dst + proto = pkt[IP].proto + if proto == 6 and TCP in pkt: + pkt_info["protocol_name"] = "TCP" + pkt_info["src_port"] = pkt[TCP].sport + pkt_info["dst_port"] = pkt[TCP].dport + elif proto == 17 and UDP in pkt: + pkt_info["protocol_name"] = "UDP" + pkt_info["src_port"] = pkt[UDP].sport + pkt_info["dst_port"] = pkt[UDP].dport + elif proto == 1 and ICMP in pkt: + pkt_info["protocol_name"] = "ICMP" + else: + pkt_info["protocol_name"] = f"IP_PROTO_{proto}" + + # IPv6 + if IPv6 in pkt: + pkt_info["src_ip"] = pkt[IPv6].src + pkt_info["dst_ip"] = pkt[IPv6].dst + nh = pkt[IPv6].nh + if nh == 6 and TCP in pkt: + pkt_info["protocol_name"] = "TCP" + pkt_info["src_port"] = pkt[TCP].sport + pkt_info["dst_port"] = pkt[TCP].dport + elif nh == 17 and UDP in pkt: + pkt_info["protocol_name"] = "UDP" + pkt_info["src_port"] = pkt[UDP].sport + pkt_info["dst_port"] = pkt[UDP].dport + elif ICMPv6Unknown in pkt: + pkt_info["protocol_name"] = "ICMPv6" + else: + pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}" + + # Raw / unknown + if Raw in pkt and not pkt_info["protocol_name"]: + pkt_info["protocol_name"] = "RAW" + asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop) -# ------------------------------------------------------------------- -# SNIFFER LOOP -# ------------------------------------------------------------------- -def sniffer_loop(iface: str, stop_event: threading.Event, bridge: str): - logger.info(f"Sniffer STARTED on {iface}") - if not check_interface_exists(iface): - logger.error(f"Interface {iface} does not exist. Exiting sniffer.") - return - if not check_interface_up(iface): - logger.error(f"Interface {iface} is DOWN. Exiting sniffer.") +# ------------------------------------------------------------------- +# Sniffer thread using AF_PACKET +# ------------------------------------------------------------------- +def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): + logger.info(f"Sniffer STARTED on {ifname}") + + if not check_interface_exists(ifname) or not check_interface_up(ifname): + logger.error(f"Interface {ifname} does not exist or is down. Stopping sniffer.") return - try: - s = socket.socket(socket.AF_PACKET, socket.SOCK_RAW, socket.ntohs(3)) - s.bind((iface, 0)) - except PermissionError: - logger.error(f"Permission denied on {iface}, need root") - return + conf.L2socket = conf.L2socket # enforce AF_PACKET usage in scapy while not stop_event.is_set(): try: - pkt, _ = s.recvfrom(65536) - handle_packet(pkt, iface, bridge) + sniff( + iface=ifname, + prn=lambda pkt: parse_packet(pkt, bridge), + store=False, + timeout=0.5, # fast stop checks + ) except Exception as e: - logger.exception(f"Error in sniffer loop on {iface}: {e}") + logger.exception(f"Sniffer error on {ifname}: {e}") + break + + logger.info(f"Sniffer STOPPED on {ifname}") - s.close() - logger.info(f"Sniffer STOPPED on {iface}") # ------------------------------------------------------------------- -# START / STOP METHODS +# Start / Stop / Status API # ------------------------------------------------------------------- def start_sniffer_thread(bridge: str): ports = get_bridge_ports(bridge) @@ -189,35 +252,39 @@ def start_sniffer_thread(bridge: str): for iface in ports: if iface in sniffer_threads: continue + stop_event = threading.Event() thread_stop_flags[iface] = stop_event thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True) sniffer_threads[iface] = thread thread.start() + return sniffer_threads + async def start_sniffing(bridge: str): start_sniffer_thread(bridge) return True + async def stop_sniffing(): - for stop_event in thread_stop_flags.values(): + for iface, stop_event in thread_stop_flags.items(): stop_event.set() - for thread in sniffer_threads.values(): + + for iface, thread in sniffer_threads.items(): thread.join(timeout=2) + sniffer_threads.clear() thread_stop_flags.clear() return True -# ------------------------------------------------------------------- -# STATUS HELPER -# ------------------------------------------------------------------- + def get_sniffer_status(): out = {} - for iface, thread in sniffer_threads.items(): + for iface, t in sniffer_threads.items(): out[iface] = { - "running": thread.is_alive(), + "running": t.is_alive(), "exists": check_interface_exists(iface), - "up": check_interface_up(iface) + "up": check_interface_up(iface), } return out diff --git a/setup_database.sh b/setup_database.sh index b728e54..a6399c4 100755 --- a/setup_database.sh +++ b/setup_database.sh @@ -51,30 +51,36 @@ CREATE TABLE IF NOT EXISTS packets ( id BIGSERIAL PRIMARY KEY, timestamp TIMESTAMPTZ DEFAULT NOW(), - -- interface metadata + -- Interface metadata iface VARCHAR(64), direction VARCHAR(16), -- Ethernet - src_mac VARCHAR(32), - dst_mac VARCHAR(32), - eth_type INTEGER, + src_mac MACADDR, + dst_mac MACADDR, + eth_type VARCHAR(16), - -- IP - src_ip VARCHAR(64), - dst_ip VARCHAR(64), - protocol INTEGER, + -- VLAN + vlan_id INTEGER, + + -- IP layer + src_ip INET, + dst_ip INET, + ip_proto VARCHAR(32), + + -- Transport layer + src_port INTEGER, + dst_port INTEGER, + + -- Packet metadata length INTEGER, - -- eBPF data (reserved for later) - ebpf_verdict VARCHAR(32), - ebpf_chain VARCHAR(64), - - -- full packet + -- Full packet dump raw BYTEA ); EOF + echo "[7] Grant privileges to user…" sudo -u postgres psql -d $DB_NAME <