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,112 +3,117 @@ from scapy.all import sniff, Ether, IP
import asyncpg
import threading
import logging
from typing import List, Dict
import os
from typing import List
# ------------------------------------------------------------------------------
# LOGGING
# ------------------------------------------------------------------------------
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("sniffer")
# ------------------------------------------------------------------------------
# CONFIG
# ------------------------------------------------------------------------------
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 = {}
bridge_ports_cache = {}
# ------------------------------------------------------------------------------
# ASYNCIO LOOP (for DB INSERTS)
# ------------------------------------------------------------------------------
# -------------------------------------------------------------------
# ASYNC LOOP FOR DATABASE INSERTS
# -------------------------------------------------------------------
async_loop = asyncio.new_event_loop()
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()
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]
# ------------------------------------------------------------------------------
# DB INSERT
# ------------------------------------------------------------------------------
base = f"/sys/class/net/{bridge}/brif/"
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):
"""
Insert a packet asynchronously into PostgreSQL.
"""
try:
conn = await asyncpg.connect(DB_DSN)
await conn.execute(
"""
await conn.execute("""
INSERT INTO packet_log(
interface, direction,
interface_ingress,
interfaces_egress,
src_mac, dst_mac, eth_type,
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)
""",
pkt_info["iface"],
pkt_info["direction"],
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"],
pkt_info.get("ebpf_verdict"),
pkt_info.get("ebpf_chain"),
pkt_info["raw"]
VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10)
""",
pkt_info["ingress"],
pkt_info["egress"],
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"],
pkt_info["raw"]
)
logger.info(
f"DB Insert OK [{pkt_info['iface']}/{pkt_info['direction']}] "
f"{pkt_info['src_ip']} → {pkt_info['dst_ip']} len={pkt_info['length']}"
)
logger.info(f"DB insert OK: ingress={pkt_info['ingress']} src={pkt_info['src_ip']} dst={pkt_info['dst_ip']}")
except Exception as e:
logger.exception(f"DB insert failed: {e}")
finally:
try:
if 'conn' in locals():
await conn.close()
except:
pass
# ------------------------------------------------------------------------------
# PACKET HANDLER
# ------------------------------------------------------------------------------
def detect_direction(iface: str) -> str:
"""
Return ingress/egress direction for this interface.
"""
return iface_directions.get(iface, "unknown")
def handle_packet(pkt, iface):
"""
Called inside Scapy sniff thread.
"""
direction = detect_direction(iface)
# -------------------------------------------------------------------
# PACKET HANDLER
# -------------------------------------------------------------------
def handle_packet(pkt, bridge: str):
pkt_iface = getattr(pkt, "sniffed_on", None)
if not pkt_iface:
logger.warning("Packet missing sniffed_on metadata")
return
ingress, egress = determine_direction(pkt_iface, bridge)
pkt_info = {
"iface": iface,
"direction": direction,
"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,
@@ -116,65 +121,37 @@ def handle_packet(pkt, iface):
"dst_ip": pkt[IP].dst if IP in pkt else None,
"protocol": pkt[IP].proto if IP in pkt else None,
"length": len(pkt),
"ebpf_verdict": None,
"ebpf_chain": None,
"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)
# ------------------------------------------------------------------------------
# SNIFFER THREAD
# ------------------------------------------------------------------------------
def start_sniffer_thread(iface: str):
def sniff_blocking():
logger.info(f"Starting sniffer on {iface}")
sniff(prn=lambda x: handle_packet(x, iface), iface=iface, store=False)
# -------------------------------------------------------------------
# SNIFFING
# -------------------------------------------------------------------
def start_sniffer_thread(bridge: str):
ports = get_bridge_ports(bridge)
thread = threading.Thread(target=sniff_blocking, daemon=True)
thread.start()
return thread
def sniff_blocking(iface):
logger.info(f"Starting sniff on {iface}")
sniff(iface=iface, prn=lambda x: handle_packet(x, bridge), store=False)
# ------------------------------------------------------------------------------
# PUBLIC API
# ------------------------------------------------------------------------------
for iface in ports:
if iface not in sniffer_threads:
thread = threading.Thread(target=sniff_blocking, args=(iface,), daemon=True)
thread.start()
sniffer_threads[iface] = thread
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.
"""
return sniffer_threads
# Save direction mapping
iface_directions.clear()
iface_directions[ingress_iface] = "ingress"
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 start_sniffing(bridge: str):
get_bridge_ports(bridge)
start_sniffer_thread(bridge)
async def stop_sniffing():
"""
Scapy sniff cannot stop easily.
We simply forget threads (they are daemon threads).
"""
logger.warning("Stopping all sniffers (threads will exit on process stop)")
# Scapy can't stop sniff(), so we drop the thread references
sniffer_threads.clear()
return True

View File

@@ -4,22 +4,17 @@ from src.network_sniffer import start_sniffing, stop_sniffing
router = APIRouter()
sniffer_running_interfaces: List[str] = []
@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.
"""
global sniffer_running_interfaces
sniffer_running_interfaces = interfaces
await start_sniffing(interfaces)
await start_sniffing(bridge)
return {
"status": "ok",
"started_on": interfaces
"started_on": bridge
}