Files
mitm-webserver/backend/src/utilities/database.py
malmert d9cb501879
All checks were successful
Build and Deploy MITM Webserver / traffic_target (push) Successful in 1s
Build and Deploy MITM Webserver / build (push) Successful in 10s
fix db query
2026-03-30 21:46:13 +02:00

995 lines
43 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.etherType import EtherTypeEnum, ethertype_from_int
from src.Models.ip_protocol import protocol_from_number
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 _analysis_protocol_name(app_protocol: Any, ip_proto_raw: Any, eth_type_raw: Any) -> str:
if app_protocol not in (None, ""):
return str(app_protocol)
if ip_proto_raw is not None:
try:
return str(protocol_from_number(int(ip_proto_raw)))
except Exception:
return f"IP_PROTO_{ip_proto_raw}"
if eth_type_raw is not None:
try:
return str(ethertype_from_int(int(eth_type_raw)))
except Exception:
return f"ETH_TYPE_{eth_type_raw}"
return "UNKNOWN"
def _serialize_row_for_broadcast(row: Dict[str, Any]) -> Dict[str, Any]:
serialized = dict(row)
_normalize_json_fields(serialized)
_attach_derived_fields(serialized)
serialized.pop("capture_session_id", None)
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]:
session_id = payload.get("capture_session_id")
session_prefix = f"{session_id}:" if session_id not in (None, "") else ""
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"{session_prefix}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"{session_prefix}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"{session_prefix}tcp:{tcp_stream}"
udp_stream = tshark_meta.get("udp_stream")
if udp_stream not in (None, "", []):
return f"{session_prefix}udp:{udp_stream}"
return None
def _attach_derived_fields(payload: Dict[str, Any]) -> None:
flow_id = _derive_flow_id(payload)
if flow_id is not None:
current_flow_id = payload.get("flow_id")
if current_flow_id in (None, ""):
payload["flow_id"] = flow_id
else:
session_id = payload.get("capture_session_id")
if session_id not in (None, "") and current_flow_id == flow_id.split(":", 1)[1]:
payload["flow_id"] = flow_id
if payload.get("raw_present") is None:
payload["raw_present"] = payload.get("raw") is not None
if payload.get("eth_type") in (None, "") and payload.get("eth_type_raw") is not None:
try:
payload["eth_type"] = ethertype_from_int(int(payload["eth_type_raw"]))
except Exception:
payload["eth_type"] = EtherTypeEnum.UNKNOWN
if payload.get("ip_proto") in (None, "") and payload.get("ip_proto_raw") is not None:
try:
payload["ip_proto"] = protocol_from_number(int(payload["ip_proto_raw"]))
except Exception:
payload["ip_proto"] = int(payload["ip_proto_raw"])
if payload.get("app_master_protocol") in (None, "") and payload.get("app_protocol") not in (None, ""):
payload["app_master_protocol"] = payload.get("app_protocol")
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()
_normalize_json_fields(pkt_info)
_attach_derived_fields(pkt_info)
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,
flow_id,
capture_session_id,
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,
vlan_id,
src_ip,
dst_ip,
ip_proto_raw,
src_port,
dst_port,
length,
capture_sources,
app_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::jsonb,$36::jsonb,$37::jsonb,$38
)
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),
flow_id = COALESCE(EXCLUDED.flow_id, packets.flow_id),
capture_session_id = COALESCE(EXCLUDED.capture_session_id, packets.capture_session_id),
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),
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),
src_port = COALESCE(EXCLUDED.src_port, packets.src_port),
dst_port = COALESCE(EXCLUDED.dst_port, packets.dst_port),
length = COALESCE(EXCLUDED.length, packets.length),
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_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("flow_id"),
pkt_info.get("capture_session_id"),
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"),
pkt_info.get("vlan_id"),
pkt_info.get("src_ip"),
pkt_info.get("dst_ip"),
pkt_info.get("protocol_raw"),
pkt_info.get("src_port"),
pkt_info.get("dst_port"),
pkt_info.get("length"),
pkt_info.get("capture_sources"),
pkt_info.get("app_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(
"""
WITH candidate_rows AS (
SELECT id
FROM packets
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 (
app_protocol IS NULL
OR app_category IS NULL
OR app_confidence IS NULL
OR app_hostname IS NULL
OR app_is_encrypted IS NULL
OR ($16::jsonb IS NOT NULL)
OR (COALESCE(array_length($17::text[], 1), 0) > 0)
)
ORDER BY id
FOR UPDATE SKIP LOCKED
)
UPDATE packets
SET
updated_at = NOW(),
app_protocol = COALESCE(packets.app_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,
flow_id = COALESCE(
packets.flow_id,
CASE
WHEN COALESCE($16::jsonb -> 'tcp' ->> 'stream', $16::jsonb -> 'tshark' ->> 'tcp_stream') IS NOT NULL
AND COALESCE($16::jsonb -> 'tcp' ->> 'stream', $16::jsonb -> 'tshark' ->> 'tcp_stream') <> '' THEN
COALESCE(NULLIF(packets.capture_session_id, '') || ':', '') ||
'tcp:' || COALESCE($16::jsonb -> 'tcp' ->> 'stream', $16::jsonb -> 'tshark' ->> 'tcp_stream')
WHEN COALESCE($16::jsonb -> 'udp' ->> 'stream', $16::jsonb -> 'tshark' ->> 'udp_stream') IS NOT NULL
AND COALESCE($16::jsonb -> 'udp' ->> 'stream', $16::jsonb -> 'tshark' ->> 'udp_stream') <> '' THEN
COALESCE(NULLIF(packets.capture_session_id, '') || ':', '') ||
'udp:' || COALESCE($16::jsonb -> 'udp' ->> 'stream', $16::jsonb -> 'tshark' ->> 'udp_stream')
ELSE NULL
END
),
capture_sources = (
SELECT ARRAY(
SELECT DISTINCT source
FROM unnest(
COALESCE(packets.capture_sources, ARRAY[]::text[]) ||
COALESCE($17::text[], ARRAY[]::text[])
) AS source
)
)
FROM candidate_rows
WHERE packets.id = candidate_rows.id
RETURNING packets.*
""",
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_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(
"""
WITH candidate_rows AS (
SELECT id
FROM packets
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(dpi_metadata -> 'tcp' ->> 'stream', '')
WHEN $3 = 'udp' THEN COALESCE(dpi_metadata -> 'udp' ->> 'stream', '')
ELSE ''
END
) = $4
AND (
app_protocol IS NULL
OR app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH')
OR app_category IS NULL
OR app_category IN ('Transport', 'Network', 'Protocol')
OR app_confidence IS NULL
OR app_hostname IS NULL
OR app_is_encrypted IS NULL
OR (COALESCE(array_length($12::text[], 1), 0) > 0)
)
ORDER BY id
FOR UPDATE SKIP LOCKED
)
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_category = CASE
WHEN packets.app_category IS NULL OR packets.app_category IN ('Transport', 'Network', 'Protocol')
THEN COALESCE($8, 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($9, packets.app_confidence)
ELSE packets.app_confidence
END,
app_hostname = COALESCE(packets.app_hostname, $10),
app_is_encrypted = COALESCE(packets.app_is_encrypted, $11),
flow_id = COALESCE(
packets.flow_id,
COALESCE(NULLIF(packets.capture_session_id, '') || ':', '') || $3 || ':' || $4
),
capture_sources = (
SELECT ARRAY(
SELECT DISTINCT source
FROM unnest(
COALESCE(packets.capture_sources, ARRAY[]::text[]) ||
COALESCE($12::text[], ARRAY[]::text[])
) AS source
)
)
FROM candidate_rows
WHERE packets.id = candidate_rows.id
RETURNING packets.*
""",
protocol,
iface,
stream_kind,
str(stream_id),
lower_bound,
upper_bound,
enrichment.get("app_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 infer_interface_hosts(
self,
*,
since: Optional[datetime] = None,
limit_per_interface: int = 100,
) -> List[Dict[str, Any]]:
"""Infer which IP/MAC endpoints are attached to each observed interface."""
if self._pool is None:
await self.init_pool()
try:
async with self._pool.acquire() as conn:
rows = await conn.fetch(
"""
WITH observations AS (
SELECT
ingress_if AS iface,
src_ip::text AS ip_address,
src_mac::text AS mac_address,
timestamp,
'source_on_ingress' AS evidence
FROM packets
WHERE ingress_if IS NOT NULL
AND (src_ip IS NOT NULL OR src_mac IS NOT NULL)
AND ($1::timestamptz IS NULL OR timestamp >= $1)
UNION ALL
SELECT
egress_if AS iface,
dst_ip::text AS ip_address,
dst_mac::text AS mac_address,
timestamp,
'destination_on_egress' AS evidence
FROM packets
WHERE egress_if IS NOT NULL
AND (dst_ip IS NOT NULL OR dst_mac IS NOT NULL)
AND ($1::timestamptz IS NULL OR timestamp >= $1)
),
filtered AS (
SELECT *
FROM observations
WHERE iface IS NOT NULL
AND COALESCE(mac_address, '') <> 'ff:ff:ff:ff:ff:ff'
AND (
COALESCE(ip_address, '') <> ''
OR COALESCE(mac_address, '') <> ''
)
),
aggregated AS (
SELECT
iface,
ip_address,
mac_address,
COUNT(*) AS packet_count,
MAX(timestamp) AS last_seen,
SUM(CASE WHEN evidence = 'source_on_ingress' THEN 1 ELSE 0 END) AS source_on_ingress_count,
SUM(CASE WHEN evidence = 'destination_on_egress' THEN 1 ELSE 0 END) AS destination_on_egress_count
FROM filtered
GROUP BY iface, ip_address, mac_address
),
ranked AS (
SELECT
*,
ROW_NUMBER() OVER (
PARTITION BY iface
ORDER BY packet_count DESC, last_seen DESC, ip_address, mac_address
) AS row_num
FROM aggregated
)
SELECT
iface,
ip_address,
mac_address,
packet_count,
last_seen,
source_on_ingress_count,
destination_on_egress_count
FROM ranked
WHERE row_num <= $2
ORDER BY iface, packet_count DESC, last_seen DESC, ip_address, mac_address
""",
since,
limit_per_interface,
)
except Exception:
logger.exception("DB interface-host analysis failed")
raise
grouped: Dict[str, List[Dict[str, Any]]] = {}
for row in rows:
record = dict(row)
iface = str(record.pop("iface"))
if hasattr(record.get("last_seen"), "isoformat"):
record["last_seen"] = record["last_seen"].isoformat()
grouped.setdefault(iface, []).append(record)
return [
{
"interface": iface,
"hosts": hosts,
}
for iface, hosts in sorted(grouped.items())
]
async def infer_interface_host_protocols(
self,
*,
since: Optional[datetime] = None,
limit_per_interface: int = 50,
limit_protocols_per_host: int = 12,
) -> List[Dict[str, Any]]:
"""Infer interface-host attachment and aggregate observed protocols and verdicts."""
if self._pool is None:
await self.init_pool()
try:
async with self._pool.acquire() as conn:
rows = await conn.fetch(
"""
WITH observations AS (
SELECT
ingress_if AS iface,
src_ip::text AS ip_address,
src_mac::text AS mac_address,
NULLIF(app_protocol::text, '') AS app_protocol_name,
ip_proto_raw,
eth_type_raw,
COALESCE(NULLIF(verdict::text, ''), 'unknown') AS verdict_name,
timestamp,
'source_on_ingress' AS evidence
FROM packets
WHERE ingress_if IS NOT NULL
AND (src_ip IS NOT NULL OR src_mac IS NOT NULL)
AND ($1::timestamptz IS NULL OR timestamp >= $1)
UNION ALL
SELECT
egress_if AS iface,
dst_ip::text AS ip_address,
dst_mac::text AS mac_address,
NULLIF(app_protocol::text, '') AS app_protocol_name,
ip_proto_raw,
eth_type_raw,
COALESCE(NULLIF(verdict::text, ''), 'unknown') AS verdict_name,
timestamp,
'destination_on_egress' AS evidence
FROM packets
WHERE egress_if IS NOT NULL
AND (dst_ip IS NOT NULL OR dst_mac IS NOT NULL)
AND ($1::timestamptz IS NULL OR timestamp >= $1)
),
filtered AS (
SELECT *
FROM observations
WHERE iface IS NOT NULL
AND COALESCE(mac_address, '') <> 'ff:ff:ff:ff:ff:ff'
AND (
COALESCE(ip_address, '') <> ''
OR COALESCE(mac_address, '') <> ''
)
),
host_aggregated AS (
SELECT
iface,
ip_address,
mac_address,
COUNT(*) AS packet_count,
MAX(timestamp) AS last_seen,
SUM(CASE WHEN evidence = 'source_on_ingress' THEN 1 ELSE 0 END) AS source_on_ingress_count,
SUM(CASE WHEN evidence = 'destination_on_egress' THEN 1 ELSE 0 END) AS destination_on_egress_count
FROM filtered
GROUP BY iface, ip_address, mac_address
),
selected_hosts AS (
SELECT *
FROM (
SELECT
*,
ROW_NUMBER() OVER (
PARTITION BY iface
ORDER BY packet_count DESC, last_seen DESC, ip_address, mac_address
) AS row_num
FROM host_aggregated
) ranked_hosts
WHERE row_num <= $2
),
protocol_aggregated AS (
SELECT
filtered.iface,
filtered.ip_address,
filtered.mac_address,
filtered.app_protocol_name,
filtered.ip_proto_raw,
filtered.eth_type_raw,
COUNT(*) AS packet_count,
MAX(filtered.timestamp) AS last_seen,
SUM(CASE WHEN filtered.verdict_name = 'accept' THEN 1 ELSE 0 END) AS accept_count,
SUM(CASE WHEN filtered.verdict_name = 'drop' THEN 1 ELSE 0 END) AS drop_count,
SUM(CASE WHEN filtered.verdict_name = 'reject' THEN 1 ELSE 0 END) AS reject_count,
SUM(CASE WHEN filtered.verdict_name NOT IN ('accept', 'drop', 'reject') THEN 1 ELSE 0 END) AS unknown_count
FROM filtered
INNER JOIN selected_hosts
ON selected_hosts.iface = filtered.iface
AND selected_hosts.ip_address IS NOT DISTINCT FROM filtered.ip_address
AND selected_hosts.mac_address IS NOT DISTINCT FROM filtered.mac_address
GROUP BY
filtered.iface,
filtered.ip_address,
filtered.mac_address,
filtered.app_protocol_name,
filtered.ip_proto_raw,
filtered.eth_type_raw
),
ranked_protocols AS (
SELECT *
FROM (
SELECT
*,
ROW_NUMBER() OVER (
PARTITION BY iface, ip_address, mac_address
ORDER BY
packet_count DESC,
last_seen DESC,
app_protocol_name NULLS LAST,
ip_proto_raw NULLS LAST,
eth_type_raw NULLS LAST
) AS row_num
FROM protocol_aggregated
) ranked
WHERE row_num <= $3
)
SELECT
selected_hosts.iface,
selected_hosts.ip_address,
selected_hosts.mac_address,
selected_hosts.packet_count AS host_packet_count,
selected_hosts.last_seen AS host_last_seen,
selected_hosts.source_on_ingress_count,
selected_hosts.destination_on_egress_count,
ranked_protocols.app_protocol_name,
ranked_protocols.ip_proto_raw,
ranked_protocols.eth_type_raw,
ranked_protocols.packet_count AS protocol_packet_count,
ranked_protocols.last_seen AS protocol_last_seen,
ranked_protocols.accept_count,
ranked_protocols.drop_count,
ranked_protocols.reject_count,
ranked_protocols.unknown_count
FROM selected_hosts
LEFT JOIN ranked_protocols
ON ranked_protocols.iface = selected_hosts.iface
AND ranked_protocols.ip_address IS NOT DISTINCT FROM selected_hosts.ip_address
AND ranked_protocols.mac_address IS NOT DISTINCT FROM selected_hosts.mac_address
ORDER BY
selected_hosts.iface,
selected_hosts.packet_count DESC,
selected_hosts.last_seen DESC,
selected_hosts.ip_address,
selected_hosts.mac_address,
ranked_protocols.packet_count DESC NULLS LAST,
ranked_protocols.last_seen DESC NULLS LAST,
ranked_protocols.app_protocol_name NULLS LAST,
ranked_protocols.ip_proto_raw NULLS LAST,
ranked_protocols.eth_type_raw NULLS LAST
""",
since,
limit_per_interface,
limit_protocols_per_host,
)
except Exception:
logger.exception("DB interface-host-protocol analysis failed")
raise
grouped: Dict[str, Dict[str, Dict[str, Any]]] = {}
for row in rows:
record = dict(row)
iface = str(record["iface"])
host_key = f"{record.get('ip_address') or 'no-ip'}|{record.get('mac_address') or 'no-mac'}"
iface_hosts = grouped.setdefault(iface, {})
host_record = iface_hosts.get(host_key)
if host_record is None:
host_record = {
"ip_address": record.get("ip_address"),
"mac_address": record.get("mac_address"),
"packet_count": int(record.get("host_packet_count") or 0),
"last_seen": record["host_last_seen"].isoformat() if hasattr(record.get("host_last_seen"), "isoformat") else record.get("host_last_seen"),
"source_on_ingress_count": int(record.get("source_on_ingress_count") or 0),
"destination_on_egress_count": int(record.get("destination_on_egress_count") or 0),
"protocols": [],
}
iface_hosts[host_key] = host_record
protocol_name = _analysis_protocol_name(
record.get("app_protocol_name"),
record.get("ip_proto_raw"),
record.get("eth_type_raw"),
)
if protocol_name not in (None, ""):
host_record["protocols"].append(
{
"protocol": str(protocol_name),
"packet_count": int(record.get("protocol_packet_count") or 0),
"last_seen": record["protocol_last_seen"].isoformat()
if hasattr(record.get("protocol_last_seen"), "isoformat")
else record.get("protocol_last_seen"),
"accept_count": int(record.get("accept_count") or 0),
"drop_count": int(record.get("drop_count") or 0),
"reject_count": int(record.get("reject_count") or 0),
"unknown_count": int(record.get("unknown_count") or 0),
}
)
result: List[Dict[str, Any]] = []
for iface, hosts in sorted(grouped.items()):
sorted_hosts = sorted(
hosts.values(),
key=lambda item: (-int(item.get("packet_count") or 0), str(item.get("last_seen") or ""), str(item.get("ip_address") or ""), str(item.get("mac_address") or "")),
)
result.append({"interface": iface, "hosts": sorted_hosts})
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