simplify sniffer
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-12-03 19:21:12 +01:00
parent fb33aadc79
commit 309aeff2f1

View File

@@ -1,11 +1,14 @@
import asyncio import asyncio
import logging import logging
import threading import threading
import os
import time import time
from typing import List, Dict, Optional import socket
import selectors
import errno
import struct
from typing import Dict, List, Optional, Any
import asyncpg import asyncpg
from src.utilities.interface_bridge_helpers import check_interface_exists, check_interface_up, get_bridge_ports_once
from scapy.all import ( from scapy.all import (
Ether, Ether,
ARP, ARP,
@@ -18,73 +21,73 @@ from scapy.all import (
Dot1Q, Dot1Q,
Raw, Raw,
) )
import socket
import selectors
import errno
import struct
from src.utilities.interface_bridge_helpers import (
check_interface_exists,
check_interface_up,
get_bridge_ports_once,
)
from src.Models.etherType import EtherTypeEnum, ethertype_from_int from src.Models.etherType import EtherTypeEnum, ethertype_from_int
from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number
# ---- Logging ---------------------------------------------------------- # ---- 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) ------------------------ # ---- Config -----------------------------------------------------------
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
# ---- Global state ----------------------------------------------------- # ---- Globals (kept minimal) -------------------------------------------
# AF_PACKET sockets keyed by interface name
af_sockets: Dict[str, socket.socket] = {} 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_thread: Optional[threading.Thread] = None
af_stop_event: Optional[threading.Event] = None af_stop_event: Optional[threading.Event] = None
# Snapshot of bridge -> ports taken once when sniffer starts
fixed_bridge_ports: Dict[str, List[str]] = {} fixed_bridge_ports: Dict[str, List[str]] = {}
# Remember which bridge the sniffer is using (single active sniffer model)
current_bridge: Optional[str] = None current_bridge: Optional[str] = None
# --------------------------------------------------------------------- # Background asyncio loop used to run DB tasks
# 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: asyncio.AbstractEventLoop) -> None: 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). # Start background loop in daemon thread immediately
threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start() threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
# ------------------------------------------------------------------- # -------------------------
# Database insertion # Database helper (pool)
# ------------------------------------------------------------------- # -------------------------
async def db_insert_packet(pkt_info: dict, bridge: str) -> None: class DB:
_pool: Optional[asyncpg.pool.Pool] = None
@classmethod
async def init_pool(cls) -> None:
if cls._pool is None:
cls._pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=5)
logger.info("DB pool initialized")
@classmethod
async def close_pool(cls) -> None:
if cls._pool:
await cls._pool.close()
cls._pool = None
logger.info("DB pool closed")
@classmethod
async def insert_packet(cls, pkt_info: Dict[str, Any]) -> None:
""" """
Asynchronously insert parsed packet information into the database. Insert packet metadata into DB. This preserves the exact columns/values in the original code.
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 if cls._pool is None:
# defensive: try to initialize if not ready
await cls.init_pool()
try: try:
conn = await asyncpg.connect(DB_DSN) async with cls._pool.acquire() as conn:
await conn.execute( await conn.execute(
""" """
INSERT INTO packets( INSERT INTO packets(
@@ -103,9 +106,9 @@ async def db_insert_packet(pkt_info: dict, bridge: str) -> None:
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12) ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12)
""", """,
pkt_info["iface"], pkt_info["iface"],
pkt_info["src_mac"], pkt_info.get("src_mac"),
pkt_info["dst_mac"], pkt_info.get("dst_mac"),
pkt_info["eth_type"], pkt_info.get("eth_type"),
pkt_info.get("vlan_id"), pkt_info.get("vlan_id"),
pkt_info.get("src_ip"), pkt_info.get("src_ip"),
pkt_info.get("dst_ip"), pkt_info.get("dst_ip"),
@@ -115,128 +118,104 @@ async def db_insert_packet(pkt_info: dict, bridge: str) -> None:
pkt_info["length"], pkt_info["length"],
pkt_info["raw"], pkt_info["raw"],
) )
except Exception as e: except Exception:
# Log any database errors but do not re-raise (sniffer should keep running) # keep sniffer alive: log and continue
logger.exception("DB insert failed: %s", e) logger.exception("DB insert failed")
finally:
if conn:
await conn.close() # schedule pool init on background loop (fire-and-forget)
asyncio.run_coroutine_threadsafe(DB.init_pool(), async_loop)
# -------------------------
# Packet parsing
# -------------------------
def _safe_get_attr(layer, attr: str):
try:
return getattr(layer, attr, None)
except Exception:
return None
# -------------------------------------------------------------------
# Packet parsing callback used by scapy-like parser
# -------------------------------------------------------------------
def parse_packet(pkt, bridge: str) -> None: def parse_packet(pkt, bridge: str) -> None:
""" """
Parse a scapy packet object and collect a normalized dict of metadata Parse a scapy Packet object into a normalized dict and schedule DB insert.
which is then scheduled to be written to the database asynchronously. Keeps all information from the original implementation (fields, VLAN handling,
IP/IPv6/TCP/UDP/ARP handling and enum mapping).
Enhancements:
- Uses EtherType and IP Protocol enums (human-readable) for API documentation.
- Records both raw numeric values and enum descriptions.
- Proper VLAN handling (Dot1Q inner ethertype).
- Keeps `protocol_name` (string) for backward compatibility.
""" """
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 # can't determine interface context
return
logger.debug("Packet captured on %s, bridge %s", pkt_iface, bridge) logger.debug("Packet captured on %s (bridge %s)", pkt_iface, bridge)
# Basic normalization structure for DB insertion. pkt_info: Dict[str, Any] = {
pkt_info = {
"iface": pkt_iface, "iface": pkt_iface,
"length": len(pkt), "length": len(pkt),
"raw": bytes(pkt), "raw": bytes(pkt),
"src_mac": None, "src_mac": None,
"dst_mac": None, "dst_mac": None,
# eth types: both raw numeric and enum/description
"eth_type_raw": None, "eth_type_raw": None,
"eth_type": EtherTypeEnum.UNKNOWN, # enum / human description "eth_type": EtherTypeEnum.UNKNOWN,
"vlan_id": None, "vlan_id": None,
# IP protocol: raw numeric and enum/description
"protocol_raw": None, "protocol_raw": None,
"protocol": IPProtocolEnum.UNKNOWN, "protocol": IPProtocolEnum.UNKNOWN,
# backward-compatible string label (older code)
"protocol_name": None, "protocol_name": None,
# L3/L4 fields
"src_ip": None, "src_ip": None,
"dst_ip": None, "dst_ip": None,
"src_port": None, "src_port": None,
"dst_port": None, "dst_port": None,
} }
# --- Layer extraction --- # Ethernet layer
# Ethernet layer + EtherType
if Ether in pkt: if Ether in pkt:
try: eth = pkt[Ether]
pkt_info["src_mac"] = pkt[Ether].src pkt_info["src_mac"] = _safe_get_attr(eth, "src")
except Exception: pkt_info["dst_mac"] = _safe_get_attr(eth, "dst")
pkt_info["src_mac"] = None
try:
pkt_info["dst_mac"] = pkt[Ether].dst
except Exception:
pkt_info["dst_mac"] = None
# Base ethertype (may be 0x8100 for VLAN) # Base ethertype
eth_type_raw: Optional[int] = None
try: try:
eth_type_raw = int(pkt[Ether].type) eth_type_raw = int(eth.type)
except Exception: except Exception:
eth_type_raw = None eth_type_raw = None
# VLAN (Dot1Q) may contain the inner ethertype # VLAN inner ethertype and vlan id if Dot1Q exists
if Dot1Q in pkt: if Dot1Q in pkt:
try: try:
# Dot1Q.type is the encapsulated ethertype
inner = int(pkt[Dot1Q].type) inner = int(pkt[Dot1Q].type)
if inner: if inner:
eth_type_raw = inner eth_type_raw = inner
except Exception: except Exception:
# ignore and keep whatever eth_type_raw was
pass pass
# capture vlan id if present
try: try:
pkt_info["vlan_id"] = int(pkt[Dot1Q].vlan) pkt_info["vlan_id"] = int(pkt[Dot1Q].vlan)
except Exception: except Exception:
pkt_info["vlan_id"] = None pkt_info["vlan_id"] = None
# Fill eth_type fields (raw + enum)
if eth_type_raw is not None: if eth_type_raw is not None:
pkt_info["eth_type_raw"] = eth_type_raw pkt_info["eth_type_raw"] = eth_type_raw
try: try:
pkt_info["eth_type"] = ethertype_from_int(eth_type_raw) pkt_info["eth_type"] = ethertype_from_int(eth_type_raw)
except Exception: except Exception:
pkt_info["eth_type"] = EtherTypeEnum.UNKNOWN pkt_info["eth_type"] = EtherTypeEnum.UNKNOWN
else:
# Unknown / missing ethertype
pkt_info["eth_type_raw"] = None
pkt_info["eth_type"] = EtherTypeEnum.UNKNOWN
# ARP (layer 2/3) # ARP
if ARP in pkt: if ARP in pkt:
arp = pkt[ARP]
pkt_info["protocol_name"] = "ARP" pkt_info["protocol_name"] = "ARP"
pkt_info["src_ip"] = getattr(pkt[ARP], "psrc", None) pkt_info["src_ip"] = _safe_get_attr(arp, "psrc")
pkt_info["dst_ip"] = getattr(pkt[ARP], "pdst", None) pkt_info["dst_ip"] = _safe_get_attr(arp, "pdst")
pkt_info["src_port"] = None pkt_info["src_port"] = None
pkt_info["dst_port"] = None pkt_info["dst_port"] = None
# ARP is a L2 protocol — leave protocol_raw/protocol as UNKNOWN (or set to a sentinel if desired)
# IPv4 # IPv4
if IP in pkt: if IP in pkt:
try: ip = pkt[IP]
pkt_info["src_ip"] = pkt[IP].src pkt_info["src_ip"] = _safe_get_attr(ip, "src")
except Exception: pkt_info["dst_ip"] = _safe_get_attr(ip, "dst")
pkt_info["src_ip"] = None
try:
pkt_info["dst_ip"] = pkt[IP].dst
except Exception:
pkt_info["dst_ip"] = None
# protocol number (IPv4 'protocol' field)
try: try:
proto_num = int(pkt[IP].proto) proto_num = int(_safe_get_attr(ip, "proto"))
except Exception: except Exception:
proto_num = None proto_num = None
@@ -247,36 +226,28 @@ def parse_packet(pkt, bridge: str) -> None:
except Exception: except Exception:
pkt_info["protocol"] = IPProtocolEnum.UNKNOWN pkt_info["protocol"] = IPProtocolEnum.UNKNOWN
# Common transports
if proto_num == 6 and TCP in pkt: if proto_num == 6 and TCP in pkt:
pkt_info["protocol_name"] = "TCP" pkt_info["protocol_name"] = "TCP"
pkt_info["src_port"] = getattr(pkt[TCP], "sport", None) pkt_info["src_port"] = _safe_get_attr(pkt[TCP], "sport")
pkt_info["dst_port"] = getattr(pkt[TCP], "dport", None) pkt_info["dst_port"] = _safe_get_attr(pkt[TCP], "dport")
elif proto_num == 17 and UDP in pkt: elif proto_num == 17 and UDP in pkt:
pkt_info["protocol_name"] = "UDP" pkt_info["protocol_name"] = "UDP"
pkt_info["src_port"] = getattr(pkt[UDP], "sport", None) pkt_info["src_port"] = _safe_get_attr(pkt[UDP], "sport")
pkt_info["dst_port"] = getattr(pkt[UDP], "dport", None) pkt_info["dst_port"] = _safe_get_attr(pkt[UDP], "dport")
elif proto_num == 1 and ICMP in pkt: elif proto_num == 1 and ICMP in pkt:
pkt_info["protocol_name"] = "ICMP" pkt_info["protocol_name"] = "ICMP"
else: else:
# leave protocol_name as numeric fallback if not matched
if pkt_info["protocol_name"] is None: if pkt_info["protocol_name"] is None:
pkt_info["protocol_name"] = f"IP_PROTO_{proto_num}" if proto_num is not None else None pkt_info["protocol_name"] = f"IP_PROTO_{proto_num}" if proto_num is not None else None
# IPv6 # IPv6
if IPv6 in pkt: if IPv6 in pkt:
try: ip6 = pkt[IPv6]
pkt_info["src_ip"] = pkt[IPv6].src pkt_info["src_ip"] = _safe_get_attr(ip6, "src")
except Exception: pkt_info["dst_ip"] = _safe_get_attr(ip6, "dst")
pkt_info["src_ip"] = None
try:
pkt_info["dst_ip"] = pkt[IPv6].dst
except Exception:
pkt_info["dst_ip"] = None
# next header / nh value
try: try:
nh = int(pkt[IPv6].nh) nh = int(_safe_get_attr(ip6, "nh"))
except Exception: except Exception:
nh = None nh = None
@@ -289,70 +260,65 @@ def parse_packet(pkt, bridge: str) -> None:
if nh == 6 and TCP in pkt: if nh == 6 and TCP in pkt:
pkt_info["protocol_name"] = "TCP" pkt_info["protocol_name"] = "TCP"
pkt_info["src_port"] = getattr(pkt[TCP], "sport", None) pkt_info["src_port"] = _safe_get_attr(pkt[TCP], "sport")
pkt_info["dst_port"] = getattr(pkt[TCP], "dport", None) pkt_info["dst_port"] = _safe_get_attr(pkt[TCP], "dport")
elif nh == 17 and UDP in pkt: elif nh == 17 and UDP in pkt:
pkt_info["protocol_name"] = "UDP" pkt_info["protocol_name"] = "UDP"
pkt_info["src_port"] = getattr(pkt[UDP], "sport", None) pkt_info["src_port"] = _safe_get_attr(pkt[UDP], "sport")
pkt_info["dst_port"] = getattr(pkt[UDP], "dport", None) pkt_info["dst_port"] = _safe_get_attr(pkt[UDP], "dport")
elif ICMPv6Unknown in pkt: elif ICMPv6Unknown in pkt:
pkt_info["protocol_name"] = "ICMPv6" pkt_info["protocol_name"] = "ICMPv6"
else: else:
if pkt_info["protocol_name"] is None: if pkt_info["protocol_name"] is None:
pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}" if nh is not None else None pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}" if nh is not None else None
# Raw payload / fallback protocol label # Raw fallback 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 # Submit DB insert to background loop (non-blocking)
asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info, bridge), async_loop) try:
asyncio.run_coroutine_threadsafe(DB.insert_packet(pkt_info), async_loop)
except Exception:
logger.exception("Failed to schedule DB insert")
# ------------------------- # -------------------------
# AF_PACKET optimized reader # AF_PACKET socket utilities
# ------------------------- # -------------------------
def _create_af_packet_socket(ifname: str, rx_buf_bytes: int = 4 * 1024 * 1024) -> Optional[socket.socket]: def _create_af_packet_socket(ifname: str, rx_buf: int = 4 * 1024 * 1024) -> Optional[socket.socket]:
""" """
Create and bind an AF_PACKET raw socket to the given interface. Create and bind an AF_PACKET raw socket to interface.
Returns the socket or None on failure. Non-blocking socket returned or None on failure.
We configure:
- large SO_RCVBUF to reduce packet drops,
- non-blocking mode,
- best-effort: set PACKET_VERSION = TPACKET_V3 if available.
""" """
try: try:
s = socket.socket(socket.AF_PACKET, socket.SOCK_RAW, socket.htons(0x0003)) # ETH_P_ALL s = socket.socket(socket.AF_PACKET, socket.SOCK_RAW, socket.htons(0x0003)) # ETH_P_ALL
except PermissionError: except PermissionError:
logger.exception("Permission denied creating AF_PACKET socket (need CAP_NET_RAW / root).") logger.exception("Permission denied creating AF_PACKET socket for %s", ifname)
return None return None
except Exception as e:
logger.exception("Failed to create AF_PACKET socket for %s: %s", ifname, e)
return None
# set a large recv buffer
try:
s.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, rx_buf_bytes)
except Exception: except Exception:
# non-fatal logger.exception("Failed creating AF_PACKET socket for %s", ifname)
logger.debug("Failed to set SO_RCVBUF on %s", ifname) return None
# Try to enable TPACKET_V3 (best-effort). Not available on all Python/platforms.
try: try:
SOL_PACKET = getattr(socket, "SOL_PACKET", 263) # fallback constant s.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, rx_buf)
except Exception:
logger.debug("SO_RCVBUF set failed for %s (non-fatal)", ifname)
# Try TPACKET_V3 best-effort
try:
SOL_PACKET = getattr(socket, "SOL_PACKET", 263)
PACKET_VERSION = getattr(socket, "PACKET_VERSION", 10) PACKET_VERSION = getattr(socket, "PACKET_VERSION", 10)
TPACKET_V3 = 3 TPACKET_V3 = 3
s.setsockopt(SOL_PACKET, PACKET_VERSION, struct.pack("I", TPACKET_V3)) s.setsockopt(SOL_PACKET, PACKET_VERSION, struct.pack("I", TPACKET_V3))
logger.debug("Requested TPACKET_V3 on %s", ifname) logger.debug("Requested TPACKET_V3 on %s", ifname)
except Exception: except Exception:
# ignore if unsupported logger.debug("TPACKETv3 not available for %s", ifname)
logger.debug("TPACKETv3 not available / not enabled for %s", ifname)
# Bind to interface index; binding works even if interface is down
try: try:
s.bind((ifname, 0)) s.bind((ifname, 0))
except OSError as e: except OSError as e:
logger.exception("Failed to bind AF_PACKET socket to %s: %s", ifname, e) logger.exception("Bind failed for %s: %s", ifname, e)
s.close() s.close()
return None return None
@@ -361,9 +327,6 @@ def _create_af_packet_socket(ifname: str, rx_buf_bytes: int = 4 * 1024 * 1024) -
def _close_socket(ifname: str) -> None: def _close_socket(ifname: str) -> None:
"""
Close and remove socket for given interface if present.
"""
s = af_sockets.pop(ifname, None) s = af_sockets.pop(ifname, None)
if s: if s:
try: try:
@@ -372,74 +335,58 @@ def _close_socket(ifname: str) -> None:
pass pass
def _ensure_sockets_for_bridge_from_snapshot(bridge: str) -> None: # -------------------------
""" # Main reader thread
Create sockets for ports listed in the fixed snapshot. # -------------------------
Only creates sockets for ports that don't already have one.
"""
ports = fixed_bridge_ports.get(bridge, [])
for p in ports:
if p in af_sockets:
continue
if not check_interface_exists(p):
logger.warning("Snapshot port %s does not exist in sysfs (skipping)", p)
continue
s = _create_af_packet_socket(p)
if s:
af_sockets[p] = s
def afpacket_reader_loop(bridge: str, stop_event: threading.Event) -> None: def afpacket_reader_loop(bridge: str, stop_event: threading.Event) -> None:
""" """
Single reader thread that multiplexes all AF_PACKET sockets for the bridge Multiplex AF_PACKET sockets using a selector and hand packets to parse_packet.
using a selector. When data arrives, we parse into a Scapy packet and call parse_packet(). Uses the fixed_bridge_ports snapshot to decide which interfaces to open.
""" """
global af_selector
logger.info("AF_PACKET reader starting for bridge %s", bridge) logger.info("AF_PACKET reader starting for bridge %s", bridge)
af_selector = selectors.DefaultSelector() sel = selectors.DefaultSelector()
# Ensure sockets exist based on the fixed snapshot (created at start) # create sockets for fixed snapshot (if any)
_ensure_sockets_for_bridge_from_snapshot(bridge) for iface in fixed_bridge_ports.get(bridge, []):
if iface in af_sockets:
continue
if not check_interface_exists(iface):
logger.warning("Snapshot port %s missing, skipping", iface)
continue
s = _create_af_packet_socket(iface)
if s:
af_sockets[iface] = s
# register sockets we have # register existing sockets
for ifname, s in list(af_sockets.items()): for iface, s in list(af_sockets.items()):
try: try:
af_selector.register(s, selectors.EVENT_READ, data=ifname) sel.register(s, selectors.EVENT_READ, data=iface)
except KeyError: except Exception:
# already registered logger.debug("Register failed for %s (continuing)", iface)
pass
except Exception as e:
logger.exception("Failed to register socket for %s: %s", ifname, e)
# main loop try:
while not stop_event.is_set(): while not stop_event.is_set():
# In this static snapshot mode we only occasionally try to register # ensure newly created sockets are registered
# newly created sockets (e.g., interfaces that existed but socket creation failed earlier). for iface, s in list(af_sockets.items()):
try: try:
for ifname, s in list(af_sockets.items()): if not any(k.fileobj is s for k in sel.get_map().values()):
try: sel.register(s, selectors.EVENT_READ, data=iface)
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: except Exception:
pass pass
except Exception:
logger.exception("Error ensuring selector registrations")
# wait for events with short timeout to remain responsive
try: try:
events = af_selector.select(timeout=1.0) events = sel.select(timeout=1.0)
except Exception as e: except Exception:
logger.exception("Selector error: %s", e) logger.exception("Selector error")
time.sleep(0.1) time.sleep(0.1)
continue continue
if not events: if not events:
# no events; loop will re-check registrations
continue continue
for key, mask in events: for key, _ in events:
sock: socket.socket = key.fileobj sock: socket.socket = key.fileobj
ifname: str = key.data iface: str = key.data
try: try:
raw = sock.recv(65536) raw = sock.recv(65536)
if not raw: if not raw:
@@ -447,46 +394,43 @@ def afpacket_reader_loop(bridge: str, stop_event: threading.Event) -> None:
except BlockingIOError: except BlockingIOError:
continue continue
except OSError as e: except OSError as e:
# handle interface removal (ENODEV) or other errors: close socket and continue
if e.errno in (errno.ENODEV, errno.ENETDOWN, errno.EBADF): if e.errno in (errno.ENODEV, errno.ENETDOWN, errno.EBADF):
logger.warning("Socket error on %s: %s — closing socket", ifname, e) logger.warning("Socket error on %s: %s — closing", iface, e)
try: try:
af_selector.unregister(sock) sel.unregister(sock)
except Exception: except Exception:
pass pass
_close_socket(ifname) _close_socket(iface)
continue continue
else: else:
logger.exception("Recv error on %s: %s", ifname, e) logger.exception("Recv error on %s", iface)
continue continue
# Parse into scapy Packet (lazy parse) # parse with scapy
try: try:
pkt = Ether(raw) pkt = Ether(raw)
# attach interface metadata so parse_packet can determine ingress pkt.sniffed_on = iface
pkt.sniffed_on = ifname
# call your existing parser
parse_packet(pkt, bridge) parse_packet(pkt, bridge)
except Exception as e: except Exception:
logger.exception("Failed to parse/process packet from %s: %s", ifname, e) logger.exception("Failed to parse/process packet from %s", iface)
continue continue
finally:
# cleanup logger.info("AF_PACKET reader stopping; cleaning up sockets")
logger.info("AF_PACKET reader stopping; closing sockets") # unregister and close
try: try:
for key in list(af_selector.get_map().values()): for key in list(sel.get_map().values()):
try: try:
af_selector.unregister(key.fileobj) sel.unregister(key.fileobj)
except Exception: except Exception:
pass pass
except Exception: except Exception:
pass pass
for ifname in list(af_sockets.keys()): for iface in list(af_sockets.keys()):
_close_socket(ifname) _close_socket(iface)
try: try:
af_selector.close() sel.close()
except Exception: except Exception:
pass pass
@@ -494,19 +438,17 @@ def afpacket_reader_loop(bridge: str, stop_event: threading.Event) -> None:
# ------------------------- # -------------------------
# Public start/stop API # Public API: start/stop/status
# ------------------------- # -------------------------
def start_afpacket_sniffer(bridge: str) -> None: def start_afpacket_sniffer(bridge: str) -> None:
""" """
Start the optimized AF_PACKET sniffer for the provided bridge. Start the sniffer: snapshot ports once, spin up reader thread.
This reads bridge ports once (snapshot) and starts a single reader thread.
""" """
global af_thread, af_stop_event, current_bridge, fixed_bridge_ports global af_thread, af_stop_event, current_bridge
if af_thread and af_thread.is_alive(): if af_thread and af_thread.is_alive():
logger.info("AF_PACKET sniffer already running") logger.info("Sniffer already running")
return return
# Warm and freeze the bridge ports snapshot. We read sysfs once here.
ports = get_bridge_ports_once(bridge) ports = get_bridge_ports_once(bridge)
fixed_bridge_ports[bridge] = ports fixed_bridge_ports[bridge] = ports
current_bridge = bridge current_bridge = bridge
@@ -520,34 +462,41 @@ def start_afpacket_sniffer(bridge: str) -> None:
def stop_afpacket_sniffer() -> None: def stop_afpacket_sniffer() -> None:
""" """
Stop the AF_PACKET sniffer thread and close sockets. Clear the fixed snapshot. Stop the reader thread and clear snapshot. Also attempt to close DB pool.
""" """
global af_thread, af_stop_event, current_bridge, fixed_bridge_ports global af_thread, af_stop_event, current_bridge
if not af_thread: if not af_thread:
return return
if af_stop_event: if af_stop_event:
af_stop_event.set() af_stop_event.set()
af_thread.join(timeout=2) af_thread.join(timeout=2)
af_thread = None af_thread = None
af_stop_event = None af_stop_event = None
# clear fixed snapshot(s)
if current_bridge: if current_bridge:
fixed_bridge_ports.pop(current_bridge, None) fixed_bridge_ports.pop(current_bridge, None)
current_bridge = None current_bridge = None
# close db pool (schedule on async loop)
try:
asyncio.run_coroutine_threadsafe(DB.close_pool(), async_loop)
except Exception:
logger.exception("Failed to schedule DB pool close")
logger.info("AF_PACKET sniffer stopped") logger.info("AF_PACKET sniffer stopped")
def get_sniffer_status() -> Dict[str, Dict[str, object]]: def get_sniffer_status() -> Dict[str, Dict[str, object]]:
""" """
Return a status dictionary describing each currently managed interface. Return simple status per managed interface.
Contains running flag, exists flag, and up flag.
""" """
running = af_thread.is_alive() if af_thread else False
out: Dict[str, Dict[str, object]] = {} out: Dict[str, Dict[str, object]] = {}
for iface in list(af_sockets.keys()): for iface in list(af_sockets.keys()):
out[iface] = { out[iface] = {
"running": af_thread.is_alive() if af_thread else False, "running": running,
"exists": check_interface_exists(iface), "exists": check_interface_exists(iface),
"up": check_interface_up(iface), "up": check_interface_up(iface),
} }