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 threading
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
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
|
||||
|
||||
# ---------------------------
|
||||
# Configuration
|
||||
# ---------------------------
|
||||
logger = logging.getLogger("network_sniffer")
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger("sniffer")
|
||||
|
||||
DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db"
|
||||
DB_TABLE = "packets" # keep your current table name; change if needed
|
||||
|
||||
# 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)
|
||||
# ----- CONFIG -----
|
||||
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
|
||||
# ------------------
|
||||
|
||||
# Active sniffer threads and stop flags
|
||||
sniffer_threads: Dict[str, threading.Thread] = {}
|
||||
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]] = {}
|
||||
|
||||
# Queue for packet events (thread-safe)
|
||||
packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE)
|
||||
# Async loop + DB pool (run in background thread)
|
||||
_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)
|
||||
# -------------------------------------------------------------------
|
||||
def _start_async_loop(loop):
|
||||
# ------------------------
|
||||
# Async loop bootstrap
|
||||
# ------------------------
|
||||
def _start_async_loop(loop: asyncio.AbstractEventLoop):
|
||||
"""Run the given event loop forever (target for background thread)."""
|
||||
asyncio.set_event_loop(loop)
|
||||
_loop_started_event.set()
|
||||
loop.run_forever()
|
||||
|
||||
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# DB helper: create pool & consumer
|
||||
# -------------------------------------------------------------------
|
||||
async def _create_pool():
|
||||
global pg_pool
|
||||
if pg_pool is None:
|
||||
def ensure_async_loop():
|
||||
"""Create and start the dedicated asyncio loop once."""
|
||||
global _async_loop, _loop_thread
|
||||
if _async_loop is not None:
|
||||
return _async_loop
|
||||
|
||||
_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...")
|
||||
pg_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=8)
|
||||
return pg_pool
|
||||
_db_pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=6)
|
||||
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):
|
||||
"""
|
||||
Async consumer: drains packet_queue and does batched inserts.
|
||||
Runs inside async_loop.
|
||||
"""
|
||||
await _create_pool()
|
||||
logger.info("DB consumer started")
|
||||
stmt = f"""
|
||||
INSERT INTO {DB_TABLE}(
|
||||
iface,
|
||||
direction,
|
||||
src_mac,
|
||||
dst_mac,
|
||||
eth_type,
|
||||
src_ip,
|
||||
dst_ip,
|
||||
protocol,
|
||||
length,
|
||||
ebpf_verdict,
|
||||
raw
|
||||
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
|
||||
"""
|
||||
|
||||
buffer: List[tuple] = []
|
||||
last_flush = time.time()
|
||||
|
||||
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:
|
||||
while not stop_event.is_set():
|
||||
# collect up to BATCH_SIZE from the queue with small wait
|
||||
try:
|
||||
# block briefly for the first item
|
||||
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")
|
||||
return fut.result(timeout=10)
|
||||
except Exception as e:
|
||||
logger.exception("Failed to create DB pool: %s", e)
|
||||
raise
|
||||
|
||||
|
||||
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
|
||||
# -------------------------------------------------------------------
|
||||
# ------------------------
|
||||
# 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 Exception:
|
||||
except FileNotFoundError:
|
||||
return False
|
||||
|
||||
|
||||
def get_bridge_ports(bridge: str) -> List[str]:
|
||||
"""
|
||||
Return member interfaces of the bridge. Uses /sys/class/net/<bridge>/brif/.
|
||||
Returns [] if bridge missing or on error.
|
||||
"""
|
||||
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")
|
||||
logger.error("Bridge %s does not exist (or brif missing)", bridge)
|
||||
return []
|
||||
|
||||
try:
|
||||
ports = os.listdir(base)
|
||||
ports = [p for p in os.listdir(base) if os.path.isdir(f"/sys/class/net/{p}")]
|
||||
except PermissionError:
|
||||
logger.error(f"No permission to read bridge ports for '{bridge}'")
|
||||
logger.error("Permission denied reading bridge ports for %s", bridge)
|
||||
return []
|
||||
|
||||
ok_ports = []
|
||||
bridge_ports_cache[bridge] = ports
|
||||
# register iface->bridge mapping
|
||||
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")
|
||||
iface_to_bridge[p] = bridge
|
||||
logger.info("Bridge %s ports: %s", bridge, ports)
|
||||
return ports
|
||||
|
||||
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):
|
||||
"""
|
||||
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
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Producer (sniffer threads)
|
||||
# -------------------------------------------------------------------
|
||||
def _make_pkt_tuple(pkt, ingress_iface: str, egress_list):
|
||||
|
||||
# ------------------------
|
||||
# DB insert
|
||||
# ------------------------
|
||||
async def _db_insert_packet_coroutine(pkt_info: dict):
|
||||
"""
|
||||
Convert packet info to a tuple matching DB insert order.
|
||||
We store only a small raw snapshot to keep DB small.
|
||||
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,
|
||||
src_mac,
|
||||
dst_mac,
|
||||
eth_type,
|
||||
src_ip,
|
||||
dst_ip,
|
||||
ip_protocol,
|
||||
packet_len,
|
||||
raw_packet
|
||||
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10)
|
||||
""",
|
||||
pkt_info.get("ingress"),
|
||||
pkt_info.get("direction", "unknown"),
|
||||
pkt_info.get("src_mac"),
|
||||
pkt_info.get("dst_mac"),
|
||||
pkt_info.get("eth_type"),
|
||||
pkt_info.get("src_ip"),
|
||||
pkt_info.get("dst_ip"),
|
||||
pkt_info.get("protocol"),
|
||||
pkt_info.get("length"),
|
||||
pkt_info.get("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:
|
||||
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
|
||||
ensure_db_pool()
|
||||
except Exception:
|
||||
src_mac = dst_mac = eth_type = None
|
||||
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:
|
||||
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
|
||||
if IP in pkt:
|
||||
return pkt[IP].src, pkt[IP].dst, pkt[IP].proto
|
||||
except Exception:
|
||||
src_ip = dst_ip = protocol = None
|
||||
pass
|
||||
return None, None, 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
|
||||
)
|
||||
|
||||
def handle_packet(pkt, bridge: str):
|
||||
"""
|
||||
Callback running inside Scapy sniff thread.
|
||||
Places a lightweight tuple in packet_queue.
|
||||
Called in sniffing thread. Build pkt_info and schedule DB insert.
|
||||
"""
|
||||
pkt_iface = getattr(pkt, "sniffed_on", None)
|
||||
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
|
||||
|
||||
ingress, egress = determine_direction(pkt_iface, bridge)
|
||||
tup = _make_pkt_tuple(pkt, ingress, egress)
|
||||
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)")
|
||||
src_ip, dst_ip, proto = _safe_get_src_dst(pkt)
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Threaded sniff loop
|
||||
# -------------------------------------------------------------------
|
||||
pkt_info = {
|
||||
"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):
|
||||
logger.info(f"Sniffer STARTED on {ifname}")
|
||||
logger.info("Sniffer thread starting for %s (bridge=%s)", ifname, bridge)
|
||||
|
||||
if not check_interface_exists(ifname):
|
||||
logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}")
|
||||
return
|
||||
if not check_interface_up(ifname):
|
||||
logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}")
|
||||
logger.error("Interface %s does not exist; stopping sniffer thread", ifname)
|
||||
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():
|
||||
try:
|
||||
sniff(
|
||||
iface=ifname,
|
||||
prn=lambda pkt: handle_packet(pkt, bridge),
|
||||
store=False,
|
||||
timeout=1 # short timeout to check stop_event frequently
|
||||
timeout=1 # returns periodically so we can check stop_event
|
||||
)
|
||||
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
|
||||
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
|
||||
except Exception as e:
|
||||
logger.exception(f"Unexpected sniffer error on {ifname}: {e}")
|
||||
except Exception:
|
||||
logger.exception("Unexpected error in sniffer loop for %s", ifname)
|
||||
stop_event.wait(0.5)
|
||||
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)
|
||||
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 {}
|
||||
|
||||
started = {}
|
||||
for iface in ports:
|
||||
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
|
||||
|
||||
stop_event = threading.Event()
|
||||
thread_stop_flags[iface] = stop_event
|
||||
|
||||
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True)
|
||||
sniffer_threads[iface] = thread
|
||||
stop_evt = threading.Event()
|
||||
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_evt, bridge), daemon=True)
|
||||
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.
|
||||
Safe to call multiple times.
|
||||
Async entrypoint you can call from FastAPI startup/route.
|
||||
"""
|
||||
global _consumer_stop_event, _consumer_future_accessor
|
||||
logger.info(f"Request to start sniffing on 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)
|
||||
ensure_async_loop()
|
||||
ensure_db_pool()
|
||||
start_sniffer_thread_for_bridge(bridge)
|
||||
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
|
||||
for iface, ev in list(thread_stop_flags.items()):
|
||||
ev.set()
|
||||
async def stop_sniffing() -> bool:
|
||||
"""
|
||||
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()):
|
||||
logger.info("Joining thread for %s", iface)
|
||||
thread.join(timeout=2)
|
||||
|
||||
sniffer_threads.clear()
|
||||
thread_stop_flags.clear()
|
||||
|
||||
# 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")
|
||||
logger.info("All sniffer threads stopped")
|
||||
return True
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Status helper
|
||||
# -------------------------------------------------------------------
|
||||
def get_sniffer_status():
|
||||
|
||||
def get_sniffer_status() -> Dict[str, Dict[str, object]]:
|
||||
"""
|
||||
Return a dict of interface -> status information.
|
||||
Return status dict: iface -> { running: bool, exists: bool, up: bool, bridge: Optional[str] }
|
||||
"""
|
||||
out = {}
|
||||
for iface, t in sniffer_threads.items():
|
||||
out[iface] = {
|
||||
"running": t.is_alive(),
|
||||
status = {}
|
||||
# include known interfaces (from cache) and active threads
|
||||
known_ifaces = set(list(sniffer_threads.keys()) + list(iface_to_bridge.keys()))
|
||||
for iface in known_ifaces:
|
||||
thread = sniffer_threads.get(iface)
|
||||
status[iface] = {
|
||||
"running": bool(thread and thread.is_alive()),
|
||||
"exists": check_interface_exists(iface),
|
||||
"up": check_interface_up(iface),
|
||||
"bridge": iface_to_bridge.get(iface)
|
||||
}
|
||||
out["queue_size"] = packet_queue.qsize()
|
||||
return out
|
||||
return status
|
||||
|
||||
|
||||
# ------------------------
|
||||
# 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