file structure and comments unified
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s
Some checks failed
Build and Deploy MITM Webserver / build (push) Failing after 3s
This commit is contained in:
@@ -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};"
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user