diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index 39c64f6..1e8aa41 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -2,7 +2,8 @@ import asyncio import logging import threading import os -from typing import List, Dict +import time +from typing import List, Dict, Optional import asyncpg from scapy.all import ( @@ -16,41 +17,70 @@ from scapy.all import ( ICMPv6Unknown, Dot1Q, Raw, - sniff, - conf, ) +# ---- New imports for AF_PACKET optimized reader ---------------------- +import socket +import selectors +import errno +import struct + +# ---- Logging ---------------------------------------------------------- logging.basicConfig(level=logging.INFO) logger = logging.getLogger("af_packet_sniffer") +# ---- Database DSN (change for your environment) ------------------------ DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" -sniffer_threads: Dict[str, threading.Thread] = {} -thread_stop_flags: Dict[str, threading.Event] = {} +# ---- Global state ----------------------------------------------------- +# AF_PACKET sockets keyed by interface name +af_sockets: Dict[str, socket.socket] = {} +# Selector for multiplexing sockets efficiently +af_selector: Optional[selectors.BaseSelector] = None +# Reader thread + stop event +af_thread: Optional[threading.Thread] = None +af_stop_event: Optional[threading.Event] = None + +# Cache of bridge -> ports bridge_ports_cache: Dict[str, List[str]] = {} -# ------------------------------------------------------------------- -# Async loop for DB inserts -# ------------------------------------------------------------------- +# --------------------------------------------------------------------- +# Async loop used to schedule DB inserts from packet callback threads. +# We create a dedicated event loop running in a background thread and +# submit coroutine tasks to it using run_coroutine_threadsafe(). +# --------------------------------------------------------------------- async_loop = asyncio.new_event_loop() -def start_async_loop(loop): +def start_async_loop(loop: asyncio.AbstractEventLoop) -> None: + """ + Entry point for the background thread running the asyncio loop. + This sets the event loop for the thread and runs it forever. + """ asyncio.set_event_loop(loop) loop.run_forever() +# Start the background asyncio loop thread (daemon so it doesn't block process exit). threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start() - -# ------------------------------------------------------------------- -# Interface / bridge checks -# ------------------------------------------------------------------- +# ------------------------- +# Interface / bridge helpers +# ------------------------- def check_interface_exists(iface: str) -> bool: + """ + Check for the presence of a network interface by testing sysfs. + Returns True if /sys/class/net/ exists. + """ return os.path.isdir(f"/sys/class/net/{iface}") def check_interface_up(iface: str) -> bool: + """ + Check whether the given interface is administratively/operationally up + by reading /sys/class/net//operstate. + Returns False if the path does not exist. + """ try: with open(f"/sys/class/net/{iface}/operstate", "r") as f: return f.read().strip() == "up" @@ -59,22 +89,39 @@ def check_interface_up(iface: str) -> bool: def get_bridge_ports(bridge: str) -> List[str]: - if bridge in bridge_ports_cache: - return bridge_ports_cache[bridge] - + """ + Read the bridge member interfaces from sysfs (/sys/class/net//brif/). + This function refreshes the cache each call (cache entry is updated), + which keeps behavior predictable for dynamic topologies. + If the bridge does not exist, an empty list is returned. + """ 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", bridge) return [] - ports = [p for p in os.listdir(base) if check_interface_exists(p)] + try: + ports = [p for p in os.listdir(base) if check_interface_exists(p)] + except Exception as e: + logger.exception("Error listing bridge ports for %s: %s", bridge, e) + ports = [] + + # update local cache and log discovered ports bridge_ports_cache[bridge] = ports - logger.info(f"Bridge {bridge} ports: {ports}") + logger.info("Bridge %s ports: %s", bridge, ports) return ports def determine_direction(pkt_iface: str, bridge: str): - ports = get_bridge_ports(bridge) + """ + Determine ingress and egress information for a packet based on the + interface it was sniffed on and the bridge port list. + + Returns: + ingress: the interface where the packet was observed + egress: list of other bridge ports (possible egress ports) + """ + ports = bridge_ports_cache.get(bridge) or get_bridge_ports(bridge) ingress = pkt_iface egress = [p for p in ports if p != pkt_iface] return ingress, egress @@ -83,7 +130,16 @@ def determine_direction(pkt_iface: str, bridge: str): # ------------------------------------------------------------------- # Database insertion # ------------------------------------------------------------------- -async def db_insert_packet(pkt_info: dict): +async def db_insert_packet(pkt_info: dict) -> None: + """ + Asynchronously insert parsed packet information into the database. + This function is designed to be scheduled on the background asyncio loop + via asyncio.run_coroutine_threadsafe() from other threads. + + pkt_info keys (expected): + - ingress, egress, src_mac, dst_mac, eth_type, vlan_id, + src_ip, dst_ip, protocol_name, src_port, dst_port, length, raw + """ conn = None try: conn = await asyncpg.connect(DB_DSN) @@ -106,7 +162,7 @@ async def db_insert_packet(pkt_info: dict): ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13) """, pkt_info["ingress"], - "unknown", + "unknown", # direction placeholder; matching/annotation can be done later pkt_info["src_mac"], pkt_info["dst_mac"], pkt_info["eth_type"], @@ -120,24 +176,35 @@ async def db_insert_packet(pkt_info: dict): pkt_info["raw"], ) except Exception as e: - logger.exception(f"DB insert failed: {e}") + # Log any database errors but do not re-raise (sniffer should keep running) + logger.exception("DB insert failed: %s", e) finally: if conn: await conn.close() # ------------------------------------------------------------------- -# Packet parsing: AF_PACKET / full stack +# Packet parsing using your existing parse_packet logic # ------------------------------------------------------------------- -def parse_packet(pkt, bridge: str): +def parse_packet(pkt, bridge: str) -> None: + """ + Parse a scapy packet object and collect a normalized dict of metadata + which is then scheduled to be written to the database asynchronously. + + The function expects that 'pkt' is a Scapy Packet and that we set + 'pkt.sniffed_on' before calling this function. + """ pkt_iface = getattr(pkt, "sniffed_on", None) if not pkt_iface: + # If sniffed_on is missing we cannot determine the interface context; + # skip this packet. return - - logger.info(f"Packet captured on {pkt_iface}, bridge {bridge}") + + logger.debug("Packet captured on %s, bridge %s", pkt_iface, bridge) ingress, egress = determine_direction(pkt_iface, bridge) + # Basic normalization structure for DB insertion. pkt_info = { "ingress": ingress, "egress": egress, @@ -154,13 +221,14 @@ def parse_packet(pkt, bridge: str): "dst_port": None, } - # Ethernet + # --- Layer extraction --- + # Ethernet layer if Ether in pkt: pkt_info["src_mac"] = pkt[Ether].src pkt_info["dst_mac"] = pkt[Ether].dst pkt_info["eth_type"] = hex(pkt[Ether].type) - # VLAN + # VLAN (802.1Q) if Dot1Q in pkt: pkt_info["vlan_id"] = pkt[Dot1Q].vlan @@ -208,84 +276,271 @@ def parse_packet(pkt, bridge: str): else: pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}" - # Raw / unknown + # Raw payload / fallback protocol label if Raw in pkt and not pkt_info["protocol_name"]: pkt_info["protocol_name"] = "RAW" + # Submit DB insert to background asyncio loop from this thread. asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop) -# ------------------------------------------------------------------- -# Sniffer thread using AF_PACKET -# ------------------------------------------------------------------- -def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): - logger.info(f"Sniffer STARTED on {ifname}") +# ------------------------- +# AF_PACKET optimized reader +# ------------------------- +def _create_af_packet_socket(ifname: str, rx_buf_bytes: int = 4 * 1024 * 1024) -> Optional[socket.socket]: + """ + Create and bind an AF_PACKET raw socket to the given interface. + Returns the socket or None on failure. - if not check_interface_exists(ifname) or not check_interface_up(ifname): - logger.error(f"Interface {ifname} does not exist or is down. Stopping sniffer.") - return + We configure: + - large SO_RCVBUF to reduce packet drops, + - non-blocking mode, + - best-effort: set PACKET_VERSION = TPACKET_V3 if available. + """ + try: + s = socket.socket(socket.AF_PACKET, socket.SOCK_RAW, socket.htons(0x0003)) # ETH_P_ALL + except PermissionError: + logger.exception("Permission denied creating AF_PACKET socket (need CAP_NET_RAW / root).") + return None + except Exception as e: + logger.exception("Failed to create AF_PACKET socket for %s: %s", ifname, e) + return None - conf.L2socket = conf.L2socket # enforce AF_PACKET usage in scapy + # set a large recv buffer + try: + s.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, rx_buf_bytes) + except Exception: + # non-fatal + logger.debug("Failed to set SO_RCVBUF on %s", ifname) - while not stop_event.is_set(): + # Try to enable TPACKET_V3 (best-effort). Not available on all Python/platforms. + try: + SOL_PACKET = getattr(socket, "SOL_PACKET", 263) # fallback constant + PACKET_VERSION = getattr(socket, "PACKET_VERSION", 10) + TPACKET_V3 = 3 + s.setsockopt(SOL_PACKET, PACKET_VERSION, struct.pack("I", TPACKET_V3)) + logger.debug("Requested TPACKET_V3 on %s", ifname) + except Exception: + # ignore if unsupported + logger.debug("TPACKETv3 not available / not enabled for %s", ifname) + + # Bind to interface index; binding works even if interface is down + try: + s.bind((ifname, 0)) + except OSError as e: + logger.exception("Failed to bind AF_PACKET socket to %s: %s", ifname, e) + s.close() + return None + + s.setblocking(False) + return s + + +def _close_socket(ifname: str) -> None: + """ + Close and remove socket for given interface if present. + """ + s = af_sockets.pop(ifname, None) + if s: try: - sniff( - iface=ifname, - prn=lambda pkt: parse_packet(pkt, bridge), - store=False, - timeout=0.5, # fast stop checks - ) - except Exception as e: - logger.exception(f"Sniffer error on {ifname}: {e}") - break - - logger.info(f"Sniffer STOPPED on {ifname}") + s.close() + except Exception: + pass -# ------------------------------------------------------------------- -# Start / Stop / Status API -# ------------------------------------------------------------------- -def start_sniffer_thread(bridge: str): +def _ensure_sockets_for_bridge(bridge: str) -> None: + """ + Ensure there is an AF_PACKET socket bound for every current bridge port. + We keep sockets for interfaces even if they are DOWN (they will start receiving when link comes up). + """ ports = get_bridge_ports(bridge) - if not ports: - logger.error(f"No valid ports found for bridge {bridge}") - return {} + for p in ports: + if p in af_sockets: + continue + if not check_interface_exists(p): + continue + s = _create_af_packet_socket(p) + if s: + af_sockets[p] = s + # register later in selector by reader thread - for iface in ports: - if iface in sniffer_threads: + +def _rebind_if_needed(ifname: str) -> None: + """ + Try to re-create a socket for an interface if it is missing (e.g. after deletion). + """ + if ifname in af_sockets: + return + if not check_interface_exists(ifname): + return + s = _create_af_packet_socket(ifname) + if s: + af_sockets[ifname] = s + if af_selector: + try: + af_selector.register(s, selectors.EVENT_READ, data=ifname) + except Exception: + # ignore duplicate registration / race + pass + + +def afpacket_reader_loop(bridge: str, stop_event: threading.Event) -> None: + """ + Single reader thread that multiplexes all AF_PACKET sockets for the bridge + using a selector. When data arrives, we parse into a Scapy packet and call parse_packet(). + + This design uses a single thread + selector rather than one thread per interface. + """ + global af_selector + logger.info("AF_PACKET reader starting for bridge %s", bridge) + af_selector = selectors.DefaultSelector() + + # ensure sockets exist for current ports + _ensure_sockets_for_bridge(bridge) + + # register sockets we have + for ifname, s in list(af_sockets.items()): + try: + af_selector.register(s, selectors.EVENT_READ, data=ifname) + except KeyError: + # already registered + pass + except Exception as e: + logger.exception("Failed to register socket for %s: %s", ifname, e) + + # main loop + while not stop_event.is_set(): + # refresh sockets for any new bridge ports (cheap) + try: + _ensure_sockets_for_bridge(bridge) + # register any new sockets with selector + for ifname, s in list(af_sockets.items()): + try: + # register only if not registered + if not any(k.fileobj is s for k in af_selector.get_map().values()): + af_selector.register(s, selectors.EVENT_READ, data=ifname) + except Exception: + # ignore races + pass + except Exception: + logger.exception("Error ensuring sockets") + + # wait for events with short timeout to remain responsive + try: + events = af_selector.select(timeout=1.0) + except Exception as e: + logger.exception("Selector error: %s", e) + time.sleep(0.1) 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 - thread.start() + if not events: + # no events; loop will re-ensure sockets again + continue - return sniffer_threads + for key, mask in events: + sock: socket.socket = key.fileobj + ifname: str = key.data + try: + # read raw frame + # using a single large buffer; AF_PACKET will give full frame + raw = sock.recv(65536) + if not raw: + continue + except BlockingIOError: + continue + except OSError as e: + # handle interface removal (ENODEV) or other errors: close socket and attempt rebind later + if e.errno in (errno.ENODEV, errno.ENETDOWN, errno.EBADF): + logger.warning("Socket error on %s: %s — closing socket and will attempt to rebind later", ifname, e) + try: + af_selector.unregister(sock) + except Exception: + pass + _close_socket(ifname) + # schedule rebind attempt on next loop iteration + continue + else: + logger.exception("Recv error on %s: %s", ifname, e) + continue + + # Parse into scapy Packet (lazy parse) + try: + pkt = Ether(raw) + # attach interface metadata so parse_packet can determine ingress + pkt.sniffed_on = ifname + # call your existing parser + parse_packet(pkt, bridge) + except Exception as e: + logger.exception("Failed to parse/process packet from %s: %s", ifname, e) + continue + + # cleanup + logger.info("AF_PACKET reader stopping; closing sockets") + try: + for key in list(af_selector.get_map().values()): + try: + af_selector.unregister(key.fileobj) + except Exception: + pass + except Exception: + pass + + for ifname in list(af_sockets.keys()): + _close_socket(ifname) + + try: + af_selector.close() + except Exception: + pass + + logger.info("AF_PACKET reader stopped") -async def start_sniffing(bridge: str): - start_sniffer_thread(bridge) - return True +# ------------------------- +# Public start/stop API +# ------------------------- +def start_afpacket_sniffer(bridge: str) -> None: + """ + Start the optimized AF_PACKET sniffer for the provided bridge. + This will bind sockets to all current bridge ports and start a single reader thread. + """ + global af_thread, af_stop_event + if af_thread and af_thread.is_alive(): + logger.info("AF_PACKET sniffer already running") + return + + # ensure we have initial ports + get_bridge_ports(bridge) + af_stop_event = threading.Event() + af_thread = threading.Thread(target=afpacket_reader_loop, args=(bridge, af_stop_event), daemon=True) + af_thread.start() + logger.info("AF_PACKET sniffer started") -async def stop_sniffing(): - 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() - return True +def stop_afpacket_sniffer() -> None: + """ + Stop the AF_PACKET sniffer thread and close sockets. + """ + global af_thread, af_stop_event + if not af_thread: + return + if af_stop_event: + af_stop_event.set() + af_thread.join(timeout=2) + af_thread = None + af_stop_event = None + logger.info("AF_PACKET sniffer stopped") -def get_sniffer_status(): - out = {} - for iface, t in sniffer_threads.items(): +def get_sniffer_status() -> Dict[str, Dict[str, object]]: + """ + Return a status dictionary describing each currently managed interface. + Contains running flag, exists flag, and up flag. + """ + out: Dict[str, Dict[str, object]] = {} + for iface in list(af_sockets.keys()): out[iface] = { - "running": t.is_alive(), + "running": af_thread.is_alive() if af_thread else False, "exists": check_interface_exists(iface), "up": check_interface_up(iface), } diff --git a/backend/src/routes_sniffer.py b/backend/src/routes_sniffer.py index 052da3e..6aa5f71 100644 --- a/backend/src/routes_sniffer.py +++ b/backend/src/routes_sniffer.py @@ -1,19 +1,99 @@ -from fastapi import APIRouter -from typing import List -from src.network_sniffer import get_sniffer_status, start_sniffing, stop_sniffing +from fastapi import APIRouter, HTTPException +from pydantic import BaseModel, Field +from typing import Dict, Any, List, Optional + +from src.network_sniffer import ( + get_sniffer_status, + start_afpacket_sniffer, + stop_afpacket_sniffer, +) router = APIRouter() -@router.post("/sniffer/start") -async def api_start(bridge: str): - ok = await start_sniffing(bridge) - return {"started": ok} +# ------------------------------ +# Pydantic Models +# ------------------------------ -@router.post("/sniffer/stop") -async def api_stop(): - ok = await stop_sniffing() - return {"stopped": ok} +class SnifferStartRequest(BaseModel): + """ + Request model for starting the sniffer on a specific bridge. + """ + bridge: str = Field(..., example="br0", description="Name of the Linux bridge to sniff on") -@router.get("/sniffer/status") -def api_status(): - return get_sniffer_status() \ No newline at end of file + +class SnifferStartResponse(BaseModel): + """ + Response model returned when sniffer starts successfully. + """ + started: bool = Field(..., description="Whether the sniffer was started successfully") + bridge: str = Field(..., description="Bridge where the sniffer was started") + + +class SnifferStopResponse(BaseModel): + """ + Response model returned when the sniffer stops successfully. + """ + stopped: bool = Field(..., description="Whether the sniffer was stopped successfully") + + +class InterfaceSnifferStatus(BaseModel): + """ + Status of an individual interface monitored by the AF_PACKET sniffer. + """ + running: bool = Field(..., description="Whether the sniffer thread is active") + exists: bool = Field(..., description="Whether the interface exists in /sys/class/net") + up: bool = Field(..., description="Whether the interface is operationally UP") + + +class SnifferStatusResponse(BaseModel): + """ + Response model for the sniffer status endpoint. + """ + interfaces: Dict[str, InterfaceSnifferStatus] = Field( + ..., description="Map of interface names to their sniffer status" + ) + + +# ------------------------------ +# Endpoints +# ------------------------------ + +@router.post("/sniffer/start", response_model=SnifferStartResponse) +def sniffer_start(req: SnifferStartRequest): + """ + Start the AF_PACKET sniffer for the given bridge. + """ + try: + start_afpacket_sniffer(req.bridge) + return SnifferStartResponse(started=True, bridge=req.bridge) + except Exception as exc: + raise HTTPException(status_code=500, detail=f"Failed to start sniffer: {exc}") + + +@router.post("/sniffer/stop", response_model=SnifferStopResponse) +def sniffer_stop(): + """ + Stop the AF_PACKET sniffer (if running). + """ + try: + stop_afpacket_sniffer() + return SnifferStopResponse(stopped=True) + except Exception as exc: + raise HTTPException(status_code=500, detail=f"Failed to stop sniffer: {exc}") + + +@router.get("/sniffer/status", response_model=SnifferStatusResponse) +def sniffer_status(): + """ + Return the sniffer status information. + """ + try: + raw = get_sniffer_status() + # Convert raw dict → typed model + typed = { + k: InterfaceSnifferStatus(**v) + for k, v in raw.items() + } + return SnifferStatusResponse(interfaces=typed) + except Exception as exc: + raise HTTPException(status_code=500, detail=f"Failed to query sniffer status: {exc}")