All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 10s
436 lines
18 KiB
Python
436 lines
18 KiB
Python
"""Aggregate AF_PACKET observations and kernel telemetry into one packet record."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
import threading
|
|
import time
|
|
from datetime import datetime, timezone
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
import src.shared_objects as shared_objects
|
|
from src.Models.etherType import EtherTypeEnum, ethertype_from_int
|
|
from src.Models.ip_protocol import protocol_from_number
|
|
from src.config import settings
|
|
from src.utilities.packet_identity import build_packet_uid
|
|
from src.utilities.packet_mark import packet_id_from_mark, verdict_from_mark
|
|
|
|
logger = logging.getLogger("packet_tracker")
|
|
|
|
|
|
def _utcnow() -> datetime:
|
|
return datetime.now(timezone.utc)
|
|
|
|
|
|
class PacketTracker:
|
|
"""Deduplicate packet observations and persist one upserted row per packet."""
|
|
|
|
def __init__(
|
|
self,
|
|
finalize_delay_seconds: float = 0.25,
|
|
retention_seconds: float = 10.0,
|
|
min_flush_interval_seconds: float = 0.05,
|
|
):
|
|
self._finalize_delay_seconds = finalize_delay_seconds
|
|
self._retention_seconds = retention_seconds
|
|
self._min_flush_interval_seconds = min_flush_interval_seconds
|
|
self._entries: Dict[str, Dict[str, Any]] = {}
|
|
self._stats: Dict[str, int] = {
|
|
"persisted_total": 0,
|
|
"persisted_capture_only": 0,
|
|
"persisted_telemetry_only": 0,
|
|
"persisted_merged": 0,
|
|
"persisted_with_raw": 0,
|
|
"persisted_without_raw": 0,
|
|
"persisted_kernel_mark": 0,
|
|
"persisted_legacy_hash": 0,
|
|
}
|
|
self._lock = threading.Lock()
|
|
self._stop_event = threading.Event()
|
|
self._thread = threading.Thread(target=self._run, daemon=True, name="packet-tracker")
|
|
self._thread.start()
|
|
|
|
def stop(self) -> None:
|
|
self._stop_event.set()
|
|
self._thread.join(timeout=settings.packet_tracker_stop_join_timeout_seconds)
|
|
|
|
def observe_packet(self, pkt_info: Dict[str, Any]) -> str:
|
|
"""Merge parsed packet information into a pending packet entry."""
|
|
now_ts = time.time()
|
|
correlation_key = self._ensure_correlation(pkt_info)
|
|
pkt_info["raw_present"] = pkt_info.get("raw") is not None
|
|
pkt_info["capture_sources"] = [pkt_info.get("capture_source") or "af_packet"]
|
|
|
|
with self._lock:
|
|
entry = self._entries.get(correlation_key)
|
|
if entry is None:
|
|
entry = self._new_entry(correlation_key, now_ts)
|
|
self._entries[correlation_key] = entry
|
|
|
|
self._merge_packet_info(entry, pkt_info, now_ts)
|
|
self._maybe_promote_reject_from_reply(pkt_info, now_ts)
|
|
self._maybe_mark_complete(entry)
|
|
return correlation_key
|
|
|
|
def observe_telemetry(self, event: Dict[str, Any]) -> Optional[str]:
|
|
"""Merge ingress/egress/verdict telemetry into a pending packet entry."""
|
|
correlation_key = self._ensure_correlation(event)
|
|
if not correlation_key:
|
|
logger.debug("Telemetry event missing packet identity: %s", event)
|
|
return None
|
|
|
|
now_ts = time.time()
|
|
with self._lock:
|
|
entry = self._entries.get(correlation_key)
|
|
if entry is None:
|
|
entry = self._new_entry(correlation_key, now_ts)
|
|
self._entries[correlation_key] = entry
|
|
|
|
payload = entry["payload"]
|
|
payload["correlation_key"] = correlation_key
|
|
payload["packet_id"] = event.get("packet_id") or payload.get("packet_id")
|
|
payload["packet_uid"] = event.get("packet_uid") or payload.get("packet_uid")
|
|
payload["correlation_source"] = event.get("correlation_source") or payload.get("correlation_source")
|
|
payload["skb_mark"] = event.get("skb_mark") or payload.get("skb_mark")
|
|
payload["telemetry_metadata"] = event
|
|
payload["last_observed_at"] = now_ts
|
|
self._add_capture_source(payload, "telemetry")
|
|
for key, value in event.items():
|
|
if value is None or key in {"event_type", "reason", "reason_code", "iface", "packet_uid", "correlation_key"}:
|
|
continue
|
|
if payload.get(key) is None:
|
|
payload[key] = value
|
|
|
|
if payload.get("eth_type") is 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("protocol") is None and payload.get("protocol_raw") is not None:
|
|
try:
|
|
payload["protocol"] = protocol_from_number(int(payload["protocol_raw"]))
|
|
except Exception:
|
|
payload["protocol"] = int(payload["protocol_raw"])
|
|
event_type = event.get("event_type")
|
|
iface = event.get("iface")
|
|
verdict_hint = event.get("verdict_hint")
|
|
|
|
if event_type == "ingress":
|
|
payload["ingress_if"] = iface
|
|
payload["ingress_seen_at"] = _utcnow()
|
|
elif event_type == "egress":
|
|
payload["egress_if"] = iface
|
|
payload["egress_seen_at"] = _utcnow()
|
|
payload["verdict"] = "accept"
|
|
payload["verdict_reason"] = "egress-observed"
|
|
payload["verdict_confidence"] = "high"
|
|
payload["verdict_seen_at"] = _utcnow()
|
|
elif event_type == "drop":
|
|
payload["verdict"] = verdict_hint or "drop"
|
|
payload["verdict_reason"] = event.get("reason") or ("mark-verdict" if verdict_hint else "kfree_skb")
|
|
payload["verdict_confidence"] = "high"
|
|
payload["verdict_seen_at"] = _utcnow()
|
|
elif event_type == "reject" or verdict_hint == "reject":
|
|
payload["verdict"] = "reject"
|
|
payload["verdict_reason"] = event.get("reason") or "netfilter-reject"
|
|
payload["verdict_confidence"] = event.get("verdict_confidence") or "medium"
|
|
payload["verdict_seen_at"] = _utcnow()
|
|
|
|
entry["last_observed_at"] = now_ts
|
|
entry["dirty"] = True
|
|
self._maybe_mark_complete(entry)
|
|
return correlation_key
|
|
|
|
def _new_entry(self, correlation_key: str, now_ts: float) -> Dict[str, Any]:
|
|
return {
|
|
"correlation_key": correlation_key,
|
|
"payload": {
|
|
"correlation_key": correlation_key,
|
|
"correlation_source": None,
|
|
"packet_id": None,
|
|
"packet_uid": None,
|
|
"skb_mark": None,
|
|
"verdict": "pending",
|
|
"verdict_reason": None,
|
|
"verdict_confidence": None,
|
|
"raw_present": False,
|
|
"capture_sources": [],
|
|
"capture_metadata": None,
|
|
"telemetry_metadata": None,
|
|
},
|
|
"persisted": False,
|
|
"dirty": True,
|
|
"finalized": False,
|
|
"stats_recorded": False,
|
|
"created_at": now_ts,
|
|
"last_observed_at": now_ts,
|
|
"last_persisted_at": 0.0,
|
|
}
|
|
|
|
def _ensure_correlation(self, payload: Dict[str, Any]) -> Optional[str]:
|
|
skb_mark = payload.get("skb_mark")
|
|
if payload.get("packet_id") is None and skb_mark is not None:
|
|
payload["packet_id"] = packet_id_from_mark(skb_mark)
|
|
if payload.get("verdict_hint") is None and skb_mark is not None:
|
|
payload["verdict_hint"] = verdict_from_mark(skb_mark)
|
|
|
|
packet_id = payload.get("packet_id")
|
|
if packet_id not in (None, ""):
|
|
packet_id = str(packet_id)
|
|
payload["packet_id"] = packet_id
|
|
payload["correlation_key"] = f"pid:{packet_id}"
|
|
if not payload.get("correlation_source"):
|
|
payload["correlation_source"] = "kernel_mark"
|
|
return payload["correlation_key"]
|
|
|
|
packet_uid = payload.get("packet_uid")
|
|
if packet_uid in (None, ""):
|
|
try:
|
|
packet_uid = build_packet_uid(payload)
|
|
except Exception:
|
|
return None
|
|
payload["packet_uid"] = packet_uid
|
|
payload["correlation_key"] = f"uid:{packet_uid}"
|
|
if not payload.get("correlation_source"):
|
|
payload["correlation_source"] = "legacy_hash"
|
|
return payload["correlation_key"]
|
|
|
|
def _add_capture_source(self, payload: Dict[str, Any], source: str) -> None:
|
|
capture_sources = payload.setdefault("capture_sources", [])
|
|
if source not in capture_sources:
|
|
capture_sources.append(source)
|
|
|
|
def _merge_packet_info(self, entry: Dict[str, Any], pkt_info: Dict[str, Any], now_ts: float) -> None:
|
|
payload = entry["payload"]
|
|
changed = False
|
|
self._add_capture_source(payload, pkt_info.get("capture_source") or "af_packet")
|
|
for key, value in pkt_info.items():
|
|
if key in {"iface", "capture_source"}:
|
|
continue
|
|
if value is None:
|
|
continue
|
|
if key == "raw" and payload.get("raw") is not None:
|
|
continue
|
|
if payload.get(key) == value:
|
|
continue
|
|
payload[key] = value
|
|
changed = True
|
|
|
|
iface = pkt_info.get("iface")
|
|
if iface and pkt_info.get("capture_iface") and not payload.get("capture_iface"):
|
|
payload["capture_iface"] = pkt_info.get("capture_iface")
|
|
changed = True
|
|
elif iface and pkt_info.get("capture_metadata") and not payload.get("capture_iface"):
|
|
payload["capture_iface"] = iface
|
|
changed = True
|
|
|
|
if iface and not pkt_info.get("capture_metadata") and not payload.get("ingress_if"):
|
|
payload["ingress_if"] = iface
|
|
payload["ingress_seen_at"] = _utcnow()
|
|
changed = True
|
|
|
|
payload["last_observed_at"] = now_ts
|
|
entry["last_observed_at"] = now_ts
|
|
entry["dirty"] = entry["dirty"] or changed
|
|
|
|
def _maybe_promote_reject_from_reply(self, pkt_info: Dict[str, Any], now_ts: float) -> None:
|
|
reject_reason = None
|
|
|
|
tcp_flags = pkt_info.get("tcp_flags")
|
|
if pkt_info.get("protocol_raw") == 6 and tcp_flags is not None and int(tcp_flags) & 0x04:
|
|
reject_reason = "tcp-rst-observed"
|
|
match = self._find_recent_drop(
|
|
src_ip=pkt_info.get("dst_ip"),
|
|
dst_ip=pkt_info.get("src_ip"),
|
|
protocol_raw=6,
|
|
src_port=pkt_info.get("dst_port"),
|
|
dst_port=pkt_info.get("src_port"),
|
|
)
|
|
elif pkt_info.get("protocol_raw") == 1 and pkt_info.get("icmp_type") == 3:
|
|
reject_reason = "icmp-unreachable-observed"
|
|
match = self._find_recent_drop(
|
|
src_ip=pkt_info.get("icmp_embedded_src_ip"),
|
|
dst_ip=pkt_info.get("icmp_embedded_dst_ip"),
|
|
protocol_raw=pkt_info.get("icmp_embedded_protocol"),
|
|
src_port=pkt_info.get("icmp_embedded_src_port"),
|
|
dst_port=pkt_info.get("icmp_embedded_dst_port"),
|
|
)
|
|
else:
|
|
return
|
|
|
|
if match is None:
|
|
return
|
|
|
|
payload = match["payload"]
|
|
if payload.get("verdict") != "drop":
|
|
return
|
|
|
|
payload["verdict"] = "reject"
|
|
payload["verdict_reason"] = reject_reason
|
|
payload["verdict_confidence"] = "medium"
|
|
payload["verdict_seen_at"] = _utcnow()
|
|
match["last_observed_at"] = now_ts
|
|
match["dirty"] = True
|
|
match["finalized"] = True
|
|
|
|
def _find_recent_drop(
|
|
self,
|
|
src_ip: Any,
|
|
dst_ip: Any,
|
|
protocol_raw: Any,
|
|
src_port: Any,
|
|
dst_port: Any,
|
|
) -> Optional[Dict[str, Any]]:
|
|
if not src_ip or not dst_ip or protocol_raw is None:
|
|
return None
|
|
|
|
cutoff = time.time() - settings.packet_tracker_reject_correlation_window_seconds
|
|
for entry in self._entries.values():
|
|
payload = entry["payload"]
|
|
if entry["last_observed_at"] < cutoff:
|
|
continue
|
|
if payload.get("verdict") != "drop":
|
|
continue
|
|
if payload.get("src_ip") != src_ip or payload.get("dst_ip") != dst_ip:
|
|
continue
|
|
if payload.get("protocol_raw") != protocol_raw:
|
|
continue
|
|
if payload.get("src_port") != src_port or payload.get("dst_port") != dst_port:
|
|
continue
|
|
return entry
|
|
return None
|
|
|
|
def _maybe_mark_complete(self, entry: Dict[str, Any]) -> None:
|
|
payload = entry["payload"]
|
|
if payload.get("verdict") in {"accept", "drop", "reject"}:
|
|
entry["finalized"] = True
|
|
|
|
def _run(self) -> None:
|
|
while not self._stop_event.is_set():
|
|
time.sleep(0.05)
|
|
due_entries: List[Dict[str, Any]] = []
|
|
expired_keys: List[str] = []
|
|
now_ts = time.time()
|
|
|
|
with self._lock:
|
|
for correlation_key, entry in list(self._entries.items()):
|
|
age = now_ts - entry["last_observed_at"]
|
|
if not entry["finalized"] and age >= self._finalize_delay_seconds:
|
|
entry["finalized"] = True
|
|
if entry["payload"].get("verdict") == "pending":
|
|
entry["payload"]["verdict"] = "unknown"
|
|
entry["payload"]["verdict_reason"] = "timeout"
|
|
entry["payload"]["verdict_confidence"] = "low"
|
|
entry["payload"]["verdict_seen_at"] = _utcnow()
|
|
entry["dirty"] = True
|
|
|
|
should_flush = entry["dirty"] and (
|
|
not entry["persisted"]
|
|
or entry["finalized"]
|
|
or (now_ts - entry["last_persisted_at"]) >= self._min_flush_interval_seconds
|
|
)
|
|
if should_flush:
|
|
due_entries.append(
|
|
{
|
|
"correlation_key": entry["correlation_key"],
|
|
"payload": dict(entry["payload"]),
|
|
}
|
|
)
|
|
elif entry["persisted"] and age >= self._retention_seconds:
|
|
if entry["finalized"] and not entry["stats_recorded"]:
|
|
self._record_stats(entry["payload"])
|
|
entry["stats_recorded"] = True
|
|
expired_keys.append(correlation_key)
|
|
|
|
for correlation_key in expired_keys:
|
|
self._entries.pop(correlation_key, None)
|
|
|
|
for entry in due_entries:
|
|
self._persist(entry)
|
|
|
|
def _persist(self, entry: Dict[str, Any]) -> None:
|
|
payload = dict(entry["payload"])
|
|
|
|
web_loop = getattr(shared_objects, "web_loop", None)
|
|
web_db = getattr(shared_objects, "db", None)
|
|
if web_loop is None or web_db is None:
|
|
return
|
|
|
|
try:
|
|
fut = asyncio.run_coroutine_threadsafe(web_db.upsert_packet(payload), web_loop)
|
|
fut.result(timeout=settings.packet_tracker_persist_timeout_seconds)
|
|
with self._lock:
|
|
current = self._entries.get(entry["correlation_key"])
|
|
if current is not None:
|
|
current["persisted"] = True
|
|
current["dirty"] = False
|
|
current["last_persisted_at"] = time.time()
|
|
except Exception:
|
|
logger.exception("Failed to persist packet %s", entry["correlation_key"])
|
|
|
|
def _record_stats(self, payload: Dict[str, Any]) -> None:
|
|
capture_sources = set(payload.get("capture_sources") or [])
|
|
self._stats["persisted_total"] += 1
|
|
if payload.get("raw_present"):
|
|
self._stats["persisted_with_raw"] += 1
|
|
else:
|
|
self._stats["persisted_without_raw"] += 1
|
|
if payload.get("correlation_source") == "kernel_mark":
|
|
self._stats["persisted_kernel_mark"] = self._stats.get("persisted_kernel_mark", 0) + 1
|
|
else:
|
|
self._stats["persisted_legacy_hash"] = self._stats.get("persisted_legacy_hash", 0) + 1
|
|
|
|
if capture_sources and capture_sources != {"telemetry"} and "telemetry" not in capture_sources:
|
|
self._stats["persisted_capture_only"] += 1
|
|
elif capture_sources == {"telemetry"}:
|
|
self._stats["persisted_telemetry_only"] += 1
|
|
else:
|
|
self._stats["persisted_merged"] += 1
|
|
|
|
def get_debug_snapshot(self) -> Dict[str, Any]:
|
|
with self._lock:
|
|
active_entries = list(self._entries.values())
|
|
stats = dict(self._stats)
|
|
|
|
active_total = len(active_entries)
|
|
active_with_raw = sum(1 for entry in active_entries if entry["payload"].get("raw_present"))
|
|
active_without_raw = active_total - active_with_raw
|
|
active_capture_only = 0
|
|
active_telemetry_only = 0
|
|
active_merged = 0
|
|
active_kernel_mark = 0
|
|
active_legacy_hash = 0
|
|
for entry in active_entries:
|
|
capture_sources = set(entry["payload"].get("capture_sources") or [])
|
|
if capture_sources and capture_sources != {"telemetry"} and "telemetry" not in capture_sources:
|
|
active_capture_only += 1
|
|
elif capture_sources == {"telemetry"}:
|
|
active_telemetry_only += 1
|
|
else:
|
|
active_merged += 1
|
|
if entry["payload"].get("correlation_source") == "kernel_mark":
|
|
active_kernel_mark += 1
|
|
else:
|
|
active_legacy_hash += 1
|
|
|
|
return {
|
|
"active_total": active_total,
|
|
"active_with_raw": active_with_raw,
|
|
"active_without_raw": active_without_raw,
|
|
"active_capture_only": active_capture_only,
|
|
"active_telemetry_only": active_telemetry_only,
|
|
"active_merged": active_merged,
|
|
"active_kernel_mark": active_kernel_mark,
|
|
"active_legacy_hash": active_legacy_hash,
|
|
"cumulative": stats,
|
|
}
|
|
|
|
|
|
packet_tracker = PacketTracker(
|
|
finalize_delay_seconds=settings.packet_tracker_finalize_delay_seconds,
|
|
retention_seconds=settings.packet_tracker_retention_seconds,
|
|
min_flush_interval_seconds=settings.packet_tracker_min_flush_interval_seconds,
|
|
)
|