fix: refactor network sniffer for improved async handling and logging
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-11-27 17:34:01 +01:00
parent 2d68718bd0
commit c3bef89fa4

View File

@@ -1,426 +1,401 @@
"""
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 asyncio
import threading import threading
import logging import logging
import os import os
import time
from typing import List, Dict, Optional from typing import List, Dict, Optional
from queue import Queue, Full, Empty
from scapy.all import sniff, Ether, IP # keep Scapy usage minimal from scapy.all import sniff, Ether, IP # scapy must be installed
import asyncpg import asyncpg
# --------------------------- logger = logging.getLogger("network_sniffer")
# Configuration
# ---------------------------
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("sniffer")
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" # ----- CONFIG -----
DB_TABLE = "packets" # keep your current table name; change if needed DB_DSN = os.getenv("MITM_DB_DSN", "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db")
TABLE_NAME = os.getenv("MITM_TABLE", "packet_log") # change if you use a different table
# Producer/consumer settings # ------------------
QUEUE_MAXSIZE = 20000 # max packets buffered in memory
BATCH_SIZE = 200 # how many rows to insert at once
FLUSH_INTERVAL = 0.5 # seconds max before flushing partial batch
SNAPSHOT_MAX_BYTES = 256 # how many bytes of the packet to store (truncate raw)
# 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] = {}
# Cache for bridge -> ports # Bridge→ports cache
bridge_ports_cache: Dict[str, List[str]] = {} bridge_ports_cache: Dict[str, List[str]] = {}
# Queue for packet events (thread-safe) # Async loop + DB pool (run in background thread)
packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE) _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 event loop + background consumer task
async_loop = asyncio.new_event_loop()
consumer_task: Optional[asyncio.Task] = None
pg_pool: Optional[asyncpg.Pool] = None
# ------------------------------------------------------------------- # ------------------------
# Async loop bootstrap (runs in background thread) # Async loop bootstrap
# ------------------------------------------------------------------- # ------------------------
def _start_async_loop(loop): def _start_async_loop(loop: asyncio.AbstractEventLoop):
"""Run the given event loop forever (target for background thread)."""
asyncio.set_event_loop(loop) asyncio.set_event_loop(loop)
_loop_started_event.set()
loop.run_forever() loop.run_forever()
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
# ------------------------------------------------------------------- def ensure_async_loop():
# DB helper: create pool & consumer """Create and start the dedicated asyncio loop once."""
# ------------------------------------------------------------------- global _async_loop, _loop_thread
async def _create_pool(): if _async_loop is not None:
global pg_pool return _async_loop
if pg_pool is None:
_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...") logger.info("Creating asyncpg pool...")
pg_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=8) _db_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=6)
return pg_pool logger.info("DB pool created")
return _db_pool
async def _close_pool():
global pg_pool
if pg_pool:
await pg_pool.close()
pg_pool = None
async def _db_consumer_loop(stop_event: asyncio.Event):
"""
Async consumer: drains packet_queue and does batched inserts.
Runs inside async_loop.
"""
await _create_pool()
logger.info("DB consumer started")
stmt = f"""
INSERT INTO {DB_TABLE}(
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)
"""
buffer: List[tuple] = []
last_flush = time.time()
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: try:
while not stop_event.is_set(): return fut.result(timeout=10)
# collect up to BATCH_SIZE from the queue with small wait except Exception as e:
try: logger.exception("Failed to create DB pool: %s", e)
# block briefly for the first item raise
item = packet_queue.get(timeout=FLUSH_INTERVAL)
except Empty:
item = None
if item:
buffer.append(item)
# try to drain additional items without blocking
while len(buffer) < BATCH_SIZE:
try:
item = packet_queue.get_nowait()
buffer.append(item)
except Empty:
break
now = time.time()
# flush if we have enough buffered or timed out
if buffer and (len(buffer) >= BATCH_SIZE or (now - last_flush) >= FLUSH_INTERVAL):
try:
async with pg_pool.acquire() as conn:
async with conn.transaction():
# executemany pattern
await conn.executemany(stmt, buffer)
logger.debug(f"Inserted {len(buffer)} packets")
except Exception as e:
logger.exception(f"DB batch insert failed: {e}")
# in case of failure, drop or requeue? We'll drop to avoid blocking.
# Optionally write to disk or metrics.
finally:
buffer.clear()
last_flush = now
# When stop requested, flush remaining
if buffer:
try:
async with pg_pool.acquire() as conn:
async with conn.transaction():
await conn.executemany(stmt, buffer)
logger.info(f"Flushed final {len(buffer)} packets on shutdown")
except Exception:
logger.exception("Failed flushing final packet buffer on shutdown")
finally:
buffer.clear()
finally:
logger.info("DB consumer stopped")
def _start_consumer(): # ------------------------
"""Schedules the DB consumer to run in the shared async_loop and returns the stop Event.""" # Bridge helpers
stop_event = asyncio.Event() # ------------------------
# create consumer task in async_loop
def _start():
global consumer_task
consumer_task = asyncio.run_coroutine_threadsafe(_db_consumer_loop(stop_event), async_loop)
threading.Thread(target=_start, daemon=True).start()
return stop_event, lambda: consumer_task # returns stop_event and accessor to future
async def _stop_consumer(stop_event: asyncio.Event, consumer_future_accessor):
"""Signal stop_event and wait for consumer to finish."""
logger.info("Stopping DB consumer...")
stop_event.set()
# wait for the consumer coroutine to finish
future = consumer_future_accessor()
if future:
try:
future.result(timeout=5)
except Exception as e:
logger.debug(f"Consumer future termination: {e}")
await _close_pool()
# -------------------------------------------------------------------
# Utility: Bridge port handling & direction
# -------------------------------------------------------------------
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:
return f.read().strip() == "up" return f.read().strip() == "up"
except Exception: 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(f"Bridge '{bridge}' does not exist") logger.error("Bridge %s does not exist (or brif missing)", bridge)
return [] return []
try: try:
ports = os.listdir(base) ports = [p for p in os.listdir(base) if os.path.isdir(f"/sys/class/net/{p}")]
except PermissionError: except PermissionError:
logger.error(f"No permission to read bridge ports for '{bridge}'") logger.error("Permission denied reading bridge ports for %s", bridge)
return [] return []
ok_ports = [] bridge_ports_cache[bridge] = ports
# register iface->bridge mapping
for p in ports: for p in ports:
if check_interface_exists(p): iface_to_bridge[p] = bridge
ok_ports.append(p) logger.info("Bridge %s ports: %s", bridge, ports)
else: return ports
logger.warning(f"Bridge member {p} listed but interface missing in /sys/class/net")
bridge_ports_cache[bridge] = ok_ports
logger.info(f"Bridge {bridge} ports: {ok_ports}")
return ok_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
# -------------------------------------------------------------------
# Producer (sniffer threads) # ------------------------
# ------------------------------------------------------------------- # DB insert
def _make_pkt_tuple(pkt, ingress_iface: str, egress_list): # ------------------------
async def _db_insert_packet_coroutine(pkt_info: dict):
""" """
Convert packet info to a tuple matching DB insert order. Runs in event loop; uses pool to insert.
We store only a small raw snapshot to keep DB small. 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
try:
async with pool.acquire() as conn:
await conn.execute(f"""
INSERT INTO {TABLE_NAME}(
interface,
direction,
src_mac,
dst_mac,
eth_type,
src_ip,
dst_ip,
ip_protocol,
packet_len,
raw_packet
) 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
def schedule_db_insert(pkt_info: dict):
"""
Thread-safe helper: schedule DB insert into background loop.
""" """
try: try:
src_mac = pkt[Ether].src if Ether in pkt else None ensure_db_pool()
dst_mac = pkt[Ether].dst if Ether in pkt else None
eth_type = pkt[Ether].type if Ether in pkt else None
except Exception: except Exception:
src_mac = dst_mac = eth_type = None logger.error("DB pool not available; dropping packet")
return
# 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)
# ------------------------
# Packet handling
# ------------------------
def _safe_get_src_dst(pkt):
"""Return (src_ip, dst_ip, proto) safely if IP present."""
try: try:
src_ip = pkt[IP].src if IP in pkt else None if IP in pkt:
dst_ip = pkt[IP].dst if IP in pkt else None return pkt[IP].src, pkt[IP].dst, pkt[IP].proto
protocol = pkt[IP].proto if IP in pkt else None
except Exception: except Exception:
src_ip = dst_ip = protocol = None pass
return None, None, None
raw = bytes(pkt)[:SNAPSHOT_MAX_BYTES] if pkt is not None else b""
# ebpf_verdict field currently used to store egress as text (you can change schema)
ebpf_verdict_text = str(egress_list) if egress_list else None
return (
ingress_iface,
"unknown", # direction
src_mac,
dst_mac,
str(eth_type) if eth_type is not None else None,
src_ip,
dst_ip,
str(protocol) if protocol is not None else None,
len(raw),
ebpf_verdict_text,
raw
)
def handle_packet(pkt, bridge: str): def handle_packet(pkt, bridge: str):
""" """
Callback running inside Scapy sniff thread. Called in sniffing thread. Build pkt_info and schedule DB insert.
Places a lightweight tuple in packet_queue.
""" """
pkt_iface = getattr(pkt, "sniffed_on", None) pkt_iface = getattr(pkt, "sniffed_on", None)
if not pkt_iface: if not pkt_iface:
logger.debug("Packet without sniffed_on metadata, ignoring") # scapy sometimes doesn't set sniffed_on; skip if unknown
logger.warning("Packet without sniffed_on - ignoring")
return return
ingress, egress = determine_direction(pkt_iface, bridge) ingress, egress = determine_direction(pkt_iface, bridge)
tup = _make_pkt_tuple(pkt, ingress, egress) src_ip, dst_ip, proto = _safe_get_src_dst(pkt)
try:
packet_queue.put_nowait(tup)
except Full:
# queue full -> drop packet and log throttle event
logger.warning("Packet queue full, dropping packet (producer side)")
# ------------------------------------------------------------------- pkt_info = {
# Threaded sniff loop "ingress": ingress,
# ------------------------------------------------------------------- "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_ip": src_ip,
"dst_ip": dst_ip,
"protocol": proto,
"length": len(pkt),
# store raw as binary; use bytes(pkt)
"raw": bytes(pkt)
}
# schedule asynchronous DB insert
schedule_db_insert(pkt_info)
# ------------------------
# Sniffer thread loop
# ------------------------
def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
logger.info(f"Sniffer STARTED on {ifname}") logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge)
if not check_interface_exists(ifname): if not check_interface_exists(ifname):
logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}") logger.error("Interface %s does not exist; stopping sniffer thread", ifname)
return
if not check_interface_up(ifname):
logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}")
return return
if not check_interface_up(ifname):
logger.warning("Interface %s is down; sniffing might still work but check interface state", ifname)
# 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( sniff(
iface=ifname, iface=ifname,
prn=lambda pkt: handle_packet(pkt, bridge), prn=lambda pkt: handle_packet(pkt, bridge),
store=False, store=False,
timeout=1 # short timeout to check stop_event frequently timeout=1 # returns periodically so we can check stop_event
) )
except PermissionError: except PermissionError:
logger.error(f"Permission denied sniffing on {ifname}. Run as root.") logger.exception("Permission denied sniffing on %s - run as root or give CAP_NET_RAW", ifname)
break break
except OSError as e: except OSError as e:
logger.error(f"Sniffer OSError on {ifname}: {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 break
except Exception as e: except Exception:
logger.exception(f"Unexpected sniffer error on {ifname}: {e}") logger.exception("Unexpected error in sniffer loop for %s", ifname)
stop_event.wait(0.5)
break break
logger.info(f"Sniffer STOPPED on {ifname}") logger.info("Sniffer thread exiting for %s", ifname)
def start_sniffer_thread(bridge: str):
# ------------------------
# 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).
"""
ports = get_bridge_ports(bridge) ports = get_bridge_ports(bridge)
if not ports: if not ports:
logger.error(f"No ports found for bridge {bridge}, not starting sniffers.") logger.error("No ports for bridge %s - nothing to start", 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 and sniffer_threads[iface].is_alive():
logger.info(f"Sniffer already running on {iface}") 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)
sniffer_threads[iface] = thread
thread.start() thread.start()
sniffer_threads[iface] = thread
thread_stop_flags[iface] = stop_evt
started[iface] = thread
logger.info("Started sniffer on %s (bridge=%s)", iface, bridge)
return sniffer_threads return started
# -------------------------------------------------------------------
# Public API for starting/stopping sniffing (to be called from FastAPI)
# -------------------------------------------------------------------
_consumer_stop_event: Optional[asyncio.Event] = None
_consumer_future_accessor = None
async def start_sniffing(bridge: str): async def start_sniffing(bridge: str) -> bool:
""" """
Start sniffing on all bridge ports and start the DB consumer. Async entrypoint you can call from FastAPI startup/route.
Safe to call multiple times.
""" """
global _consumer_stop_event, _consumer_future_accessor ensure_async_loop()
logger.info(f"Request to start sniffing on bridge {bridge}") ensure_db_pool()
start_sniffer_thread_for_bridge(bridge)
# start consumer if not running
if _consumer_stop_event is None:
# create asyncio.Event in async_loop
fut = asyncio.run_coroutine_threadsafe(asyncio.sleep(0), async_loop)
# schedule creation of event and consumer
_consumer_stop_event = asyncio.run_coroutine_threadsafe(asyncio.Event(), async_loop).result()
# start consumer in async_loop as a future
def schedule_consumer():
nonlocal _consumer_stop_event
# schedule consumer coroutine directly
future = asyncio.run_coroutine_threadsafe(_db_consumer_loop(_consumer_stop_event), async_loop)
return future
# store accessor to the Future to allow waiting for termination on stop
_consumer_future_accessor = schedule_consumer
# start sniffer threads
start_sniffer_thread(bridge)
return True return True
async def stop_sniffing():
"""
Stop all sniffer threads and the DB consumer.
"""
global _consumer_stop_event, _consumer_future_accessor
logger.info("Stopping all sniffers and consumer")
# stop sniffer threads async def stop_sniffing() -> bool:
for iface, ev in list(thread_stop_flags.items()): """
ev.set() 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()): 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")
# stop consumer
if _consumer_stop_event is not None:
# signal consumer in async_loop
def set_stop():
_consumer_stop_event.set()
asyncio.run_coroutine_threadsafe(asyncio.to_thread(set_stop), async_loop).result(timeout=2)
# wait for consumer future
if _consumer_future_accessor:
fut = _consumer_future_accessor()
try:
fut.result(timeout=5)
except Exception as e:
logger.debug(f"Consumer future join error: {e}")
_consumer_stop_event = None
_consumer_future_accessor = None
# flush queue attempt (best-effort)
logger.info("Flushing packet queue (best-effort)")
# Not blocking: drop packets on stop
while not packet_queue.empty():
try:
packet_queue.get_nowait()
except Empty:
break
# close pool cleanly from async loop
try:
asyncio.run_coroutine_threadsafe(_close_pool(), async_loop).result(timeout=5)
except Exception:
logger.exception("Failed closing pg pool cleanly")
logger.info("All stopped")
return True return True
# -------------------------------------------------------------------
# Status helper def get_sniffer_status() -> Dict[str, Dict[str, object]]:
# -------------------------------------------------------------------
def get_sniffer_status():
""" """
Return a dict of interface -> status information. Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] }
""" """
out = {} status = {}
for iface, t 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": t.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)
} }
out["queue_size"] = packet_queue.qsize() 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)