add utility for database
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s

This commit is contained in:
2025-12-03 19:34:42 +01:00
parent 309aeff2f1
commit 47b95e84dd
2 changed files with 85 additions and 66 deletions

View File

@@ -9,6 +9,7 @@ import struct
from typing import Dict, List, Optional, Any
import asyncpg
from backend.src.utilities.database import DB
from scapy.all import (
Ether,
ARP,
@@ -59,72 +60,11 @@ threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).star
# -------------------------
# Database helper (pool)
# Database init (pool)
# -------------------------
class DB:
_pool: Optional[asyncpg.pool.Pool] = None
@classmethod
async def init_pool(cls) -> None:
if cls._pool is None:
cls._pool = await asyncpg.create_pool(dsn=DB_DSN, min_size=1, max_size=5)
logger.info("DB pool initialized")
@classmethod
async def close_pool(cls) -> None:
if cls._pool:
await cls._pool.close()
cls._pool = None
logger.info("DB pool closed")
@classmethod
async def insert_packet(cls, pkt_info: Dict[str, Any]) -> None:
"""
Insert packet metadata into DB. This preserves the exact columns/values in the original code.
"""
if cls._pool is None:
# defensive: try to initialize if not ready
await cls.init_pool()
try:
async with cls._pool.acquire() as conn:
await conn.execute(
"""
INSERT INTO packets(
iface,
src_mac,
dst_mac,
eth_type,
vlan_id,
src_ip,
dst_ip,
ip_proto,
src_port,
dst_port,
length,
raw
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12)
""",
pkt_info["iface"],
pkt_info.get("src_mac"),
pkt_info.get("dst_mac"),
pkt_info.get("eth_type"),
pkt_info.get("vlan_id"),
pkt_info.get("src_ip"),
pkt_info.get("dst_ip"),
pkt_info.get("protocol"),
pkt_info.get("src_port"),
pkt_info.get("dst_port"),
pkt_info["length"],
pkt_info["raw"],
)
except Exception:
# keep sniffer alive: log and continue
logger.exception("DB insert failed")
db = DB(DB_DSN)
# schedule pool init on background loop (fire-and-forget)
asyncio.run_coroutine_threadsafe(DB.init_pool(), async_loop)
asyncio.run_coroutine_threadsafe(db.init_pool(), async_loop)
# -------------------------
@@ -278,7 +218,7 @@ def parse_packet(pkt, bridge: str) -> None:
# Submit DB insert to background loop (non-blocking)
try:
asyncio.run_coroutine_threadsafe(DB.insert_packet(pkt_info), async_loop)
asyncio.run_coroutine_threadsafe(db.insert_packet(pkt_info), async_loop)
except Exception:
logger.exception("Failed to schedule DB insert")
@@ -481,7 +421,7 @@ def stop_afpacket_sniffer() -> None:
# close db pool (schedule on async loop)
try:
asyncio.run_coroutine_threadsafe(DB.close_pool(), async_loop)
asyncio.run_coroutine_threadsafe(db.close_pool(), async_loop)
except Exception:
logger.exception("Failed to schedule DB pool close")

View File

@@ -0,0 +1,79 @@
import logging
from typing import Dict, List, Optional, Any
import asyncpg
# ---- Logging ----------------------------------------------------------
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("af_packet_sniffer")
class DB:
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._min_size = min_size
self._max_size = max_size
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")
async def close_pool(self) -> None:
"""Gracefully close connection pool."""
if self._pool:
await self._pool.close()
self._pool = None
logger.info("DB pool closed")
async def insert_packet(self, pkt_info: Dict[str, Any]) -> None:
"""
Insert packet metadata. Same fields + same SQL statement as before.
"""
if self._pool is None:
await self.init_pool() # lazy init
try:
async with self._pool.acquire() as conn:
await conn.execute(
"""
INSERT INTO packets(
iface,
src_mac,
dst_mac,
eth_type,
vlan_id,
src_ip,
dst_ip,
ip_proto,
src_port,
dst_port,
length,
raw
) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12)
""",
pkt_info["iface"],
pkt_info.get("src_mac"),
pkt_info.get("dst_mac"),
pkt_info.get("eth_type"),
pkt_info.get("vlan_id"),
pkt_info.get("src_ip"),
pkt_info.get("dst_ip"),
pkt_info.get("protocol"),
pkt_info.get("src_port"),
pkt_info.get("dst_port"),
pkt_info["length"],
pkt_info["raw"],
)
except Exception:
logger.exception("DB insert failed")