"""Database helper for packet persistence and retrieval.""" import asyncio import base64 import json import logging from datetime import datetime, timezone from typing import Any, Dict, List, Optional import asyncpg from asyncpg.pool import Pool from pydantic import ValidationError from src.Models.packets import PacketDBModel logger = logging.getLogger("af_packet_sniffer") def _db_text(value: Any) -> Any: if value is None: return None if isinstance(value, (str, int)): return value if hasattr(value, "value"): return getattr(value, "value") return str(value) def _serialize_row_for_broadcast(row: Dict[str, Any]) -> Dict[str, Any]: serialized = dict(row) _normalize_json_fields(serialized) raw_val = serialized.get("raw") if isinstance(raw_val, (bytes, bytearray)): serialized["raw_b64"] = base64.b64encode(raw_val).decode("ascii") serialized.pop("raw", None) for key, value in list(serialized.items()): if hasattr(value, "isoformat"): serialized[key] = value.isoformat() return serialized def _normalize_json_fields(payload: Dict[str, Any]) -> None: for key in ("dpi_metadata", "capture_metadata", "telemetry_metadata"): value = payload.get(key) if isinstance(value, str): try: parsed = json.loads(value) except json.JSONDecodeError: continue if isinstance(parsed, dict): payload[key] = parsed class DatabasePool: """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 self._pool: Optional[Pool] = None self._min_size = min_size self._max_size = max_size self.broadcaster = None self._init_lock: Optional[asyncio.Lock] = None async def init_pool(self) -> None: """Initialize the connection pool once per process.""" if self._pool is not None: return if self._init_lock is None: self._init_lock = asyncio.Lock() async with self._init_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: """Close the pool if present.""" if self._pool is None: return try: await self._pool.close() 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 one packet record and publish it to subscribers.""" await self.upsert_packet(pkt_info) async def upsert_packet(self, pkt_info: Dict[str, Any]) -> None: """Insert or update one packet record and publish it to subscribers.""" if self._pool is None: await self.init_pool() dpi_metadata = pkt_info.get("dpi_metadata") telemetry_metadata = pkt_info.get("telemetry_metadata") try: async with self._pool.acquire() as conn: row = await conn.fetchrow( """ INSERT INTO packets ( correlation_key, packet_id, packet_uid, correlation_source, skb_mark, capture_iface, ingress_if, egress_if, verdict, verdict_reason, verdict_confidence, ingress_seen_at, egress_seen_at, verdict_seen_at, src_mac, dst_mac, eth_type_raw, eth_type, vlan_id, src_ip, dst_ip, ip_proto_raw, ip_proto, src_port, dst_port, length, raw_present, capture_sources, app_protocol, app_master_protocol, app_category, app_confidence, app_hostname, app_is_encrypted, app_risk_score, dpi_metadata, capture_metadata, telemetry_metadata, raw ) VALUES( $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20, $21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36::jsonb, $37::jsonb,$38::jsonb,$39 ) ON CONFLICT (correlation_key) DO UPDATE SET updated_at = NOW(), packet_id = COALESCE(EXCLUDED.packet_id, packets.packet_id), packet_uid = COALESCE(EXCLUDED.packet_uid, packets.packet_uid), correlation_source = COALESCE(EXCLUDED.correlation_source, packets.correlation_source), skb_mark = COALESCE(EXCLUDED.skb_mark, packets.skb_mark), capture_iface = COALESCE(EXCLUDED.capture_iface, packets.capture_iface), ingress_if = COALESCE(EXCLUDED.ingress_if, packets.ingress_if), egress_if = COALESCE(EXCLUDED.egress_if, packets.egress_if), verdict = COALESCE(EXCLUDED.verdict, packets.verdict), verdict_reason = COALESCE(EXCLUDED.verdict_reason, packets.verdict_reason), verdict_confidence = COALESCE(EXCLUDED.verdict_confidence, packets.verdict_confidence), ingress_seen_at = COALESCE(EXCLUDED.ingress_seen_at, packets.ingress_seen_at), egress_seen_at = COALESCE(EXCLUDED.egress_seen_at, packets.egress_seen_at), verdict_seen_at = COALESCE(EXCLUDED.verdict_seen_at, packets.verdict_seen_at), src_mac = COALESCE(EXCLUDED.src_mac, packets.src_mac), dst_mac = COALESCE(EXCLUDED.dst_mac, packets.dst_mac), eth_type_raw = COALESCE(EXCLUDED.eth_type_raw, packets.eth_type_raw), eth_type = COALESCE(EXCLUDED.eth_type, packets.eth_type), vlan_id = COALESCE(EXCLUDED.vlan_id, packets.vlan_id), src_ip = COALESCE(EXCLUDED.src_ip, packets.src_ip), dst_ip = COALESCE(EXCLUDED.dst_ip, packets.dst_ip), ip_proto_raw = COALESCE(EXCLUDED.ip_proto_raw, packets.ip_proto_raw), ip_proto = COALESCE(EXCLUDED.ip_proto, packets.ip_proto), src_port = COALESCE(EXCLUDED.src_port, packets.src_port), dst_port = COALESCE(EXCLUDED.dst_port, packets.dst_port), length = COALESCE(EXCLUDED.length, packets.length), raw_present = COALESCE(EXCLUDED.raw_present, FALSE) OR COALESCE(packets.raw_present, FALSE), capture_sources = ( SELECT ARRAY( SELECT DISTINCT source FROM unnest( COALESCE(packets.capture_sources, ARRAY[]::text[]) || COALESCE(EXCLUDED.capture_sources, ARRAY[]::text[]) ) AS source ) ), app_protocol = COALESCE(EXCLUDED.app_protocol, packets.app_protocol), app_master_protocol = COALESCE(EXCLUDED.app_master_protocol, packets.app_master_protocol), app_category = COALESCE(EXCLUDED.app_category, packets.app_category), app_confidence = COALESCE(EXCLUDED.app_confidence, packets.app_confidence), app_hostname = COALESCE(EXCLUDED.app_hostname, packets.app_hostname), app_is_encrypted = COALESCE(EXCLUDED.app_is_encrypted, packets.app_is_encrypted), app_risk_score = COALESCE(EXCLUDED.app_risk_score, packets.app_risk_score), dpi_metadata = COALESCE(EXCLUDED.dpi_metadata, packets.dpi_metadata), capture_metadata = COALESCE(EXCLUDED.capture_metadata, packets.capture_metadata), telemetry_metadata = COALESCE(EXCLUDED.telemetry_metadata, packets.telemetry_metadata), raw = COALESCE(EXCLUDED.raw, packets.raw) RETURNING * """, pkt_info["correlation_key"], pkt_info.get("packet_id"), pkt_info.get("packet_uid"), pkt_info.get("correlation_source"), pkt_info.get("skb_mark"), pkt_info.get("capture_iface"), pkt_info.get("ingress_if"), pkt_info.get("egress_if"), pkt_info.get("verdict"), pkt_info.get("verdict_reason"), pkt_info.get("verdict_confidence"), pkt_info.get("ingress_seen_at"), pkt_info.get("egress_seen_at"), pkt_info.get("verdict_seen_at"), pkt_info.get("src_mac"), pkt_info.get("dst_mac"), pkt_info.get("eth_type_raw"), _db_text(pkt_info.get("eth_type")), pkt_info.get("vlan_id"), pkt_info.get("src_ip"), pkt_info.get("dst_ip"), pkt_info.get("protocol_raw"), _db_text(pkt_info.get("protocol_name") or pkt_info.get("protocol")), pkt_info.get("src_port"), pkt_info.get("dst_port"), pkt_info.get("length"), pkt_info.get("raw_present"), pkt_info.get("capture_sources"), pkt_info.get("app_protocol"), pkt_info.get("app_master_protocol"), pkt_info.get("app_category"), pkt_info.get("app_confidence"), pkt_info.get("app_hostname"), pkt_info.get("app_is_encrypted"), pkt_info.get("app_risk_score"), json.dumps(dpi_metadata) if dpi_metadata is not None else None, json.dumps(pkt_info.get("capture_metadata")) if pkt_info.get("capture_metadata") is not None else None, json.dumps(telemetry_metadata) if telemetry_metadata is not None else None, pkt_info.get("raw"), ) except Exception: logger.exception("DB upsert failed") return if row: persisted = _serialize_row_for_broadcast(dict(row)) pkt_info.update(persisted) if self.broadcaster: try: self.broadcaster.sync_publish(pkt_info) except Exception: logger.exception("Failed to publish pkt_info to broadcaster") async def backfill_packet_metadata( self, *, iface: str, src_ip: str, dst_ip: str, src_port: int, dst_port: int, protocol: int, length: int, observed_at_ms: int, enrichment: Dict[str, Any], window_ms: int, ) -> int: """Update recent packet rows after tshark metadata arrives asynchronously.""" if self._pool is None: await self.init_pool() lower_bound = datetime.fromtimestamp(max(observed_at_ms - window_ms, 0) / 1000.0, tz=timezone.utc) upper_bound = datetime.fromtimestamp(max(observed_at_ms + window_ms, 0) / 1000.0, tz=timezone.utc) dpi_metadata = enrichment.get("dpi_metadata") capture_sources = [str(source) for source in enrichment.get("capture_sources", []) if source] try: async with self._pool.acquire() as conn: rows = await conn.fetch( """ UPDATE packets SET updated_at = NOW(), app_protocol = COALESCE(packets.app_protocol, $10), app_master_protocol = COALESCE(packets.app_master_protocol, $11), app_category = COALESCE(packets.app_category, $12), app_confidence = COALESCE(packets.app_confidence, $13), app_hostname = COALESCE(packets.app_hostname, $14), app_is_encrypted = COALESCE(packets.app_is_encrypted, $15), dpi_metadata = CASE WHEN $16::jsonb IS NULL THEN packets.dpi_metadata WHEN packets.dpi_metadata IS NULL THEN $16::jsonb ELSE packets.dpi_metadata || $16::jsonb END, capture_sources = ( SELECT ARRAY( SELECT DISTINCT source FROM unnest( COALESCE(packets.capture_sources, ARRAY[]::text[]) || COALESCE($17::text[], ARRAY[]::text[]) ) AS source ) ) WHERE ip_proto_raw = $1 AND (capture_iface = $2 OR ingress_if = $2 OR egress_if = $2) AND src_ip = $3::inet AND dst_ip = $4::inet AND src_port = $5 AND dst_port = $6 AND length = $7 AND timestamp BETWEEN $8 AND $9 AND ( packets.app_protocol IS NULL OR packets.app_master_protocol IS NULL OR packets.app_category IS NULL OR packets.app_confidence IS NULL OR packets.app_hostname IS NULL OR packets.app_is_encrypted IS NULL OR ($16::jsonb IS NOT NULL) OR (COALESCE(array_length($17::text[], 1), 0) > 0) ) RETURNING * """, protocol, iface, src_ip, dst_ip, src_port, dst_port, length, lower_bound, upper_bound, enrichment.get("app_protocol"), enrichment.get("app_master_protocol"), enrichment.get("app_category"), enrichment.get("app_confidence"), enrichment.get("app_hostname"), enrichment.get("app_is_encrypted"), json.dumps(dpi_metadata) if dpi_metadata is not None else None, capture_sources if capture_sources else None, ) except Exception: logger.exception("DB packet metadata backfill failed") return 0 if not rows: return 0 updated_count = 0 for row in rows: serialized = _serialize_row_for_broadcast(dict(row)) updated_count += 1 if self.broadcaster: try: self.broadcaster.sync_publish(serialized) except Exception: logger.exception("Failed to publish tshark-enriched packet row") return updated_count async def fetch_latest(self, limit: int) -> List[PacketDBModel]: """Fetch newest packet rows as validated `PacketDBModel` instances.""" if self._pool is None: await self.init_pool() async with self._pool.acquire() as conn: rows = await conn.fetch( """ SELECT * FROM packets ORDER BY updated_at DESC, id DESC LIMIT $1 """, limit, ) result: List[PacketDBModel] = [] for row in rows: data = dict(row) _normalize_json_fields(data) raw_val = data.get("raw") if isinstance(raw_val, (bytes, bytearray)): data["raw_b64"] = base64.b64encode(raw_val).decode("ascii") data.pop("raw", None) try: packet_model = PacketDBModel(**data) except ValidationError as exc: logger.warning( "Skipping DB row that failed PacketDBModel validation (id=%s): %s", data.get("id"), exc, ) continue result.append(packet_model) return result async def clear_all_packets(self, reset_identity: bool = True) -> bool: """Truncate the packet table and optionally reset identity counters.""" if self._pool is None: await self.init_pool() restart_clause = "RESTART IDENTITY" if reset_identity else "" query = f"TRUNCATE TABLE packets {restart_clause};" try: async with self._pool.acquire() as conn: await conn.execute(query) logger.info("Successfully cleared all packets from the database (reset_id=%s)", reset_identity) return True except Exception: logger.exception("Failed to clear packets table") return False