use shared objects like loops db access etc
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s

This commit is contained in:
2025-12-03 20:49:22 +01:00
parent 50225a6dec
commit a18576a61e
8 changed files with 526 additions and 46 deletions

View File

@@ -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

View File

@@ -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