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