fix: enhance network sniffer with status endpoint and improve thread management
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s
This commit is contained in:
@@ -4,18 +4,22 @@ import asyncpg
|
||||
import threading
|
||||
import logging
|
||||
import os
|
||||
from typing import List
|
||||
from typing import List, Dict
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("sniffer")
|
||||
|
||||
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
|
||||
|
||||
sniffer_threads = {}
|
||||
# Active sniffer threads and stop flags
|
||||
sniffer_threads: Dict[str, threading.Thread] = {}
|
||||
thread_stop_flags: Dict[str, threading.Event] = {}
|
||||
|
||||
# Cache for bridge → ports
|
||||
bridge_ports_cache = {}
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# ASYNC LOOP FOR DATABASE INSERTS
|
||||
# ASYNC LOOP FOR DB INSERTS
|
||||
# -------------------------------------------------------------------
|
||||
async_loop = asyncio.new_event_loop()
|
||||
|
||||
@@ -27,19 +31,15 @@ threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# BRIDGE PORT DETECTION
|
||||
# BRIDGE PORT HANDLING
|
||||
# -------------------------------------------------------------------
|
||||
def get_bridge_ports(bridge: str) -> List[str]:
|
||||
"""
|
||||
Return all interfaces that belong to this Linux bridge.
|
||||
Cached for performance.
|
||||
"""
|
||||
if bridge in bridge_ports_cache:
|
||||
return bridge_ports_cache[bridge]
|
||||
|
||||
base = f"/sys/class/net/{bridge}/brif/"
|
||||
if not os.path.isdir(base):
|
||||
logger.error(f"Bridge {bridge} does not exist or brif folder missing")
|
||||
logger.error(f"Bridge {bridge} does not exist")
|
||||
return []
|
||||
|
||||
ports = os.listdir(base)
|
||||
@@ -49,15 +49,9 @@ def get_bridge_ports(bridge: str) -> List[str]:
|
||||
|
||||
|
||||
def determine_direction(pkt_iface: str, bridge: str):
|
||||
"""
|
||||
Determine ingress + egress:
|
||||
- ingress = interface where packet was received
|
||||
- egress = other bridge ports (the packet will be forwarded to)
|
||||
"""
|
||||
ports = get_bridge_ports(bridge)
|
||||
ingress = pkt_iface
|
||||
egress = [p for p in ports if p != pkt_iface]
|
||||
|
||||
return ingress, egress
|
||||
|
||||
|
||||
@@ -67,6 +61,7 @@ def determine_direction(pkt_iface: str, bridge: str):
|
||||
async def db_insert_packet(pkt_info: dict):
|
||||
try:
|
||||
conn = await asyncpg.connect(DB_DSN)
|
||||
|
||||
await conn.execute("""
|
||||
INSERT INTO packets(
|
||||
iface,
|
||||
@@ -80,10 +75,11 @@ async def db_insert_packet(pkt_info: dict):
|
||||
length,
|
||||
ebpf_verdict,
|
||||
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)
|
||||
""",
|
||||
pkt_info["ingress"], # interface where packet arrived
|
||||
"unknown", # direction (can update later)
|
||||
pkt_info["ingress"],
|
||||
"unknown",
|
||||
pkt_info["src_mac"],
|
||||
pkt_info["dst_mac"],
|
||||
pkt_info["eth_type"],
|
||||
@@ -91,15 +87,12 @@ async def db_insert_packet(pkt_info: dict):
|
||||
pkt_info["dst_ip"],
|
||||
pkt_info["protocol"],
|
||||
pkt_info["length"],
|
||||
str(pkt_info.get("egress")), # optional egress info as text
|
||||
str(pkt_info["egress"]),
|
||||
pkt_info["raw"]
|
||||
)
|
||||
|
||||
logger.info(f"DB insert OK: ingress={pkt_info['ingress']} src={pkt_info['src_ip']} dst={pkt_info['dst_ip']}")
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"DB insert failed: {e}")
|
||||
|
||||
finally:
|
||||
if 'conn' in locals():
|
||||
await conn.close()
|
||||
@@ -110,9 +103,7 @@ async def db_insert_packet(pkt_info: dict):
|
||||
# -------------------------------------------------------------------
|
||||
def handle_packet(pkt, bridge: str):
|
||||
pkt_iface = getattr(pkt, "sniffed_on", None)
|
||||
|
||||
if not pkt_iface:
|
||||
logger.warning("Packet missing sniffed_on metadata")
|
||||
return
|
||||
|
||||
ingress, egress = determine_direction(pkt_iface, bridge)
|
||||
@@ -139,25 +130,65 @@ def handle_packet(pkt, bridge: str):
|
||||
def start_sniffer_thread(bridge: str):
|
||||
ports = get_bridge_ports(bridge)
|
||||
|
||||
def sniff_blocking(iface):
|
||||
logger.info(f"Starting sniff on {iface}")
|
||||
sniff(iface=iface, prn=lambda x: handle_packet(x, bridge), store=False)
|
||||
|
||||
for iface in ports:
|
||||
if iface not in sniffer_threads:
|
||||
thread = threading.Thread(target=sniff_blocking, args=(iface,), daemon=True)
|
||||
thread.start()
|
||||
sniffer_threads[iface] = thread
|
||||
if iface in sniffer_threads:
|
||||
continue
|
||||
|
||||
stop_event = threading.Event()
|
||||
thread_stop_flags[iface] = stop_event
|
||||
|
||||
def sniff_blocking(ifname=iface):
|
||||
logger.info(f"Sniffer STARTED on {ifname}")
|
||||
|
||||
sniff(
|
||||
iface=ifname,
|
||||
prn=lambda pkt: handle_packet(pkt, bridge),
|
||||
store=False,
|
||||
stop_filter=lambda _: stop_event.is_set(),
|
||||
timeout=1 # ensure periodic stop checks
|
||||
)
|
||||
|
||||
logger.info(f"Sniffer STOPPED on {ifname}")
|
||||
|
||||
thread = threading.Thread(target=sniff_blocking, daemon=True)
|
||||
sniffer_threads[iface] = thread
|
||||
thread.start()
|
||||
|
||||
return sniffer_threads
|
||||
|
||||
|
||||
async def start_sniffing(bridge: str):
|
||||
get_bridge_ports(bridge)
|
||||
start_sniffer_thread(bridge)
|
||||
return True
|
||||
|
||||
|
||||
async def stop_sniffing():
|
||||
# Scapy can't stop sniff(), so we drop the thread references
|
||||
logger.info("Stopping all sniffers…")
|
||||
|
||||
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()
|
||||
|
||||
logger.info("All sniffers stopped.")
|
||||
return True
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# STATUS HELPER (for API endpoint)
|
||||
# -------------------------------------------------------------------
|
||||
def get_sniffer_status():
|
||||
"""
|
||||
Return a dict of interface → running/stopped status.
|
||||
Useful for an API /status endpoint.
|
||||
"""
|
||||
status = {}
|
||||
|
||||
for iface, t in sniffer_threads.items():
|
||||
status[iface] = "running" if t.is_alive() else "stopped"
|
||||
|
||||
return status
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastapi import APIRouter
|
||||
from typing import List
|
||||
from src.network_sniffer import start_sniffing, stop_sniffing
|
||||
from src.network_sniffer import get_sniffer_status, start_sniffing, stop_sniffing
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -22,3 +22,7 @@ async def start_sniffer(bridge: str):
|
||||
async def stop_sniffer_api():
|
||||
await stop_sniffing()
|
||||
return {"status": "stopped"}
|
||||
|
||||
@router.get("/status")
|
||||
async def status():
|
||||
return get_sniffer_status()
|
||||
Reference in New Issue
Block a user