use shared objects like loops db access etc
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 8s
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user