diff --git a/backend/src/config.py b/backend/src/config.py index 0122ed0..0cef55e 100644 --- a/backend/src/config.py +++ b/backend/src/config.py @@ -43,6 +43,11 @@ class BackendSettings: packet_tracker_retention_seconds: float packet_tracker_min_flush_interval_seconds: float packet_tracker_persist_timeout_seconds: float + packet_tracker_persist_retry_backoff_seconds: float + packet_tracker_persist_retry_backoff_max_seconds: float + packet_tracker_error_log_interval_seconds: float + packet_tracker_flush_batch_size: int + packet_tracker_max_entries: int packet_tracker_stop_join_timeout_seconds: float packet_tracker_reject_correlation_window_seconds: float sniffer_buffer_capacity: int @@ -52,6 +57,12 @@ class BackendSettings: sniffer_buffer_drain_interval_seconds: float sniffer_thread_join_timeout_seconds: float bridge_bpf_build_dir: str + bridge_telemetry_raw_sample_every: int + bridge_telemetry_meta_sample_every: int + bridge_telemetry_ingress_perf_pages: int + bridge_telemetry_meta_perf_pages: int + bridge_telemetry_event_queue_maxsize: int + bridge_telemetry_drop_log_interval_seconds: float bridge_link_state_thread_join_timeout_seconds: float bridge_link_state_failure_holdoff_seconds: float bridge_link_state_recovery_holdoff_seconds: float @@ -78,6 +89,17 @@ def load_settings() -> BackendSettings: packet_tracker_retention_seconds=_env_float("BACKEND_PACKET_TRACKER_RETENTION_SECONDS", 10.0), packet_tracker_min_flush_interval_seconds=_env_float("BACKEND_PACKET_TRACKER_MIN_FLUSH_INTERVAL_SECONDS", 0.05), packet_tracker_persist_timeout_seconds=_env_float("BACKEND_PACKET_TRACKER_PERSIST_TIMEOUT_SECONDS", 2.0), + packet_tracker_persist_retry_backoff_seconds=_env_float( + "BACKEND_PACKET_TRACKER_PERSIST_RETRY_BACKOFF_SECONDS", + 0.25, + ), + packet_tracker_persist_retry_backoff_max_seconds=_env_float( + "BACKEND_PACKET_TRACKER_PERSIST_RETRY_BACKOFF_MAX_SECONDS", + 5.0, + ), + packet_tracker_error_log_interval_seconds=_env_float("BACKEND_PACKET_TRACKER_ERROR_LOG_INTERVAL_SECONDS", 5.0), + packet_tracker_flush_batch_size=max(1, _env_int("BACKEND_PACKET_TRACKER_FLUSH_BATCH_SIZE", 500)), + packet_tracker_max_entries=max(1, _env_int("BACKEND_PACKET_TRACKER_MAX_ENTRIES", 50_000)), packet_tracker_stop_join_timeout_seconds=_env_float("BACKEND_PACKET_TRACKER_STOP_JOIN_TIMEOUT_SECONDS", 2.0), packet_tracker_reject_correlation_window_seconds=_env_float( "BACKEND_PACKET_TRACKER_REJECT_CORRELATION_WINDOW_SECONDS", @@ -90,6 +112,12 @@ def load_settings() -> BackendSettings: sniffer_buffer_drain_interval_seconds=_env_float("BACKEND_SNIFFER_BUFFER_DRAIN_INTERVAL_SECONDS", 5.0), sniffer_thread_join_timeout_seconds=_env_float("BACKEND_SNIFFER_THREAD_JOIN_TIMEOUT_SECONDS", 2.0), bridge_bpf_build_dir=_env_str("BACKEND_BRIDGE_BPF_BUILD_DIR", "/tmp/mitm-bpf"), + bridge_telemetry_raw_sample_every=max(0, _env_int("BACKEND_BRIDGE_TELEMETRY_RAW_SAMPLE_EVERY", 1)), + bridge_telemetry_meta_sample_every=max(0, _env_int("BACKEND_BRIDGE_TELEMETRY_META_SAMPLE_EVERY", 1)), + bridge_telemetry_ingress_perf_pages=max(1, _env_int("BACKEND_BRIDGE_TELEMETRY_INGRESS_PERF_PAGES", 256)), + bridge_telemetry_meta_perf_pages=max(1, _env_int("BACKEND_BRIDGE_TELEMETRY_META_PERF_PAGES", 128)), + bridge_telemetry_event_queue_maxsize=max(1, _env_int("BACKEND_BRIDGE_TELEMETRY_EVENT_QUEUE_MAXSIZE", 20_000)), + bridge_telemetry_drop_log_interval_seconds=_env_float("BACKEND_BRIDGE_TELEMETRY_DROP_LOG_INTERVAL_SECONDS", 5.0), bridge_link_state_thread_join_timeout_seconds=_env_float( "BACKEND_BRIDGE_LINK_STATE_THREAD_JOIN_TIMEOUT_SECONDS", 2.0, diff --git a/backend/src/utilities/bridge_telemetry.py b/backend/src/utilities/bridge_telemetry.py index e0a5593..28f1202 100644 --- a/backend/src/utilities/bridge_telemetry.py +++ b/backend/src/utilities/bridge_telemetry.py @@ -6,10 +6,12 @@ import base64 import json import logging import os +import queue import signal import subprocess import sys import threading +import time from pathlib import Path from typing import Iterable, Mapping, Optional @@ -27,6 +29,18 @@ class BridgeTelemetryManager: self._session_ids_by_interface: dict[str, tuple[str, ...]] = {} self._process: Optional[subprocess.Popen[str]] = None self._reader_thread: Optional[threading.Thread] = None + self._event_queue: queue.Queue = queue.Queue(maxsize=settings.bridge_telemetry_event_queue_maxsize) + self._worker_thread = threading.Thread( + target=self._event_worker_loop, + daemon=True, + name="bridge-telemetry-worker", + ) + self._worker_thread.start() + self._dropped_events = 0 + self._dropped_raw_payloads = 0 + self._last_drop_log_at = 0.0 + self._suppressed_collector_messages = 0 + self._last_collector_warning_at = 0.0 self._lock = threading.Lock() def update_sessions(self, session_interfaces: Mapping[str, Iterable[str]]) -> None: @@ -80,8 +94,21 @@ class BridgeTelemetryManager: ",".join(sorted(self._interfaces)), "--build-dir", settings.bridge_bpf_build_dir, + "--raw-sample-every", + str(settings.bridge_telemetry_raw_sample_every), + "--meta-sample-every", + str(settings.bridge_telemetry_meta_sample_every), + "--ingress-pages", + str(settings.bridge_telemetry_ingress_perf_pages), + "--meta-pages", + str(settings.bridge_telemetry_meta_perf_pages), ] - logger.info("Starting bridge telemetry collector for interfaces=%s", sorted(self._interfaces)) + logger.info( + "Starting bridge telemetry collector for interfaces=%s raw_sample_every=%s meta_sample_every=%s", + sorted(self._interfaces), + settings.bridge_telemetry_raw_sample_every, + settings.bridge_telemetry_meta_sample_every, + ) try: self._process = subprocess.Popen( cmd, @@ -158,6 +185,90 @@ class BridgeTelemetryManager: except Exception: logger.exception("Failed to process ingress raw packet event") + def _event_worker_loop(self) -> None: + while True: + event = self._event_queue.get() + try: + if not isinstance(event, dict): + continue + + if event.get("event_type") == "ingress": + self._handle_ingress_packet(event) + + try: + packet_tracker.observe_telemetry(event) + except Exception: + logger.exception("Failed to process telemetry event: %s", event) + finally: + self._event_queue.task_done() + + def _enqueue_event(self, event: dict[str, object]) -> None: + try: + self._event_queue.put_nowait(event) + return + except queue.Full: + pass + + raw_b64 = event.pop("raw_b64", None) + dropped_raw = isinstance(raw_b64, str) and bool(raw_b64) + + try: + self._event_queue.get_nowait() + self._event_queue.task_done() + except queue.Empty: + pass + + try: + self._event_queue.put_nowait(event) + except queue.Full: + with self._lock: + self._dropped_events += 1 + self._log_drop_summary() + return + + with self._lock: + self._dropped_events += 1 + if dropped_raw: + self._dropped_raw_payloads += 1 + self._log_drop_summary() + + def _log_drop_summary(self) -> None: + now_ts = time.time() + if now_ts - self._last_drop_log_at < settings.bridge_telemetry_drop_log_interval_seconds: + return + + self._last_drop_log_at = now_ts + with self._lock: + dropped_events = self._dropped_events + dropped_raw_payloads = self._dropped_raw_payloads + queue_size = self._event_queue.qsize() + + logger.warning( + "Bridge telemetry is overloaded; queue=%s dropped_events=%s dropped_raw_payloads=%s", + queue_size, + dropped_events, + dropped_raw_payloads, + ) + + def _log_collector_message(self, text: str) -> None: + lower_text = text.lower() + is_loss_message = "lost" in lower_text and "sample" in lower_text + if not is_loss_message: + logger.info("bridge-telemetry: %s", text) + return + + now_ts = time.time() + if now_ts - self._last_collector_warning_at < settings.bridge_telemetry_drop_log_interval_seconds: + with self._lock: + self._suppressed_collector_messages += 1 + return + + with self._lock: + suppressed = self._suppressed_collector_messages + self._suppressed_collector_messages = 0 + self._last_collector_warning_at = now_ts + logger.warning("bridge-telemetry: %s (suppressed similar messages=%s)", text, suppressed) + def _read_loop(self, process: subprocess.Popen[str]) -> None: stdout = process.stdout if stdout is None: @@ -170,20 +281,14 @@ class BridgeTelemetryManager: try: event = json.loads(text) except json.JSONDecodeError: - logger.info("bridge-telemetry: %s", text) + self._log_collector_message(text) continue if "event_type" not in event: logger.info("bridge-telemetry: %s", event) continue - if event.get("event_type") == "ingress": - self._handle_ingress_packet(event) - - try: - packet_tracker.observe_telemetry(event) - except Exception: - logger.exception("Failed to process telemetry event: %s", event) + self._enqueue_event(event) rc = process.poll() if rc not in (0, None): diff --git a/backend/src/utilities/ebpf_bridge_events.py b/backend/src/utilities/ebpf_bridge_events.py index a5f13b5..c7b5ce2 100644 --- a/backend/src/utilities/ebpf_bridge_events.py +++ b/backend/src/utilities/ebpf_bridge_events.py @@ -76,6 +76,8 @@ BPF_SOURCE = r""" #define EVENT_INGRESS 1 #define EVENT_EGRESS 2 #define EVENT_DROP 3 +#define RAW_SAMPLE_EVERY __RAW_SAMPLE_EVERY__ +#define META_SAMPLE_EVERY __META_SAMPLE_EVERY__ struct vlan_hdr_t { __be16 h_vlan_TCI; @@ -112,6 +114,16 @@ struct event_t { BPF_PERF_OUTPUT(ingress_events); BPF_PERF_OUTPUT(meta_events); +static __always_inline int should_emit_sample(__u32 packet_mark, __u32 every) { + if (every == 0) { + return 0; + } + if (every == 1) { + return 1; + } + return (packet_mark % every) == 0; +} + static __always_inline __u32 ensure_packet_mark(struct __sk_buff *skb) { __u32 next = skb->mark; @@ -346,7 +358,11 @@ int handle_ingress(struct __sk_buff *skb) { return TC_ACT_OK; } - ingress_events.perf_submit_skb(skb, skb->len, &event, sizeof(event)); + if (should_emit_sample(event.skb_mark, RAW_SAMPLE_EVERY)) { + ingress_events.perf_submit_skb(skb, skb->len, &event, sizeof(event)); + } else if (should_emit_sample(event.skb_mark, META_SAMPLE_EVERY)) { + meta_events.perf_submit(skb, &event, sizeof(event)); + } return TC_ACT_OK; } @@ -368,7 +384,9 @@ int handle_egress(struct __sk_buff *skb) { return TC_ACT_OK; } - meta_events.perf_submit(skb, &event, sizeof(event)); + if (should_emit_sample(event.skb_mark, META_SAMPLE_EVERY)) { + meta_events.perf_submit(skb, &event, sizeof(event)); + } return TC_ACT_OK; } @@ -397,7 +415,9 @@ TRACEPOINT_PROBE(skb, kfree_skb) { return 0; } - meta_events.perf_submit(args, &event, sizeof(event)); + if (should_emit_sample(event.skb_mark, META_SAMPLE_EVERY)) { + meta_events.perf_submit(args, &event, sizeof(event)); + } return 0; } """ @@ -539,10 +559,41 @@ def _emit_meta_event(cpu: int, data: int, size: int) -> None: print(json.dumps(payload, separators=(",", ":")), flush=True) +def _build_bpf_source(raw_sample_every: int, meta_sample_every: int) -> str: + return ( + BPF_SOURCE.replace("__RAW_SAMPLE_EVERY__", str(max(0, raw_sample_every))) + .replace("__META_SAMPLE_EVERY__", str(max(0, meta_sample_every))) + ) + + def _parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="tc/eBPF bridge telemetry collector") parser.add_argument("--ifaces", required=True, help="Comma-separated list of interfaces to instrument") parser.add_argument("--build-dir", required=True, help="Directory for compiled tc BPF objects") + parser.add_argument( + "--raw-sample-every", + type=int, + default=1, + help="Emit full raw ingress packets every Nth marked packet. Use 0 to disable raw packet export.", + ) + parser.add_argument( + "--meta-sample-every", + type=int, + default=1, + help="Emit metadata events every Nth marked packet. Use 0 to disable metadata-only events.", + ) + parser.add_argument( + "--ingress-pages", + type=int, + default=256, + help="Perf-buffer page count for raw ingress packet events.", + ) + parser.add_argument( + "--meta-pages", + type=int, + default=128, + help="Perf-buffer page count for metadata events.", + ) return parser.parse_args() @@ -621,6 +672,11 @@ def _cleanup_tc(ifaces: Iterable[str]) -> None: def main() -> int: args = _parse_args() + raw_sample_every = max(0, args.raw_sample_every) + meta_sample_every = max(0, args.meta_sample_every) + ingress_pages = max(1, args.ingress_pages) + meta_pages = max(1, args.meta_pages) + global TARGET_INTERFACES TARGET_INTERFACES = {iface.strip() for iface in args.ifaces.split(",") if iface.strip()} if not TARGET_INTERFACES: @@ -630,7 +686,7 @@ def main() -> int: signal.signal(signal.SIGTERM, _sigterm) signal.signal(signal.SIGINT, _sigterm) - bpf = BPF(text=BPF_SOURCE) + bpf = BPF(text=_build_bpf_source(raw_sample_every, meta_sample_every)) ingress_prog_name = "" egress_prog_name = "" try: @@ -643,14 +699,18 @@ def main() -> int: "ingress_program": _json_safe(ingress_prog_name), "egress_program": _json_safe(egress_prog_name), "build_dir": str(args.build_dir), + "raw_sample_every": raw_sample_every, + "meta_sample_every": meta_sample_every, + "ingress_pages": ingress_pages, + "meta_pages": meta_pages, }, separators=(",", ":"), ), flush=True, ) - bpf["ingress_events"].open_perf_buffer(_emit_ingress_event, page_cnt=256) - bpf["meta_events"].open_perf_buffer(_emit_meta_event, page_cnt=128) + bpf["ingress_events"].open_perf_buffer(_emit_ingress_event, page_cnt=ingress_pages) + bpf["meta_events"].open_perf_buffer(_emit_meta_event, page_cnt=meta_pages) while True: bpf.perf_buffer_poll() except KeyboardInterrupt: diff --git a/backend/src/utilities/packet_tracker.py b/backend/src/utilities/packet_tracker.py index de724a9..abb3150 100644 --- a/backend/src/utilities/packet_tracker.py +++ b/backend/src/utilities/packet_tracker.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import concurrent.futures import logging import threading import time @@ -174,7 +175,13 @@ class PacketTracker: "persisted_without_raw": 0, "persisted_kernel_mark": 0, "persisted_legacy_hash": 0, + "persist_failed_total": 0, + "persist_timeout_total": 0, + "persist_failure_log_suppressed": 0, + "evicted_persisted_total": 0, + "evicted_unpersisted_total": 0, } + self._last_persist_error_log_at = 0.0 self._lock = threading.Lock() self._stop_event = threading.Event() self._thread = threading.Thread(target=self._run, daemon=True, name="packet-tracker") @@ -233,6 +240,7 @@ class PacketTracker: self._merge_packet_info(entry, pkt_info, now_ts) self._maybe_promote_reject_from_reply(pkt_info, now_ts) self._maybe_mark_complete(entry) + self._enforce_entry_limit_locked() return correlation_key def observe_telemetry(self, event: Dict[str, Any]) -> Optional[str]: @@ -315,6 +323,7 @@ class PacketTracker: entry["last_observed_at"] = now_ts entry["dirty"] = True self._maybe_mark_complete(entry) + self._enforce_entry_limit_locked() return correlation_key def _new_entry(self, correlation_key: str, now_ts: float) -> Dict[str, Any]: @@ -344,8 +353,35 @@ class PacketTracker: "created_at": now_ts, "last_observed_at": now_ts, "last_persisted_at": 0.0, + "last_persist_attempt_at": 0.0, + "persist_failures": 0, } + def _enforce_entry_limit_locked(self) -> None: + overflow = len(self._entries) - settings.packet_tracker_max_entries + if overflow <= 0: + return + eviction_chunk = max(1, min(settings.packet_tracker_flush_batch_size, settings.packet_tracker_max_entries)) + evict_count = min(len(self._entries), max(overflow, eviction_chunk)) + + candidates = sorted( + self._entries.items(), + key=lambda item: ( + 0 if item[1].get("persisted") else 1, + 0 if item[1].get("finalized") else 1, + float(item[1].get("last_observed_at") or 0.0), + ), + ) + for correlation_key, entry in candidates[:evict_count]: + if entry.get("persisted"): + self._stats["evicted_persisted_total"] += 1 + if entry.get("finalized") and not entry.get("stats_recorded"): + self._record_stats(entry["payload"]) + entry["stats_recorded"] = True + else: + self._stats["evicted_unpersisted_total"] += 1 + self._entries.pop(correlation_key, None) + 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: @@ -635,13 +671,15 @@ class PacketTracker: 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"]), - } - ) + if should_flush and self._persist_backoff_elapsed(entry, now_ts): + if len(due_entries) < settings.packet_tracker_flush_batch_size: + due_entries.append( + { + "correlation_key": entry["correlation_key"], + "payload": dict(entry["payload"]), + } + ) + entry["last_persist_attempt_at"] = now_ts elif entry["persisted"] and age >= self._retention_seconds: if entry["finalized"] and not entry["stats_recorded"]: self._record_stats(entry["payload"]) @@ -654,6 +692,21 @@ class PacketTracker: for entry in due_entries: self._persist(entry) + def _persist_backoff_elapsed(self, entry: Dict[str, Any], now_ts: float) -> bool: + failures = int(entry.get("persist_failures") or 0) + if failures <= 0: + return True + + base = max(0.0, settings.packet_tracker_persist_retry_backoff_seconds) + if base <= 0: + return True + + backoff = min( + settings.packet_tracker_persist_retry_backoff_max_seconds, + base * (2 ** min(failures - 1, 6)), + ) + return now_ts - float(entry.get("last_persist_attempt_at") or 0.0) >= backoff + def _persist(self, entry: Dict[str, Any]) -> None: payload = dict(entry["payload"]) @@ -671,8 +724,34 @@ class PacketTracker: current["persisted"] = True current["dirty"] = False current["last_persisted_at"] = time.time() - except Exception: - logger.exception("Failed to persist packet %s", entry["correlation_key"]) + current["persist_failures"] = 0 + except Exception as exc: + try: + fut.cancel() + except Exception: + pass + + is_timeout = isinstance(exc, (TimeoutError, concurrent.futures.TimeoutError, asyncio.TimeoutError)) + with self._lock: + self._stats["persist_failed_total"] += 1 + if is_timeout: + self._stats["persist_timeout_total"] += 1 + current = self._entries.get(entry["correlation_key"]) + if current is not None: + current["dirty"] = True + current["persist_failures"] = int(current.get("persist_failures") or 0) + 1 + + now_ts = time.time() + if now_ts - self._last_persist_error_log_at >= settings.packet_tracker_error_log_interval_seconds: + self._last_persist_error_log_at = now_ts + logger.warning( + "Packet persistence is overloaded; failed to persist %s (%s). Further errors are rate-limited.", + entry["correlation_key"], + type(exc).__name__, + ) + else: + with self._lock: + self._stats["persist_failure_log_suppressed"] += 1 def _record_stats(self, payload: Dict[str, Any]) -> None: capture_sources = set(payload.get("capture_sources") or [])