Files
mitm-webserver/backend/src/utilities/database.py
malmert d2179b2813
All checks were successful
Build and Deploy MITM Webserver / traffic_target (push) Successful in 0s
Build and Deploy MITM Webserver / build (push) Successful in 10s
test fix timestampt to capture time not db insert time
2026-03-09 14:33:43 +01:00

603 lines
26 KiB
Python

"""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)
_attach_derived_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
def _derive_flow_id(payload: Dict[str, Any]) -> Optional[str]:
dpi_metadata = payload.get("dpi_metadata")
if not isinstance(dpi_metadata, dict):
return None
tcp_meta = dpi_metadata.get("tcp")
if isinstance(tcp_meta, dict):
stream = tcp_meta.get("stream")
if stream not in (None, "", []):
return f"tcp:{stream}"
udp_meta = dpi_metadata.get("udp")
if isinstance(udp_meta, dict):
stream = udp_meta.get("stream")
if stream not in (None, "", []):
return f"udp:{stream}"
tshark_meta = dpi_metadata.get("tshark")
if isinstance(tshark_meta, dict):
tcp_stream = tshark_meta.get("tcp_stream")
if tcp_stream not in (None, "", []):
return f"tcp:{tcp_stream}"
udp_stream = tshark_meta.get("udp_stream")
if udp_stream not in (None, "", []):
return f"udp:{udp_stream}"
return None
def _attach_derived_fields(payload: Dict[str, Any]) -> None:
if payload.get("flow_id") in (None, ""):
flow_id = _derive_flow_id(payload)
if flow_id is not None:
payload["flow_id"] = flow_id
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 (
timestamp,
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,$37::jsonb,
$38::jsonb,$39::jsonb,$40
)
ON CONFLICT (correlation_key) DO UPDATE SET
updated_at = NOW(),
timestamp = CASE
WHEN packets.timestamp IS NULL THEN EXCLUDED.timestamp
WHEN EXCLUDED.timestamp IS NULL THEN packets.timestamp
ELSE LEAST(packets.timestamp, EXCLUDED.timestamp)
END,
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.get("timestamp"),
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,
eth_type_raw: Optional[int],
src_ip: str,
dst_ip: str,
src_port: int,
dst_port: int,
protocol: Optional[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, $11),
app_master_protocol = COALESCE(packets.app_master_protocol, $12),
app_category = COALESCE(packets.app_category, $13),
app_confidence = COALESCE(packets.app_confidence, $14),
app_hostname = COALESCE(packets.app_hostname, $15),
app_is_encrypted = COALESCE(packets.app_is_encrypted, $16),
dpi_metadata = CASE
WHEN $17::jsonb IS NULL THEN packets.dpi_metadata
WHEN packets.dpi_metadata IS NULL THEN $17::jsonb
ELSE packets.dpi_metadata || $17::jsonb
END,
capture_sources = (
SELECT ARRAY(
SELECT DISTINCT source
FROM unnest(
COALESCE(packets.capture_sources, ARRAY[]::text[]) ||
COALESCE($18::text[], ARRAY[]::text[])
) AS source
)
)
WHERE
(
($1::int IS NOT NULL AND ip_proto_raw = $1)
OR
($1::int IS NULL AND $2::int IS NOT NULL AND eth_type_raw = $2)
)
AND (capture_iface = $3 OR ingress_if = $3 OR egress_if = $3)
AND src_ip = $4::inet
AND dst_ip = $5::inet
AND COALESCE(src_port, 0) = $6
AND COALESCE(dst_port, 0) = $7
AND length = $8
AND timestamp BETWEEN $9 AND $10
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 ($17::jsonb IS NOT NULL)
OR (COALESCE(array_length($18::text[], 1), 0) > 0)
)
RETURNING *
""",
protocol,
eth_type_raw,
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 backfill_stream_metadata(
self,
*,
iface: str,
protocol: int,
stream_kind: str,
stream_id: int,
observed_at_ms: int,
enrichment: Dict[str, Any],
window_ms: int,
) -> int:
"""Propagate tshark stream context across packets already tagged with the same stream id."""
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)
capture_sources = [str(source) for source in enrichment.get("capture_sources", []) if source]
stream_kind = str(stream_kind).lower()
if stream_kind not in {"tcp", "udp"}:
return 0
try:
async with self._pool.acquire() as conn:
rows = await conn.fetch(
"""
UPDATE packets
SET
updated_at = NOW(),
app_protocol = CASE
WHEN packets.app_protocol IS NULL OR packets.app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH')
THEN COALESCE($7, packets.app_protocol)
ELSE packets.app_protocol
END,
app_master_protocol = CASE
WHEN packets.app_master_protocol IS NULL OR packets.app_master_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH')
THEN COALESCE($8, packets.app_master_protocol)
ELSE packets.app_master_protocol
END,
app_category = CASE
WHEN packets.app_category IS NULL OR packets.app_category IN ('Transport', 'Network', 'Protocol')
THEN COALESCE($9, packets.app_category)
ELSE packets.app_category
END,
app_confidence = CASE
WHEN packets.app_protocol IS NULL OR packets.app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH')
THEN COALESCE($10, packets.app_confidence)
ELSE packets.app_confidence
END,
app_hostname = COALESCE(packets.app_hostname, $11),
app_is_encrypted = COALESCE(packets.app_is_encrypted, $12),
capture_sources = (
SELECT ARRAY(
SELECT DISTINCT source
FROM unnest(
COALESCE(packets.capture_sources, ARRAY[]::text[]) ||
COALESCE($13::text[], ARRAY[]::text[])
) AS source
)
)
WHERE
ip_proto_raw = $1
AND (capture_iface = $2 OR ingress_if = $2 OR egress_if = $2)
AND timestamp BETWEEN $5 AND $6
AND (
CASE
WHEN $3 = 'tcp' THEN COALESCE(packets.dpi_metadata -> 'tcp' ->> 'stream', '')
WHEN $3 = 'udp' THEN COALESCE(packets.dpi_metadata -> 'udp' ->> 'stream', '')
ELSE ''
END
) = $4
AND (
packets.app_protocol IS NULL
OR packets.app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH')
OR packets.app_master_protocol IS NULL
OR packets.app_master_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH')
OR packets.app_category IS NULL
OR packets.app_category IN ('Transport', 'Network', 'Protocol')
OR packets.app_confidence IS NULL
OR packets.app_hostname IS NULL
OR packets.app_is_encrypted IS NULL
OR (COALESCE(array_length($13::text[], 1), 0) > 0)
)
RETURNING *
""",
protocol,
iface,
stream_kind,
str(stream_id),
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"),
capture_sources if capture_sources else None,
)
except Exception:
logger.exception("DB stream 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 stream-context 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 timestamp DESC, id DESC
LIMIT $1
""",
limit,
)
result: List[PacketDBModel] = []
for row in rows:
data = dict(row)
_normalize_json_fields(data)
_attach_derived_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