fix: refactor sniffer API routes for consistency and clarity
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,170 +1,284 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
from scapy.all import sniff, Ether, IP
|
|
||||||
import asyncpg
|
|
||||||
import threading
|
import threading
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import socket
|
import time
|
||||||
from typing import List, Dict
|
from typing import List, Dict, Optional
|
||||||
|
from queue import Queue, Full, Empty
|
||||||
|
|
||||||
|
from scapy.all import sniff, Ether, IP # keep Scapy usage minimal
|
||||||
|
import asyncpg
|
||||||
|
|
||||||
|
# ---------------------------
|
||||||
|
# Configuration
|
||||||
|
# ---------------------------
|
||||||
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"
|
||||||
|
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)
|
||||||
|
|
||||||
# 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] = {}
|
||||||
|
|
||||||
# Cache for bridge → ports
|
# Cache for bridge -> ports
|
||||||
bridge_ports_cache = {}
|
bridge_ports_cache: Dict[str, List[str]] = {}
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
# Queue for packet events (thread-safe)
|
||||||
# ASYNC LOOP FOR DB INSERTS
|
packet_queue: Queue = Queue(maxsize=QUEUE_MAXSIZE)
|
||||||
# -------------------------------------------------------------------
|
|
||||||
|
# Async event loop + background consumer task
|
||||||
async_loop = asyncio.new_event_loop()
|
async_loop = asyncio.new_event_loop()
|
||||||
|
consumer_task: Optional[asyncio.Task] = None
|
||||||
|
pg_pool: Optional[asyncpg.Pool] = None
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------
|
||||||
def start_async_loop(loop):
|
# Async loop bootstrap (runs in background thread)
|
||||||
|
# -------------------------------------------------------------------
|
||||||
|
def _start_async_loop(loop):
|
||||||
asyncio.set_event_loop(loop)
|
asyncio.set_event_loop(loop)
|
||||||
loop.run_forever()
|
loop.run_forever()
|
||||||
|
|
||||||
|
threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).start()
|
||||||
threading.Thread(target=start_async_loop, args=(async_loop,), daemon=True).start()
|
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
# BRIDGE PORT HANDLING
|
# DB helper: create pool & consumer
|
||||||
|
# -------------------------------------------------------------------
|
||||||
|
async def _create_pool():
|
||||||
|
global pg_pool
|
||||||
|
if pg_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
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
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")
|
||||||
|
|
||||||
|
|
||||||
|
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:
|
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:
|
||||||
return f.read().strip() == "up"
|
return f.read().strip() == "up"
|
||||||
except FileNotFoundError:
|
except Exception:
|
||||||
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]
|
||||||
|
|
||||||
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"[ERROR] Bridge '{bridge}' does not exist.")
|
logger.error(f"Bridge '{bridge}' does not exist")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
try:
|
try:
|
||||||
ports = os.listdir(base)
|
ports = os.listdir(base)
|
||||||
except PermissionError:
|
except PermissionError:
|
||||||
logger.error(f"[ERROR] No permissions to read bridge ports for '{bridge}'.")
|
logger.error(f"No permission to read bridge ports for '{bridge}'")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
ok_ports = []
|
ok_ports = []
|
||||||
for p in ports:
|
for p in ports:
|
||||||
if not check_interface_exists(p):
|
if check_interface_exists(p):
|
||||||
logger.warning(f"[WARN] Port '{p}' in bridge but does not exist in /sys/class/net")
|
ok_ports.append(p)
|
||||||
continue
|
else:
|
||||||
ok_ports.append(p)
|
logger.warning(f"Bridge member {p} listed but interface missing in /sys/class/net")
|
||||||
|
|
||||||
bridge_ports_cache[bridge] = ok_ports
|
bridge_ports_cache[bridge] = ok_ports
|
||||||
logger.info(f"[INFO] Bridge {bridge} ports: {ok_ports}")
|
logger.info(f"Bridge {bridge} ports: {ok_ports}")
|
||||||
return ok_ports
|
return ok_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
|
# Producer (sniffer threads)
|
||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
async def db_insert_packet(pkt_info: dict):
|
def _make_pkt_tuple(pkt, ingress_iface: str, egress_list):
|
||||||
conn = None
|
"""
|
||||||
|
Convert packet info to a tuple matching DB insert order.
|
||||||
|
We store only a small raw snapshot to keep DB small.
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
conn = await asyncpg.connect(DB_DSN)
|
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
|
||||||
|
|
||||||
await conn.execute("""
|
try:
|
||||||
INSERT INTO packets(
|
src_ip = pkt[IP].src if IP in pkt else None
|
||||||
iface,
|
dst_ip = pkt[IP].dst if IP in pkt else None
|
||||||
direction,
|
protocol = pkt[IP].proto if IP in pkt else None
|
||||||
src_mac,
|
except Exception:
|
||||||
dst_mac,
|
src_ip = dst_ip = protocol = None
|
||||||
eth_type,
|
|
||||||
src_ip,
|
|
||||||
dst_ip,
|
|
||||||
protocol,
|
|
||||||
length,
|
|
||||||
ebpf_verdict,
|
|
||||||
raw
|
|
||||||
)
|
|
||||||
VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)
|
|
||||||
""",
|
|
||||||
pkt_info["ingress"],
|
|
||||||
"unknown",
|
|
||||||
pkt_info["src_mac"],
|
|
||||||
pkt_info["dst_mac"],
|
|
||||||
pkt_info["eth_type"],
|
|
||||||
pkt_info["src_ip"],
|
|
||||||
pkt_info["dst_ip"],
|
|
||||||
pkt_info["protocol"],
|
|
||||||
pkt_info["length"],
|
|
||||||
str(pkt_info["egress"]),
|
|
||||||
pkt_info["raw"]
|
|
||||||
)
|
|
||||||
|
|
||||||
except (asyncpg.PostgresError, ConnectionError, OSError) as e:
|
raw = bytes(pkt)[:SNAPSHOT_MAX_BYTES] if pkt is not None else b""
|
||||||
logger.exception(f"DB insert failed: {e}")
|
|
||||||
|
|
||||||
finally:
|
# ebpf_verdict field currently used to store egress as text (you can change schema)
|
||||||
if conn:
|
ebpf_verdict_text = str(egress_list) if egress_list else None
|
||||||
await conn.close()
|
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
# PACKET HANDLER
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
def handle_packet(pkt, bridge: str):
|
def handle_packet(pkt, bridge: str):
|
||||||
|
"""
|
||||||
|
Callback running inside Scapy sniff thread.
|
||||||
|
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")
|
||||||
return
|
return
|
||||||
|
|
||||||
ingress, egress = determine_direction(pkt_iface, bridge)
|
ingress, egress = determine_direction(pkt_iface, bridge)
|
||||||
|
tup = _make_pkt_tuple(pkt, ingress, egress)
|
||||||
pkt_info = {
|
try:
|
||||||
"ingress": ingress,
|
packet_queue.put_nowait(tup)
|
||||||
"egress": egress,
|
except Full:
|
||||||
"src_mac": pkt[Ether].src if Ether in pkt else None,
|
# queue full -> drop packet and log throttle event
|
||||||
"dst_mac": pkt[Ether].dst if Ether in pkt else None,
|
logger.warning("Packet queue full, dropping packet (producer side)")
|
||||||
"eth_type": pkt[Ether].type if Ether in pkt else None,
|
|
||||||
"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,
|
|
||||||
"length": len(pkt),
|
|
||||||
"raw": pkt.json()
|
|
||||||
}
|
|
||||||
|
|
||||||
asyncio.run_coroutine_threadsafe(db_insert_packet(pkt_info), async_loop)
|
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
# CLEAN STOPPING SNIFFERS
|
# Threaded sniff 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(f"Sniffer STARTED on {ifname}")
|
||||||
|
|
||||||
if not check_interface_exists(ifname):
|
if not check_interface_exists(ifname):
|
||||||
logger.error(f"[ERROR] Interface {ifname} does not exist. Stopping sniffer.")
|
logger.error(f"Interface {ifname} does not exist. Exiting sniffer for {ifname}")
|
||||||
return
|
return
|
||||||
|
|
||||||
if not check_interface_up(ifname):
|
if not check_interface_up(ifname):
|
||||||
logger.error(f"[ERROR] Interface {ifname} is DOWN. Stopping sniffer.")
|
logger.error(f"Interface {ifname} is DOWN. Exiting sniffer for {ifname}")
|
||||||
return
|
return
|
||||||
|
|
||||||
while not stop_event.is_set():
|
while not stop_event.is_set():
|
||||||
@@ -173,13 +287,13 @@ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
|
|||||||
iface=ifname,
|
iface=ifname,
|
||||||
prn=lambda pkt: handle_packet(pkt, bridge),
|
prn=lambda pkt: handle_packet(pkt, bridge),
|
||||||
store=False,
|
store=False,
|
||||||
timeout=1 # periodic return so we can check stop_event
|
timeout=1 # short timeout to check stop_event frequently
|
||||||
)
|
)
|
||||||
except PermissionError:
|
except PermissionError:
|
||||||
logger.error(f"[ERROR] Permission denied sniffing on {ifname}. Run as root.")
|
logger.error(f"Permission denied sniffing on {ifname}. Run as root.")
|
||||||
break
|
break
|
||||||
except OSError as e:
|
except OSError as e:
|
||||||
logger.error(f"[ERROR] Sniffer error on {ifname}: {e}")
|
logger.error(f"Sniffer OSError on {ifname}: {e}")
|
||||||
break
|
break
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Unexpected sniffer error on {ifname}: {e}")
|
logger.exception(f"Unexpected sniffer error on {ifname}: {e}")
|
||||||
@@ -187,67 +301,126 @@ def sniffer_loop(ifname: str, stop_event: threading.Event, bridge: str):
|
|||||||
|
|
||||||
logger.info(f"Sniffer STOPPED on {ifname}")
|
logger.info(f"Sniffer STOPPED on {ifname}")
|
||||||
|
|
||||||
|
|
||||||
def start_sniffer_thread(bridge: str):
|
def start_sniffer_thread(bridge: str):
|
||||||
ports = get_bridge_ports(bridge)
|
ports = get_bridge_ports(bridge)
|
||||||
if not ports:
|
if not ports:
|
||||||
logger.error(f"[ERROR] Could not start sniffer: No valid ports found for bridge {bridge}")
|
logger.error(f"No ports found for bridge {bridge}, not starting sniffers.")
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
for iface in ports:
|
for iface in ports:
|
||||||
if iface in sniffer_threads:
|
if iface in sniffer_threads and sniffer_threads[iface].is_alive():
|
||||||
logger.info(f"[INFO] Sniffer on {iface} is already running")
|
logger.info(f"Sniffer already running on {iface}")
|
||||||
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(
|
thread = threading.Thread(target=sniffer_loop, args=(iface, stop_event, bridge), daemon=True)
|
||||||
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
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------
|
||||||
|
# Public API for starting/stopping sniffing (to be called from FastAPI)
|
||||||
|
# -------------------------------------------------------------------
|
||||||
|
_consumer_stop_event: Optional[asyncio.Event] = None
|
||||||
|
_consumer_future_accessor = None
|
||||||
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
# PUBLIC ASYNC START/STOP METHODS
|
|
||||||
# -------------------------------------------------------------------
|
|
||||||
async def start_sniffing(bridge: str):
|
async def start_sniffing(bridge: str):
|
||||||
logger.info(f"Starting sniffing for bridge {bridge}")
|
"""
|
||||||
|
Start sniffing on all bridge ports and start the DB consumer.
|
||||||
|
Safe to call multiple times.
|
||||||
|
"""
|
||||||
|
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)
|
start_sniffer_thread(bridge)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
async def stop_sniffing():
|
async def stop_sniffing():
|
||||||
logger.info("Stopping all sniffers…")
|
"""
|
||||||
|
Stop all sniffer threads and the DB consumer.
|
||||||
|
"""
|
||||||
|
global _consumer_stop_event, _consumer_future_accessor
|
||||||
|
logger.info("Stopping all sniffers and consumer")
|
||||||
|
|
||||||
for iface, stop_event in thread_stop_flags.items():
|
# stop sniffer threads
|
||||||
stop_event.set()
|
for iface, ev in list(thread_stop_flags.items()):
|
||||||
|
ev.set()
|
||||||
|
|
||||||
for iface, thread in sniffer_threads.items():
|
for iface, thread in list(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()
|
||||||
|
|
||||||
logger.info("All sniffers 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
|
# Status helper
|
||||||
# -------------------------------------------------------------------
|
# -------------------------------------------------------------------
|
||||||
def get_sniffer_status():
|
def get_sniffer_status():
|
||||||
|
"""
|
||||||
|
Return a dict of interface -> status information.
|
||||||
|
"""
|
||||||
out = {}
|
out = {}
|
||||||
for iface, t in sniffer_threads.items():
|
for iface, t in sniffer_threads.items():
|
||||||
out[iface] = {
|
out[iface] = {
|
||||||
"running": t.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),
|
||||||
}
|
}
|
||||||
|
out["queue_size"] = packet_queue.qsize()
|
||||||
return out
|
return out
|
||||||
|
|||||||
@@ -4,25 +4,16 @@ from src.network_sniffer import get_sniffer_status, start_sniffing, stop_sniffin
|
|||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
@router.post("/start_sniffer")
|
@router.post("/sniffer/start")
|
||||||
async def start_sniffer(bridge: str):
|
async def api_start(bridge: str):
|
||||||
"""
|
ok = await start_sniffing(bridge)
|
||||||
Start packet sniffing on bridge member interfaces.
|
return {"started": ok}
|
||||||
"""
|
|
||||||
|
|
||||||
await start_sniffing(bridge)
|
@router.post("/sniffer/stop")
|
||||||
|
async def api_stop():
|
||||||
|
ok = await stop_sniffing()
|
||||||
|
return {"stopped": ok}
|
||||||
|
|
||||||
return {
|
@router.get("/sniffer/status")
|
||||||
"status": "ok",
|
def api_status():
|
||||||
"started_on": bridge
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/stop_sniffer")
|
|
||||||
async def stop_sniffer_api():
|
|
||||||
await stop_sniffing()
|
|
||||||
return {"status": "stopped"}
|
|
||||||
|
|
||||||
@router.get("/status")
|
|
||||||
async def status():
|
|
||||||
return get_sniffer_status()
|
return get_sniffer_status()
|
||||||
Reference in New Issue
Block a user