fix: enhance database schema for packet metadata and improve data types
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s

This commit is contained in:
2025-11-27 18:36:26 +01:00
parent ef946bec3c
commit a2603d4690
2 changed files with 185 additions and 112 deletions

View File

@@ -1,36 +1,55 @@
import socket
import struct
import threading
import asyncio import asyncio
import asyncpg
import os
import logging import logging
from typing import Dict, List import threading
import os
from typing import List, Dict
import asyncpg
from scapy.all import (
Ether,
ARP,
IP,
IPv6,
TCP,
UDP,
ICMP,
ICMPv6Unknown,
Dot1Q,
Raw,
sniff,
conf,
)
logging.basicConfig(level=logging.INFO) logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("afpacket_sniffer") logger = logging.getLogger("af_packet_sniffer")
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
# Active sniffer threads and stop flags
sniffer_threads: Dict[str, threading.Thread] = {} sniffer_threads: Dict[str, threading.Thread] = {}
thread_stop_flags: Dict[str, threading.Event] = {} thread_stop_flags: Dict[str, threading.Event] = {}
# Cache for bridge -> ports
bridge_ports_cache: Dict[str, List[str]] = {} bridge_ports_cache: Dict[str, List[str]] = {}
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# ASYNC LOOP FOR DB INSERTS # Async loop for DB inserts
# ------------------------------------------------------------------- # -------------------------------------------------------------------
async_loop = asyncio.new_event_loop() async_loop = asyncio.new_event_loop()
threading.Thread(target=lambda: async_loop.run_forever(), daemon=True).start()
def start_async_loop(loop):
asyncio.set_event_loop(loop)
loop.run_forever()
threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start()
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# BRIDGE PORT HANDLING # Interface / bridge checks
# ------------------------------------------------------------------- # -------------------------------------------------------------------
def check_interface_exists(iface: str) -> bool: def check_interface_exists(iface: str) -> bool:
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:
try: try:
with open(f"/sys/class/net/{iface}/operstate", "r") as f: with open(f"/sys/class/net/{iface}/operstate", "r") as f:
@@ -38,6 +57,7 @@ def check_interface_up(iface: str) -> bool:
except FileNotFoundError: except FileNotFoundError:
return False return False
def get_bridge_ports(bridge: str) -> List[str]: def get_bridge_ports(bridge: str) -> List[str]:
if bridge in bridge_ports_cache: if bridge in bridge_ports_cache:
return bridge_ports_cache[bridge] return bridge_ports_cache[bridge]
@@ -47,59 +67,57 @@ def get_bridge_ports(bridge: str) -> List[str]:
logger.error(f"Bridge '{bridge}' does not exist") logger.error(f"Bridge '{bridge}' does not exist")
return [] return []
ports = [] ports = [p for p in os.listdir(base) if check_interface_exists(p)]
try:
for p in os.listdir(base):
if check_interface_exists(p):
ports.append(p)
else:
logger.warning(f"Port '{p}' listed in bridge but does not exist")
except PermissionError:
logger.error(f"No permission to read bridge '{bridge}' ports")
bridge_ports_cache[bridge] = ports bridge_ports_cache[bridge] = ports
logger.info(f"Bridge {bridge} ports: {ports}") logger.info(f"Bridge {bridge} ports: {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) ports = 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
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# DATABASE INSERTION # Database insertion
# ------------------------------------------------------------------- # -------------------------------------------------------------------
async def db_insert_packet(pkt_info: dict): async def db_insert_packet(pkt_info: dict):
conn = None conn = None
try: try:
conn = await asyncpg.connect(DB_DSN) conn = await asyncpg.connect(DB_DSN)
await conn.execute(""" await conn.execute(
"""
INSERT INTO packets( INSERT INTO packets(
iface, iface,
direction, direction,
src_mac, src_mac,
dst_mac, dst_mac,
eth_type, eth_type,
vlan_id,
src_ip, src_ip,
dst_ip, dst_ip,
protocol, ip_proto,
src_port,
dst_port,
length, length,
ebpf_verdict,
raw raw
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) ) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13)
""", """,
pkt_info["iface"], pkt_info["ingress"],
"unknown", "unknown",
pkt_info["src_mac"], pkt_info["src_mac"],
pkt_info["dst_mac"], pkt_info["dst_mac"],
pkt_info["eth_type"], pkt_info["eth_type"],
pkt_info["src_ip"], pkt_info.get("vlan_id"),
pkt_info["dst_ip"], pkt_info.get("src_ip"),
pkt_info["protocol"], pkt_info.get("dst_ip"),
pkt_info.get("protocol_name"),
pkt_info.get("src_port"),
pkt_info.get("dst_port"),
pkt_info["length"], pkt_info["length"],
str(pkt_info["egress"]), pkt_info["raw"],
pkt_info["raw"]
) )
except Exception as e: except Exception as e:
logger.exception(f"DB insert failed: {e}") logger.exception(f"DB insert failed: {e}")
@@ -107,78 +125,123 @@ async def db_insert_packet(pkt_info: dict):
if conn: if conn:
await conn.close() await conn.close()
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# PACKET HANDLER # Packet parsing: AF_PACKET / full stack
# ------------------------------------------------------------------- # -------------------------------------------------------------------
def handle_packet(pkt_bytes: bytes, iface: str, bridge: str): def parse_packet(pkt, bridge: str):
# Ethernet header pkt_iface = getattr(pkt, "sniffed_on", None)
if len(pkt_bytes) < 14: if not pkt_iface:
return return
eth_header = pkt_bytes[:14]
dst_mac, src_mac, eth_type = struct.unpack("!6s6sH", eth_header)
dst_mac = ':'.join('%02x' % b for b in dst_mac)
src_mac = ':'.join('%02x' % b for b in src_mac)
eth_type = socket.ntohs(eth_type)
# IP header ingress, egress = determine_direction(pkt_iface, bridge)
src_ip = dst_ip = None
protocol = None
if eth_type == 0x0800 and len(pkt_bytes) >= 34:
ip_header = pkt_bytes[14:34]
iph = struct.unpack('!BBHHHBBH4s4s', ip_header)
src_ip = socket.inet_ntoa(iph[8])
dst_ip = socket.inet_ntoa(iph[9])
protocol = iph[6]
ingress, egress = determine_direction(iface, bridge)
pkt_info = { pkt_info = {
"iface": iface, "ingress": ingress,
"egress": egress, "egress": egress,
"src_mac": src_mac, "length": len(pkt),
"dst_mac": dst_mac, "raw": bytes(pkt),
"eth_type": hex(eth_type), "src_mac": None,
"src_ip": src_ip, "dst_mac": None,
"dst_ip": dst_ip, "eth_type": None,
"protocol": protocol, "vlan_id": None,
"length": len(pkt_bytes), "src_ip": None,
"raw": pkt_bytes "dst_ip": None,
"protocol_name": None,
"src_port": None,
"dst_port": None,
} }
# Ethernet
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
if Dot1Q in pkt:
pkt_info["vlan_id"] = pkt[Dot1Q].vlan
# ARP
if ARP in pkt:
pkt_info["protocol_name"] = "ARP"
pkt_info["src_ip"] = pkt[ARP].psrc
pkt_info["dst_ip"] = pkt[ARP].pdst
pkt_info["src_port"] = None
pkt_info["dst_port"] = None
# IPv4
if IP in pkt:
pkt_info["src_ip"] = pkt[IP].src
pkt_info["dst_ip"] = pkt[IP].dst
proto = pkt[IP].proto
if proto == 6 and TCP in pkt:
pkt_info["protocol_name"] = "TCP"
pkt_info["src_port"] = pkt[TCP].sport
pkt_info["dst_port"] = pkt[TCP].dport
elif proto == 17 and UDP in pkt:
pkt_info["protocol_name"] = "UDP"
pkt_info["src_port"] = pkt[UDP].sport
pkt_info["dst_port"] = pkt[UDP].dport
elif proto == 1 and ICMP in pkt:
pkt_info["protocol_name"] = "ICMP"
else:
pkt_info["protocol_name"] = f"IP_PROTO_{proto}"
# IPv6
if IPv6 in pkt:
pkt_info["src_ip"] = pkt[IPv6].src
pkt_info["dst_ip"] = pkt[IPv6].dst
nh = pkt[IPv6].nh
if nh == 6 and TCP in pkt:
pkt_info["protocol_name"] = "TCP"
pkt_info["src_port"] = pkt[TCP].sport
pkt_info["dst_port"] = pkt[TCP].dport
elif nh == 17 and UDP in pkt:
pkt_info["protocol_name"] = "UDP"
pkt_info["src_port"] = pkt[UDP].sport
pkt_info["dst_port"] = pkt[UDP].dport
elif ICMPv6Unknown in pkt:
pkt_info["protocol_name"] = "ICMPv6"
else:
pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}"
# Raw / unknown
if Raw in pkt and not pkt_info["protocol_name"]:
pkt_info["protocol_name"] = "RAW"
asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop) asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop)
# -------------------------------------------------------------------
# SNIFFER LOOP
# -------------------------------------------------------------------
def sniffer_loop(iface: str, stop_event: threading.Event, bridge: str):
logger.info(f"Sniffer STARTED on {iface}")
if not check_interface_exists(iface): # -------------------------------------------------------------------
logger.error(f"Interface {iface} does not exist. Exiting sniffer.") # Sniffer thread using AF_PACKET
return # -------------------------------------------------------------------
if not check_interface_up(iface): def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
logger.error(f"Interface {iface} is DOWN. Exiting sniffer.") logger.info(f"Sniffer STARTED on {ifname}")
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 return
try: conf.L2socket = conf.L2socket # enforce AF_PACKET usage in scapy
s = socket.socket(socket.AF_PACKET, socket.SOCK_RAW, socket.ntohs(3))
s.bind((iface, 0))
except PermissionError:
logger.error(f"Permission denied on {iface}, need root")
return
while not stop_event.is_set(): while not stop_event.is_set():
try: try:
pkt, _ = s.recvfrom(65536) sniff(
handle_packet(pkt, iface, bridge) iface=ifname,
prn=lambda pkt: parse_packet(pkt, bridge),
store=False,
timeout=0.5, # fast stop checks
)
except Exception as e: except Exception as e:
logger.exception(f"Error in sniffer loop on {iface}: {e}") logger.exception(f"Sniffer error on {ifname}: {e}")
break
logger.info(f"Sniffer STOPPED on {ifname}")
s.close()
logger.info(f"Sniffer STOPPED on {iface}")
# ------------------------------------------------------------------- # -------------------------------------------------------------------
# START / STOP METHODS # Start / Stop / Status API
# ------------------------------------------------------------------- # -------------------------------------------------------------------
def start_sniffer_thread(bridge: str): def start_sniffer_thread(bridge: str):
ports = get_bridge_ports(bridge) ports = get_bridge_ports(bridge)
@@ -189,35 +252,39 @@ def start_sniffer_thread(bridge: str):
for iface in ports: for iface in ports:
if iface in sniffer_threads: if iface in sniffer_threads:
continue continue
stop_event = threading.Event() stop_event = threading.Event()
thread_stop_flags[iface] = stop_event thread_stop_flags[iface] = stop_event
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True) thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True)
sniffer_threads[iface] = thread sniffer_threads[iface] = thread
thread.start() thread.start()
return sniffer_threads return sniffer_threads
async def start_sniffing(bridge: str): async def start_sniffing(bridge: str):
start_sniffer_thread(bridge) start_sniffer_thread(bridge)
return True return True
async def stop_sniffing(): async def stop_sniffing():
for stop_event in thread_stop_flags.values(): for iface, stop_event in thread_stop_flags.items():
stop_event.set() stop_event.set()
for thread in sniffer_threads.values():
for iface, thread in sniffer_threads.items():
thread.join(timeout=2) thread.join(timeout=2)
sniffer_threads.clear() sniffer_threads.clear()
thread_stop_flags.clear() thread_stop_flags.clear()
return True return True
# -------------------------------------------------------------------
# STATUS HELPER
# -------------------------------------------------------------------
def get_sniffer_status(): def get_sniffer_status():
out = {} out = {}
for iface, thread in sniffer_threads.items(): for iface, t in sniffer_threads.items():
out[iface] = { out[iface] = {
"running": thread.is_alive(), "running": t.is_alive(),
"exists": check_interface_exists(iface), "exists": check_interface_exists(iface),
"up": check_interface_up(iface) "up": check_interface_up(iface),
} }
return out return out

