feat: implement AF_PACKET sniffer API with start, stop, and status endpoints
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-11-29 13:25:44 +01:00
parent 9ca8ed531b
commit 24a47fdf37
2 changed files with 432 additions and 97 deletions

View File

@@ -2,7 +2,8 @@ import asyncio
import logging import logging
import threading import threading
import os import os
from typing import List, Dict import time
from typing import List, Dict, Optional
import asyncpg import asyncpg
from scapy.all import ( from scapy.all import (
@@ -16,41 +17,70 @@ from scapy.all import (
ICMPv6Unknown, ICMPv6Unknown,
Dot1Q, Dot1Q,
Raw, Raw,
sniff,
conf,
) )
# ---- New imports for AF_PACKET optimized reader ----------------------
import socket
import selectors
import errno
import struct
# ---- Logging ----------------------------------------------------------
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("af_packet_sniffer") logger = logging.getLogger("af_packet_sniffer")
# ---- Database DSN (change for your environment) ------------------------
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
sniffer_threads: Dict[str, threading.Thread] = {} # ---- Global state -----------------------------------------------------
thread_stop_flags: Dict[str, threading.Event] = {} # 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]] = {} 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() 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) asyncio.set_event_loop(loop)
loop.run_forever() 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() threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start()
# -------------------------
# ------------------------------------------------------------------- # Interface / bridge helpers
# Interface / bridge checks # -------------------------
# -------------------------------------------------------------------
def check_interface_exists(iface: str) -> bool: def check_interface_exists(iface: str) -> bool:
"""
Check for the presence of a network interface by testing sysfs.
Returns True if /sys/class/net/<iface> exists.
"""
return os.path.isdir(f"/sys/class/net/{iface}") return os.path.isdir(f"/sys/class/net/{iface}")
def check_interface_up(iface: str) -> bool: def check_interface_up(iface: str) -> bool:
"""
Check whether the given interface is administratively/operationally up
by reading /sys/class/net/<iface>/operstate.
Returns False if the path does not exist.
"""
try: try:
with open(f"/sys/class/net/{iface}/operstate", "r") as f: with open(f"/sys/class/net/{iface}/operstate", "r") as f:
return f.read().strip() == "up" return f.read().strip() == "up"
@@ -59,22 +89,39 @@ def check_interface_up(iface: str) -> bool:
def get_bridge_ports(bridge: str) -> List[str]: 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/<bridge>/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/" base = f"/sys/class/net/{bridge}/brif/"
if not os.path.isdir(base): if not os.path.isdir(base):
logger.error(f"Bridge '{bridge}' does not exist") logger.error("Bridge '%s' does not exist", bridge)
return [] return []
try:
ports = [p for p in os.listdir(base) if check_interface_exists(p)] 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 bridge_ports_cache[bridge] = ports
logger.info(f"Bridge {bridge} ports: {ports}") logger.info("Bridge %s ports: %s", bridge, ports)
return ports return ports
def determine_direction(pkt_iface: str, bridge: str): 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 ingress = pkt_iface
egress = [p for p in ports if p != pkt_iface] egress = [p for p in ports if p != pkt_iface]
return ingress, egress return ingress, egress
@@ -83,7 +130,16 @@ def determine_direction(pkt_iface: str, bridge: str):
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# Database insertion # 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 conn = None
try: try:
conn = await asyncpg.connect(DB_DSN) 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) ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13)
""", """,
pkt_info["ingress"], pkt_info["ingress"],
"unknown", "unknown", # direction placeholder; matching/annotation can be done later
pkt_info["src_mac"], pkt_info["src_mac"],
pkt_info["dst_mac"], pkt_info["dst_mac"],
pkt_info["eth_type"], pkt_info["eth_type"],
@@ -120,24 +176,35 @@ async def db_insert_packet(pkt_info: dict):
pkt_info["raw"], pkt_info["raw"],
) )
except Exception as e: 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: finally:
if conn: if conn:
await conn.close() 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) pkt_iface = getattr(pkt, "sniffed_on", None)
if not pkt_iface: if not pkt_iface:
# If sniffed_on is missing we cannot determine the interface context;
# skip this packet.
return 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) ingress, egress = determine_direction(pkt_iface, bridge)
# Basic normalization structure for DB insertion.
pkt_info = { pkt_info = {
"ingress": ingress, "ingress": ingress,
"egress": egress, "egress": egress,
@@ -154,13 +221,14 @@ def parse_packet(pkt, bridge: str):
"dst_port": None, "dst_port": None,
} }
# Ethernet # --- Layer extraction ---
# Ethernet layer
if Ether in pkt: if Ether in pkt:
pkt_info["src_mac"] = pkt[Ether].src pkt_info["src_mac"] = pkt[Ether].src
pkt_info["dst_mac"] = pkt[Ether].dst pkt_info["dst_mac"] = pkt[Ether].dst
pkt_info["eth_type"] = hex(pkt[Ether].type) pkt_info["eth_type"] = hex(pkt[Ether].type)
# VLAN # VLAN (802.1Q)
if Dot1Q in pkt: if Dot1Q in pkt:
pkt_info["vlan_id"] = pkt[Dot1Q].vlan pkt_info["vlan_id"] = pkt[Dot1Q].vlan
@@ -208,84 +276,271 @@ def parse_packet(pkt, bridge: str):
else: else:
pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}" 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"]: if Raw in pkt and not pkt_info["protocol_name"]:
pkt_info["protocol_name"] = "RAW" 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) asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop)
# ------------------------------------------------------------------- # -------------------------
# Sniffer thread using AF_PACKET # AF_PACKET optimized reader
# ------------------------------------------------------------------- # -------------------------
def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str): def _create_af_packet_socket(ifname: str, rx_buf_bytes: int = 4 * 1024 * 1024) -> Optional[socket.socket]:
logger.info(f"Sniffer STARTED on {ifname}") """
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): We configure:
logger.error(f"Interface {ifname} does not exist or is down. Stopping sniffer.") - large SO_RCVBUF to reduce packet drops,
return - non-blocking mode,
- best-effort: set PACKET_VERSION = TPACKET_V3 if available.
conf.L2socket = conf.L2socket # enforce AF_PACKET usage in scapy """
while not stop_event.is_set():
try: try:
sniff( s = socket.socket(socket.AF_PACKET, socket.SOCK_RAW, socket.htons(0x0003)) # ETH_P_ALL
iface=ifname, except PermissionError:
prn=lambda pkt: parse_packet(pkt, bridge), logger.exception("Permission denied creating AF_PACKET socket (need CAP_NET_RAW / root).")
store=False, return None
timeout=0.5, # fast stop checks
)
except Exception as e: except Exception as e:
logger.exception(f"Sniffer error on {ifname}: {e}") logger.exception("Failed to create AF_PACKET socket for %s: %s", ifname, e)
break return None
logger.info(f"Sniffer STOPPED on {ifname}") # 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)
# 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:
# Start / Stop / Status API """
# ------------------------------------------------------------------- Close and remove socket for given interface if present.
def start_sniffer_thread(bridge: str): """
s = af_sockets.pop(ifname, None)
if s:
try:
s.close()
except Exception:
pass
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) ports = get_bridge_ports(bridge)
if not ports: for p in ports:
logger.error(f"No valid ports found for bridge {bridge}") if p in af_sockets:
return {} 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 continue
stop_event = threading.Event() if not events:
thread_stop_flags[iface] = stop_event # no events; loop will re-ensure sockets again
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True) continue
sniffer_threads[iface] = thread
thread.start()
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) # Public start/stop API
return True # -------------------------
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(): def stop_afpacket_sniffer() -> None:
for iface, stop_event in thread_stop_flags.items(): """
stop_event.set() Stop the AF_PACKET sniffer thread and close sockets.
"""
for iface, thread in sniffer_threads.items(): global af_thread, af_stop_event
thread.join(timeout=2) if not af_thread:
return
sniffer_threads.clear() if af_stop_event:
thread_stop_flags.clear() af_stop_event.set()
return True af_thread.join(timeout=2)
af_thread = None
af_stop_event = None
logger.info("AF_PACKET sniffer stopped")
def get_sniffer_status(): def get_sniffer_status() -> Dict[str, Dict[str, object]]:
out = {} """
for iface, t in sniffer_threads.items(): 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] = { out[iface] = {
"running": t.is_alive(), "running": af_thread.is_alive() if af_thread else False,
"exists": check_interface_exists(iface), "exists": check_interface_exists(iface),
"up": check_interface_up(iface), "up": check_interface_up(iface),
} }

View File

@@ -1,19 +1,99 @@
from fastapi import APIRouter from fastapi import APIRouter, HTTPException
from typing import List from pydantic import BaseModel, Field
from src.network_sniffer import get_sniffer_status, start_sniffing, stop_sniffing from typing import Dict, Any, List, Optional
from src.network_sniffer import (
get_sniffer_status,
start_afpacket_sniffer,
stop_afpacket_sniffer,
)
router = APIRouter() router = APIRouter()
@router.post("/sniffer/start") # ------------------------------
async def api_start(bridge: str): # Pydantic Models
ok = await start_sniffing(bridge) # ------------------------------
return {"started": ok}
@router.post("/sniffer/stop") class SnifferStartRequest(BaseModel):
async def api_stop(): """
ok = await stop_sniffing() Request model for starting the sniffer on a specific bridge.
return {"stopped": ok} """
bridge: str = Field(..., example="br0", description="Name of the Linux bridge to sniff on")
@router.get("/sniffer/status")
def api_status(): class SnifferStartResponse(BaseModel):
return get_sniffer_status() """
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}")