fix: refactor sniffer API routes for consistency and clarity
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-11-27 16:17:03 +01:00
parent f3d2fca9d2
commit 2d68718bd0
2 changed files with 297 additions and 133 deletions

View File

@@ -1,170 +1,284 @@
import asyncio
from scapy.all import sniff, Ether, IP
import asyncpg
import threading
import logging
import os
import socket
from typing import List, Dict
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
import asyncpg
# ---------------------------
# Configuration
# ---------------------------
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)
# Active sniffer threads and stop flags
sniffer_threads: Dict[str, threading.Thread] = {}
thread_stop_flags: Dict[str, threading.Event] = {}
# Cache for bridge → ports
bridge_ports_cache = {}
# Cache for bridge -> ports
bridge_ports_cache: Dict[str, List[str]] = {}
# -------------------------------------------------------------------
# ASYNC LOOP FOR DB INSERTS
# -------------------------------------------------------------------
# Queue for packet events (thread-safe)
packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE)
# Async event loop + background consumer task
async_loop = asyncio.new_event_loop()
consumer_task: Optional[asyncio.Task] = None
pg_pool: Optional[asyncpg.Pool] = None
def start_async_loop(loop):
# -------------------------------------------------------------------
# Async loop bootstrap (runs in background thread)
# -------------------------------------------------------------------
def _start_async_loop(loop):
asyncio.set_event_loop(loop)
loop.run_forever()
threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start()
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
# -------------------------------------------------------------------
# BRIDGE PORT HANDLING
# DB helper: create pool & consumer
# -------------------------------------------------------------------
async def _create_pool():
global pg_pool
if pg_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
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()
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 FileNotFoundError:
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"[ERROR] Bridge '{bridge}' does not exist.")
logger.error(f"Bridge '{bridge}' does not exist")
return []
try:
ports = os.listdir(base)
except PermissionError:
logger.error(f"[ERROR] No permissions to read bridge ports for '{bridge}'.")
logger.error(f"No permission to read bridge ports for '{bridge}'")
return []
ok_ports = []
for p in ports:
if not check_interface_exists(p):
logger.warning(f"[WARN] Port '{p}' in bridge but does not exist in /sys/class/net")
continue
ok_ports.append(p)
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"[INFO] Bridge {bridge} ports: {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
# -------------------------------------------------------------------
# DATABASE INSERTION
# Producer (sniffer threads)
# -------------------------------------------------------------------
async def db_insert_packet(pkt_info: dict):
conn = None
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:
conn = await asyncpg.connect(DB_DSN)
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
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["ingress"],
"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"]
)
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
except (asyncpg.PostgresError, ConnectionError, OSError) as e:
logger.exception(f"DB insert failed: {e}")
raw = bytes(pkt)[:SNAPSHOT_MAX_BYTES] if pkt is not None else b""
finally:
if conn:
await conn.close()
# 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
)
# -------------------------------------------------------------------
# PACKET HANDLER
# -------------------------------------------------------------------
def handle_packet(pkt, bridge: str):
"""
Callback running inside Scapy sniff thread.
Places a lightweight tuple in packet_queue.
"""
pkt_iface = getattr(pkt, "sniffed_on", None)
if not pkt_iface:
logger.debug("Packet without sniffed_on metadata, ignoring")
return
ingress, egress = determine_direction(pkt_iface, bridge)
pkt_info = {
"ingress": ingress,
"egress": egress,
"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": 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,
"length": len(pkt),
"raw": pkt.json()
}
asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop)
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)")
# -------------------------------------------------------------------
# CLEAN STOPPING SNIFFERS
# Threaded sniff loop
# -------------------------------------------------------------------
def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
logger.info(f"Sniffer STARTED on {ifname}")
if not check_interface_exists(ifname):
logger.error(f"[ERROR] Interface {ifname} does not exist. Stopping sniffer.")
logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}")
return
if not check_interface_up(ifname):
logger.error(f"[ERROR] Interface {ifname} is DOWN. Stopping sniffer.")
logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}")
return
while not stop_event.is_set():
@@ -173,13 +287,13 @@ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
iface=ifname,
prn=lambda pkt: handle_packet(pkt, bridge),
store=False,
timeout=1 # periodic return so we can check stop_event
timeout=1 # short timeout to check stop_event frequently
)
except PermissionError:
logger.error(f"[ERROR] Permission denied sniffing on {ifname}. Run as root.")
logger.error(f"Permission denied sniffing on {ifname}. Run as root.")
break
except OSError as e:
logger.error(f"[ERROR] Sniffer error on {ifname}: {e}")
logger.error(f"Sniffer OSError on {ifname}: {e}")
break
except Exception as e:
logger.exception(f"Unexpected sniffer error on {ifname}: {e}")
@@ -187,67 +301,126 @@ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
logger.info(f"Sniffer STOPPED on {ifname}")
def start_sniffer_thread(bridge: str):
ports = get_bridge_ports(bridge)
if not ports:
logger.error(f"[ERROR] Could not start sniffer: No valid ports found for bridge {bridge}")
logger.error(f"No ports found for bridge {bridge}, not starting sniffers.")
return {}
for iface in ports:
if iface in sniffer_threads:
logger.info(f"[INFO] Sniffer on {iface} is already running")
if iface in sniffer_threads and sniffer_threads[iface].is_alive():
logger.info(f"Sniffer already running on {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
)
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True)
sniffer_threads[iface] = thread
thread.start()
return sniffer_threads
# -------------------------------------------------------------------
# Public API for starting/stopping sniffing (to be called from FastAPI)
# -------------------------------------------------------------------
_consumer_stop_event: Optional[asyncio.Event] = None
_consumer_future_accessor = None
# -------------------------------------------------------------------
# PUBLIC ASYNC START/STOP METHODS
# -------------------------------------------------------------------
async def start_sniffing(bridge: str):
logger.info(f"Starting sniffing for bridge {bridge}")
"""
Start sniffing on all bridge ports and start the DB consumer.
Safe to call multiple times.
"""
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)
return True
async def stop_sniffing():
logger.info("Stopping all sniffers…")
"""
Stop all sniffer threads and the DB consumer.
"""
global _consumer_stop_event, _consumer_future_accessor
logger.info("Stopping all sniffers and consumer")
for iface, stop_event in thread_stop_flags.items():
stop_event.set()
# stop sniffer threads
for iface, ev in list(thread_stop_flags.items()):
ev.set()
for iface, thread in sniffer_threads.items():
for iface, thread in list(sniffer_threads.items()):
thread.join(timeout=2)
sniffer_threads.clear()
thread_stop_flags.clear()
logger.info("All sniffers 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
# -------------------------------------------------------------------
# STATUS HELPER
# Status helper
# -------------------------------------------------------------------
def get_sniffer_status():
"""
Return a dict of interface -> status information.
"""
out = {}
for iface, t in sniffer_threads.items():
out[iface] = {
"running": t.is_alive(),
"exists": check_interface_exists(iface),
"up": check_interface_up(iface)
"up": check_interface_up(iface),
}
out["queue_size"] = packet_queue.qsize()
return out