View File

@@ -51,30 +51,36 @@ CREATE TABLE IF NOT EXISTS packets (
id BIGSERIAL PRIMARY KEY, id BIGSERIAL PRIMARY KEY,
timestamp TIMESTAMPTZ DEFAULT NOW(), timestamp TIMESTAMPTZ DEFAULT NOW(),
-- interface metadata -- Interface metadata
iface VARCHAR(64), iface VARCHAR(64),
direction VARCHAR(16), direction VARCHAR(16),
-- Ethernet -- Ethernet
src_mac VARCHAR(32), src_mac MACADDR,
dst_mac VARCHAR(32), dst_mac MACADDR,
eth_type INTEGER, eth_type VARCHAR(16),
-- IP -- VLAN
src_ip VARCHAR(64), vlan_id INTEGER,
dst_ip VARCHAR(64),
protocol INTEGER, -- IP layer
src_ip INET,
dst_ip INET,
ip_proto VARCHAR(32),
-- Transport layer
src_port INTEGER,
dst_port INTEGER,
-- Packet metadata
length INTEGER, length INTEGER,
-- eBPF data (reserved for later) -- Full packet dump
ebpf_verdict VARCHAR(32),
ebpf_chain VARCHAR(64),
-- full packet
raw BYTEA raw BYTEA
); );
EOF EOF
echo "[7] Grant privileges to user…" echo "[7] Grant privileges to user…"
sudo -u postgres psql -d $DB_NAME <<EOF sudo -u postgres psql -d $DB_NAME <<EOF
-- Grant all table privileges -- Grant all table privileges