diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index da81b23..4107719 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -1,170 +1,284 @@ import asyncio -from scapy.all import sniff, Ether, IP -import asyncpg import threading import logging import os -import socket -from typing import List, Dict +import time +from typing import List, Dict, Optional +from queue import Queue, Full, Empty +from scapy.all import sniff, Ether, IP # keep Scapy usage minimal +import asyncpg + +# --------------------------- +# Configuration +# --------------------------- logging.basicConfig(level=logging.INFO) logger = logging.getLogger("sniffer") DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" +DB_TABLE = "packets" # keep your current table name; change if needed + +# Producer/consumer settings +QUEUE_MAXSIZE = 20000 # max packets buffered in memory +BATCH_SIZE = 200 # how many rows to insert at once +FLUSH_INTERVAL = 0.5 # seconds max before flushing partial batch +SNAPSHOT_MAX_BYTES = 256 # how many bytes of the packet to store (truncate raw) # 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 = {} +# Cache for bridge -> ports +bridge_ports_cache: Dict[str, List[str]] = {} -# ------------------------------------------------------------------- -# ASYNC LOOP FOR DB INSERTS -# ------------------------------------------------------------------- +# Queue for packet events (thread-safe) +packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE) + +# Async event loop + background consumer task async_loop = asyncio.new_event_loop() +consumer_task: Optional[asyncio.Task] = None +pg_pool: Optional[asyncpg.Pool] = None - -def start_async_loop(loop): +# ------------------------------------------------------------------- +# Async loop bootstrap (runs in background thread) +# ------------------------------------------------------------------- +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() # ------------------------------------------------------------------- -# BRIDGE PORT HANDLING +# DB helper: create pool & consumer +# ------------------------------------------------------------------- +async def _create_pool(): + global pg_pool + if pg_pool is None: + logger.info("Creating asyncpg pool...") + pg_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=8) + return pg_pool + +async def _close_pool(): + global pg_pool + if pg_pool: + await pg_pool.close() + pg_pool = None + +async def _db_consumer_loop(stop_event: asyncio.Event): + """ + Async consumer: drains packet_queue and does batched inserts. + Runs inside async_loop. + """ + await _create_pool() + logger.info("DB consumer started") + stmt = f""" + INSERT INTO {DB_TABLE}( + iface, + direction, + src_mac, + dst_mac, + eth_type, + src_ip, + dst_ip, + protocol, + length, + ebpf_verdict, + raw + ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) + """ + + buffer: List[tuple] = [] + last_flush = time.time() + + try: + while not stop_event.is_set(): + # collect up to BATCH_SIZE from the queue with small wait + try: + # block briefly for the first item + item = packet_queue.get(timeout=FLUSH_INTERVAL) + except Empty: + item = None + + if item: + buffer.append(item) + # try to drain additional items without blocking + while len(buffer) < BATCH_SIZE: + try: + item = packet_queue.get_nowait() + buffer.append(item) + except Empty: + break + + now = time.time() + # flush if we have enough buffered or timed out + if buffer and (len(buffer) >= BATCH_SIZE or (now - last_flush) >= FLUSH_INTERVAL): + try: + async with pg_pool.acquire() as conn: + async with conn.transaction(): + # executemany pattern + await conn.executemany(stmt, buffer) + logger.debug(f"Inserted {len(buffer)} packets") + except Exception as e: + logger.exception(f"DB batch insert failed: {e}") + # in case of failure, drop or requeue? We'll drop to avoid blocking. + # Optionally write to disk or metrics. + finally: + buffer.clear() + last_flush = now + + # When stop requested, flush remaining + if buffer: + try: + async with pg_pool.acquire() as conn: + async with conn.transaction(): + await conn.executemany(stmt, buffer) + logger.info(f"Flushed final {len(buffer)} packets on shutdown") + except Exception: + logger.exception("Failed flushing final packet buffer on shutdown") + finally: + buffer.clear() + + finally: + logger.info("DB consumer stopped") + + +def _start_consumer(): + """Schedules the DB consumer to run in the shared async_loop and returns the stop Event.""" + stop_event = asyncio.Event() + # create consumer task in async_loop + def _start(): + global consumer_task + consumer_task = asyncio.run_coroutine_threadsafe(_db_consumer_loop(stop_event), async_loop) + threading.Thread(target=_start, daemon=True).start() + return stop_event, lambda: consumer_task # returns stop_event and accessor to future + +async def _stop_consumer(stop_event: asyncio.Event, consumer_future_accessor): + """Signal stop_event and wait for consumer to finish.""" + logger.info("Stopping DB consumer...") + stop_event.set() + # wait for the consumer coroutine to finish + future = consumer_future_accessor() + if future: + try: + future.result(timeout=5) + except Exception as e: + logger.debug(f"Consumer future termination: {e}") + await _close_pool() + +# ------------------------------------------------------------------- +# Utility: Bridge port handling & direction # ------------------------------------------------------------------- def check_interface_exists(iface: str) -> bool: return os.path.isdir(f"/sys/class/net/{iface}") - def check_interface_up(iface: str) -> bool: try: with open(f"/sys/class/net/{iface}/operstate", "r") as f: return f.read().strip() == "up" - except FileNotFoundError: + except Exception: return False - def get_bridge_ports(bridge: str) -> List[str]: 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"[ERROR] Bridge '{bridge}' does not exist.") + logger.error(f"Bridge '{bridge}' does not exist") return [] try: ports = os.listdir(base) except PermissionError: - logger.error(f"[ERROR] No permissions to read bridge ports for '{bridge}'.") + logger.error(f"No permission to read bridge ports for '{bridge}'") return [] ok_ports = [] for p in ports: - if not check_interface_exists(p): - logger.warning(f"[WARN] Port '{p}' in bridge but does not exist in /sys/class/net") - continue - ok_ports.append(p) + if check_interface_exists(p): + ok_ports.append(p) + else: + logger.warning(f"Bridge member {p} listed but interface missing in /sys/class/net") bridge_ports_cache[bridge] = ok_ports - logger.info(f"[INFO] Bridge {bridge} ports: {ok_ports}") + logger.info(f"Bridge {bridge} ports: {ok_ports}") return ok_ports - def determine_direction(pkt_iface: str, bridge: str): ports = get_bridge_ports(bridge) ingress = pkt_iface egress = [p for p in ports if p != pkt_iface] return ingress, egress - # ------------------------------------------------------------------- -# DATABASE INSERTION +# Producer (sniffer threads) # ------------------------------------------------------------------- -async def db_insert_packet(pkt_info: dict): - conn = None +def _make_pkt_tuple(pkt, ingress_iface: str, egress_list): + """ + Convert packet info to a tuple matching DB insert order. + We store only a small raw snapshot to keep DB small. + """ try: - conn = await asyncpg.connect(DB_DSN) + 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 + except Exception: + src_mac = dst_mac = eth_type = None - await conn.execute(""" - INSERT INTO packets( - iface, - direction, - src_mac, - dst_mac, - eth_type, - src_ip, - dst_ip, - protocol, - length, - ebpf_verdict, - raw - ) - VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) - """, - pkt_info["ingress"], - "unknown", - 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"], - str(pkt_info["egress"]), - pkt_info["raw"] - ) + try: + src_ip = pkt[IP].src 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 + except Exception: + src_ip = dst_ip = protocol = None - except (asyncpg.PostgresError, ConnectionError, OSError) as e: - logger.exception(f"DB insert failed: {e}") + raw = bytes(pkt)[:SNAPSHOT_MAX_BYTES] if pkt is not None else b"" - finally: - if conn: - await conn.close() + # ebpf_verdict field currently used to store egress as text (you can change schema) + ebpf_verdict_text = str(egress_list) if egress_list else None + return ( + ingress_iface, + "unknown", # direction + src_mac, + dst_mac, + str(eth_type) if eth_type is not None else None, + src_ip, + dst_ip, + str(protocol) if protocol is not None else None, + len(raw), + ebpf_verdict_text, + raw + ) -# ------------------------------------------------------------------- -# PACKET HANDLER -# ------------------------------------------------------------------- def handle_packet(pkt, bridge: str): + """ + Callback running inside Scapy sniff thread. + Places a lightweight tuple in packet_queue. + """ pkt_iface = getattr(pkt, "sniffed_on", None) if not pkt_iface: + logger.debug("Packet without sniffed_on metadata, ignoring") return ingress, egress = determine_direction(pkt_iface, bridge) - - pkt_info = { - "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, - "src_ip": pkt[IP].src 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, - "length": len(pkt), - "raw": pkt.json() - } - - asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop) - + tup = _make_pkt_tuple(pkt, ingress, egress) + try: + packet_queue.put_nowait(tup) + except Full: + # queue full -> drop packet and log throttle event + logger.warning("Packet queue full, dropping packet (producer side)") # ------------------------------------------------------------------- -# CLEAN STOPPING SNIFFERS +# Threaded sniff loop # ------------------------------------------------------------------- def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): logger.info(f"Sniffer STARTED on {ifname}") - if not check_interface_exists(ifname): - logger.error(f"[ERROR] Interface {ifname} does not exist. Stopping sniffer.") + logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}") return - if not check_interface_up(ifname): - logger.error(f"[ERROR] Interface {ifname} is DOWN. Stopping sniffer.") + logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}") return while not stop_event.is_set(): @@ -173,13 +287,13 @@ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): iface=ifname, prn=lambda pkt: handle_packet(pkt, bridge), store=False, - timeout=1 # periodic return so we can check stop_event + timeout=1 # short timeout to check stop_event frequently ) except PermissionError: - logger.error(f"[ERROR] Permission denied sniffing on {ifname}. Run as root.") + logger.error(f"Permission denied sniffing on {ifname}. Run as root.") break except OSError as e: - logger.error(f"[ERROR] Sniffer error on {ifname}: {e}") + logger.error(f"Sniffer OSError on {ifname}: {e}") break except Exception as e: logger.exception(f"Unexpected sniffer error on {ifname}: {e}") @@ -187,67 +301,126 @@ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): logger.info(f"Sniffer STOPPED on {ifname}") - def start_sniffer_thread(bridge: str): ports = get_bridge_ports(bridge) if not ports: - logger.error(f"[ERROR] Could not start sniffer: No valid ports found for bridge {bridge}") + logger.error(f"No ports found for bridge {bridge}, not starting sniffers.") return {} for iface in ports: - if iface in sniffer_threads: - logger.info(f"[INFO] Sniffer on {iface} is already running") + if iface in sniffer_threads and sniffer_threads[iface].is_alive(): + logger.info(f"Sniffer already running on {iface}") continue stop_event = threading.Event() thread_stop_flags[iface] = stop_event - thread = threading.Thread( - target=sniffer_loop, - args=(iface, stop_event, bridge), - daemon=True - ) - + thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True) sniffer_threads[iface] = thread thread.start() return sniffer_threads +# ------------------------------------------------------------------- +# Public API for starting/stopping sniffing (to be called from FastAPI) +# ------------------------------------------------------------------- +_consumer_stop_event: Optional[asyncio.Event] = None +_consumer_future_accessor = None -# ------------------------------------------------------------------- -# PUBLIC ASYNC START/STOP METHODS -# ------------------------------------------------------------------- async def start_sniffing(bridge: str): - logger.info(f"Starting sniffing for bridge {bridge}") + """ + Start sniffing on all bridge ports and start the DB consumer. + Safe to call multiple times. + """ + global _consumer_stop_event, _consumer_future_accessor + logger.info(f"Request to start sniffing on bridge {bridge}") + + # start consumer if not running + if _consumer_stop_event is None: + # create asyncio.Event in async_loop + fut = asyncio.run_coroutine_threadsafe(asyncio.sleep(0), async_loop) + # schedule creation of event and consumer + _consumer_stop_event = asyncio.run_coroutine_threadsafe(asyncio.Event(), async_loop).result() + # start consumer in async_loop as a future + def schedule_consumer(): + nonlocal _consumer_stop_event + # schedule consumer coroutine directly + future = asyncio.run_coroutine_threadsafe(_db_consumer_loop(_consumer_stop_event), async_loop) + return future + + # store accessor to the Future to allow waiting for termination on stop + _consumer_future_accessor = schedule_consumer + + # start sniffer threads start_sniffer_thread(bridge) return True - async def stop_sniffing(): - logger.info("Stopping all sniffers…") + """ + Stop all sniffer threads and the DB consumer. + """ + global _consumer_stop_event, _consumer_future_accessor + logger.info("Stopping all sniffers and consumer") - for iface, stop_event in thread_stop_flags.items(): - stop_event.set() + # stop sniffer threads + for iface, ev in list(thread_stop_flags.items()): + ev.set() - for iface, thread in sniffer_threads.items(): + for iface, thread in list(sniffer_threads.items()): thread.join(timeout=2) sniffer_threads.clear() thread_stop_flags.clear() - logger.info("All sniffers stopped.") + # stop consumer + if _consumer_stop_event is not None: + # signal consumer in async_loop + def set_stop(): + _consumer_stop_event.set() + asyncio.run_coroutine_threadsafe(asyncio.to_thread(set_stop), async_loop).result(timeout=2) + + # wait for consumer future + if _consumer_future_accessor: + fut = _consumer_future_accessor() + try: + fut.result(timeout=5) + except Exception as e: + logger.debug(f"Consumer future join error: {e}") + + _consumer_stop_event = None + _consumer_future_accessor = None + + # flush queue attempt (best-effort) + logger.info("Flushing packet queue (best-effort)") + # Not blocking: drop packets on stop + while not packet_queue.empty(): + try: + packet_queue.get_nowait() + except Empty: + break + + # close pool cleanly from async loop + try: + asyncio.run_coroutine_threadsafe(_close_pool(), async_loop).result(timeout=5) + except Exception: + logger.exception("Failed closing pg pool cleanly") + + logger.info("All stopped") return True - # ------------------------------------------------------------------- -# STATUS HELPER +# Status helper # ------------------------------------------------------------------- def get_sniffer_status(): + """ + Return a dict of interface -> status information. + """ out = {} for iface, t in sniffer_threads.items(): out[iface] = { "running": t.is_alive(), "exists": check_interface_exists(iface), - "up": check_interface_up(iface) + "up": check_interface_up(iface), } + out["queue_size"] = packet_queue.qsize() return out diff --git a/backend/src/routes_sniffer.py b/backend/src/routes_sniffer.py index 959c1db..052da3e 100644 --- a/backend/src/routes_sniffer.py +++ b/backend/src/routes_sniffer.py @@ -4,25 +4,16 @@ from src.network_sniffer import get_sniffer_status, start_sniffing, stop_sniffin router = APIRouter() -@router.post("/start_sniffer") -async def start_sniffer(bridge: str): - """ - Start packet sniffing on bridge member interfaces. - """ +@router.post("/sniffer/start") +async def api_start(bridge: str): + ok = await start_sniffing(bridge) + return {"started": ok} - await start_sniffing(bridge) +@router.post("/sniffer/stop") +async def api_stop(): + ok = await stop_sniffing() + return {"stopped": ok} - return { - "status": "ok", - "started_on": bridge - } - - -@router.post("/stop_sniffer") -async def stop_sniffer_api(): - await stop_sniffing() - return {"status": "stopped"} - -@router.get("/status") -async def status(): +@router.get("/sniffer/status") +def api_status(): return get_sniffer_status() \ No newline at end of file