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 @@
"""
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/<bridge>/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