fix: refactor network sniffer for improved structure and performance
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user