fix: refactor network sniffer for improved async handling and logging
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:
@@ -1,426 +1,401 @@
|
|||||||
|
"""
|
||||||
|
network_sniffer.py
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
- import start_sniffing, stop_sniffing, get_sniffer_status from this module in your FastAPI routes.
|
||||||
|
- start_sniffing(bridge_name) will spawn per-interface sniffing threads for all bridge ports.
|
||||||
|
- stop_sniffing() will stop all running sniffer threads cleanly.
|
||||||
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import threading
|
import threading
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import time
|
|
||||||
from typing import List, Dict, Optional
|
from typing import List, Dict, Optional
|
||||||
from queue import Queue, Full, Empty
|
|
||||||
|
|
||||||
from scapy.all import sniff, Ether, IP # keep Scapy usage minimal
|
from scapy.all import sniff, Ether, IP # scapy must be installed
|
||||||
import asyncpg
|
import asyncpg
|
||||||
|
|
||||||
# ---------------------------
|
logger = logging.getLogger("network_sniffer")
|
||||||
# Configuration
|
|
||||||
# ---------------------------
|
|
||||||
logging.basicConfig(level=logging.INFO)
|
logging.basicConfig(level=logging.INFO)
|
||||||
logger = logging.getLogger("sniffer")
|
|
||||||
|
|
||||||
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
|
# ----- CONFIG -----
|
||||||
DB_TABLE = "packets" # keep your current table name; change if needed
|
DB_DSN = os.getenv("MITM_DB_DSN", "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db")
|
||||||
|
TABLE_NAME = os.getenv("MITM_TABLE", "packet_log") # change if you use a different table
|
||||||
# Producer/consumer settings
|
# ------------------
|
||||||
QUEUE_MAXSIZE = 20000 # max packets buffered in memory
|
|
||||||
BATCH_SIZE = 200 # how many rows to insert at once
|
|
||||||
FLUSH_INTERVAL = 0.5 # seconds max before flushing partial batch
|
|
||||||
SNAPSHOT_MAX_BYTES = 256 # how many bytes of the packet to store (truncate raw)
|
|
||||||
|
|
||||||
# Active sniffer threads and stop flags
|
# 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] = {}
|
||||||
|
# map which interface belongs to which bridge (reverse mapping)
|
||||||
|
iface_to_bridge: Dict[str, str] = {}
|
||||||
|
|
||||||
# Cache for bridge -> ports
|
# Bridge→ports cache
|
||||||
bridge_ports_cache: Dict[str, List[str]] = {}
|
bridge_ports_cache: Dict[str, List[str]] = {}
|
||||||
|
|
||||||
# Queue for packet events (thread-safe)
|
# Async loop + DB pool (run in background thread)
|
||||||
packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE)
|
_async_loop: Optional[asyncio.AbstractEventLoop] = None
|
||||||
|
_db_pool: Optional[asyncpg.pool.Pool] = None
|
||||||
|
_loop_thread: Optional[threading.Thread] = None
|
||||||
|
_loop_started_event = threading.Event()
|
||||||
|
|
||||||
# Async event loop + background consumer task
|
|
||||||
async_loop = asyncio.new_event_loop()
|
|
||||||
consumer_task: Optional[asyncio.Task] = None
|
|
||||||
pg_pool: Optional[asyncpg.Pool] = None
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
# ------------------------
|
||||||
# Async loop bootstrap (runs in background thread)
|
# Async loop bootstrap
|
||||||
# -------------------------------------------------------------------
|
# ------------------------
|
||||||
def _start_async_loop(loop):
|
def _start_async_loop(loop: asyncio.AbstractEventLoop):
|
||||||
|
"""Run the given event loop forever (target for background thread)."""
|
||||||
asyncio.set_event_loop(loop)
|
asyncio.set_event_loop(loop)
|
||||||
|
_loop_started_event.set()
|
||||||
loop.run_forever()
|
loop.run_forever()
|
||||||
|
|
||||||
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
def ensure_async_loop():
|
||||||
# DB helper: create pool & consumer
|
"""Create and start the dedicated asyncio loop once."""
|
||||||
# -------------------------------------------------------------------
|
global _async_loop, _loop_thread
|
||||||
async def _create_pool():
|
if _async_loop is not None:
|
||||||
global pg_pool
|
return _async_loop
|
||||||
if pg_pool is None:
|
|
||||||
|
_async_loop = asyncio.new_event_loop()
|
||||||
|
_loop_thread = threading.Thread(target=_start_async_loop, args=(_async_loop,), daemon=True)
|
||||||
|
_loop_thread.start()
|
||||||
|
# wait for loop to be set in thread
|
||||||
|
_loop_started_event.wait(timeout=5)
|
||||||
|
if not _loop_started_event.is_set():
|
||||||
|
raise RuntimeError("Failed to start async loop thread")
|
||||||
|
logger.info("Async loop started in background thread")
|
||||||
|
return _async_loop
|
||||||
|
|
||||||
|
|
||||||
|
async def _create_db_pool():
|
||||||
|
"""Coroutine that creates asyncpg pool. Run this in the background loop."""
|
||||||
|
global _db_pool
|
||||||
|
if _db_pool is None:
|
||||||
logger.info("Creating asyncpg pool...")
|
logger.info("Creating asyncpg pool...")
|
||||||
pg_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=8)
|
_db_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=6)
|
||||||
return pg_pool
|
logger.info("DB pool created")
|
||||||
|
return _db_pool
|
||||||
|
|
||||||
async def _close_pool():
|
|
||||||
global pg_pool
|
|
||||||
if pg_pool:
|
|
||||||
await pg_pool.close()
|
|
||||||
pg_pool = None
|
|
||||||
|
|
||||||
async def _db_consumer_loop(stop_event: asyncio.Event):
|
def ensure_db_pool():
|
||||||
|
"""Ensure pool exists by scheduling creation on background loop."""
|
||||||
|
ensure_async_loop()
|
||||||
|
# schedule coroutine and wait for it to complete
|
||||||
|
fut = asyncio.run_coroutine_threadsafe(_create_db_pool(), _async_loop)
|
||||||
|
try:
|
||||||
|
return fut.result(timeout=10)
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception("Failed to create DB pool: %s", e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
# ------------------------
|
||||||
|
# Bridge helpers
|
||||||
|
# ------------------------
|
||||||
|
def check_interface_exists(iface: str) -> bool:
|
||||||
|
return os.path.isdir(f"/sys/class/net/{iface}")
|
||||||
|
|
||||||
|
|
||||||
|
def check_interface_up(iface: str) -> bool:
|
||||||
|
try:
|
||||||
|
with open(f"/sys/class/net/{iface}/operstate", "r") as f:
|
||||||
|
return f.read().strip() == "up"
|
||||||
|
except FileNotFoundError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def get_bridge_ports(bridge: str) -> List[str]:
|
||||||
"""
|
"""
|
||||||
Async consumer: drains packet_queue and does batched inserts.
|
Return member interfaces of the bridge. Uses /sys/class/net/<bridge>/brif/.
|
||||||
Runs inside async_loop.
|
Returns [] if bridge missing or on error.
|
||||||
"""
|
"""
|
||||||
await _create_pool()
|
if bridge in bridge_ports_cache:
|
||||||
logger.info("DB consumer started")
|
return bridge_ports_cache[bridge]
|
||||||
stmt = f"""
|
|
||||||
INSERT INTO {DB_TABLE}(
|
base = f"/sys/class/net/{bridge}/brif/"
|
||||||
iface,
|
if not os.path.isdir(base):
|
||||||
|
logger.error("Bridge %s does not exist (or brif missing)", bridge)
|
||||||
|
return []
|
||||||
|
|
||||||
|
try:
|
||||||
|
ports = [p for p in os.listdir(base) if os.path.isdir(f"/sys/class/net/{p}")]
|
||||||
|
except PermissionError:
|
||||||
|
logger.error("Permission denied reading bridge ports for %s", bridge)
|
||||||
|
return []
|
||||||
|
|
||||||
|
bridge_ports_cache[bridge] = ports
|
||||||
|
# register iface->bridge mapping
|
||||||
|
for p in ports:
|
||||||
|
iface_to_bridge[p] = bridge
|
||||||
|
logger.info("Bridge %s ports: %s", bridge, ports)
|
||||||
|
return ports
|
||||||
|
|
||||||
|
|
||||||
|
def determine_direction(pkt_iface: str, bridge: str):
|
||||||
|
"""
|
||||||
|
ingress = interface where packet was captured
|
||||||
|
egress = list of other bridge ports (where it would be forwarded)
|
||||||
|
"""
|
||||||
|
ports = get_bridge_ports(bridge)
|
||||||
|
ingress = pkt_iface
|
||||||
|
egress = [p for p in ports if p != pkt_iface]
|
||||||
|
return ingress, egress
|
||||||
|
|
||||||
|
|
||||||
|
# ------------------------
|
||||||
|
# DB insert
|
||||||
|
# ------------------------
|
||||||
|
async def _db_insert_packet_coroutine(pkt_info: dict):
|
||||||
|
"""
|
||||||
|
Runs in event loop; uses pool to insert.
|
||||||
|
Matches the schema of packet_log with minimal fields:
|
||||||
|
interface, direction, src_mac, dst_mac, eth_type,
|
||||||
|
src_ip, dst_ip, ip_protocol, packet_len, raw_packet
|
||||||
|
Adjust columns if your table differs.
|
||||||
|
"""
|
||||||
|
pool = await _create_db_pool() # ensure pool
|
||||||
|
conn = None
|
||||||
|
try:
|
||||||
|
async with pool.acquire() as conn:
|
||||||
|
await conn.execute(f"""
|
||||||
|
INSERT INTO {TABLE_NAME}(
|
||||||
|
interface,
|
||||||
direction,
|
direction,
|
||||||
src_mac,
|
src_mac,
|
||||||
dst_mac,
|
dst_mac,
|
||||||
eth_type,
|
eth_type,
|
||||||
src_ip,
|
src_ip,
|
||||||
dst_ip,
|
dst_ip,
|
||||||
protocol,
|
ip_protocol,
|
||||||
length,
|
packet_len,
|
||||||
ebpf_verdict,
|
raw_packet
|
||||||
raw
|
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10)
|
||||||
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
|
""",
|
||||||
"""
|
pkt_info.get("ingress"),
|
||||||
|
pkt_info.get("direction", "unknown"),
|
||||||
buffer: List[tuple] = []
|
pkt_info.get("src_mac"),
|
||||||
last_flush = time.time()
|
pkt_info.get("dst_mac"),
|
||||||
|
pkt_info.get("eth_type"),
|
||||||
try:
|
pkt_info.get("src_ip"),
|
||||||
while not stop_event.is_set():
|
pkt_info.get("dst_ip"),
|
||||||
# collect up to BATCH_SIZE from the queue with small wait
|
pkt_info.get("protocol"),
|
||||||
try:
|
pkt_info.get("length"),
|
||||||
# block briefly for the first item
|
pkt_info.get("raw")
|
||||||
item = packet_queue.get(timeout=FLUSH_INTERVAL)
|
|
||||||
except Empty:
|
|
||||||
item = None
|
|
||||||
|
|
||||||
if item:
|
|
||||||
buffer.append(item)
|
|
||||||
# try to drain additional items without blocking
|
|
||||||
while len(buffer) < BATCH_SIZE:
|
|
||||||
try:
|
|
||||||
item = packet_queue.get_nowait()
|
|
||||||
buffer.append(item)
|
|
||||||
except Empty:
|
|
||||||
break
|
|
||||||
|
|
||||||
now = time.time()
|
|
||||||
# flush if we have enough buffered or timed out
|
|
||||||
if buffer and (len(buffer) >= BATCH_SIZE or (now - last_flush) >= FLUSH_INTERVAL):
|
|
||||||
try:
|
|
||||||
async with pg_pool.acquire() as conn:
|
|
||||||
async with conn.transaction():
|
|
||||||
# executemany pattern
|
|
||||||
await conn.executemany(stmt, buffer)
|
|
||||||
logger.debug(f"Inserted {len(buffer)} packets")
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"DB batch insert failed: {e}")
|
|
||||||
# in case of failure, drop or requeue? We'll drop to avoid blocking.
|
|
||||||
# Optionally write to disk or metrics.
|
|
||||||
finally:
|
|
||||||
buffer.clear()
|
|
||||||
last_flush = now
|
|
||||||
|
|
||||||
# When stop requested, flush remaining
|
|
||||||
if buffer:
|
|
||||||
try:
|
|
||||||
async with pg_pool.acquire() as conn:
|
|
||||||
async with conn.transaction():
|
|
||||||
await conn.executemany(stmt, buffer)
|
|
||||||
logger.info(f"Flushed final {len(buffer)} packets on shutdown")
|
|
||||||
except Exception:
|
|
||||||
logger.exception("Failed flushing final packet buffer on shutdown")
|
|
||||||
finally:
|
|
||||||
buffer.clear()
|
|
||||||
|
|
||||||
finally:
|
|
||||||
logger.info("DB consumer stopped")
|
|
||||||
|
|
||||||
|
|
||||||
def _start_consumer():
|
|
||||||
"""Schedules the DB consumer to run in the shared async_loop and returns the stop Event."""
|
|
||||||
stop_event = asyncio.Event()
|
|
||||||
# create consumer task in async_loop
|
|
||||||
def _start():
|
|
||||||
global consumer_task
|
|
||||||
consumer_task = asyncio.run_coroutine_threadsafe(_db_consumer_loop(stop_event), async_loop)
|
|
||||||
threading.Thread(target=_start, daemon=True).start()
|
|
||||||
return stop_event, lambda: consumer_task # returns stop_event and accessor to future
|
|
||||||
|
|
||||||
async def _stop_consumer(stop_event: asyncio.Event, consumer_future_accessor):
|
|
||||||
"""Signal stop_event and wait for consumer to finish."""
|
|
||||||
logger.info("Stopping DB consumer...")
|
|
||||||
stop_event.set()
|
|
||||||
# wait for the consumer coroutine to finish
|
|
||||||
future = consumer_future_accessor()
|
|
||||||
if future:
|
|
||||||
try:
|
|
||||||
future.result(timeout=5)
|
|
||||||
except Exception as e:
|
|
||||||
logger.debug(f"Consumer future termination: {e}")
|
|
||||||
await _close_pool()
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
# Utility: Bridge port handling & direction
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
def check_interface_exists(iface: str) -> bool:
|
|
||||||
return os.path.isdir(f"/sys/class/net/{iface}")
|
|
||||||
|
|
||||||
def check_interface_up(iface: str) -> bool:
|
|
||||||
try:
|
|
||||||
with open(f"/sys/class/net/{iface}/operstate", "r") as f:
|
|
||||||
return f.read().strip() == "up"
|
|
||||||
except Exception:
|
|
||||||
return False
|
|
||||||
|
|
||||||
def get_bridge_ports(bridge: str) -> List[str]:
|
|
||||||
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")
|
|
||||||
return []
|
|
||||||
|
|
||||||
try:
|
|
||||||
ports = os.listdir(base)
|
|
||||||
except PermissionError:
|
|
||||||
logger.error(f"No permission to read bridge ports for '{bridge}'")
|
|
||||||
return []
|
|
||||||
|
|
||||||
ok_ports = []
|
|
||||||
for p in ports:
|
|
||||||
if check_interface_exists(p):
|
|
||||||
ok_ports.append(p)
|
|
||||||
else:
|
|
||||||
logger.warning(f"Bridge member {p} listed but interface missing in /sys/class/net")
|
|
||||||
|
|
||||||
bridge_ports_cache[bridge] = ok_ports
|
|
||||||
logger.info(f"Bridge {bridge} ports: {ok_ports}")
|
|
||||||
return ok_ports
|
|
||||||
|
|
||||||
def determine_direction(pkt_iface: str, bridge: str):
|
|
||||||
ports = get_bridge_ports(bridge)
|
|
||||||
ingress = pkt_iface
|
|
||||||
egress = [p for p in ports if p != pkt_iface]
|
|
||||||
return ingress, egress
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
# Producer (sniffer threads)
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
def _make_pkt_tuple(pkt, ingress_iface: str, egress_list):
|
|
||||||
"""
|
|
||||||
Convert packet info to a tuple matching DB insert order.
|
|
||||||
We store only a small raw snapshot to keep DB small.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
src_mac = pkt[Ether].src if Ether in pkt else None
|
|
||||||
dst_mac = pkt[Ether].dst if Ether in pkt else None
|
|
||||||
eth_type = pkt[Ether].type if Ether in pkt else None
|
|
||||||
except Exception:
|
|
||||||
src_mac = dst_mac = eth_type = None
|
|
||||||
|
|
||||||
try:
|
|
||||||
src_ip = pkt[IP].src if IP in pkt else None
|
|
||||||
dst_ip = pkt[IP].dst if IP in pkt else None
|
|
||||||
protocol = pkt[IP].proto if IP in pkt else None
|
|
||||||
except Exception:
|
|
||||||
src_ip = dst_ip = protocol = None
|
|
||||||
|
|
||||||
raw = bytes(pkt)[:SNAPSHOT_MAX_BYTES] if pkt is not None else b""
|
|
||||||
|
|
||||||
# ebpf_verdict field currently used to store egress as text (you can change schema)
|
|
||||||
ebpf_verdict_text = str(egress_list) if egress_list else None
|
|
||||||
|
|
||||||
return (
|
|
||||||
ingress_iface,
|
|
||||||
"unknown", # direction
|
|
||||||
src_mac,
|
|
||||||
dst_mac,
|
|
||||||
str(eth_type) if eth_type is not None else None,
|
|
||||||
src_ip,
|
|
||||||
dst_ip,
|
|
||||||
str(protocol) if protocol is not None else None,
|
|
||||||
len(raw),
|
|
||||||
ebpf_verdict_text,
|
|
||||||
raw
|
|
||||||
)
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("DB insert failed")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def schedule_db_insert(pkt_info: dict):
|
||||||
|
"""
|
||||||
|
Thread-safe helper: schedule DB insert into background loop.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
ensure_db_pool()
|
||||||
|
except Exception:
|
||||||
|
logger.error("DB pool not available; dropping packet")
|
||||||
|
return
|
||||||
|
|
||||||
|
# schedule coroutine
|
||||||
|
loop = ensure_async_loop()
|
||||||
|
fut = asyncio.run_coroutine_threadsafe(_db_insert_packet_coroutine(pkt_info), loop)
|
||||||
|
# optional: attach callback to log failures asynchronously
|
||||||
|
def _on_done(f):
|
||||||
|
try:
|
||||||
|
f.result()
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Async DB insert failed")
|
||||||
|
fut.add_done_callback(_on_done)
|
||||||
|
|
||||||
|
|
||||||
|
# ------------------------
|
||||||
|
# Packet handling
|
||||||
|
# ------------------------
|
||||||
|
def _safe_get_src_dst(pkt):
|
||||||
|
"""Return (src_ip, dst_ip, proto) safely if IP present."""
|
||||||
|
try:
|
||||||
|
if IP in pkt:
|
||||||
|
return pkt[IP].src, pkt[IP].dst, pkt[IP].proto
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None, None, None
|
||||||
|
|
||||||
|
|
||||||
def handle_packet(pkt, bridge: str):
|
def handle_packet(pkt, bridge: str):
|
||||||
"""
|
"""
|
||||||
Callback running inside Scapy sniff thread.
|
Called in sniffing thread. Build pkt_info and schedule DB insert.
|
||||||
Places a lightweight tuple in packet_queue.
|
|
||||||
"""
|
"""
|
||||||
pkt_iface = getattr(pkt, "sniffed_on", None)
|
pkt_iface = getattr(pkt, "sniffed_on", None)
|
||||||
if not pkt_iface:
|
if not pkt_iface:
|
||||||
logger.debug("Packet without sniffed_on metadata, ignoring")
|
# scapy sometimes doesn't set sniffed_on; skip if unknown
|
||||||
|
logger.warning("Packet without sniffed_on - ignoring")
|
||||||
return
|
return
|
||||||
|
|
||||||
ingress, egress = determine_direction(pkt_iface, bridge)
|
ingress, egress = determine_direction(pkt_iface, bridge)
|
||||||
tup = _make_pkt_tuple(pkt, ingress, egress)
|
src_ip, dst_ip, proto = _safe_get_src_dst(pkt)
|
||||||
try:
|
|
||||||
packet_queue.put_nowait(tup)
|
|
||||||
except Full:
|
|
||||||
# queue full -> drop packet and log throttle event
|
|
||||||
logger.warning("Packet queue full, dropping packet (producer side)")
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
pkt_info = {
|
||||||
# Threaded sniff loop
|
"ingress": ingress,
|
||||||
# -------------------------------------------------------------------
|
"egress": egress,
|
||||||
|
"direction": "ingress", # we store 'ingress' here (packet was received on ingress)
|
||||||
|
"src_mac": pkt[Ether].src if Ether in pkt else None,
|
||||||
|
"dst_mac": pkt[Ether].dst if Ether in pkt else None,
|
||||||
|
"eth_type": pkt[Ether].type if Ether in pkt else None,
|
||||||
|
"src_ip": src_ip,
|
||||||
|
"dst_ip": dst_ip,
|
||||||
|
"protocol": proto,
|
||||||
|
"length": len(pkt),
|
||||||
|
# store raw as binary; use bytes(pkt)
|
||||||
|
"raw": bytes(pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
# schedule asynchronous DB insert
|
||||||
|
schedule_db_insert(pkt_info)
|
||||||
|
|
||||||
|
|
||||||
|
# ------------------------
|
||||||
|
# Sniffer thread loop
|
||||||
|
# ------------------------
|
||||||
def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
|
def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
|
||||||
logger.info(f"Sniffer STARTED on {ifname}")
|
logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge)
|
||||||
|
|
||||||
if not check_interface_exists(ifname):
|
if not check_interface_exists(ifname):
|
||||||
logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}")
|
logger.error("Interface %s does not exist; stopping sniffer thread", ifname)
|
||||||
return
|
|
||||||
if not check_interface_up(ifname):
|
|
||||||
logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
|
if not check_interface_up(ifname):
|
||||||
|
logger.warning("Interface %s is down; sniffing might still work but check interface state", ifname)
|
||||||
|
|
||||||
|
# loop with short timeout so stop_event is checked often
|
||||||
while not stop_event.is_set():
|
while not stop_event.is_set():
|
||||||
try:
|
try:
|
||||||
sniff(
|
sniff(
|
||||||
iface=ifname,
|
iface=ifname,
|
||||||
prn=lambda pkt: handle_packet(pkt, bridge),
|
prn=lambda pkt: handle_packet(pkt, bridge),
|
||||||
store=False,
|
store=False,
|
||||||
timeout=1 # short timeout to check stop_event frequently
|
timeout=1 # returns periodically so we can check stop_event
|
||||||
)
|
)
|
||||||
except PermissionError:
|
except PermissionError:
|
||||||
logger.error(f"Permission denied sniffing on {ifname}. Run as root.")
|
logger.exception("Permission denied sniffing on %s - run as root or give CAP_NET_RAW", ifname)
|
||||||
break
|
break
|
||||||
except OSError as e:
|
except OSError as e:
|
||||||
logger.error(f"Sniffer OSError on {ifname}: {e}")
|
logger.exception("OS error while sniffing on %s: %s", ifname, e)
|
||||||
|
# brief sleep to avoid tight loop on repeated errors
|
||||||
|
stop_event.wait(1)
|
||||||
break
|
break
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.exception(f"Unexpected sniffer error on {ifname}: {e}")
|
logger.exception("Unexpected error in sniffer loop for %s", ifname)
|
||||||
|
stop_event.wait(0.5)
|
||||||
break
|
break
|
||||||
|
|
||||||
logger.info(f"Sniffer STOPPED on {ifname}")
|
logger.info("Sniffer thread exiting for %s", ifname)
|
||||||
|
|
||||||
def start_sniffer_thread(bridge: str):
|
|
||||||
|
# ------------------------
|
||||||
|
# Control API (importable)
|
||||||
|
# ------------------------
|
||||||
|
def start_sniffer_thread_for_bridge(bridge: str) -> Dict[str, threading.Thread]:
|
||||||
|
"""
|
||||||
|
Start sniff threads for all ports of given bridge.
|
||||||
|
Returns mapping iface->thread for started threads (existing threads are left running).
|
||||||
|
"""
|
||||||
ports = get_bridge_ports(bridge)
|
ports = get_bridge_ports(bridge)
|
||||||
if not ports:
|
if not ports:
|
||||||
logger.error(f"No ports found for bridge {bridge}, not starting sniffers.")
|
logger.error("No ports for bridge %s - nothing to start", bridge)
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
started = {}
|
||||||
for iface in ports:
|
for iface in ports:
|
||||||
if iface in sniffer_threads and sniffer_threads[iface].is_alive():
|
if iface in sniffer_threads and sniffer_threads[iface].is_alive():
|
||||||
logger.info(f"Sniffer already running on {iface}")
|
logger.info("Sniffer already running on %s", iface)
|
||||||
|
started[iface] = sniffer_threads[iface]
|
||||||
continue
|
continue
|
||||||
|
|
||||||
stop_event = threading.Event()
|
stop_evt = threading.Event()
|
||||||
thread_stop_flags[iface] = stop_event
|
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_evt, bridge), daemon=True)
|
||||||
|
|
||||||
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True)
|
|
||||||
sniffer_threads[iface] = thread
|
|
||||||
thread.start()
|
thread.start()
|
||||||
|
sniffer_threads[iface] = thread
|
||||||
|
thread_stop_flags[iface] = stop_evt
|
||||||
|
started[iface] = thread
|
||||||
|
logger.info("Started sniffer on %s (bridge=%s)", iface, bridge)
|
||||||
|
|
||||||
return sniffer_threads
|
return started
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
# Public API for starting/stopping sniffing (to be called from FastAPI)
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
_consumer_stop_event: Optional[asyncio.Event] = None
|
|
||||||
_consumer_future_accessor = None
|
|
||||||
|
|
||||||
async def start_sniffing(bridge: str):
|
async def start_sniffing(bridge: str) -> bool:
|
||||||
"""
|
"""
|
||||||
Start sniffing on all bridge ports and start the DB consumer.
|
Async entrypoint you can call from FastAPI startup/route.
|
||||||
Safe to call multiple times.
|
|
||||||
"""
|
"""
|
||||||
global _consumer_stop_event, _consumer_future_accessor
|
ensure_async_loop()
|
||||||
logger.info(f"Request to start sniffing on bridge {bridge}")
|
ensure_db_pool()
|
||||||
|
start_sniffer_thread_for_bridge(bridge)
|
||||||
# start consumer if not running
|
|
||||||
if _consumer_stop_event is None:
|
|
||||||
# create asyncio.Event in async_loop
|
|
||||||
fut = asyncio.run_coroutine_threadsafe(asyncio.sleep(0), async_loop)
|
|
||||||
# schedule creation of event and consumer
|
|
||||||
_consumer_stop_event = asyncio.run_coroutine_threadsafe(asyncio.Event(), async_loop).result()
|
|
||||||
# start consumer in async_loop as a future
|
|
||||||
def schedule_consumer():
|
|
||||||
nonlocal _consumer_stop_event
|
|
||||||
# schedule consumer coroutine directly
|
|
||||||
future = asyncio.run_coroutine_threadsafe(_db_consumer_loop(_consumer_stop_event), async_loop)
|
|
||||||
return future
|
|
||||||
|
|
||||||
# store accessor to the Future to allow waiting for termination on stop
|
|
||||||
_consumer_future_accessor = schedule_consumer
|
|
||||||
|
|
||||||
# start sniffer threads
|
|
||||||
start_sniffer_thread(bridge)
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def stop_sniffing():
|
|
||||||
"""
|
|
||||||
Stop all sniffer threads and the DB consumer.
|
|
||||||
"""
|
|
||||||
global _consumer_stop_event, _consumer_future_accessor
|
|
||||||
logger.info("Stopping all sniffers and consumer")
|
|
||||||
|
|
||||||
# stop sniffer threads
|
async def stop_sniffing() -> bool:
|
||||||
for iface, ev in list(thread_stop_flags.items()):
|
"""
|
||||||
ev.set()
|
Stop all sniffers and wait briefly for threads to exit.
|
||||||
|
"""
|
||||||
|
logger.info("Stopping sniffers (all interfaces)...")
|
||||||
|
for iface, evt in list(thread_stop_flags.items()):
|
||||||
|
evt.set()
|
||||||
|
|
||||||
|
# join threads with timeout
|
||||||
for iface, thread in list(sniffer_threads.items()):
|
for iface, thread in list(sniffer_threads.items()):
|
||||||
|
logger.info("Joining thread for %s", iface)
|
||||||
thread.join(timeout=2)
|
thread.join(timeout=2)
|
||||||
|
|
||||||
sniffer_threads.clear()
|
sniffer_threads.clear()
|
||||||
thread_stop_flags.clear()
|
thread_stop_flags.clear()
|
||||||
|
logger.info("All sniffer threads stopped")
|
||||||
# stop consumer
|
|
||||||
if _consumer_stop_event is not None:
|
|
||||||
# signal consumer in async_loop
|
|
||||||
def set_stop():
|
|
||||||
_consumer_stop_event.set()
|
|
||||||
asyncio.run_coroutine_threadsafe(asyncio.to_thread(set_stop), async_loop).result(timeout=2)
|
|
||||||
|
|
||||||
# wait for consumer future
|
|
||||||
if _consumer_future_accessor:
|
|
||||||
fut = _consumer_future_accessor()
|
|
||||||
try:
|
|
||||||
fut.result(timeout=5)
|
|
||||||
except Exception as e:
|
|
||||||
logger.debug(f"Consumer future join error: {e}")
|
|
||||||
|
|
||||||
_consumer_stop_event = None
|
|
||||||
_consumer_future_accessor = None
|
|
||||||
|
|
||||||
# flush queue attempt (best-effort)
|
|
||||||
logger.info("Flushing packet queue (best-effort)")
|
|
||||||
# Not blocking: drop packets on stop
|
|
||||||
while not packet_queue.empty():
|
|
||||||
try:
|
|
||||||
packet_queue.get_nowait()
|
|
||||||
except Empty:
|
|
||||||
break
|
|
||||||
|
|
||||||
# close pool cleanly from async loop
|
|
||||||
try:
|
|
||||||
asyncio.run_coroutine_threadsafe(_close_pool(), async_loop).result(timeout=5)
|
|
||||||
except Exception:
|
|
||||||
logger.exception("Failed closing pg pool cleanly")
|
|
||||||
|
|
||||||
logger.info("All stopped")
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
# Status helper
|
def get_sniffer_status() -> Dict[str, Dict[str, object]]:
|
||||||
# -------------------------------------------------------------------
|
|
||||||
def get_sniffer_status():
|
|
||||||
"""
|
"""
|
||||||
Return a dict of interface -> status information.
|
Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] }
|
||||||
"""
|
"""
|
||||||
out = {}
|
status = {}
|
||||||
for iface, t in sniffer_threads.items():
|
# include known interfaces (from cache) and active threads
|
||||||
out[iface] = {
|
known_ifaces = set(list(sniffer_threads.keys()) + list(iface_to_bridge.keys()))
|
||||||
"running": t.is_alive(),
|
for iface in known_ifaces:
|
||||||
|
thread = sniffer_threads.get(iface)
|
||||||
|
status[iface] = {
|
||||||
|
"running": bool(thread and thread.is_alive()),
|
||||||
"exists": check_interface_exists(iface),
|
"exists": check_interface_exists(iface),
|
||||||
"up": check_interface_up(iface),
|
"up": check_interface_up(iface),
|
||||||
|
"bridge": iface_to_bridge.get(iface)
|
||||||
}
|
}
|
||||||
out["queue_size"] = packet_queue.qsize()
|
return status
|
||||||
return out
|
|
||||||
|
|
||||||
|
# ------------------------
|
||||||
|
# Optional convenience: stop all on process exit
|
||||||
|
# ------------------------
|
||||||
|
def _cleanup_on_exit():
|
||||||
|
try:
|
||||||
|
# stop sniffers
|
||||||
|
for evt in thread_stop_flags.values():
|
||||||
|
evt.set()
|
||||||
|
for t in sniffer_threads.values():
|
||||||
|
t.join(timeout=1)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
# shutdown db pool and loop
|
||||||
|
if _db_pool is not None:
|
||||||
|
try:
|
||||||
|
fut = asyncio.run_coroutine_threadsafe(_db_pool.close(), _async_loop)
|
||||||
|
fut.result(timeout=5)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if _async_loop is not None:
|
||||||
|
try:
|
||||||
|
_async_loop.call_soon_threadsafe(_async_loop.stop)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
# register cleanup (best-effort)
|
||||||
|
import atexit
|
||||||
|
atexit.register(_cleanup_on_exit)
|
||||||
|
|||||||
Reference in New Issue
Block a user