diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index 0abef56..d070417 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -4,18 +4,22 @@ import asyncpg import threading import logging import os -from typing import List +from typing import List, Dict logging.basicConfig(level=logging.INFO) logger = logging.getLogger("sniffer") 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 = {} # ------------------------------------------------------------------- -# ASYNC LOOP FOR DATABASE INSERTS +# ASYNC LOOP FOR DB INSERTS # ------------------------------------------------------------------- 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]: - """ - 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/" 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 [] ports = os.listdir(base) @@ -49,15 +49,9 @@ def get_bridge_ports(bridge: str) -> List[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) ingress = pkt_iface egress = [p for p in ports if p != pkt_iface] - return ingress, egress @@ -67,6 +61,7 @@ def determine_direction(pkt_iface: str, bridge: str): async def db_insert_packet(pkt_info: dict): try: conn = await asyncpg.connect(DB_DSN) + await conn.execute(""" INSERT INTO packets( iface, @@ -80,10 +75,11 @@ async def db_insert_packet(pkt_info: dict): length, ebpf_verdict, 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 - "unknown", # direction (can update later) + pkt_info["ingress"], + "unknown", pkt_info["src_mac"], pkt_info["dst_mac"], pkt_info["eth_type"], @@ -91,15 +87,12 @@ async def db_insert_packet(pkt_info: dict): pkt_info["dst_ip"], pkt_info["protocol"], pkt_info["length"], - str(pkt_info.get("egress")), # optional egress info as text + str(pkt_info["egress"]), 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: logger.exception(f"DB insert failed: {e}") - finally: if 'conn' in locals(): await conn.close() @@ -110,9 +103,7 @@ async def db_insert_packet(pkt_info: dict): # ------------------------------------------------------------------- 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) @@ -139,25 +130,65 @@ def handle_packet(pkt, bridge: str): def start_sniffer_thread(bridge: str): 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: - if iface not in sniffer_threads: - thread = threading.Thread(target=sniff_blocking, args=(iface,), daemon=True) - thread.start() - sniffer_threads[iface] = thread + if iface in sniffer_threads: + continue + + 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 + thread.start() return sniffer_threads async def start_sniffing(bridge: str): - get_bridge_ports(bridge) start_sniffer_thread(bridge) + return True 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() + thread_stop_flags.clear() + + logger.info("All sniffers stopped.") 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 diff --git a/backend/src/routes_sniffer.py b/backend/src/routes_sniffer.py index eb4b51d..959c1db 100644 --- a/backend/src/routes_sniffer.py +++ b/backend/src/routes_sniffer.py @@ -1,6 +1,6 @@ from fastapi import APIRouter 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() @@ -22,3 +22,7 @@ async def start_sniffer(bridge: str): async def stop_sniffer_api(): await stop_sniffing() return {"status": "stopped"} + +@router.get("/status") +async def status(): + return get_sniffer_status() \ No newline at end of file