fix: enhance network sniffer with status endpoint and improve thread management
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-11-27 15:43:35 +01:00
parent 3140d3435f
commit 4cc7d30de1
2 changed files with 70 additions and 35 deletions

View File

@@ -4,18 +4,22 @@ import asyncpg
import threading import threading
import logging import logging
import os import os
from typing import List from typing import List, Dict
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("sniffer") logger = logging.getLogger("sniffer")
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
sniffer_threads = {} # 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 = {} bridge_ports_cache = {}
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# ASYNC LOOP FOR DATABASE INSERTS # ASYNC LOOP FOR DB INSERTS
# ------------------------------------------------------------------- # -------------------------------------------------------------------
async_loop = asyncio.new_event_loop() async_loop = asyncio.new_event_loop()
@@ -27,19 +31,15 @@ threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# BRIDGE PORT DETECTION # BRIDGE PORT HANDLING
# ------------------------------------------------------------------- # -------------------------------------------------------------------
def get_bridge_ports(bridge: str) -> List[str]: def get_bridge_ports(bridge: str) -> List[str]:
"""
Return all interfaces that belong to this Linux bridge.
Cached for performance.
"""
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 or brif folder missing") logger.error(f"Bridge {bridge} does not exist")
return [] return []
ports = os.listdir(base) ports = os.listdir(base)
@@ -49,15 +49,9 @@ def get_bridge_ports(bridge: str) -> List[str]:
def determine_direction(pkt_iface: str, bridge: str): def determine_direction(pkt_iface: str, bridge: str):
"""
Determine ingress + egress:
- ingress = interface where packet was received
- egress = other bridge ports (the packet will be forwarded to)
"""
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
@@ -67,6 +61,7 @@ def determine_direction(pkt_iface: str, bridge: str):
async def db_insert_packet(pkt_info: dict): async def db_insert_packet(pkt_info: dict):
try: try:
conn = await asyncpg.connect(DB_DSN) conn = await asyncpg.connect(DB_DSN)
await conn.execute(""" await conn.execute("""
INSERT INTO packets( INSERT INTO packets(
iface, iface,
@@ -80,10 +75,11 @@ async def db_insert_packet(pkt_info: dict):
length, length,
ebpf_verdict, ebpf_verdict,
raw raw
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) )
VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
""", """,
pkt_info["ingress"], # interface where packet arrived pkt_info["ingress"],
"unknown", # direction (can update later) "unknown",
pkt_info["src_mac"], pkt_info["src_mac"],
pkt_info["dst_mac"], pkt_info["dst_mac"],
pkt_info["eth_type"], pkt_info["eth_type"],
@@ -91,15 +87,12 @@ async def db_insert_packet(pkt_info: dict):
pkt_info["dst_ip"], pkt_info["dst_ip"],
pkt_info["protocol"], pkt_info["protocol"],
pkt_info["length"], pkt_info["length"],
str(pkt_info.get("egress")), # optional egress info as text str(pkt_info["egress"]),
pkt_info["raw"] pkt_info["raw"]
) )
logger.info(f"DB insert OK: ingress={pkt_info['ingress']} src={pkt_info['src_ip']} dst={pkt_info['dst_ip']}")
except Exception as e: except Exception as e:
logger.exception(f"DB insert failed: {e}") logger.exception(f"DB insert failed: {e}")
finally: finally:
if 'conn' in locals(): if 'conn' in locals():
await conn.close() await conn.close()
@@ -110,9 +103,7 @@ async def db_insert_packet(pkt_info: dict):
# ------------------------------------------------------------------- # -------------------------------------------------------------------
def handle_packet(pkt, bridge: str): def handle_packet(pkt, bridge: str):
pkt_iface = getattr(pkt, "sniffed_on", None) pkt_iface = getattr(pkt, "sniffed_on", None)
if not pkt_iface: if not pkt_iface:
logger.warning("Packet missing sniffed_on metadata")
return return
ingress, egress = determine_direction(pkt_iface, bridge) ingress, egress = determine_direction(pkt_iface, bridge)
@@ -139,25 +130,65 @@ def handle_packet(pkt, bridge: str):
def start_sniffer_thread(bridge: str): def start_sniffer_thread(bridge: str):
ports = get_bridge_ports(bridge) ports = get_bridge_ports(bridge)
def sniff_blocking(iface):
logger.info(f"Starting sniff on {iface}")
sniff(iface=iface, prn=lambda x: handle_packet(x, bridge), store=False)
for iface in ports: for iface in ports:
if iface not in sniffer_threads: if iface in sniffer_threads:
thread = threading.Thread(target=sniff_blocking, args=(iface,), daemon=True) continue
thread.start()
stop_event = threading.Event()
thread_stop_flags[iface] = stop_event
def sniff_blocking(ifname=iface):
logger.info(f"Sniffer STARTED on {ifname}")
sniff(
iface=ifname,
prn=lambda pkt: handle_packet(pkt, bridge),
store=False,
stop_filter=lambda _: stop_event.is_set(),
timeout=1 # ensure periodic stop checks
)
logger.info(f"Sniffer STOPPED on {ifname}")
thread = threading.Thread(target=sniff_blocking, daemon=True)
sniffer_threads[iface] = thread sniffer_threads[iface] = thread
thread.start()
return sniffer_threads return sniffer_threads
async def start_sniffing(bridge: str): async def start_sniffing(bridge: str):
get_bridge_ports(bridge)
start_sniffer_thread(bridge) start_sniffer_thread(bridge)
return True
async def stop_sniffing(): async def stop_sniffing():
# Scapy can't stop sniff(), so we drop the thread references logger.info("Stopping all sniffers…")
for iface, stop_event in thread_stop_flags.items():
stop_event.set()
for iface, thread in sniffer_threads.items():
thread.join(timeout=2)
sniffer_threads.clear() sniffer_threads.clear()
thread_stop_flags.clear()
logger.info("All sniffers stopped.")
return True return True
# -------------------------------------------------------------------
# STATUS HELPER (for API endpoint)
# -------------------------------------------------------------------
def get_sniffer_status():
"""
Return a dict of interface → running/stopped status.
Useful for an API /status endpoint.
"""
status = {}
for iface, t in sniffer_threads.items():
status[iface] = "running" if t.is_alive() else "stopped"
return status

View File

@@ -1,6 +1,6 @@
from fastapi import APIRouter from fastapi import APIRouter
from typing import List from typing import List
from src.network_sniffer import start_sniffing, stop_sniffing from src.network_sniffer import get_sniffer_status, start_sniffing, stop_sniffing
router = APIRouter() router = APIRouter()
@@ -22,3 +22,7 @@ async def start_sniffer(bridge: str):
async def stop_sniffer_api(): async def stop_sniffer_api():
await stop_sniffing() await stop_sniffing()
return {"status": "stopped"} return {"status": "stopped"}
@router.get("/status")
async def status():
return get_sniffer_status()