From 47b95e84dd008d52d5df1dd0bf1cda3764c65bb0 Mon Sep 17 00:00:00 2001 From: malmert Date: Wed, 3 Dec 2025 19:34:42 +0100 Subject: [PATCH] add utility for database --- backend/src/network_sniffer.py | 72 +++------------------------- backend/src/utilities/database.py | 79 +++++++++++++++++++++++++++++++ 2 files changed, 85 insertions(+), 66 deletions(-) create mode 100644 backend/src/utilities/database.py diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index 4228ce3..a13f678 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -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") diff --git a/backend/src/utilities/database.py b/backend/src/utilities/database.py new file mode 100644 index 0000000..bf6d79a --- /dev/null +++ b/backend/src/utilities/database.py @@ -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")