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};"