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 threading
import logging
import os
import time
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
# ---------------------------
# Configuration
# ---------------------------
logger = logging.getLogger("network_sniffer")
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("sniffer")
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
DB_TABLE = "packets" # keep your current table name; change if needed
# 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)
# ----- CONFIG -----
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
# ------------------
# 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] = {}
# Cache for bridge -> ports
# Bridge→ports cache
bridge_ports_cache: Dict[str, List[str]] = {}
# Queue for packet events (thread-safe)
packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE)
# 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 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)
# -------------------------------------------------------------------
def _start_async_loop(loop):
# ------------------------
# 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()
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
# -------------------------------------------------------------------
# DB helper: create pool & consumer
# -------------------------------------------------------------------
async def _create_pool():
global pg_pool
if pg_pool is None:
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...")
pg_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=8)
return pg_pool
_db_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=6)
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):
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:
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:
return f.read().strip() == "up"
except FileNotFoundError:
return False
def get_bridge_ports(bridge: str) -> List[str]:
"""
Async consumer: drains packet_queue and does batched inserts.
Runs inside async_loop.
Return member interfaces of the bridge. Uses /sys/class/net/<bridge>/brif/.
Returns [] if bridge missing or on error.
"""
await _create_pool()
logger.info("DB consumer started")
stmt = f"""
INSERT INTO {DB_TABLE}(
iface,
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)
return []
try:
ports = [p for p in os.listdir(base) if os.path.isdir(f"/sys/class/net/{p}")]
except PermissionError:
logger.error("Permission denied reading bridge ports for %s", bridge)
return []
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)
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
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,
protocol,
length,
ebpf_verdict,
raw
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
"""
buffer: List[tuple] = []
last_flush = time.time()
try:
while not stop_event.is_set():
# collect up to BATCH_SIZE from the queue with small wait
try:
# block briefly for the first item
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."""
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:
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:
return f.read().strip() == "up"
except Exception:
return False
def get_bridge_ports(bridge: str) -> List[str]:
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(f"Bridge '{bridge}' does not exist")
return []
try:
ports = os.listdir(base)
except PermissionError:
logger.error(f"No permission to read bridge ports for '{bridge}'")
return []
ok_ports = []
for p in ports:
if check_interface_exists(p):
ok_ports.append(p)
else:
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):
ports = get_bridge_ports(bridge)
ingress = pkt_iface
egress = [p for p in ports if p != pkt_iface]
return ingress, egress
# -------------------------------------------------------------------
# Producer (sniffer threads)
# -------------------------------------------------------------------
def _make_pkt_tuple(pkt, ingress_iface: str, egress_list):
"""
Convert packet info to a tuple matching DB insert order.
We store only a small raw snapshot to keep DB small.
"""
try:
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
except Exception:
src_mac = dst_mac = eth_type = None
try:
src_ip = pkt[IP].src if IP in pkt else None
dst_ip = pkt[IP].dst if IP in pkt else None
protocol = pkt[IP].proto if IP in pkt else None
except Exception:
src_ip = dst_ip = protocol = 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
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:
ensure_db_pool()
except Exception:
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:
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):
"""
Callback running inside Scapy sniff thread.
Places a lightweight tuple in packet_queue.
Called in sniffing thread. Build pkt_info and schedule DB insert.
"""
pkt_iface = getattr(pkt, "sniffed_on", None)
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
ingress, egress = determine_direction(pkt_iface, bridge)
tup = _make_pkt_tuple(pkt, ingress, egress)
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)")
src_ip, dst_ip, proto = _safe_get_src_dst(pkt)
# -------------------------------------------------------------------
# Threaded sniff loop
# -------------------------------------------------------------------
pkt_info = {
"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):
logger.info(f"Sniffer STARTED on {ifname}")
logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge)
if not check_interface_exists(ifname):
logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}")
return
if not check_interface_up(ifname):
logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}")
logger.error("Interface %s does not exist; stopping sniffer thread", ifname)
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():
try:
sniff(
iface=ifname,
prn=lambda pkt: handle_packet(pkt, bridge),
store=False,
timeout=1 # short timeout to check stop_event frequently
timeout=1 # returns periodically so we can check stop_event
)
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
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
except Exception as e:
logger.exception(f"Unexpected sniffer error on {ifname}: {e}")
except Exception:
logger.exception("Unexpected error in sniffer loop for %s", ifname)
stop_event.wait(0.5)
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)
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 {}
started = {}
for iface in ports:
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
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
stop_evt = threading.Event()
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_evt, bridge), daemon=True)
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.
Safe to call multiple times.
Async entrypoint you can call from FastAPI startup/route.
"""
global _consumer_stop_event, _consumer_future_accessor
logger.info(f"Request to start sniffing on 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)
ensure_async_loop()
ensure_db_pool()
start_sniffer_thread_for_bridge(bridge)
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
for iface, ev in list(thread_stop_flags.items()):
ev.set()
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)
thread.join(timeout=2)
sniffer_threads.clear()
thread_stop_flags.clear()
# 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")
logger.info("All sniffer threads stopped")
return True
# -------------------------------------------------------------------
# Status helper
# -------------------------------------------------------------------
def get_sniffer_status():
def get_sniffer_status() -> Dict[str, Dict[str, object]]:
"""
Return a dict of interface -> status information.
Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] }
"""
out = {}
for iface, t in sniffer_threads.items():
out[iface] = {
"running": t.is_alive(),
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()),
"exists": check_interface_exists(iface),
"up": check_interface_up(iface),
"bridge": iface_to_bridge.get(iface)
}
out["queue_size"] = packet_queue.qsize()
return out
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)