fix: refactor network sniffer for improved structure and performance
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s

This commit is contained in:
2025-11-27 18:18:31 +01:00
parent bdf6f53033
commit ef946bec3c

View File

@@ -1,101 +1,36 @@
""" import socket
network_sniffer.py import struct
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 threading import threading
import logging import asyncio
import os
from typing import List, Dict, Optional
from scapy.all import sniff, Ether, IP # scapy must be installed
import asyncpg import asyncpg
import os
import logging
from typing import Dict, List
logger = logging.getLogger("network_sniffer")
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("afpacket_sniffer")
# ----- CONFIG ----- DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
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
# ------------------
# Active sniffer threads and stop flags # Active sniffer threads and stop flags
sniffer_threads: Dict[str, threading.Thread] = {} sniffer_threads: Dict[str, threading.Thread] = {}
thread_stop_flags: Dict[str, threading.Event] = {} 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]] = {} bridge_ports_cache: Dict[str, List[str]] = {}
# Async loop + DB pool (run in background thread) # -------------------------------------------------------------------
_async_loop: Optional[asyncio.AbstractEventLoop] = None # ASYNC LOOP FOR DB INSERTS
_db_pool: Optional[asyncpg.pool.Pool] = None # -------------------------------------------------------------------
_loop_thread: Optional[threading.Thread] = None async_loop = asyncio.new_event_loop()
_loop_started_event = threading.Event() threading.Thread(target=lambda: async_loop.run_forever(), daemon=True).start()
# -------------------------------------------------------------------
# ------------------------ # BRIDGE PORT HANDLING
# 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
# ------------------------
def check_interface_exists(iface: str) -> bool: def check_interface_exists(iface: str) -> bool:
return os.path.isdir(f"/sys/class/net/{iface}") return os.path.isdir(f"/sys/class/net/{iface}")
def check_interface_up(iface: str) -> bool: def check_interface_up(iface: str) -> bool:
try: try:
with open(f"/sys/class/net/{iface}/operstate", "r") as f: with open(f"/sys/class/net/{iface}/operstate", "r") as f:
@@ -103,62 +38,44 @@ def check_interface_up(iface: str) -> bool:
except FileNotFoundError: except FileNotFoundError:
return False return False
def get_bridge_ports(bridge: str) -> List[str]: def get_bridge_ports(bridge: str) -> List[str]:
"""
Return member interfaces of the bridge. Uses /sys/class/net/<bridge>/brif/.
Returns [] if bridge missing or on error.
"""
if bridge in bridge_ports_cache: if bridge in bridge_ports_cache:
return bridge_ports_cache[bridge] return bridge_ports_cache[bridge]
base = f"/sys/class/net/{bridge}/brif/" base = f"/sys/class/net/{bridge}/brif/"
if not os.path.isdir(base): 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 [] return []
ports = []
try: 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: except PermissionError:
logger.error("Permission denied reading bridge ports for %s", bridge) logger.error(f"No permission to read bridge '{bridge}' ports")
return []
bridge_ports_cache[bridge] = ports bridge_ports_cache[bridge] = ports
# register iface->bridge mapping logger.info(f"Bridge {bridge} ports: {ports}")
for p in ports:
iface_to_bridge[p] = bridge
logger.info("Bridge %s ports: %s", bridge, ports)
return ports return ports
def determine_direction(pkt_iface: str, bridge: str): 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) ports = get_bridge_ports(bridge)
ingress = pkt_iface ingress = pkt_iface
egress = [p for p in ports if p != pkt_iface] egress = [p for p in ports if p != pkt_iface]
return ingress, egress return ingress, egress
# -------------------------------------------------------------------
# ------------------------ # DATABASE INSERTION
# DB insert # -------------------------------------------------------------------
# ------------------------ async def db_insert_packet(pkt_info: dict):
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
conn = None conn = None
try: try:
async with pool.acquire() as conn: conn = await asyncpg.connect(DB_DSN)
await conn.execute(f""" await conn.execute("""
INSERT INTO {TABLE_NAME}( INSERT INTO packets(
iface, iface,
direction, direction,
src_mac, src_mac,
@@ -168,234 +85,139 @@ async def _db_insert_packet_coroutine(pkt_info: dict):
dst_ip, dst_ip,
protocol, protocol,
length, length,
ebpf_verdict,
raw raw
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
""", """,
pkt_info.get("ingress"), pkt_info["iface"],
pkt_info.get("direction", "unknown"), "unknown",
pkt_info.get("src_mac"), pkt_info["src_mac"],
pkt_info.get("dst_mac"), pkt_info["dst_mac"],
pkt_info.get("eth_type"), pkt_info["eth_type"],
pkt_info.get("src_ip"), pkt_info["src_ip"],
pkt_info.get("dst_ip"), pkt_info["dst_ip"],
pkt_info.get("protocol"), pkt_info["protocol"],
pkt_info.get("length"), pkt_info["length"],
pkt_info.get("raw") str(pkt_info["egress"]),
pkt_info["raw"]
) )
except Exception: except Exception as e:
logger.exception("DB insert failed") logger.exception(f"DB insert failed: {e}")
raise finally:
if conn:
await conn.close()
# -------------------------------------------------------------------
def schedule_db_insert(pkt_info: dict): # PACKET HANDLER
""" # -------------------------------------------------------------------
Thread-safe helper: schedule DB insert into background loop. def handle_packet(pkt_bytes: bytes, iface: str, bridge: str):
""" # Ethernet header
try: if len(pkt_bytes) < 14:
ensure_db_pool()
except Exception:
logger.error("DB pool not available; dropping packet")
return 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 # IP header
loop = ensure_async_loop() src_ip = dst_ip = None
fut = asyncio.run_coroutine_threadsafe(_db_insert_packet_coroutine(pkt_info), loop) protocol = None
# optional: attach callback to log failures asynchronously if eth_type == 0x0800 and len(pkt_bytes) >= 34:
def _on_done(f): ip_header = pkt_bytes[14:34]
try: iph = struct.unpack('!BBHHHBBH4s4s', ip_header)
f.result() src_ip = socket.inet_ntoa(iph[8])
except Exception: dst_ip = socket.inet_ntoa(iph[9])
logger.exception("Async DB insert failed") protocol = iph[6]
fut.add_done_callback(_on_done)
ingress, egress = determine_direction(iface, bridge)
# ------------------------
# 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)
pkt_info = { pkt_info = {
"ingress": ingress, "iface": iface,
"egress": egress, "egress": egress,
"direction": "ingress", # we store 'ingress' here (packet was received on ingress) "src_mac": src_mac,
"src_mac": pkt[Ether].src if Ether in pkt else None, "dst_mac": dst_mac,
"dst_mac": pkt[Ether].dst if Ether in pkt else None, "eth_type": hex(eth_type),
"eth_type": pkt[Ether].type if Ether in pkt else None,
"src_ip": src_ip, "src_ip": src_ip,
"dst_ip": dst_ip, "dst_ip": dst_ip,
"protocol": proto, "protocol": protocol,
"length": len(pkt), "length": len(pkt_bytes),
# store raw as binary; use bytes(pkt) "raw": pkt_bytes
"raw": bytes(pkt)
} }
# schedule asynchronous DB insert asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop)
schedule_db_insert(pkt_info)
# -------------------------------------------------------------------
# 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):
# Sniffer thread loop logger.error(f"Interface {iface} does not exist. Exiting sniffer.")
# ------------------------ return
def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): if not check_interface_up(iface):
logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge) logger.error(f"Interface {iface} is DOWN. Exiting sniffer.")
if not check_interface_exists(ifname):
logger.error("Interface %s does not exist; stopping sniffer thread", ifname)
return return
if not check_interface_up(ifname): try:
logger.warning("Interface %s is down; sniffing might still work but check interface state", ifname) 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(): while not stop_event.is_set():
try: try:
sniff( pkt, _ = s.recvfrom(65536)
iface=ifname, handle_packet(pkt, iface, bridge)
prn=lambda pkt: handle_packet(pkt, bridge), except Exception as e:
store=False, logger.exception(f"Error in sniffer loop on {iface}: {e}")
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
logger.info("Sniffer thread exiting for %s", ifname) s.close()
logger.info(f"Sniffer STOPPED on {iface}")
# -------------------------------------------------------------------
# ------------------------ # START / STOP METHODS
# Control API (importable) # -------------------------------------------------------------------
# ------------------------ def start_sniffer_thread(bridge: str):
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).
"""
ports = get_bridge_ports(bridge) ports = get_bridge_ports(bridge)
if not ports: 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 {} return {}
started = {}
for iface in ports: for iface in ports:
if iface in sniffer_threads and sniffer_threads[iface].is_alive(): if iface in sniffer_threads:
logger.info("Sniffer already running on %s", iface)
started[iface] = sniffer_threads[iface]
continue continue
stop_event = threading.Event()
stop_evt = threading.Event() thread_stop_flags[iface] = stop_event
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_evt, bridge), daemon=True) thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True)
thread.start()
sniffer_threads[iface] = thread sniffer_threads[iface] = thread
thread_stop_flags[iface] = stop_evt thread.start()
started[iface] = thread return sniffer_threads
logger.info("Started sniffer on %s (bridge=%s)", iface, bridge)
return started async def start_sniffing(bridge: str):
start_sniffer_thread(bridge)
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)
return True return True
async def stop_sniffing():
async def stop_sniffing() -> bool: for stop_event in thread_stop_flags.values():
""" stop_event.set()
Stop all sniffers and wait briefly for threads to exit. for thread in sniffer_threads.values():
"""
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)
thread.join(timeout=2) thread.join(timeout=2)
sniffer_threads.clear() sniffer_threads.clear()
thread_stop_flags.clear() thread_stop_flags.clear()
logger.info("All sniffer threads stopped")
return True return True
# -------------------------------------------------------------------
def get_sniffer_status() -> Dict[str, Dict[str, object]]: # STATUS HELPER
""" # -------------------------------------------------------------------
Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] } def get_sniffer_status():
""" out = {}
status = {} for iface, thread in sniffer_threads.items():
# include known interfaces (from cache) and active threads out[iface] = {
known_ifaces = set(list(sniffer_threads.keys()) + list(iface_to_bridge.keys())) "running": thread.is_alive(),
for iface in known_ifaces:
thread = sniffer_threads.get(iface)
status[iface] = {
"running": bool(thread and thread.is_alive()),
"exists": check_interface_exists(iface), "exists": check_interface_exists(iface),
"up": check_interface_up(iface), "up": check_interface_up(iface)
"bridge": iface_to_bridge.get(iface)
} }
return status return out
# ------------------------
# 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)