add utility for database
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 7s
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user