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 threading
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import List
|
from typing import List, Dict
|
||||||
|
|
||||||
logging.basicConfig(level=logging.INFO)
|
logging.basicConfig(level=logging.INFO)
|
||||||
logger = logging.getLogger("sniffer")
|
logger = logging.getLogger("sniffer")
|
||||||
|
|
||||||
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
|
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 = {}
|
bridge_ports_cache = {}
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
# ASYNC LOOP FOR DATABASE INSERTS
|
# ASYNC LOOP FOR DB INSERTS
|
||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
async_loop = asyncio.new_event_loop()
|
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]:
|
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:
|
if bridge in bridge_ports_cache:
|
||||||
return bridge_ports_cache[bridge]
|
return bridge_ports_cache[bridge]
|
||||||
|
|
||||||
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 or brif folder missing")
|
logger.error(f"Bridge {bridge} does not exist")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
ports = os.listdir(base)
|
ports = os.listdir(base)
|
||||||
@@ -49,15 +49,9 @@ def get_bridge_ports(bridge: str) -> List[str]:
|
|||||||
|
|
||||||
|
|
||||||
def determine_direction(pkt_iface: str, bridge: 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)
|
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
|
||||||
|
|
||||||
|
|
||||||
@@ -67,6 +61,7 @@ def determine_direction(pkt_iface: str, bridge: str):
|
|||||||
async def db_insert_packet(pkt_info: dict):
|
async def db_insert_packet(pkt_info: dict):
|
||||||
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,
|
||||||
@@ -80,10 +75,11 @@ async def db_insert_packet(pkt_info: dict):
|
|||||||
length,
|
length,
|
||||||
ebpf_verdict,
|
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)
|
||||||
""",
|
""",
|
||||||
pkt_info["ingress"], # interface where packet arrived
|
pkt_info["ingress"],
|
||||||
"unknown", # direction (can update later)
|
"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"],
|
||||||
@@ -91,15 +87,12 @@ async def db_insert_packet(pkt_info: dict):
|
|||||||
pkt_info["dst_ip"],
|
pkt_info["dst_ip"],
|
||||||
pkt_info["protocol"],
|
pkt_info["protocol"],
|
||||||
pkt_info["length"],
|
pkt_info["length"],
|
||||||
str(pkt_info.get("egress")), # optional egress info as text
|
str(pkt_info["egress"]),
|
||||||
pkt_info["raw"]
|
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:
|
except Exception as e:
|
||||||
logger.exception(f"DB insert failed: {e}")
|
logger.exception(f"DB insert failed: {e}")
|
||||||
|
|
||||||
finally:
|
finally:
|
||||||
if 'conn' in locals():
|
if 'conn' in locals():
|
||||||
await conn.close()
|
await conn.close()
|
||||||
@@ -110,9 +103,7 @@ async def db_insert_packet(pkt_info: dict):
|
|||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
def handle_packet(pkt, bridge: str):
|
def handle_packet(pkt, bridge: str):
|
||||||
pkt_iface = getattr(pkt, "sniffed_on", None)
|
pkt_iface = getattr(pkt, "sniffed_on", None)
|
||||||
|
|
||||||
if not pkt_iface:
|
if not pkt_iface:
|
||||||
logger.warning("Packet missing sniffed_on metadata")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
ingress, egress = determine_direction(pkt_iface, bridge)
|
ingress, egress = determine_direction(pkt_iface, bridge)
|
||||||
@@ -139,25 +130,65 @@ def handle_packet(pkt, bridge: str):
|
|||||||
def start_sniffer_thread(bridge: str):
|
def start_sniffer_thread(bridge: str):
|
||||||
ports = get_bridge_ports(bridge)
|
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:
|
for iface in ports:
|
||||||
if iface not in sniffer_threads:
|
if iface in sniffer_threads:
|
||||||
thread = threading.Thread(target=sniff_blocking, args=(iface,), daemon=True)
|
continue
|
||||||
thread.start()
|
|
||||||
sniffer_threads[iface] = thread
|
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
|
return sniffer_threads
|
||||||
|
|
||||||
|
|
||||||
async def start_sniffing(bridge: str):
|
async def start_sniffing(bridge: str):
|
||||||
get_bridge_ports(bridge)
|
|
||||||
start_sniffer_thread(bridge)
|
start_sniffer_thread(bridge)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
async def stop_sniffing():
|
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()
|
sniffer_threads.clear()
|
||||||
|
thread_stop_flags.clear()
|
||||||
|
|
||||||
|
logger.info("All sniffers stopped.")
|
||||||
return True
|
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 fastapi import APIRouter
|
||||||
from typing import List
|
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()
|
router = APIRouter()
|
||||||
|
|
||||||
@@ -22,3 +22,7 @@ async def start_sniffer(bridge: str):
|
|||||||
async def stop_sniffer_api():
|
async def stop_sniffer_api():
|
||||||
await stop_sniffing()
|
await stop_sniffing()
|
||||||
return {"status": "stopped"}
|
return {"status": "stopped"}
|
||||||
|
|
||||||
|
@router.get("/status")
|
||||||
|
async def status():
|
||||||
|
return get_sniffer_status()
|
||||||
Reference in New Issue
Block a user