diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index d793be6..65748ac 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -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 diff --git a/backend/src/routes_sniffer.py b/backend/src/routes_sniffer.py index f0d140d..9d7829c 100644 --- a/backend/src/routes_sniffer.py +++ b/backend/src/routes_sniffer.py @@ -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 }