diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index 4107719..773544c 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -1,426 +1,401 @@ +""" +network_sniffer.py + +Usage: +- import start_sniffing, stop_sniffing, get_sniffer_status from this module in your FastAPI routes. +- start_sniffing(bridge_name) will spawn per-interface sniffing threads for all bridge ports. +- stop_sniffing() will stop all running sniffer threads cleanly. +""" + import asyncio import threading import logging import os -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 +from scapy.all import sniff, Ether, IP # scapy must be installed import asyncpg -# --------------------------- -# Configuration -# --------------------------- +logger = logging.getLogger("network_sniffer") 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) +# ----- CONFIG ----- +DB_DSN = os.getenv("MITM_DB_DSN", "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db") +TABLE_NAME = os.getenv("MITM_TABLE", "packet_log") # change if you use a different table +# ------------------ # Active sniffer threads and stop flags sniffer_threads: Dict[str, threading.Thread] = {} thread_stop_flags: Dict[str, threading.Event] = {} +# map which interface belongs to which bridge (reverse mapping) +iface_to_bridge: Dict[str, str] = {} -# Cache for bridge -> ports +# Bridge→ports cache bridge_ports_cache: Dict[str, List[str]] = {} -# Queue for packet events (thread-safe) -packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE) +# Async loop + DB pool (run in background thread) +_async_loop: Optional[asyncio.AbstractEventLoop] = None +_db_pool: Optional[asyncpg.pool.Pool] = None +_loop_thread: Optional[threading.Thread] = None +_loop_started_event = threading.Event() -# Async event loop + background consumer task -async_loop = asyncio.new_event_loop() -consumer_task: Optional[asyncio.Task] = None -pg_pool: Optional[asyncpg.Pool] = None -# ------------------------------------------------------------------- -# Async loop bootstrap (runs in background thread) -# ------------------------------------------------------------------- -def _start_async_loop(loop): +# ------------------------ +# Async loop bootstrap +# ------------------------ +def _start_async_loop(loop: asyncio.AbstractEventLoop): + """Run the given event loop forever (target for background thread).""" asyncio.set_event_loop(loop) + _loop_started_event.set() loop.run_forever() -threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start() -# ------------------------------------------------------------------- -# DB helper: create pool & consumer -# ------------------------------------------------------------------- -async def _create_pool(): - global pg_pool - if pg_pool is None: +def ensure_async_loop(): + """Create and start the dedicated asyncio loop once.""" + global _async_loop, _loop_thread + if _async_loop is not None: + return _async_loop + + _async_loop = asyncio.new_event_loop() + _loop_thread = threading.Thread(target=_start_async_loop, args=(_async_loop,), daemon=True) + _loop_thread.start() + # wait for loop to be set in thread + _loop_started_event.wait(timeout=5) + if not _loop_started_event.is_set(): + raise RuntimeError("Failed to start async loop thread") + logger.info("Async loop started in background thread") + return _async_loop + + +async def _create_db_pool(): + """Coroutine that creates asyncpg pool. Run this in the background loop.""" + global _db_pool + if _db_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 + _db_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=6) + logger.info("DB pool created") + return _db_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() +def ensure_db_pool(): + """Ensure pool exists by scheduling creation on background loop.""" + ensure_async_loop() + # schedule coroutine and wait for it to complete + fut = asyncio.run_coroutine_threadsafe(_create_db_pool(), _async_loop) 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") + return fut.result(timeout=10) + except Exception as e: + logger.exception("Failed to create DB pool: %s", e) + raise -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 -# ------------------------------------------------------------------- +# ------------------------ +# Bridge helpers +# ------------------------ 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 Exception: + except FileNotFoundError: return False + def get_bridge_ports(bridge: str) -> List[str]: + """ + Return member interfaces of the bridge. Uses /sys/class/net//brif/. + Returns [] if bridge missing or on error. + """ 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") + logger.error("Bridge %s does not exist (or brif missing)", bridge) return [] try: - ports = os.listdir(base) + ports = [p for p in os.listdir(base) if os.path.isdir(f"/sys/class/net/{p}")] except PermissionError: - logger.error(f"No permission to read bridge ports for '{bridge}'") + logger.error("Permission denied reading bridge ports for %s", bridge) return [] - ok_ports = [] + bridge_ports_cache[bridge] = ports + # register iface->bridge mapping for p in ports: - if check_interface_exists(p): - ok_ports.append(p) - else: - logger.warning(f"Bridge member {p} listed but interface missing in /sys/class/net") + iface_to_bridge[p] = bridge + logger.info("Bridge %s ports: %s", bridge, ports) + return ports - bridge_ports_cache[bridge] = ok_ports - logger.info(f"Bridge {bridge} ports: {ok_ports}") - return ok_ports def determine_direction(pkt_iface: str, bridge: str): + """ + ingress = interface where packet was captured + egress = list of other bridge ports (where it would be forwarded) + """ ports = get_bridge_ports(bridge) ingress = pkt_iface egress = [p for p in ports if p != pkt_iface] return ingress, egress -# ------------------------------------------------------------------- -# Producer (sniffer threads) -# ------------------------------------------------------------------- -def _make_pkt_tuple(pkt, ingress_iface: str, egress_list): + +# ------------------------ +# DB insert +# ------------------------ +async def _db_insert_packet_coroutine(pkt_info: dict): """ - Convert packet info to a tuple matching DB insert order. - We store only a small raw snapshot to keep DB small. + Runs in event loop; uses pool to insert. + Matches the schema of packet_log with minimal fields: + interface, direction, src_mac, dst_mac, eth_type, + src_ip, dst_ip, ip_protocol, packet_len, raw_packet + Adjust columns if your table differs. + """ + pool = await _create_db_pool() # ensure pool + conn = None + try: + async with pool.acquire() as conn: + await conn.execute(f""" + INSERT INTO {TABLE_NAME}( + interface, + direction, + src_mac, + dst_mac, + eth_type, + src_ip, + dst_ip, + ip_protocol, + packet_len, + raw_packet + ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) + """, + pkt_info.get("ingress"), + pkt_info.get("direction", "unknown"), + pkt_info.get("src_mac"), + pkt_info.get("dst_mac"), + pkt_info.get("eth_type"), + pkt_info.get("src_ip"), + pkt_info.get("dst_ip"), + pkt_info.get("protocol"), + pkt_info.get("length"), + pkt_info.get("raw") + ) + except Exception: + logger.exception("DB insert failed") + raise + + +def schedule_db_insert(pkt_info: dict): + """ + Thread-safe helper: schedule DB insert into background loop. """ try: - 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 + ensure_db_pool() except Exception: - src_mac = dst_mac = eth_type = None + logger.error("DB pool not available; dropping packet") + return + # schedule coroutine + loop = ensure_async_loop() + fut = asyncio.run_coroutine_threadsafe(_db_insert_packet_coroutine(pkt_info), loop) + # optional: attach callback to log failures asynchronously + def _on_done(f): + try: + f.result() + except Exception: + logger.exception("Async DB insert failed") + fut.add_done_callback(_on_done) + + +# ------------------------ +# Packet handling +# ------------------------ +def _safe_get_src_dst(pkt): + """Return (src_ip, dst_ip, proto) safely if IP present.""" 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 + if IP in pkt: + return pkt[IP].src, pkt[IP].dst, pkt[IP].proto except Exception: - src_ip = dst_ip = protocol = None + pass + return None, None, None - raw = bytes(pkt)[:SNAPSHOT_MAX_BYTES] if pkt is not None else b"" - - # 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 - ) def handle_packet(pkt, bridge: str): """ - Callback running inside Scapy sniff thread. - Places a lightweight tuple in packet_queue. + Called in sniffing thread. Build pkt_info and schedule DB insert. """ pkt_iface = getattr(pkt, "sniffed_on", None) if not pkt_iface: - logger.debug("Packet without sniffed_on metadata, ignoring") + # scapy sometimes doesn't set sniffed_on; skip if unknown + logger.warning("Packet without sniffed_on - ignoring") return ingress, egress = determine_direction(pkt_iface, bridge) - 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)") + src_ip, dst_ip, proto = _safe_get_src_dst(pkt) -# ------------------------------------------------------------------- -# Threaded sniff loop -# ------------------------------------------------------------------- + pkt_info = { + "ingress": ingress, + "egress": egress, + "direction": "ingress", # we store 'ingress' here (packet was received on ingress) + "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": src_ip, + "dst_ip": dst_ip, + "protocol": proto, + "length": len(pkt), + # store raw as binary; use bytes(pkt) + "raw": bytes(pkt) + } + + # schedule asynchronous DB insert + schedule_db_insert(pkt_info) + + +# ------------------------ +# Sniffer thread loop +# ------------------------ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): - logger.info(f"Sniffer STARTED on {ifname}") + logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge) + if not check_interface_exists(ifname): - logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}") - return - if not check_interface_up(ifname): - logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}") + logger.error("Interface %s does not exist; stopping sniffer thread", ifname) return + if not check_interface_up(ifname): + logger.warning("Interface %s is down; sniffing might still work but check interface state", ifname) + + # loop with short timeout so stop_event is checked often while not stop_event.is_set(): try: sniff( iface=ifname, prn=lambda pkt: handle_packet(pkt, bridge), store=False, - timeout=1 # short timeout to check stop_event frequently + timeout=1 # returns periodically so we can check stop_event ) except PermissionError: - logger.error(f"Permission denied sniffing on {ifname}. Run as root.") + logger.exception("Permission denied sniffing on %s - run as root or give CAP_NET_RAW", ifname) break except OSError as e: - logger.error(f"Sniffer OSError on {ifname}: {e}") + logger.exception("OS error while sniffing on %s: %s", ifname, e) + # brief sleep to avoid tight loop on repeated errors + stop_event.wait(1) break - except Exception as e: - logger.exception(f"Unexpected sniffer error on {ifname}: {e}") + except Exception: + logger.exception("Unexpected error in sniffer loop for %s", ifname) + stop_event.wait(0.5) break - logger.info(f"Sniffer STOPPED on {ifname}") + logger.info("Sniffer thread exiting for %s", ifname) -def start_sniffer_thread(bridge: str): + +# ------------------------ +# Control API (importable) +# ------------------------ +def start_sniffer_thread_for_bridge(bridge: str) -> Dict[str, threading.Thread]: + """ + Start sniff threads for all ports of given bridge. + Returns mapping iface->thread for started threads (existing threads are left running). + """ ports = get_bridge_ports(bridge) if not ports: - logger.error(f"No ports found for bridge {bridge}, not starting sniffers.") + logger.error("No ports for bridge %s - nothing to start", bridge) return {} + started = {} for iface in ports: if iface in sniffer_threads and sniffer_threads[iface].is_alive(): - logger.info(f"Sniffer already running on {iface}") + logger.info("Sniffer already running on %s", iface) + started[iface] = sniffer_threads[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) - sniffer_threads[iface] = thread + stop_evt = threading.Event() + thread = threading.Thread(target=sniffer_loop, args=(iface, stop_evt, bridge), daemon=True) thread.start() + sniffer_threads[iface] = thread + thread_stop_flags[iface] = stop_evt + started[iface] = thread + logger.info("Started sniffer on %s (bridge=%s)", iface, bridge) - return sniffer_threads + return started -# ------------------------------------------------------------------- -# Public API for starting/stopping sniffing (to be called from FastAPI) -# ------------------------------------------------------------------- -_consumer_stop_event: Optional[asyncio.Event] = None -_consumer_future_accessor = None -async def start_sniffing(bridge: str): +async def start_sniffing(bridge: str) -> bool: """ - Start sniffing on all bridge ports and start the DB consumer. - Safe to call multiple times. + Async entrypoint you can call from FastAPI startup/route. """ - 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) + ensure_async_loop() + ensure_db_pool() + start_sniffer_thread_for_bridge(bridge) return True -async def stop_sniffing(): - """ - Stop all sniffer threads and the DB consumer. - """ - global _consumer_stop_event, _consumer_future_accessor - logger.info("Stopping all sniffers and consumer") - # stop sniffer threads - for iface, ev in list(thread_stop_flags.items()): - ev.set() +async def stop_sniffing() -> bool: + """ + Stop all sniffers and wait briefly for threads to exit. + """ + logger.info("Stopping sniffers (all interfaces)...") + for iface, evt in list(thread_stop_flags.items()): + evt.set() + # join threads with timeout for iface, thread in list(sniffer_threads.items()): + logger.info("Joining thread for %s", iface) thread.join(timeout=2) sniffer_threads.clear() thread_stop_flags.clear() - - # 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") + logger.info("All sniffer threads stopped") return True -# ------------------------------------------------------------------- -# Status helper -# ------------------------------------------------------------------- -def get_sniffer_status(): + +def get_sniffer_status() -> Dict[str, Dict[str, object]]: """ - Return a dict of interface -> status information. + Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] } """ - out = {} - for iface, t in sniffer_threads.items(): - out[iface] = { - "running": t.is_alive(), + status = {} + # include known interfaces (from cache) and active threads + known_ifaces = set(list(sniffer_threads.keys()) + list(iface_to_bridge.keys())) + for iface in known_ifaces: + thread = sniffer_threads.get(iface) + status[iface] = { + "running": bool(thread and thread.is_alive()), "exists": check_interface_exists(iface), "up": check_interface_up(iface), + "bridge": iface_to_bridge.get(iface) } - out["queue_size"] = packet_queue.qsize() - return out + return status + + +# ------------------------ +# Optional convenience: stop all on process exit +# ------------------------ +def _cleanup_on_exit(): + try: + # stop sniffers + for evt in thread_stop_flags.values(): + evt.set() + for t in sniffer_threads.values(): + t.join(timeout=1) + except Exception: + pass + # shutdown db pool and loop + if _db_pool is not None: + try: + fut = asyncio.run_coroutine_threadsafe(_db_pool.close(), _async_loop) + fut.result(timeout=5) + except Exception: + pass + if _async_loop is not None: + try: + _async_loop.call_soon_threadsafe(_async_loop.stop) + except Exception: + pass + + +# register cleanup (best-effort) +import atexit +atexit.register(_cleanup_on_exit)