"""Manage optional tshark packet enrichment workers and matching.""" from __future__ import annotations import asyncio import logging import os import signal import subprocess import threading import time from typing import Any, Dict, Iterable, List, Optional, Tuple import src.shared_objects as shared_objects from scapy.all import IP, IPv6, TCP, UDP # type: ignore from src.config import settings logger = logging.getLogger("tshark_manager") _FIELDS: List[str] = [ "frame.time_epoch", "frame.interface_name", "frame.len", "_ws.col.Protocol", "_ws.col.Info", "ip.src", "ipv6.src", "ip.dst", "ipv6.dst", "ip.proto", "ipv6.nxt", "tcp.srcport", "udp.srcport", "tcp.dstport", "udp.dstport", "frame.protocols", "http.request.method", "http.request.uri", "http.host", "http.user_agent", "http.response.code", "http.response.phrase", "http.server", "http.content_type", "tls.handshake.extensions_server_name", "tls.handshake.version", "dns.flags.response", "dns.qry.name", "dns.qry.type", "dns.resp.name", "dns.a", "dns.aaaa", "dns.cname", ] def _safe_text(value: Any) -> Optional[str]: if value is None: return None text = str(value).strip() return text if text else None def _safe_int(value: Any) -> Optional[int]: try: text = _safe_text(value) if text is None: return None return int(text) except Exception: return None def _safe_float(value: Any) -> Optional[float]: try: text = _safe_text(value) if text is None: return None return float(text) except Exception: return None def _safe_bool_flag(value: Any) -> Optional[bool]: text = _safe_text(value) if text is None: return None if text in {"1", "true", "True"}: return True if text in {"0", "false", "False"}: return False return None def _jsonable(metadata: Dict[str, Any]) -> Dict[str, Any]: out: Dict[str, Any] = {} for key, value in metadata.items(): if value in (None, "", [], {}): continue out[key] = value return out def _packet_signature( iface: str, protocol: int, src_ip: str, dst_ip: str, src_port: int, dst_port: int, length: int, ) -> Tuple[str, int, str, str, int, int, int]: return (str(iface or ""), int(protocol), str(src_ip), str(dst_ip), int(src_port), int(dst_port), int(length)) def _signature_from_packet(pkt: Any, iface: str) -> Optional[Tuple[str, int, str, str, int, int, int]]: if IP in pkt: ip_layer = pkt[IP] protocol = int(getattr(ip_layer, "proto", 0) or 0) src_ip = str(getattr(ip_layer, "src", "") or "") dst_ip = str(getattr(ip_layer, "dst", "") or "") elif IPv6 in pkt: ip_layer = pkt[IPv6] protocol = int(getattr(ip_layer, "nh", 0) or 0) src_ip = str(getattr(ip_layer, "src", "") or "") dst_ip = str(getattr(ip_layer, "dst", "") or "") else: return None src_port = 0 dst_port = 0 if TCP in pkt: src_port = int(getattr(pkt[TCP], "sport", 0) or 0) dst_port = int(getattr(pkt[TCP], "dport", 0) or 0) elif UDP in pkt: src_port = int(getattr(pkt[UDP], "sport", 0) or 0) dst_port = int(getattr(pkt[UDP], "dport", 0) or 0) if not src_ip or not dst_ip: return None return _packet_signature(iface, protocol, src_ip, dst_ip, src_port, dst_port, len(pkt)) def _packet_observed_at_ms(pkt: Any) -> int: try: packet_time = float(getattr(pkt, "time", 0.0) or 0.0) if packet_time > 0: return int(packet_time * 1000) except Exception: pass return int(time.time() * 1000) def _parse_line(line: str, fallback_iface: str) -> Optional[Dict[str, Any]]: parts = line.rstrip("\n").split("\t") if len(parts) < len(_FIELDS): parts.extend([""] * (len(_FIELDS) - len(parts))) row = dict(zip(_FIELDS, parts)) timestamp = _safe_float(row["frame.time_epoch"]) length = _safe_int(row["frame.len"]) protocol_col = _safe_text(row["_ws.col.Protocol"]) info_col = _safe_text(row["_ws.col.Info"]) protocol = _safe_int(row["ip.proto"]) or _safe_int(row["ipv6.nxt"]) src_ip = _safe_text(row["ip.src"]) or _safe_text(row["ipv6.src"]) dst_ip = _safe_text(row["ip.dst"]) or _safe_text(row["ipv6.dst"]) src_port = _safe_int(row["tcp.srcport"]) or _safe_int(row["udp.srcport"]) or 0 dst_port = _safe_int(row["tcp.dstport"]) or _safe_int(row["udp.dstport"]) or 0 iface = _safe_text(row["frame.interface_name"]) or fallback_iface if timestamp is None or length is None or protocol is None or src_ip is None or dst_ip is None: return None observed_at_ms = int(timestamp * 1000) return { "iface": iface, "observed_at_ms": observed_at_ms, "length": length, "protocol": protocol, "src_ip": src_ip, "dst_ip": dst_ip, "src_port": src_port, "dst_port": dst_port, "protocol_col": protocol_col, "info_col": info_col, "frame_protocols": _safe_text(row["frame.protocols"]), "http": _jsonable( { "method": _safe_text(row["http.request.method"]), "uri": _safe_text(row["http.request.uri"]), "host": _safe_text(row["http.host"]), "user_agent": _safe_text(row["http.user_agent"]), "response_code": _safe_int(row["http.response.code"]), "response_phrase": _safe_text(row["http.response.phrase"]), "server": _safe_text(row["http.server"]), "content_type": _safe_text(row["http.content_type"]), } ), "tls": _jsonable( { "server_name": _safe_text(row["tls.handshake.extensions_server_name"]), "handshake_version": _safe_text(row["tls.handshake.version"]), } ), "dns": _jsonable( { "is_response": _safe_bool_flag(row["dns.flags.response"]), "query_name": _safe_text(row["dns.qry.name"]), "query_type": _safe_text(row["dns.qry.type"]), "response_name": _safe_text(row["dns.resp.name"]), "a": _safe_text(row["dns.a"]), "aaaa": _safe_text(row["dns.aaaa"]), "cname": _safe_text(row["dns.cname"]), } ), } def _build_enrichment(event: Dict[str, Any]) -> Dict[str, Any]: http_meta = dict(event.get("http") or {}) tls_meta = dict(event.get("tls") or {}) dns_meta = dict(event.get("dns") or {}) protocols = str(event.get("frame_protocols") or "") protocol_col = _safe_text(event.get("protocol_col")) info_col = _safe_text(event.get("info_col")) app_protocol: Optional[str] = None app_category: Optional[str] = None app_hostname: Optional[str] = None app_is_encrypted: Optional[bool] = None if http_meta: app_protocol = "HTTP" app_category = "Web" app_hostname = http_meta.get("host") app_is_encrypted = False elif dns_meta: app_protocol = "DNS" app_category = "Infrastructure" app_hostname = dns_meta.get("query_name") or dns_meta.get("response_name") app_is_encrypted = False elif tls_meta or "quic" in protocols.lower(): app_protocol = "QUIC" if "quic" in protocols.lower() and not tls_meta else "TLS" app_category = "Encrypted" app_hostname = tls_meta.get("server_name") app_is_encrypted = True elif protocol_col and protocol_col.upper() not in {"TCP", "UDP", "IP", "IPV6", "ETH", "ARP"}: app_protocol = protocol_col tshark_meta = _jsonable( { "observed_at_ms": event.get("observed_at_ms"), "protocol": protocol_col, "info": info_col, "frame_protocols": event.get("frame_protocols"), "length": event.get("length"), } ) dpi_metadata = _jsonable( { "tshark": tshark_meta, "http": http_meta, "tls": tls_meta, "dns": dns_meta, } ) return { "app_protocol": app_protocol, "app_master_protocol": app_protocol, "app_category": app_category, "app_confidence": "high" if app_protocol else None, "app_hostname": app_hostname, "app_is_encrypted": app_is_encrypted, "dpi_metadata": dpi_metadata or None, "capture_sources": ["tshark"], } def _has_useful_enrichment(enrichment: Dict[str, Any]) -> bool: if enrichment.get("app_protocol") is not None: return True if enrichment.get("app_hostname") is not None: return True dpi_metadata = enrichment.get("dpi_metadata") or {} if not isinstance(dpi_metadata, dict): return False for key in ("http", "tls", "dns"): value = dpi_metadata.get(key) if isinstance(value, dict) and value: return True return False class TsharkManager: """Own long-lived tshark subprocesses and recent packet metadata cache.""" def __init__(self) -> None: self._enabled = settings.tshark_enabled self._workers: Dict[str, Dict[str, Any]] = {} self._cache: Dict[Tuple[str, int, str, str, int, int, int], List[Dict[str, Any]]] = {} self._last_error_by_iface: Dict[str, str] = {} self._stats: Dict[str, Any] = { "raw_lines_total": 0, "events_total": 0, "events_useful": 0, "parse_failures": 0, "lookup_hits": 0, "backfill_updates_total": 0, "last_event_by_iface": {}, "last_unparsed_line_by_iface": {}, } self._lock = threading.Lock() @property def enabled(self) -> bool: return self._enabled def update_interfaces(self, interfaces: Iterable[str]) -> None: if not self._enabled: return targets = {iface.strip() for iface in interfaces if iface and iface.strip()} with self._lock: current = set(self._workers.keys()) for iface in sorted(current - targets): self._stop_worker_locked(iface) for iface in sorted(targets - current): self._start_worker_locked(iface) self._purge_cache_locked() def stop(self) -> None: with self._lock: for iface in list(self._workers.keys()): self._stop_worker_locked(iface) self._cache.clear() def lookup_packet(self, pkt: Any, iface: Optional[str]) -> Dict[str, Any]: if not self._enabled or not iface: return {} signature = _signature_from_packet(pkt, iface) if signature is None: return {} packet_time_ms = _packet_observed_at_ms(pkt) with self._lock: self._purge_cache_locked(now_ms=packet_time_ms) entries = self._cache.get(signature) if not entries: return {} window_ms = settings.tshark_match_window_ms best_index: Optional[int] = None best_delta: Optional[int] = None for index, event in enumerate(entries): delta = abs(int(event.get("observed_at_ms") or 0) - packet_time_ms) if delta > window_ms: continue if best_delta is None or delta < best_delta: best_index = index best_delta = delta if best_index is None: return {} event = entries.pop(best_index) if not entries: self._cache.pop(signature, None) self._stats["lookup_hits"] += 1 return _build_enrichment(event) def get_debug_snapshot(self) -> Dict[str, Any]: with self._lock: self._purge_cache_locked() return { "enabled": self._enabled, "workers": { iface: { "running": bool(worker.get("process") and worker["process"].poll() is None), "thread_alive": bool(worker.get("reader") and worker["reader"].is_alive()), "last_error": self._last_error_by_iface.get(iface), } for iface, worker in self._workers.items() }, "cache_entries": sum(len(entries) for entries in self._cache.values()), "last_error_by_iface": dict(self._last_error_by_iface), "stats": dict(self._stats), } def _start_worker_locked(self, iface: str) -> None: cmd = [ "tshark", "-l", "-n", "-Q", "-i", iface, "-T", "fields", "-E", "header=n", "-E", "separator=\t", "-E", "quote=n", "-E", "occurrence=f", ] if settings.tshark_try_heuristic_first: cmd.extend(["-o", "tcp.try_heuristic_first:true"]) cmd.extend(["-o", "udp.try_heuristic_first:true"]) if settings.tshark_display_filter: cmd.extend(["-Y", settings.tshark_display_filter]) for field in _FIELDS: cmd.extend(["-e", field]) logger.info("Starting tshark worker for %s", iface) try: process = subprocess.Popen( cmd, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True, bufsize=1, env=os.environ.copy(), start_new_session=True, ) except Exception: logger.exception("Failed to start tshark worker for %s", iface) self._last_error_by_iface[iface] = "failed_to_start_process" return reader = threading.Thread( target=self._read_loop, args=(iface, process), daemon=True, name=f"tshark-reader-{iface}", ) self._workers[iface] = { "process": process, "reader": reader, } reader.start() def _stop_worker_locked(self, iface: str) -> None: worker = self._workers.pop(iface, None) if worker is None: return process = worker.get("process") reader = worker.get("reader") if process is not None and process.poll() is None: try: os.killpg(os.getpgid(process.pid), signal.SIGTERM) process.wait(timeout=settings.tshark_process_stop_timeout_seconds) except subprocess.TimeoutExpired: try: os.killpg(os.getpgid(process.pid), signal.SIGKILL) except ProcessLookupError: pass except Exception: logger.exception("Failed to stop tshark worker for %s cleanly", iface) if reader is not None and reader.is_alive(): reader.join(timeout=settings.tshark_reader_join_timeout_seconds) stale_keys = [key for key in self._cache.keys() if key[0] == iface] for key in stale_keys: self._cache.pop(key, None) def _read_loop(self, iface: str, process: subprocess.Popen[str]) -> None: stdout = process.stdout if stdout is None: return for line in stdout: text = line.rstrip("\n") if not text: continue with self._lock: self._stats["raw_lines_total"] += 1 event = _parse_line(text, iface) if event is None: with self._lock: self._stats["parse_failures"] += 1 self._stats["last_unparsed_line_by_iface"][iface] = text logger.debug("Ignoring unparsable tshark line on %s: %s", iface, text) continue signature = _packet_signature( str(event["iface"]), int(event["protocol"]), str(event["src_ip"]), str(event["dst_ip"]), int(event["src_port"]), int(event["dst_port"]), int(event["length"]), ) with self._lock: entries = self._cache.setdefault(signature, []) entries.append(event) self._stats["events_total"] += 1 self._stats["last_event_by_iface"][iface] = { "observed_at_ms": event.get("observed_at_ms"), "protocol_col": event.get("protocol_col"), "info_col": event.get("info_col"), "protocol": event.get("protocol"), "src_ip": event.get("src_ip"), "dst_ip": event.get("dst_ip"), "src_port": event.get("src_port"), "dst_port": event.get("dst_port"), "frame_protocols": event.get("frame_protocols"), "http": event.get("http"), "tls": event.get("tls"), "dns": event.get("dns"), } self._purge_cache_locked(now_ms=int(event["observed_at_ms"])) enrichment = _build_enrichment(event) if _has_useful_enrichment(enrichment): with self._lock: self._stats["events_useful"] += 1 self._schedule_backfill(event, enrichment) rc = process.poll() if rc not in (0, None): logger.warning("tshark worker for %s exited with code %s", iface, rc) def _schedule_backfill(self, event: Dict[str, Any], enrichment: Dict[str, Any]) -> None: web_db = getattr(shared_objects, "db", None) web_loop = getattr(shared_objects, "web_loop", None) if web_db is None or web_loop is None: return async def _backfill() -> None: updated = await web_db.backfill_packet_metadata( iface=str(event["iface"]), src_ip=str(event["src_ip"]), dst_ip=str(event["dst_ip"]), src_port=int(event["src_port"]), dst_port=int(event["dst_port"]), protocol=int(event["protocol"]), length=int(event["length"]), observed_at_ms=int(event["observed_at_ms"]), enrichment=enrichment, window_ms=settings.tshark_match_window_ms, ) if updated: with self._lock: self._stats["backfill_updates_total"] += updated logger.debug( "Backfilled tshark metadata for %s packets on %s %s:%s -> %s:%s", updated, event["iface"], event["src_ip"], event["src_port"], event["dst_ip"], event["dst_port"], ) try: future = asyncio.run_coroutine_threadsafe(_backfill(), web_loop) future.add_done_callback(self._log_backfill_result) except Exception: logger.exception("Failed to schedule tshark metadata backfill") @staticmethod def _log_backfill_result(future: Any) -> None: try: future.result() except Exception: logger.exception("tshark metadata backfill task failed") def _purge_cache_locked(self, now_ms: Optional[int] = None) -> None: if now_ms is None: now_ms = int(time.time() * 1000) expiry_ms = int(settings.tshark_cache_ttl_seconds * 1000) stale_keys: List[Tuple[str, int, str, str, int, int, int]] = [] for key, entries in self._cache.items(): fresh_entries = [ event for event in entries if now_ms - int(event.get("observed_at_ms") or 0) <= expiry_ms ] if fresh_entries: self._cache[key] = fresh_entries else: stale_keys.append(key) for key in stale_keys: self._cache.pop(key, None) tshark_manager = TsharkManager()