"""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, )