From a18576a61e8524a5182ab522b4c1838deff2b42b Mon Sep 17 00:00:00 2001 From: malmert Date: Wed, 3 Dec 2025 20:49:22 +0100 Subject: [PATCH] use shared objects like loops db access etc --- backend/src/{ => api}/network_api.py | 0 backend/src/api/packet_api.py | 144 ++++++++++++++++++ .../{routes_sniffer.py => api/sniffer_api.py} | 0 backend/src/main.py | 95 ++++++++++-- backend/src/network_sniffer.py | 71 +++++++-- backend/src/shared_objects.py | 8 + backend/src/utilities/database.py | 114 +++++++++++--- backend/src/utilities/packet_broadcaster.py | 140 +++++++++++++++++ 8 files changed, 526 insertions(+), 46 deletions(-) rename backend/src/{ => api}/network_api.py (100%) create mode 100644 backend/src/api/packet_api.py rename backend/src/{routes_sniffer.py => api/sniffer_api.py} (100%) create mode 100644 backend/src/shared_objects.py create mode 100644 backend/src/utilities/packet_broadcaster.py diff --git a/backend/src/network_api.py b/backend/src/api/network_api.py similarity index 100% rename from backend/src/network_api.py rename to backend/src/api/network_api.py diff --git a/backend/src/api/packet_api.py b/backend/src/api/packet_api.py new file mode 100644 index 0000000..42e5f6a --- /dev/null +++ b/backend/src/api/packet_api.py @@ -0,0 +1,144 @@ +# src/routers/packets.py +import asyncio +import base64 +import json +import logging +from typing import Optional, Any, Dict, List + +from fastapi import APIRouter, Query, WebSocket, WebSocketDisconnect, HTTPException +from fastapi.responses import JSONResponse + +import src.shared_objects as shared + +logger = logging.getLogger("packets_router") +router = APIRouter() + + +def _serialize_row_for_json(row: Dict[str, Any]) -> Dict[str, Any]: + """ + Convert DB row / pkt_info to a JSON-serializable dict. + - If 'raw' is bytes, produce 'raw_b64' and drop 'raw'. + - Fallback to str() for unknown/unserializable values. + """ + out: Dict[str, Any] = {} + for k, v in row.items(): + if k == "raw" and isinstance(v, (bytes, bytearray)): + out["raw_b64"] = base64.b64encode(v).decode("ascii") + continue + # try to JSON serialize the value directly + try: + json.dumps({k: v}) + out[k] = v + except (TypeError, ValueError): + out[k] = str(v) + return out + + +async def _serialize_rows(rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + return [_serialize_row_for_json(r) for r in rows] + + +@router.get("/packets") +async def get_packets(limit: int = Query(100, ge=1, le=10000)): + """ + Return latest `limit` packets (newest first). The DB helper already converts + `raw` to `raw_b64` in fetch_latest, but we defensively re-serialize here. + """ + db = shared.db + if db is None: + logger.warning("GET /packets called but DB is not available") + raise HTTPException(status_code=503, detail="Database not available") + + try: + rows = await db.fetch_latest(limit) + serial = await _serialize_rows(rows) + return JSONResponse(content={"count": len(serial), "packets": serial}) + except Exception: + logger.exception("Failed to fetch latest packets from DB") + raise HTTPException(status_code=500, detail="Failed to fetch packets") + + +@router.websocket("/ws/packets") +async def websocket_packets(ws: WebSocket): + """ + WebSocket live feed endpoint. + + Accepts optional query param `subscribe_recent` (e.g. ?subscribe_recent=20) + which will deliver the last N packets immediately on connect. + """ + await ws.accept() + logger.debug("WebSocket connection accepted: %s", ws.client) + + db = shared.db + broadcaster = shared.broadcaster + + if db is None: + await ws.send_json({"error": "database not available"}) + await ws.close() + logger.warning("WebSocket closed: DB not available") + return + + if broadcaster is None: + await ws.send_json({"error": "broadcaster not available"}) + await ws.close() + logger.warning("WebSocket closed: broadcaster not available") + return + + # Parse subscribe_recent from query params (defensive) + try: + subscribe_recent_raw = ws.query_params.get("subscribe_recent", "0") + subscribe_recent = int(subscribe_recent_raw) + if subscribe_recent < 0: + subscribe_recent = 0 + except Exception: + subscribe_recent = 0 + + q: Optional[asyncio.Queue] = None + try: + # Optionally send recent history first + if subscribe_recent > 0: + recent = await db.fetch_latest(subscribe_recent) + recent_serial = await _serialize_rows(recent) + await ws.send_json({"type": "recent", "count": len(recent_serial), "packets": recent_serial}) + + # Subscribe to broadcaster to receive live packets + q = await broadcaster.subscribe() + logger.info("WebSocket subscribed client %s (queue maxsize=%d)", ws.client, q.maxsize) + + # Simple heartbeat: periodically ensure client is responsive (optional) + # We'll implement by awaiting q.get() which blocks until a message is published. + while True: + msg = await q.get() + # Normalize message to JSON-able dict + if isinstance(msg, dict): + payload = _serialize_row_for_json(msg) + else: + # not a dict — try to json-serialize directly + try: + json.dumps(msg) + payload = msg + except Exception: + payload = {"data": str(msg)} + + try: + await ws.send_json(payload) + except Exception: + # sending failed (client disconnected or write error) + logger.info("WebSocket send failed for client %s — unsubscribing", ws.client) + break + except WebSocketDisconnect: + logger.info("WebSocket client disconnected: %s", ws.client) + except Exception: + logger.exception("Unexpected error in websocket_packets") + finally: + # Clean up subscriber queue + if q is not None: + try: + await broadcaster.unsubscribe(q) + except Exception: + logger.exception("Failed to unsubscribe websocket queue") + try: + await ws.close() + except Exception: + pass + logger.debug("WebSocket connection closed and cleaned up for client %s", ws.client) diff --git a/backend/src/routes_sniffer.py b/backend/src/api/sniffer_api.py similarity index 100% rename from backend/src/routes_sniffer.py rename to backend/src/api/sniffer_api.py diff --git a/backend/src/main.py b/backend/src/main.py index ef5f3b7..57b9d5c 100644 --- a/backend/src/main.py +++ b/backend/src/main.py @@ -1,12 +1,23 @@ +# src/main.py +import asyncio +import logging +import os from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware -import os -import src.Models.netplan as netplan -import src.network_api as network_api -import src.routes_sniffer as sniffer_router -from asyncio import subprocess -from fastapi import HTTPException +from src.api import packet_api +from src.utilities.packet_broadcaster import PacketBroadcaster +import src.shared_objects as shared_objects +from src.utilities.database import DatabasePool +import src.api.network_api as network_api +import src.api.sniffer_api as sniffer_api + +# ---- Config ----------------------------------------------------------- +DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" + +# ---- Globals ----------------------------------------------- +# Create DatabasePool instance (pool created on startup) +shared_objects.db = DatabasePool(DB_DSN) app = FastAPI( root_path="/api", @@ -26,15 +37,76 @@ app.add_middleware( allow_headers=["*"], ) - # --------------------- # Startup / Shutdown # --------------------- +@app.on_event("startup") +async def on_startup(): + """ + Initialize DB pool and broadcaster on the FastAPI event loop and + publish them into shared_objects so other modules (sniffer, routers) + can access them. + """ + loop = asyncio.get_running_loop() + shared_objects.web_loop = loop + + # Initialize DB pool bound to this loop + try: + await shared_objects.db.init_pool() + except Exception: + logging.exception("Failed to initialize DB pool") + raise + + # Create broadcaster and attach to DB so DB.insert_packet can publish updates + try: + shared_objects.broadcaster = PacketBroadcaster(loop) + shared_objects.db.broadcaster = shared_objects.broadcaster + except Exception: + logging.exception("Failed to create/attach broadcaster") + # continue — DB is primary; broadcaster optional + + # Drain any buffered packets from the sniffer (if it started earlier) + try: + # import sniffer here to avoid circular imports at module import time + from src import sniffer + + # sniffer provides drain_buffer_to_shared_db() + try: + sniffer.drain_buffer_to_shared_db() + except Exception: + logging.exception("Failed to drain sniffer buffer") + except ImportError: + # sniffer not present or not importable; skip + logging.debug("sniffer module not importable at startup; skipping buffer drain") + @app.on_event("shutdown") -def shutdown_event(): - network_api.shutdown_network_api() +async def shutdown_event(): + """ + Shutdown actions: stop network API and close DB pool if present. + """ + # try to shut down network API components + try: + network_api.shutdown_network_api() + except Exception: + logging.exception("Error shutting down network API") + + # close DB pool if available in shared_objects + try: + web_db = getattr(shared_objects, "db", None) + if web_db is not None: + await web_db.close_pool() + except Exception: + logging.exception("Failed to close DB pool during shutdown") + + # clear shared runtime objects (optional cleanup) + try: + shared_objects.db = None + shared_objects.broadcaster = None + shared_objects.web_loop = None + except Exception: + pass # --------------------- @@ -45,11 +117,13 @@ def shutdown_event(): def hello(): return {"message": "Hello from FastAPI 🎉"} + @app.get("/versions") def versions(): message = os.popen("python --version").read().strip() return {"message": message} + @app.get("/nft/ruleset") def nft_ruleset(): message = os.popen("sudo nft --json list ruleset").read().strip() @@ -61,4 +135,5 @@ def nft_ruleset(): # --------------------- app.include_router(network_api.router, prefix="/network", tags=["network"]) -app.include_router(sniffer_router.router, prefix="/sniffer", tags=["sniffer"]) \ No newline at end of file +app.include_router(sniffer_api.router, prefix="/sniffer", tags=["sniffer"]) +app.include_router(packet_api.router, prefix="/packets", tags=["packets"]) \ No newline at end of file diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index ca5befd..e49a777 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -1,3 +1,4 @@ +# src/sniffer.py import asyncio import logging import threading @@ -8,8 +9,9 @@ import errno import struct from typing import Dict, List, Optional, Any -import asyncpg -from src.utilities.database import DB +# NOTE: ensure this path points to your shared runtime module +from backend.src import shared_objects + from scapy.all import ( Ether, ARP, @@ -35,9 +37,6 @@ from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number logging.basicConfig(level=logging.INFO) logger = logging.getLogger("af_packet_sniffer") -# ---- Config ----------------------------------------------------------- -DB_DSN = "postgresql://mitm_user:mitm_password@localhost:5432/mitm_db" - # ---- Globals (kept minimal) ------------------------------------------- af_sockets: Dict[str, socket.socket] = {} af_thread: Optional[threading.Thread] = None @@ -46,7 +45,11 @@ af_stop_event: Optional[threading.Event] = None fixed_bridge_ports: Dict[str, List[str]] = {} current_bridge: Optional[str] = None -# Background asyncio loop used to run DB tasks +# small bounded buffer for packets produced before shared_objects is ready +_PACKET_BUFFER: List[Dict[str, Any]] = [] +_BUFFER_CAPACITY = 20000 + +# Background asyncio loop used for internal tasks in this module (kept but not used for DB pool) async_loop = asyncio.new_event_loop() @@ -60,11 +63,37 @@ threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).star # ------------------------- -# Database init (pool) +# Helpers for buffer draining # ------------------------- -db = DB(DB_DSN) -# schedule pool init on background loop (fire-and-forget) -asyncio.run_coroutine_threadsafe(db.init_pool(), async_loop) +def drain_buffer_to_shared_db() -> None: + """ + Attempt to schedule buffered packets for insertion on shared_objects.web_loop. + Call this from main.py after shared_objects.db and shared_objects.web_loop are initialized. + """ + try: + web_loop = getattr(shared_objects, "web_loop", None) + web_db = getattr(shared_objects, "db", None) + if web_db is None or web_loop is None: + return + + # schedule draining on the web loop to avoid blocking this thread + def _drain() -> None: + while _PACKET_BUFFER: + pkt = _PACKET_BUFFER.pop(0) + try: + asyncio.run_coroutine_threadsafe(web_db.insert_packet(pkt), web_loop) + except Exception: + # re-buffer first element and stop to avoid busy loop + _PACKET_BUFFER.insert(0, pkt) + break + + try: + web_loop.call_soon_threadsafe(_drain) + except Exception: + # fallback: run directly (best-effort) + _drain() + except Exception: + logger.exception("Failed to drain packet buffer") # ------------------------- @@ -80,8 +109,6 @@ def _safe_get_attr(layer, attr: str): def parse_packet(pkt, bridge: str) -> None: """ Parse a scapy Packet object into a normalized dict and schedule DB insert. - Keeps all information from the original implementation (fields, VLAN handling, - IP/IPv6/TCP/UDP/ARP handling and enum mapping). """ pkt_iface = getattr(pkt, "sniffed_on", None) if not pkt_iface: @@ -216,9 +243,18 @@ def parse_packet(pkt, bridge: str) -> None: if Raw in pkt and not pkt_info["protocol_name"]: pkt_info["protocol_name"] = "RAW" - # Submit DB insert to background loop (non-blocking) + # Submit DB insert to shared web loop if available, otherwise buffer try: - asyncio.run_coroutine_threadsafe(db.insert_packet(pkt_info), async_loop) + web_loop = getattr(shared_objects, "web_loop", None) + web_db = getattr(shared_objects, "db", None) + if web_db is not None and web_loop is not None: + asyncio.run_coroutine_threadsafe(web_db.insert_packet(pkt_info), web_loop) + else: + # buffer (bounded) until the web app initializes + _PACKET_BUFFER.append(pkt_info) + if len(_PACKET_BUFFER) > _BUFFER_CAPACITY: + # drop oldest packet + _PACKET_BUFFER.pop(0) except Exception: logger.exception("Failed to schedule DB insert") @@ -419,9 +455,12 @@ def stop_afpacket_sniffer() -> None: fixed_bridge_ports.pop(current_bridge, None) current_bridge = None - # close db pool (schedule on async loop) + # close db pool if shared.web_loop is available; otherwise leave to main try: - asyncio.run_coroutine_threadsafe(db.close_pool(), async_loop) + web_loop = getattr(shared_objects, "web_loop", None) + web_db = getattr(shared_objects, "db", None) + if web_db is not None and web_loop is not None: + asyncio.run_coroutine_threadsafe(web_db.close_pool(), web_loop) except Exception: logger.exception("Failed to schedule DB pool close") diff --git a/backend/src/shared_objects.py b/backend/src/shared_objects.py new file mode 100644 index 0000000..552d362 --- /dev/null +++ b/backend/src/shared_objects.py @@ -0,0 +1,8 @@ +from typing import Optional +import asyncio + +# These are filled at FastAPI startup +# DB instance +db = None +web_loop: Optional[asyncio.AbstractEventLoop] = None +broadcaster = None \ No newline at end of file diff --git a/backend/src/utilities/database.py b/backend/src/utilities/database.py index bf6d79a..2da0ba3 100644 --- a/backend/src/utilities/database.py +++ b/backend/src/utilities/database.py @@ -1,47 +1,80 @@ - +# src/utilities/database.py import logging +import base64 +import asyncio from typing import Dict, List, Optional, Any import asyncpg +from asyncpg.pool import Pool # ---- Logging ---------------------------------------------------------- logging.basicConfig(level=logging.INFO) logger = logging.getLogger("af_packet_sniffer") -class DB: + +class DatabasePool: + """ + Lightweight asyncpg connection pool wrapper. + + - Lazy pool creation via init_pool() + - Safe against concurrent init_pool() calls via an asyncio.Lock created on first use + - insert_packet() forwards the pkt_info to an optional broadcaster after successful insert + """ + def __init__(self, dsn: str, min_size: int = 1, max_size: int = 5): - """ - Create a DB helper bound to the given DSN. - Pool is created lazily on first use unless init_pool() is explicitly called. - """ self._dsn = dsn - self._pool: Optional[asyncpg.pool.Pool] = None + self._pool: Optional[Pool] = None self._min_size = min_size self._max_size = max_size + self.broadcaster = None + # created on first init_pool() (must be created on an event loop) + self._init_lock: Optional[asyncio.Lock] = None async def init_pool(self) -> None: - """Create connection pool if not already created.""" - if self._pool is None: - self._pool = await asyncpg.create_pool( - dsn=self._dsn, - min_size=self._min_size, - max_size=self._max_size, - ) - logger.info("DB pool initialized") + """Initialize the asyncpg pool if not already initialized (idempotent).""" + if self._pool is not None: + return + + # Ensure a lock exists that is bound to the running event loop + if self._init_lock is None: + self._init_lock = asyncio.Lock() + + async with self._init_lock: + # Double-check after acquiring lock + if self._pool is not None: + return + logger.info("Initializing DB pool (dsn=%s)", self._dsn) + try: + self._pool = await asyncpg.create_pool( + dsn=self._dsn, min_size=self._min_size, max_size=self._max_size + ) + logger.info("DB pool initialized") + except Exception: + logger.exception("Failed to create DB pool") + raise async def close_pool(self) -> None: - """Gracefully close connection pool.""" - if self._pool: + """Close the pool if it exists.""" + if self._pool is None: + return + try: await self._pool.close() - self._pool = None logger.info("DB pool closed") + except Exception: + logger.exception("Error closing DB pool") + finally: + self._pool = None async def insert_packet(self, pkt_info: Dict[str, Any]) -> None: """ - Insert packet metadata. Same fields + same SQL statement as before. + Insert packet metadata into the `packets` table. + + Preserves the same columns/values as before. + After a successful insert, if a broadcaster is attached it will be + notified via broadcaster.sync_publish(pkt_info). """ if self._pool is None: - await self.init_pool() # lazy init + await self.init_pool() try: async with self._pool.acquire() as conn: @@ -77,3 +110,44 @@ class DB: ) except Exception: logger.exception("DB insert failed") + return + + # notify broadcaster (non-blocking). broadcaster is expected to be thread-safe. + if self.broadcaster: + try: + self.broadcaster.sync_publish(pkt_info) + except Exception: + logger.exception("Failed to publish pkt_info to broadcaster") + + async def fetch_latest(self, limit: int) -> List[Dict[str, Any]]: + """ + Fetch the latest `limit` packets (newest first). + + Returns a list of dicts. If the `raw` column is binary it is converted + to `raw_b64` (base64 string) and `raw` is removed. + """ + if self._pool is None: + await self.init_pool() + + async with self._pool.acquire() as conn: + # Select explicit columns to ensure predictable dict keys + rows = await conn.fetch( + """ + SELECT id, iface, src_mac, dst_mac, eth_type, vlan_id, src_ip, dst_ip, + ip_proto, src_port, dst_port, length, raw, created_at + FROM packets + ORDER BY id DESC + LIMIT $1 + """, + limit, + ) + + out: List[Dict[str, Any]] = [] + for r in rows: + d = dict(r) + raw_val = d.get("raw") + if isinstance(raw_val, (bytes, bytearray)): + d["raw_b64"] = base64.b64encode(raw_val).decode("ascii") + d.pop("raw", None) + out.append(d) + return out diff --git a/backend/src/utilities/packet_broadcaster.py b/backend/src/utilities/packet_broadcaster.py new file mode 100644 index 0000000..daa3d7a --- /dev/null +++ b/backend/src/utilities/packet_broadcaster.py @@ -0,0 +1,140 @@ +# src/utilities/packet_broadcaster.py +import asyncio +import logging +from typing import Dict, Any, List, Optional + +logger = logging.getLogger("packet_broadcaster") + + +class PacketBroadcaster: + """ + Simple in-process broadcaster: + - Maintains a set of subscriber asyncio.Queues (one per websocket connection). + - publish(msg) is run on the broadcaster's event loop. + - sync_publish(msg) is thread-safe and can be called from other threads / loops. + + Note: create this on the FastAPI event loop (e.g. in startup) so that its lock and + operations run on that same loop. + """ + + def __init__(self, loop: asyncio.AbstractEventLoop, queue_maxsize: int = 1024): + self._loop = loop + self._queue_maxsize = queue_maxsize + + # create lock and subscribers on the target loop to avoid cross-loop asyncio primitives + self._subscribers: List[asyncio.Queue] = [] + # create lock bound to the same loop by scheduling its construction on that loop + self._lock: Optional[asyncio.Lock] = None + try: + # ensure lock is created on the given loop + def _make_lock(): + self._lock = asyncio.Lock() + + loop.call_soon_threadsafe(_make_lock) + except Exception: + # fallback — create in current loop if call_soon_threadsafe fails + self._lock = asyncio.Lock() + + self._closed = False + + async def subscribe(self) -> asyncio.Queue: + """ + Create a subscriber queue and add it to the list. + Caller is expected to await on the returned queue to receive messages. + """ + if self._closed: + raise RuntimeError("PacketBroadcaster is closed") + + q: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize) + # wait until lock exists + while self._lock is None: + await asyncio.sleep(0) # yield to event loop briefly + + async with self._lock: + self._subscribers.append(q) + return q + + async def unsubscribe(self, q: asyncio.Queue) -> None: + """ + Remove a subscriber queue if present. + """ + if self._lock is None: + return + async with self._lock: + try: + self._subscribers.remove(q) + except ValueError: + pass + + async def publish(self, msg: Dict[str, Any]) -> None: + """ + Publish msg to all subscribers (must be called on the broadcaster's loop). + We use put_nowait to avoid blocking. If a subscriber queue is full we drop + that subscriber's message to avoid backpressure. + """ + if self._closed: + return + + if self._lock is None: + # not initialized yet; nothing to do + return + + async with self._lock: + subs = list(self._subscribers) + + for q in subs: + try: + q.put_nowait(msg) + except asyncio.QueueFull: + # drop message for this subscriber + continue + except Exception as exc: + logger.exception("Unexpected error when publishing to subscriber: %s", exc) + # attempt to remove broken subscriber + try: + async with self._lock: + if q in self._subscribers: + self._subscribers.remove(q) + except Exception: + pass + + def sync_publish(self, msg: Dict[str, Any]) -> None: + """ + Thread-safe publish method: schedule publish(msg) on the broadcaster's loop. + Safe to call from other threads / event loops. + + We schedule creation of the publish task on the broadcaster loop using + call_soon_threadsafe so that publish() runs on the correct loop. + """ + if self._closed: + return + + try: + # schedule the coroutine to run on the broadcaster loop + self._loop.call_soon_threadsafe(asyncio.create_task, self.publish(msg)) + except Exception as exc: + # swallow errors but log for debugging + logger.exception("sync_publish failed to schedule publish: %s", exc) + + async def close(self) -> None: + """ + Close the broadcaster: mark closed, clear subscribers, and drain queues. + """ + self._closed = True + if self._lock is None: + return + async with self._lock: + subs = list(self._subscribers) + self._subscribers.clear() + + for q in subs: + try: + # optionally notify subscribers of closure by putting None (client must handle) + # q.put_nowait(None) + while not q.empty(): + try: + q.get_nowait() + except Exception: + break + except Exception: + pass