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 logging
import threading
import os
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
from src.utilities.interface_bridge_helpers import check_interface_exists, check_interface_up, get_bridge_ports_once
from scapy.all import (
Ether,
ARP,
@@ -18,73 +21,73 @@ from scapy.all import (
Dot1Q,
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.ip_protocol import IPProtocolEnum, protocol_from_number
# ---- Logging ----------------------------------------------------------
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("af_packet_sniffer")
# ---- Database DSN (change for your environment) ------------------------
# ---- Config -----------------------------------------------------------
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
# ---- Global state -----------------------------------------------------
# AF_PACKET sockets keyed by interface name
# ---- Globals (kept minimal) -------------------------------------------
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
# Snapshot of bridge -> ports taken once when sniffer starts
fixed_bridge_ports: Dict[str, List[str]] = {}
# Remember which bridge the sniffer is using (single active sniffer model)
current_bridge: Optional[str] = None
# ---------------------------------------------------------------------
# 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().
# ---------------------------------------------------------------------
# Background asyncio loop used to run DB tasks
async_loop = asyncio.new_event_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.
"""
def _start_async_loop(loop: asyncio.AbstractEventLoop) -> None:
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()
# Start background loop in daemon thread immediately
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
# -------------------------------------------------------------------
# Database insertion
# -------------------------------------------------------------------
async def db_insert_packet(pkt_info: dict, bridge: str) -> None:
# -------------------------
# Database helper (pool)
# -------------------------
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.
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
Insert packet metadata into DB. This preserves the exact columns/values in the original code.
"""
conn = None
if cls._pool is None:
# defensive: try to initialize if not ready
await cls.init_pool()
try:
conn = await asyncpg.connect(DB_DSN)
async with cls._pool.acquire() as conn:
await conn.execute(
"""
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)
""",
pkt_info["iface"],
pkt_info["src_mac"],
pkt_info["dst_mac"],
pkt_info["eth_type"],
pkt_info.get("src_mac"),
pkt_info.get("dst_mac"),
pkt_info.get("eth_type"),
pkt_info.get("vlan_id"),
pkt_info.get("src_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["raw"],
)
except Exception as 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()
except Exception:
# keep sniffer alive: log and continue
logger.exception("DB insert failed")
# 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:
"""
Parse a scapy packet object and collect a normalized dict of metadata
which is then scheduled to be written to the database asynchronously.
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.
Parse a scapy Packet object into a normalized dict and schedule DB insert.
Keeps all information from the original implementation (fields, VLAN handling,
IP/IPv6/TCP/UDP/ARP handling and enum mapping).
"""
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
return # can't determine interface context
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 = {
pkt_info: Dict[str, Any] = {
"iface": pkt_iface,
"length": len(pkt),
"raw": bytes(pkt),
"src_mac": None,
"dst_mac": None,
# eth types: both raw numeric and enum/description
"eth_type_raw": None,
"eth_type": EtherTypeEnum.UNKNOWN, # enum / human description
"eth_type": EtherTypeEnum.UNKNOWN,
"vlan_id": None,
# IP protocol: raw numeric and enum/description
"protocol_raw": None,
"protocol": IPProtocolEnum.UNKNOWN,
# backward-compatible string label (older code)
"protocol_name": None,
# L3/L4 fields
"src_ip": None,
"dst_ip": None,
"src_port": None,
"dst_port": None,
}
# --- Layer extraction ---
# Ethernet layer + EtherType
# Ethernet layer
if Ether in pkt:
try:
pkt_info["src_mac"] = pkt[Ether].src
except Exception:
pkt_info["src_mac"] = None
try:
pkt_info["dst_mac"] = pkt[Ether].dst
except Exception:
pkt_info["dst_mac"] = None
eth = pkt[Ether]
pkt_info["src_mac"] = _safe_get_attr(eth, "src")
pkt_info["dst_mac"] = _safe_get_attr(eth, "dst")
# Base ethertype (may be 0x8100 for VLAN)
eth_type_raw: Optional[int] = None
# Base ethertype
try:
eth_type_raw = int(pkt[Ether].type)
eth_type_raw = int(eth.type)
except Exception:
eth_type_raw = None
# VLAN (Dot1Q) may contain the inner ethertype
# VLAN inner ethertype and vlan id if Dot1Q exists
if Dot1Q in pkt:
try:
# Dot1Q.type is the encapsulated ethertype
inner = int(pkt[Dot1Q].type)
if inner:
eth_type_raw = inner
except Exception:
# ignore and keep whatever eth_type_raw was
pass
# capture vlan id if present
try:
pkt_info["vlan_id"] = int(pkt[Dot1Q].vlan)
except Exception:
pkt_info["vlan_id"] = None
# Fill eth_type fields (raw + enum)
if eth_type_raw is not None:
pkt_info["eth_type_raw"] = eth_type_raw
try:
pkt_info["eth_type"] = ethertype_from_int(eth_type_raw)
except Exception:
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:
arp = pkt[ARP]
pkt_info["protocol_name"] = "ARP"
pkt_info["src_ip"] = getattr(pkt[ARP], "psrc", None)
pkt_info["dst_ip"] = getattr(pkt[ARP], "pdst", None)
pkt_info["src_ip"] = _safe_get_attr(arp, "psrc")
pkt_info["dst_ip"] = _safe_get_attr(arp, "pdst")
pkt_info["src_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
if IP in pkt:
try:
pkt_info["src_ip"] = pkt[IP].src
except Exception:
pkt_info["src_ip"] = None
try:
pkt_info["dst_ip"] = pkt[IP].dst
except Exception:
pkt_info["dst_ip"] = None
ip = pkt[IP]
pkt_info["src_ip"] = _safe_get_attr(ip, "src")
pkt_info["dst_ip"] = _safe_get_attr(ip, "dst")
# protocol number (IPv4 'protocol' field)
try:
proto_num = int(pkt[IP].proto)
proto_num = int(_safe_get_attr(ip, "proto"))
except Exception:
proto_num = None
@@ -247,36 +226,28 @@ def parse_packet(pkt, bridge: str) -> None:
except Exception:
pkt_info["protocol"] = IPProtocolEnum.UNKNOWN
# Common transports
if proto_num == 6 and TCP in pkt:
pkt_info["protocol_name"] = "TCP"
pkt_info["src_port"] = getattr(pkt[TCP], "sport", None)
pkt_info["dst_port"] = getattr(pkt[TCP], "dport", None)
pkt_info["src_port"] = _safe_get_attr(pkt[TCP], "sport")
pkt_info["dst_port"] = _safe_get_attr(pkt[TCP], "dport")
elif proto_num == 17 and UDP in pkt:
pkt_info["protocol_name"] = "UDP"
pkt_info["src_port"] = getattr(pkt[UDP], "sport", None)
pkt_info["dst_port"] = getattr(pkt[UDP], "dport", None)
pkt_info["src_port"] = _safe_get_attr(pkt[UDP], "sport")
pkt_info["dst_port"] = _safe_get_attr(pkt[UDP], "dport")
elif proto_num == 1 and ICMP in pkt:
pkt_info["protocol_name"] = "ICMP"
else:
# leave protocol_name as numeric fallback if not matched
if pkt_info["protocol_name"] is None:
pkt_info["protocol_name"] = f"IP_PROTO_{proto_num}" if proto_num is not None else None
# IPv6
if IPv6 in pkt:
try:
pkt_info["src_ip"] = pkt[IPv6].src
except Exception:
pkt_info["src_ip"] = None
try:
pkt_info["dst_ip"] = pkt[IPv6].dst
except Exception:
pkt_info["dst_ip"] = None
ip6 = pkt[IPv6]
pkt_info["src_ip"] = _safe_get_attr(ip6, "src")
pkt_info["dst_ip"] = _safe_get_attr(ip6, "dst")
# next header / nh value
try:
nh = int(pkt[IPv6].nh)
nh = int(_safe_get_attr(ip6, "nh"))
except Exception:
nh = None
@@ -289,70 +260,65 @@ def parse_packet(pkt, bridge: str) -> None:
if nh == 6 and TCP in pkt:
pkt_info["protocol_name"] = "TCP"
pkt_info["src_port"] = getattr(pkt[TCP], "sport", None)
pkt_info["dst_port"] = getattr(pkt[TCP], "dport", None)
pkt_info["src_port"] = _safe_get_attr(pkt[TCP], "sport")
pkt_info["dst_port"] = _safe_get_attr(pkt[TCP], "dport")
elif nh == 17 and UDP in pkt:
pkt_info["protocol_name"] = "UDP"
pkt_info["src_port"] = getattr(pkt[UDP], "sport", None)
pkt_info["dst_port"] = getattr(pkt[UDP], "dport", None)
pkt_info["src_port"] = _safe_get_attr(pkt[UDP], "sport")
pkt_info["dst_port"] = _safe_get_attr(pkt[UDP], "dport")
elif ICMPv6Unknown in pkt:
pkt_info["protocol_name"] = "ICMPv6"
else:
if pkt_info["protocol_name"] is 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"]:
pkt_info["protocol_name"] = "RAW"
# Submit DB insert to background asyncio loop from this thread
asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info, bridge), async_loop)
# Submit DB insert to background loop (non-blocking)
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.
Returns the socket 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.
Create and bind an AF_PACKET raw socket to interface.
Non-blocking socket returned or None on failure.
"""
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).")
logger.exception("Permission denied creating AF_PACKET socket for %s", ifname)
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:
# non-fatal
logger.debug("Failed to set SO_RCVBUF on %s", ifname)
logger.exception("Failed creating AF_PACKET socket for %s", ifname)
return None
# Try to enable TPACKET_V3 (best-effort). Not available on all Python/platforms.
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)
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)
logger.debug("TPACKETv3 not available 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)
logger.exception("Bind failed for %s: %s", ifname, e)
s.close()
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:
"""
Close and remove socket for given interface if present.
"""
s = af_sockets.pop(ifname, None)
if s:
try:
@@ -372,74 +335,58 @@ def _close_socket(ifname: str) -> None:
pass
def _ensure_sockets_for_bridge_from_snapshot(bridge: str) -> None:
"""
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
# -------------------------
# Main reader thread
# -------------------------
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().
Multiplex AF_PACKET sockets using a selector and hand packets to 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)
af_selector = selectors.DefaultSelector()
sel = selectors.DefaultSelector()
# Ensure sockets exist based on the fixed snapshot (created at start)
_ensure_sockets_for_bridge_from_snapshot(bridge)
# create sockets for fixed snapshot (if any)
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
for ifname, s in list(af_sockets.items()):
# register existing sockets
for iface, 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)
sel.register(s, selectors.EVENT_READ, data=iface)
except Exception:
logger.debug("Register failed for %s (continuing)", iface)
# main loop
try:
while not stop_event.is_set():
# In this static snapshot mode we only occasionally try to register
# newly created sockets (e.g., interfaces that existed but socket creation failed earlier).
# ensure newly created sockets are registered
for iface, s in list(af_sockets.items()):
try:
for ifname, s in list(af_sockets.items()):
try:
if not any(k.fileobj is s for k in af_selector.get_map().values()):
af_selector.register(s, selectors.EVENT_READ, data=ifname)
if not any(k.fileobj is s for k in sel.get_map().values()):
sel.register(s, selectors.EVENT_READ, data=iface)
except Exception:
pass
except Exception:
logger.exception("Error ensuring selector registrations")
# 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)
events = sel.select(timeout=1.0)
except Exception:
logger.exception("Selector error")
time.sleep(0.1)
continue
if not events:
# no events; loop will re-check registrations
continue
for key, mask in events:
for key, _ in events:
sock: socket.socket = key.fileobj
ifname: str = key.data
iface: str = key.data
try:
raw = sock.recv(65536)
if not raw:
@@ -447,46 +394,43 @@ def afpacket_reader_loop(bridge: str, stop_event: threading.Event) -> None:
except BlockingIOError:
continue
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):
logger.warning("Socket error on %s: %s — closing socket", ifname, e)
logger.warning("Socket error on %s: %s — closing", iface, e)
try:
af_selector.unregister(sock)
sel.unregister(sock)
except Exception:
pass
_close_socket(ifname)
_close_socket(iface)
continue
else:
logger.exception("Recv error on %s: %s", ifname, e)
logger.exception("Recv error on %s", iface)
continue
# Parse into scapy Packet (lazy parse)
# parse with scapy
try:
pkt = Ether(raw)
# attach interface metadata so parse_packet can determine ingress
pkt.sniffed_on = ifname
# call your existing parser
pkt.sniffed_on = iface
parse_packet(pkt, bridge)
except Exception as e:
logger.exception("Failed to parse/process packet from %s: %s", ifname, e)
except Exception:
logger.exception("Failed to parse/process packet from %s", iface)
continue
# cleanup
logger.info("AF_PACKET reader stopping; closing sockets")
finally:
logger.info("AF_PACKET reader stopping; cleaning up sockets")
# unregister and close
try:
for key in list(af_selector.get_map().values()):
for key in list(sel.get_map().values()):
try:
af_selector.unregister(key.fileobj)
sel.unregister(key.fileobj)
except Exception:
pass
except Exception:
pass
for ifname in list(af_sockets.keys()):
_close_socket(ifname)
for iface in list(af_sockets.keys()):
_close_socket(iface)
try:
af_selector.close()
sel.close()
except Exception:
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:
"""
Start the optimized AF_PACKET sniffer for the provided bridge.
This reads bridge ports once (snapshot) and starts a single reader thread.
Start the sniffer: snapshot ports once, spin up 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():
logger.info("AF_PACKET sniffer already running")
logger.info("Sniffer already running")
return
# Warm and freeze the bridge ports snapshot. We read sysfs once here.
ports = get_bridge_ports_once(bridge)
fixed_bridge_ports[bridge] = ports
current_bridge = bridge
@@ -520,34 +462,41 @@ def start_afpacket_sniffer(bridge: str) -> 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:
return
if af_stop_event:
af_stop_event.set()
af_thread.join(timeout=2)
af_thread = None
af_stop_event = None
# clear fixed snapshot(s)
if current_bridge:
fixed_bridge_ports.pop(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")
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.
Return simple status per managed interface.
"""
running = af_thread.is_alive() if af_thread else False
out: Dict[str, Dict[str, object]] = {}
for iface in list(af_sockets.keys()):
out[iface] = {
"running": af_thread.is_alive() if af_thread else False,
"running": running,
"exists": check_interface_exists(iface),
"up": check_interface_up(iface),
}