fix: refactor network sniffer and routes to streamline bridge handling and improve logging
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-11-25 16:32:16 +01:00
parent c762942f70
commit abea4c0ba4
2 changed files with 95 additions and 123 deletions

View File

@@ -3,66 +3,82 @@ from scapy.all import sniff, Ether, IP
import asyncpg import asyncpg
import threading import threading
import logging import logging
from typing import List, Dict import os
from typing import List
# ------------------------------------------------------------------------------
# LOGGING
# ------------------------------------------------------------------------------
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("sniffer") logger = logging.getLogger("sniffer")
# ------------------------------------------------------------------------------
# CONFIG
# ------------------------------------------------------------------------------
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
# interface → direction map
# Example: { "eth0": "ingress", "eth1": "egress" }
iface_directions: Dict[str, str] = {}
sniffer_threads = {} sniffer_threads = {}
bridge_ports_cache = {}
# ------------------------------------------------------------------------------ # -------------------------------------------------------------------
# ASYNCIO LOOP (for DB INSERTS) # ASYNC LOOP FOR DATABASE INSERTS
# ------------------------------------------------------------------------------ # -------------------------------------------------------------------
async_loop = asyncio.new_event_loop() async_loop = asyncio.new_event_loop()
def start_async_loop(loop): def start_async_loop(loop):
asyncio.set_event_loop(loop) asyncio.set_event_loop(loop)
loop.run_forever() 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()
logger.info("Asyncio DB event loop started") # -------------------------------------------------------------------
# BRIDGE PORT DETECTION
# -------------------------------------------------------------------
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:
return bridge_ports_cache[bridge]
# ------------------------------------------------------------------------------ base = f"/sys/class/net/{bridge}/brif/"
# DB INSERT if not os.path.isdir(base):
# ------------------------------------------------------------------------------ logger.error(f"Bridge {bridge} does not exist or brif folder missing")
return []
ports = os.listdir(base)
bridge_ports_cache[bridge] = ports
logger.info(f"Bridge {bridge} ports: {ports}")
return ports
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)
ingress = pkt_iface
egress = [p for p in ports if p != pkt_iface]
return ingress, egress
# -------------------------------------------------------------------
# DATABASE INSERTION
# -------------------------------------------------------------------
async def db_insert_packet(pkt_info: dict): async def db_insert_packet(pkt_info: dict):
"""
Insert a packet asynchronously into PostgreSQL.
"""
try: try:
conn = await asyncpg.connect(DB_DSN) conn = await asyncpg.connect(DB_DSN)
await conn.execute( await conn.execute("""
"""
INSERT INTO packet_log( INSERT INTO packet_log(
interface, direction, interface_ingress,
interfaces_egress,
src_mac, dst_mac, eth_type, src_mac, dst_mac, eth_type,
src_ip, dst_ip, ip_protocol, src_ip, dst_ip, ip_protocol,
packet_len, nft_verdict, nft_chain, raw_packet packet_len, raw_packet
) )
VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10)
""", """,
pkt_info["iface"], pkt_info["ingress"],
pkt_info["direction"], pkt_info["egress"],
pkt_info["src_mac"], pkt_info["src_mac"],
pkt_info["dst_mac"], pkt_info["dst_mac"],
pkt_info["eth_type"], pkt_info["eth_type"],
@@ -70,45 +86,34 @@ 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"],
pkt_info.get("ebpf_verdict"),
pkt_info.get("ebpf_chain"),
pkt_info["raw"] pkt_info["raw"]
) )
logger.info( logger.info(f"DB insert OK: ingress={pkt_info['ingress']} src={pkt_info['src_ip']} dst={pkt_info['dst_ip']}")
f"DB Insert OK [{pkt_info['iface']}/{pkt_info['direction']}] "
f"{pkt_info['src_ip']} → {pkt_info['dst_ip']} len={pkt_info['length']}"
)
except Exception as e: except Exception as e:
logger.exception(f"DB insert failed: {e}") logger.exception(f"DB insert failed: {e}")
finally: finally:
try: if 'conn' in locals():
await conn.close() await conn.close()
except:
pass
# ------------------------------------------------------------------------------
# -------------------------------------------------------------------
# PACKET HANDLER # PACKET HANDLER
# ------------------------------------------------------------------------------ # -------------------------------------------------------------------
def handle_packet(pkt, bridge: str):
pkt_iface = getattr(pkt, "sniffed_on", None)
def detect_direction(iface: str) -> str: if not pkt_iface:
""" logger.warning("Packet missing sniffed_on metadata")
Return ingress/egress direction for this interface. return
"""
return iface_directions.get(iface, "unknown")
ingress, egress = determine_direction(pkt_iface, bridge)
def handle_packet(pkt, iface):
"""
Called inside Scapy sniff thread.
"""
direction = detect_direction(iface)
pkt_info = { pkt_info = {
"iface": iface, "ingress": ingress,
"direction": direction, "egress": egress,
"src_mac": pkt[Ether].src if Ether in pkt else None, "src_mac": pkt[Ether].src if Ether in pkt else None,
"dst_mac": pkt[Ether].dst 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, "eth_type": pkt[Ether].type if Ether in pkt else None,
@@ -116,65 +121,37 @@ def handle_packet(pkt, iface):
"dst_ip": pkt[IP].dst 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, "protocol": pkt[IP].proto if IP in pkt else None,
"length": len(pkt), "length": len(pkt),
"ebpf_verdict": None,
"ebpf_chain": None,
"raw": bytes(pkt) "raw": bytes(pkt)
} }
# Debug per-packet
logger.debug(
f"Packet captured on {iface} ({direction}): "
f"{pkt_info['src_ip']} -> {pkt_info['dst_ip']} ({pkt_info['length']} bytes)"
)
# Submit to asyncio loop
asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop) asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop)
# ------------------------------------------------------------------------------
# SNIFFER THREAD
# ------------------------------------------------------------------------------
def start_sniffer_thread(iface: str): # -------------------------------------------------------------------
def sniff_blocking(): # SNIFFING
logger.info(f"Starting sniffer on {iface}") # -------------------------------------------------------------------
sniff(prn=lambda x: handle_packet(x, iface), iface=iface, store=False) def start_sniffer_thread(bridge: str):
ports = get_bridge_ports(bridge)
thread = threading.Thread(target=sniff_blocking, daemon=True) 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:
if iface not in sniffer_threads:
thread = threading.Thread(target=sniff_blocking, args=(iface,), daemon=True)
thread.start() thread.start()
return thread sniffer_threads[iface] = thread
# ------------------------------------------------------------------------------ return sniffer_threads
# PUBLIC API
# ------------------------------------------------------------------------------
async def start_sniffing(interfaces: List[str], ingress_iface: str, egress_iface: str):
"""
Start sniffing on a list of interfaces.
Also store ingress/egress direction info.
"""
# Save direction mapping async def start_sniffing(bridge: str):
iface_directions.clear() get_bridge_ports(bridge)
iface_directions[ingress_iface] = "ingress" start_sniffer_thread(bridge)
iface_directions[egress_iface] = "egress"
logger.info(f"Direction mapping: {iface_directions}")
for iface in interfaces:
if iface in sniffer_threads:
logger.info(f"Sniffer already active on {iface}")
continue
sniffer_threads[iface] = start_sniffer_thread(iface)
return {"status": "sniffers started", "interfaces": interfaces}
async def stop_sniffing(): async def stop_sniffing():
""" # Scapy can't stop sniff(), so we drop the thread references
Scapy sniff cannot stop easily.
We simply forget threads (they are daemon threads).
"""
logger.warning("Stopping all sniffers (threads will exit on process stop)")
sniffer_threads.clear() sniffer_threads.clear()
return True return True

View File

@@ -4,22 +4,17 @@ from src.network_sniffer import start_sniffing, stop_sniffing
router = APIRouter() router = APIRouter()
sniffer_running_interfaces: List[str] = []
@router.post("/sniffer/start") @router.post("/sniffer/start")
async def start_sniffer(bridge: str, interfaces: List[str]): async def start_sniffer(bridge: str):
""" """
Start packet sniffing on bridge member interfaces. Start packet sniffing on bridge member interfaces.
""" """
global sniffer_running_interfaces
sniffer_running_interfaces = interfaces
await start_sniffing(interfaces) await start_sniffing(bridge)
return { return {
"status": "ok", "status": "ok",
"started_on": interfaces "started_on": bridge
} }