diff --git a/backend/src/network_sniffer.py b/backend/src/network_sniffer.py index fa84fab..aeae75e 100644 --- a/backend/src/network_sniffer.py +++ b/backend/src/network_sniffer.py @@ -426,11 +426,12 @@ def parse_packet( pkt_info["icmp_embedded_src_port"] = _safe_get_attr(inner[UDP], "sport") pkt_info["icmp_embedded_dst_port"] = _safe_get_attr(inner[UDP], "dport") + pkt_info["packet_uid"] = build_packet_uid(pkt_info) + if pkt_info.get("packet_id"): pkt_info["correlation_key"] = f"pid:{pkt_info['packet_id']}" pkt_info["correlation_source"] = "kernel_mark" else: - pkt_info["packet_uid"] = build_packet_uid(pkt_info) pkt_info["correlation_key"] = f"uid:{pkt_info['packet_uid']}" pkt_info["correlation_source"] = "legacy_hash" @@ -541,16 +542,13 @@ def _ensure_socket_for_session(sockets: Dict[str, socket.socket], iface: str, br def _sync_bridge_telemetry() -> None: - interfaces = sorted( - { - iface - for session in sessions.values() - if session.get("is_bridge") - for iface in session.get("ports", []) - } - ) + bridge_session_interfaces = { + session_id: list(session.get("ports", [])) + for session_id, session in sessions.items() + if session.get("is_bridge") + } try: - bridge_telemetry_manager.update_interfaces(interfaces) + bridge_telemetry_manager.update_sessions(bridge_session_interfaces) except Exception: logger.exception("Failed to update bridge telemetry collector") diff --git a/backend/src/utilities/bridge_telemetry.py b/backend/src/utilities/bridge_telemetry.py index 7eaafeb..e0a5593 100644 --- a/backend/src/utilities/bridge_telemetry.py +++ b/backend/src/utilities/bridge_telemetry.py @@ -11,7 +11,7 @@ import subprocess import sys import threading from pathlib import Path -from typing import Iterable, Optional +from typing import Iterable, Mapping, Optional from src.config import settings from src.utilities.packet_tracker import packet_tracker @@ -24,17 +24,33 @@ class BridgeTelemetryManager: def __init__(self) -> None: self._interfaces: set[str] = set() + self._session_ids_by_interface: dict[str, tuple[str, ...]] = {} self._process: Optional[subprocess.Popen[str]] = None self._reader_thread: Optional[threading.Thread] = None self._lock = threading.Lock() - def update_interfaces(self, interfaces: Iterable[str]) -> None: + def update_sessions(self, session_interfaces: Mapping[str, Iterable[str]]) -> None: """Restart the collector when the active bridge interface set changes.""" - normalized = {iface.strip() for iface in interfaces if iface and iface.strip()} + normalized: dict[str, set[str]] = {} + for session_id, interfaces in session_interfaces.items(): + if not session_id: + continue + iface_set = {iface.strip() for iface in interfaces if iface and iface.strip()} + if iface_set: + normalized[session_id] = iface_set + + normalized_interfaces = sorted({iface for ifaces in normalized.values() for iface in ifaces}) + session_ids_by_interface = { + iface: tuple(sorted(session_id for session_id, ifaces in normalized.items() if iface in ifaces)) + for iface in normalized_interfaces + } + with self._lock: - if normalized == self._interfaces: + interfaces_changed = set(normalized_interfaces) != self._interfaces + self._interfaces = set(normalized_interfaces) + self._session_ids_by_interface = session_ids_by_interface + if not interfaces_changed: return - self._interfaces = normalized self._restart_locked() def stop(self) -> None: @@ -123,11 +139,22 @@ class BridgeTelemetryManager: "skb_mark": event.get("skb_mark"), "capture_mode": "tc_ingress", } + capture_iface = str(event.get("iface") or "") + capture_session_id: Optional[str] = None + with self._lock: + session_ids = self._session_ids_by_interface.get(capture_iface, ()) + if session_ids: + capture_session_id = session_ids[0] try: from src.network_sniffer import parse_packet_bytes - parse_packet_bytes(packet_bytes, str(event.get("iface")), capture_metadata=capture_metadata) + parse_packet_bytes( + packet_bytes, + capture_iface, + capture_metadata=capture_metadata, + capture_session_id=capture_session_id, + ) except Exception: logger.exception("Failed to process ingress raw packet event") diff --git a/backend/src/utilities/database.py b/backend/src/utilities/database.py index 53a740a..3606245 100644 --- a/backend/src/utilities/database.py +++ b/backend/src/utilities/database.py @@ -91,10 +91,15 @@ def _derive_flow_id(payload: Dict[str, Any]) -> Optional[str]: 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: + 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 @@ -172,6 +177,9 @@ class DatabasePool: 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") diff --git a/backend/src/utilities/packet_tracker.py b/backend/src/utilities/packet_tracker.py index 27db42e..87c4fd8 100644 --- a/backend/src/utilities/packet_tracker.py +++ b/backend/src/utilities/packet_tracker.py @@ -180,6 +180,7 @@ class PacketTracker: "correlation_source": None, "packet_id": None, "packet_uid": None, + "capture_session_id": None, "skb_mark": None, "verdict": "pending", "verdict_reason": None, @@ -205,6 +206,15 @@ class PacketTracker: if payload.get("verdict_hint") is None and skb_mark is not None: payload["verdict_hint"] = verdict_from_mark(skb_mark) + packet_uid = payload.get("packet_uid") + if packet_uid in (None, ""): + try: + packet_uid = build_packet_uid(payload) + except Exception: + packet_uid = None + if packet_uid not in (None, ""): + payload["packet_uid"] = packet_uid + packet_id = payload.get("packet_id") if packet_id not in (None, ""): packet_id = str(packet_id) @@ -214,13 +224,8 @@ class PacketTracker: 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 + return None payload["correlation_key"] = f"uid:{packet_uid}" if not payload.get("correlation_source"): payload["correlation_source"] = "legacy_hash"