From ef946bec3cc781d9161f4c2093b1968e70df0233 Mon Sep 17 00:00:00 2001 From: malmert Date: Thu, 27 Nov 2025 18:18:31 +0100 Subject: [PATCH] fix: refactor network sniffer for improved structure and performance --- backend/src/network_sniffer.py | 460 ++++++++++----------------------- 1 file changed, 141 insertions(+), 319 deletions(-) diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index 9e953d6..f3ebff8 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -1,101 +1,36 @@ -""" -network_sniffer.py - -Usage: -- import start_sniffing, stop_sniffing, get_sniffer_status from this module in your FastAPI routes. -- start_sniffing(bridge_name) will spawn per-interface sniffing threads for all bridge ports. -- stop_sniffing() will stop all running sniffer threads cleanly. -""" - -import asyncio +import socket +import struct import threading -import logging -import os -from typing import List, Dict, Optional - -from scapy.all import sniff, Ether, IP # scapy must be installed +import asyncio import asyncpg +import os +import logging +from typing import Dict, List -logger = logging.getLogger("network_sniffer") logging.basicConfig(level=logging.INFO) +logger = logging.getLogger("afpacket_sniffer") -# ----- CONFIG ----- -DB_DSN = os.getenv("MITM_DB_DSN", "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db") -TABLE_NAME = os.getenv("MITM_TABLE", "packets") # change if you use a different table -# ------------------ +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] = {} -# map which interface belongs to which bridge (reverse mapping) -iface_to_bridge: Dict[str, str] = {} -# Bridge→ports cache +# Cache for bridge -> ports bridge_ports_cache: Dict[str, List[str]] = {} -# Async loop + DB pool (run in background thread) -_async_loop: Optional[asyncio.AbstractEventLoop] = None -_db_pool: Optional[asyncpg.pool.Pool] = None -_loop_thread: Optional[threading.Thread] = None -_loop_started_event = threading.Event() +# ------------------------------------------------------------------- +# ASYNC LOOP FOR DB INSERTS +# ------------------------------------------------------------------- +async_loop = asyncio.new_event_loop() +threading.Thread(target=lambda: async_loop.run_forever(), daemon=True).start() - -# ------------------------ -# Async loop bootstrap -# ------------------------ -def _start_async_loop(loop: asyncio.AbstractEventLoop): - """Run the given event loop forever (target for background thread).""" - asyncio.set_event_loop(loop) - _loop_started_event.set() - loop.run_forever() - - -def ensure_async_loop(): - """Create and start the dedicated asyncio loop once.""" - global _async_loop, _loop_thread - if _async_loop is not None: - return _async_loop - - _async_loop = asyncio.new_event_loop() - _loop_thread = threading.Thread(target=_start_async_loop, args=(_async_loop,), daemon=True) - _loop_thread.start() - # wait for loop to be set in thread - _loop_started_event.wait(timeout=5) - if not _loop_started_event.is_set(): - raise RuntimeError("Failed to start async loop thread") - logger.info("Async loop started in background thread") - return _async_loop - - -async def _create_db_pool(): - """Coroutine that creates asyncpg pool. Run this in the background loop.""" - global _db_pool - if _db_pool is None: - logger.info("Creating asyncpg pool...") - _db_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=6) - logger.info("DB pool created") - return _db_pool - - -def ensure_db_pool(): - """Ensure pool exists by scheduling creation on background loop.""" - ensure_async_loop() - # schedule coroutine and wait for it to complete - fut = asyncio.run_coroutine_threadsafe(_create_db_pool(), _async_loop) - try: - return fut.result(timeout=10) - except Exception as e: - logger.exception("Failed to create DB pool: %s", e) - raise - - -# ------------------------ -# Bridge helpers -# ------------------------ +# ------------------------------------------------------------------- +# BRIDGE PORT HANDLING +# ------------------------------------------------------------------- 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: @@ -103,299 +38,186 @@ def check_interface_up(iface: str) -> bool: except FileNotFoundError: return False - def get_bridge_ports(bridge: str) -> List[str]: - """ - Return member interfaces of the bridge. Uses /sys/class/net//brif/. - Returns [] if bridge missing or on error. - """ if bridge in bridge_ports_cache: return bridge_ports_cache[bridge] base = f"/sys/class/net/{bridge}/brif/" if not os.path.isdir(base): - logger.error("Bridge %s does not exist (or brif missing)", bridge) + logger.error(f"Bridge '{bridge}' does not exist") return [] + ports = [] try: - ports = [p for p in os.listdir(base) if os.path.isdir(f"/sys/class/net/{p}")] + 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("Permission denied reading bridge ports for %s", bridge) - return [] + logger.error(f"No permission to read bridge '{bridge}' ports") bridge_ports_cache[bridge] = ports - # register iface->bridge mapping - for p in ports: - iface_to_bridge[p] = bridge - logger.info("Bridge %s ports: %s", bridge, ports) + logger.info(f"Bridge {bridge} ports: {ports}") return ports - def determine_direction(pkt_iface: str, bridge: str): - """ - ingress = interface where packet was captured - egress = list of other bridge ports (where it would be forwarded) - """ ports = get_bridge_ports(bridge) ingress = pkt_iface egress = [p for p in ports if p != pkt_iface] return ingress, egress - -# ------------------------ -# DB insert -# ------------------------ -async def _db_insert_packet_coroutine(pkt_info: dict): - """ - Runs in event loop; uses pool to insert. - Matches the schema of packet_log with minimal fields: - interface, direction, src_mac, dst_mac, eth_type, - src_ip, dst_ip, ip_protocol, packet_len, raw_packet - Adjust columns if your table differs. - """ - pool = await _create_db_pool() # ensure pool +# ------------------------------------------------------------------- +# DATABASE INSERTION +# ------------------------------------------------------------------- +async def db_insert_packet(pkt_info: dict): conn = None try: - async with pool.acquire() as conn: - await conn.execute(f""" - INSERT INTO {TABLE_NAME}( - iface, - direction, - src_mac, - dst_mac, - eth_type, - src_ip, - dst_ip, - protocol, - length, - raw - ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) - """, - pkt_info.get("ingress"), - pkt_info.get("direction", "unknown"), - pkt_info.get("src_mac"), - pkt_info.get("dst_mac"), - pkt_info.get("eth_type"), - pkt_info.get("src_ip"), - pkt_info.get("dst_ip"), - pkt_info.get("protocol"), - pkt_info.get("length"), - pkt_info.get("raw") - ) - except Exception: - logger.exception("DB insert failed") - raise + conn = await asyncpg.connect(DB_DSN) + await conn.execute(""" + INSERT INTO packets( + iface, + direction, + src_mac, + dst_mac, + eth_type, + src_ip, + dst_ip, + protocol, + 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"] + ) + except Exception as e: + logger.exception(f"DB insert failed: {e}") + finally: + if conn: + await conn.close() - -def schedule_db_insert(pkt_info: dict): - """ - Thread-safe helper: schedule DB insert into background loop. - """ - try: - ensure_db_pool() - except Exception: - logger.error("DB pool not available; dropping packet") +# ------------------------------------------------------------------- +# PACKET HANDLER +# ------------------------------------------------------------------- +def handle_packet(pkt_bytes: bytes, iface: str, bridge: str): + # Ethernet header + if len(pkt_bytes) < 14: 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) - # schedule coroutine - loop = ensure_async_loop() - fut = asyncio.run_coroutine_threadsafe(_db_insert_packet_coroutine(pkt_info), loop) - # optional: attach callback to log failures asynchronously - def _on_done(f): - try: - f.result() - except Exception: - logger.exception("Async DB insert failed") - fut.add_done_callback(_on_done) + # 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] - -# ------------------------ -# Packet handling -# ------------------------ -def _safe_get_src_dst(pkt): - """Return (src_ip, dst_ip, proto) safely if IP present.""" - try: - if IP in pkt: - return pkt[IP].src, pkt[IP].dst, pkt[IP].proto - except Exception: - pass - return None, None, None - - -def handle_packet(pkt, bridge: str): - """ - Called in sniffing thread. Build pkt_info and schedule DB insert. - """ - pkt_iface = getattr(pkt, "sniffed_on", None) - if not pkt_iface: - # scapy sometimes doesn't set sniffed_on; skip if unknown - logger.warning("Packet without sniffed_on - ignoring") - return - - ingress, egress = determine_direction(pkt_iface, bridge) - src_ip, dst_ip, proto = _safe_get_src_dst(pkt) + ingress, egress = determine_direction(iface, bridge) pkt_info = { - "ingress": ingress, + "iface": iface, "egress": egress, - "direction": "ingress", # we store 'ingress' here (packet was received on ingress) - "src_mac": pkt[Ether].src if Ether in pkt else None, - "dst_mac": pkt[Ether].dst if Ether in pkt else None, - "eth_type": pkt[Ether].type if Ether in pkt else None, + "src_mac": src_mac, + "dst_mac": dst_mac, + "eth_type": hex(eth_type), "src_ip": src_ip, "dst_ip": dst_ip, - "protocol": proto, - "length": len(pkt), - # store raw as binary; use bytes(pkt) - "raw": bytes(pkt) + "protocol": protocol, + "length": len(pkt_bytes), + "raw": pkt_bytes } - # schedule asynchronous DB insert - schedule_db_insert(pkt_info) + 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}") -# ------------------------ -# Sniffer thread loop -# ------------------------ -def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): - logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge) - - if not check_interface_exists(ifname): - logger.error("Interface %s does not exist; stopping sniffer thread", ifname) + 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.") return - if not check_interface_up(ifname): - logger.warning("Interface %s is down; sniffing might still work but check interface state", ifname) + 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 - # loop with short timeout so stop_event is checked often while not stop_event.is_set(): try: - sniff( - iface=ifname, - prn=lambda pkt: handle_packet(pkt, bridge), - store=False, - timeout=1 # returns periodically so we can check stop_event - ) - except PermissionError: - logger.exception("Permission denied sniffing on %s - run as root or give CAP_NET_RAW", ifname) - break - except OSError as e: - logger.exception("OS error while sniffing on %s: %s", ifname, e) - # brief sleep to avoid tight loop on repeated errors - stop_event.wait(1) - break - except Exception: - logger.exception("Unexpected error in sniffer loop for %s", ifname) - stop_event.wait(0.5) - break + pkt, _ = s.recvfrom(65536) + handle_packet(pkt, iface, bridge) + except Exception as e: + logger.exception(f"Error in sniffer loop on {iface}: {e}") - logger.info("Sniffer thread exiting for %s", ifname) + s.close() + logger.info(f"Sniffer STOPPED on {iface}") - -# ------------------------ -# Control API (importable) -# ------------------------ -def start_sniffer_thread_for_bridge(bridge: str) -> Dict[str, threading.Thread]: - """ - Start sniff threads for all ports of given bridge. - Returns mapping iface->thread for started threads (existing threads are left running). - """ +# ------------------------------------------------------------------- +# START / STOP METHODS +# ------------------------------------------------------------------- +def start_sniffer_thread(bridge: str): ports = get_bridge_ports(bridge) if not ports: - logger.error("No ports for bridge %s - nothing to start", bridge) + logger.error(f"No valid ports found for bridge {bridge}") return {} - started = {} for iface in ports: - if iface in sniffer_threads and sniffer_threads[iface].is_alive(): - logger.info("Sniffer already running on %s", iface) - started[iface] = sniffer_threads[iface] + if iface in sniffer_threads: continue - - stop_evt = threading.Event() - thread = threading.Thread(target=sniffer_loop, args=(iface, stop_evt, bridge), daemon=True) - thread.start() + 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_stop_flags[iface] = stop_evt - started[iface] = thread - logger.info("Started sniffer on %s (bridge=%s)", iface, bridge) + thread.start() + return sniffer_threads - return started - - -async def start_sniffing(bridge: str) -> bool: - """ - Async entrypoint you can call from FastAPI startup/route. - """ - ensure_async_loop() - ensure_db_pool() - start_sniffer_thread_for_bridge(bridge) +async def start_sniffing(bridge: str): + start_sniffer_thread(bridge) return True - -async def stop_sniffing() -> bool: - """ - Stop all sniffers and wait briefly for threads to exit. - """ - logger.info("Stopping sniffers (all interfaces)...") - for iface, evt in list(thread_stop_flags.items()): - evt.set() - - # join threads with timeout - for iface, thread in list(sniffer_threads.items()): - logger.info("Joining thread for %s", iface) +async def stop_sniffing(): + for stop_event in thread_stop_flags.values(): + stop_event.set() + for thread in sniffer_threads.values(): thread.join(timeout=2) - sniffer_threads.clear() thread_stop_flags.clear() - logger.info("All sniffer threads stopped") return True - -def get_sniffer_status() -> Dict[str, Dict[str, object]]: - """ - Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] } - """ - status = {} - # include known interfaces (from cache) and active threads - known_ifaces = set(list(sniffer_threads.keys()) + list(iface_to_bridge.keys())) - for iface in known_ifaces: - thread = sniffer_threads.get(iface) - status[iface] = { - "running": bool(thread and thread.is_alive()), +# ------------------------------------------------------------------- +# STATUS HELPER +# ------------------------------------------------------------------- +def get_sniffer_status(): + out = {} + for iface, thread in sniffer_threads.items(): + out[iface] = { + "running": thread.is_alive(), "exists": check_interface_exists(iface), - "up": check_interface_up(iface), - "bridge": iface_to_bridge.get(iface) + "up": check_interface_up(iface) } - return status - - -# ------------------------ -# Optional convenience: stop all on process exit -# ------------------------ -def _cleanup_on_exit(): - try: - # stop sniffers - for evt in thread_stop_flags.values(): - evt.set() - for t in sniffer_threads.values(): - t.join(timeout=1) - except Exception: - pass - # shutdown db pool and loop - if _db_pool is not None: - try: - fut = asyncio.run_coroutine_threadsafe(_db_pool.close(), _async_loop) - fut.result(timeout=5) - except Exception: - pass - if _async_loop is not None: - try: - _async_loop.call_soon_threadsafe(_async_loop.stop) - except Exception: - pass - - -# register cleanup (best-effort) -import atexit -atexit.register(_cleanup_on_exit) + return out