file structure and comments unified
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s

This commit is contained in:
2026-03-06 19:03:26 +01:00
parent 7b9a7d3a4b
commit b60a1d3118
20 changed files with 697 additions and 1823 deletions

View File

@@ -1,28 +1,21 @@
# src/utilities/database.py
import logging
import base64
import asyncio
from typing import Dict, List, Optional, Any
from pydantic import ValidationError
"""Database helper for packet persistence and retrieval."""
import asyncio
import base64
import logging
from typing import Any, Dict, List, Optional
import asyncpg
from asyncpg.pool import Pool
from pydantic import ValidationError
from src.Models.packets import PacketDBModel
# ---- Logging ----------------------------------------------------------
logger = logging.getLogger("af_packet_sniffer")
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
"""
"""Asyncpg connection pool wrapper used by the packet APIs."""
def __init__(self, dsn: str, min_size: int = 1, max_size: int = 5):
self._dsn = dsn
@@ -30,26 +23,26 @@ class DatabasePool:
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:
"""Initialize the asyncpg pool if not already initialized (idempotent)."""
"""Initialize the connection pool once per process."""
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
dsn=self._dsn,
min_size=self._min_size,
max_size=self._max_size,
)
logger.info("DB pool initialized")
except Exception:
@@ -57,9 +50,10 @@ class DatabasePool:
raise
async def close_pool(self) -> None:
"""Close the pool if it exists."""
"""Close the pool if present."""
if self._pool is None:
return
try:
await self._pool.close()
logger.info("DB pool closed")
@@ -69,13 +63,7 @@ class DatabasePool:
self._pool = None
async def insert_packet(self, pkt_info: Dict[str, Any]) -> None:
"""
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).
"""
"""Insert one packet record and publish it to subscribers."""
if self._pool is None:
await self.init_pool()
@@ -115,13 +103,11 @@ class DatabasePool:
except Exception:
logger.exception("DB insert failed")
return
# Update the dictionary with the DB-generated values
if new_row:
pkt_info["id"] = new_row["id"]
# Convert timestamp to ISO string for JSON serialization in WebSockets
pkt_info["timestamp"] = new_row["timestamp"].isoformat()
# notify broadcaster (non-blocking). broadcaster is expected to be thread-safe.
if self.broadcaster:
try:
self.broadcaster.sync_publish(pkt_info)
@@ -129,11 +115,7 @@ class DatabasePool:
logger.exception("Failed to publish pkt_info to broadcaster")
async def fetch_latest(self, limit: int) -> List[PacketDBModel]:
"""
Fetch the latest `limit` packets (newest first).
Returns a list of PacketDBModel. Converts raw bytes -> raw_b64 for JSON-safe output.
"""
"""Fetch newest packet rows as validated `PacketDBModel` instances."""
if self._pool is None:
await self.init_pool()
@@ -148,45 +130,34 @@ class DatabasePool:
limit,
)
out: List[PacketDBModel] = []
result: List[PacketDBModel] = []
for row in rows:
data = dict(row)
for r in rows:
d = dict(r)
# convert byte raw -> base64 string (and remove raw)
raw_val = d.get("raw")
raw_val = data.get("raw")
if isinstance(raw_val, (bytes, bytearray)):
d["raw_b64"] = base64.b64encode(raw_val).decode("ascii")
d.pop("raw", None)
data["raw_b64"] = base64.b64encode(raw_val).decode("ascii")
data.pop("raw", None)
# Validate/construct Pydantic model
try:
packet_model = PacketDBModel(**d)
except ValidationError as ve:
# Log and skip invalid rows (or handle otherwise)
packet_model = PacketDBModel(**data)
except ValidationError as exc:
logger.warning(
"Skipping DB row that failed PacketDBModel validation (id=%s): %s",
d.get("id"),
ve,
data.get("id"),
exc,
)
continue
out.append(packet_model)
result.append(packet_model)
return result
return out
async def clear_all_packets(self, reset_identity: bool = True) -> bool:
"""
Deletes all rows from the `packets` table.
If reset_identity is True, the auto-increment ID counter is reset to 1.
Returns True if successful, False otherwise.
"""
"""Truncate the packet table and optionally reset identity counters."""
if self._pool is None:
await self.init_pool()
# TRUNCATE is faster than DELETE and resets the identity counter
restart_clause = "RESTART IDENTITY" if reset_identity else ""
query = f"TRUNCATE TABLE packets {restart_clause};"

View File

@@ -1,140 +1,100 @@
# src/utilities/packet_broadcaster.py
"""In-process packet broadcaster for websocket subscribers."""
import asyncio
import logging
from typing import Dict, Any, List, Optional
from typing import Any, Dict, 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.
"""
"""Manage subscriber queues and publish packet events."""
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
self._closed = False
try:
# ensure lock is created on the given loop
def _make_lock():
def _make_lock() -> None:
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.
"""
"""Create and register a queue for one subscriber."""
if self._closed:
raise RuntimeError("PacketBroadcaster is closed")
q: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize)
# wait until lock exists
queue: asyncio.Queue = asyncio.Queue(maxsize=self._queue_maxsize)
while self._lock is None:
await asyncio.sleep(0) # yield to event loop briefly
await asyncio.sleep(0)
async with self._lock:
self._subscribers.append(q)
return q
self._subscribers.append(queue)
async def unsubscribe(self, q: asyncio.Queue) -> None:
"""
Remove a subscriber queue if present.
"""
return queue
async def unsubscribe(self, queue: asyncio.Queue) -> None:
"""Unregister a subscriber queue if it exists."""
if self._lock is None:
return
async with self._lock:
try:
self._subscribers.remove(q)
self._subscribers.remove(queue)
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
"""Publish one message to all current subscribers."""
if self._closed or self._lock is None:
return
async with self._lock:
subs = list(self._subscribers)
subscribers = list(self._subscribers)
for q in subs:
for queue in subscribers:
try:
q.put_nowait(msg)
queue.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
logger.exception("Unexpected subscriber publish error: %s", exc)
try:
async with self._lock:
if q in self._subscribers:
self._subscribers.remove(q)
if queue in self._subscribers:
self._subscribers.remove(queue)
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.
"""
"""Thread-safe wrapper that schedules `publish` on the broadcaster 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.
"""
"""Close the broadcaster and clear queued messages."""
self._closed = True
if self._lock is None:
return
async with self._lock:
subs = list(self._subscribers)
subscribers = list(self._subscribers)
self._subscribers.clear()
for q in subs:
for queue in subscribers:
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
while not queue.empty():
queue.get_nowait()
except Exception:
pass