All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 11s
- Deleted flow_identity.py and nfstream_flow_worker.py as they are no longer needed. - Removed nfstream_manager.py and its associated logic for managing NFStream workers. - Added tshark_manager.py to manage tshark packet enrichment and matching. - Updated setup_build_server.sh to include default environment variables for tshark. - Implemented packet signature generation and enrichment logic in the new TsharkManager class.
422 lines
18 KiB
Python
422 lines
18 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)
|
|
|
|
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")
|
|
|
|
try:
|
|
async with self._pool.acquire() as conn:
|
|
rows = await conn.fetch(
|
|
"""
|
|
UPDATE packets
|
|
SET
|
|
updated_at = NOW(),
|
|
app_protocol = COALESCE(packets.app_protocol, $9),
|
|
app_master_protocol = COALESCE(packets.app_master_protocol, $10),
|
|
app_category = COALESCE(packets.app_category, $11),
|
|
app_confidence = COALESCE(packets.app_confidence, $12),
|
|
app_hostname = COALESCE(packets.app_hostname, $13),
|
|
app_is_encrypted = COALESCE(packets.app_is_encrypted, $14),
|
|
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
|
|
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)
|
|
)
|
|
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,
|
|
)
|
|
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
|