diff --git a/backend/src/utilities/database.py b/backend/src/utilities/database.py index dfcd6b8..e7920e8 100644 --- a/backend/src/utilities/database.py +++ b/backend/src/utilities/database.py @@ -375,6 +375,124 @@ class DatabasePool: return updated_count + async def backfill_stream_metadata( + self, + *, + iface: str, + protocol: int, + stream_kind: str, + stream_id: int, + observed_at_ms: int, + enrichment: Dict[str, Any], + window_ms: int, + ) -> int: + """Propagate tshark stream context across packets already tagged with the same stream id.""" + if self._pool is None: + await self.init_pool() + + lower_bound = datetime.fromtimestamp(max(observed_at_ms - window_ms, 0) / 1000.0, tz=timezone.utc) + upper_bound = datetime.fromtimestamp(max(observed_at_ms + window_ms, 0) / 1000.0, tz=timezone.utc) + capture_sources = [str(source) for source in enrichment.get("capture_sources", []) if source] + stream_kind = str(stream_kind).lower() + if stream_kind not in {"tcp", "udp"}: + return 0 + + try: + async with self._pool.acquire() as conn: + rows = await conn.fetch( + """ + UPDATE packets + SET + updated_at = NOW(), + app_protocol = CASE + WHEN packets.app_protocol IS NULL OR packets.app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH') + THEN COALESCE($7, packets.app_protocol) + ELSE packets.app_protocol + END, + app_master_protocol = CASE + WHEN packets.app_master_protocol IS NULL OR packets.app_master_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH') + THEN COALESCE($8, packets.app_master_protocol) + ELSE packets.app_master_protocol + END, + app_category = CASE + WHEN packets.app_category IS NULL OR packets.app_category IN ('Transport', 'Network', 'Protocol') + THEN COALESCE($9, packets.app_category) + ELSE packets.app_category + END, + app_confidence = CASE + WHEN packets.app_protocol IS NULL OR packets.app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH') + THEN COALESCE($10, packets.app_confidence) + ELSE packets.app_confidence + END, + app_hostname = COALESCE(packets.app_hostname, $11), + app_is_encrypted = COALESCE(packets.app_is_encrypted, $12), + capture_sources = ( + SELECT ARRAY( + SELECT DISTINCT source + FROM unnest( + COALESCE(packets.capture_sources, ARRAY[]::text[]) || + COALESCE($13::text[], ARRAY[]::text[]) + ) AS source + ) + ) + WHERE + ip_proto_raw = $1 + AND (capture_iface = $2 OR ingress_if = $2 OR egress_if = $2) + AND timestamp BETWEEN $5 AND $6 + AND ( + CASE + WHEN $3 = 'tcp' THEN COALESCE(packets.dpi_metadata -> 'tcp' ->> 'stream', '') + WHEN $3 = 'udp' THEN COALESCE(packets.dpi_metadata -> 'udp' ->> 'stream', '') + ELSE '' + END + ) = $4 + AND ( + packets.app_protocol IS NULL + OR packets.app_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH') + OR packets.app_master_protocol IS NULL + OR packets.app_master_protocol IN ('TCP', 'UDP', 'IP', 'IPv6', 'ETH') + OR packets.app_category IS NULL + OR packets.app_category IN ('Transport', 'Network', 'Protocol') + OR packets.app_confidence IS NULL + OR packets.app_hostname IS NULL + OR packets.app_is_encrypted IS NULL + OR (COALESCE(array_length($13::text[], 1), 0) > 0) + ) + RETURNING * + """, + protocol, + iface, + stream_kind, + str(stream_id), + lower_bound, + upper_bound, + enrichment.get("app_protocol"), + enrichment.get("app_master_protocol"), + enrichment.get("app_category"), + enrichment.get("app_confidence"), + enrichment.get("app_hostname"), + enrichment.get("app_is_encrypted"), + capture_sources if capture_sources else None, + ) + except Exception: + logger.exception("DB stream metadata backfill failed") + return 0 + + if not rows: + return 0 + + updated_count = 0 + for row in rows: + serialized = _serialize_row_for_broadcast(dict(row)) + updated_count += 1 + if self.broadcaster: + try: + self.broadcaster.sync_publish(serialized) + except Exception: + logger.exception("Failed to publish stream-context packet row") + + return updated_count + async def fetch_latest(self, limit: int) -> List[PacketDBModel]: """Fetch newest packet rows as validated `PacketDBModel` instances.""" if self._pool is None: diff --git a/backend/src/utilities/tshark_manager.py b/backend/src/utilities/tshark_manager.py index d877bae..448febf 100644 --- a/backend/src/utilities/tshark_manager.py +++ b/backend/src/utilities/tshark_manager.py @@ -30,10 +30,27 @@ _FIELDS: List[str] = [ "ipv6.dst", "ip.proto", "ipv6.nxt", + "tcp.stream", + "udp.stream", "tcp.srcport", "udp.srcport", "tcp.dstport", "udp.dstport", + "tcp.flags", + "tcp.seq_raw", + "tcp.ack_raw", + "tcp.len", + "tcp.analysis.retransmission", + "tcp.analysis.fast_retransmission", + "tcp.analysis.spurious_retransmission", + "tcp.analysis.keep_alive", + "tcp.analysis.keep_alive_ack", + "tcp.analysis.duplicate_ack", + "arp.opcode", + "icmp.type", + "icmp.code", + "icmpv6.type", + "icmpv6.code", "frame.protocols", "http.request.method", "http.request.uri", @@ -93,6 +110,60 @@ def _safe_bool_flag(value: Any) -> Optional[bool]: return None +def _tcp_flags_summary(flags_value: Optional[int], payload_len: int) -> Optional[str]: + if flags_value is None: + return None + + syn = bool(flags_value & 0x02) + ack = bool(flags_value & 0x10) + fin = bool(flags_value & 0x01) + rst = bool(flags_value & 0x04) + psh = bool(flags_value & 0x08) + urg = bool(flags_value & 0x20) + + if syn and ack: + return "SYN-ACK" + if syn: + return "SYN" + if rst and ack: + return "RST-ACK" + if rst: + return "RST" + if fin and ack: + return "FIN-ACK" + if fin: + return "FIN" + if psh and ack: + return "PSH-ACK" + if psh: + return "PSH" + if ack and payload_len == 0: + return "ACK" + if urg and ack: + return "URG-ACK" + return None + + +def _tcp_flag_names(flags_value: Optional[int]) -> List[str]: + if flags_value is None: + return [] + + names: List[str] = [] + for bit, name in ( + (0x01, "FIN"), + (0x02, "SYN"), + (0x04, "RST"), + (0x08, "PSH"), + (0x10, "ACK"), + (0x20, "URG"), + (0x40, "ECE"), + (0x80, "CWR"), + ): + if flags_value & bit: + names.append(name) + return names + + def _jsonable(metadata: Dict[str, Any]) -> Dict[str, Any]: out: Dict[str, Any] = {} for key, value in metadata.items(): @@ -102,6 +173,17 @@ def _jsonable(metadata: Dict[str, Any]) -> Dict[str, Any]: return out +def _stream_key(event: Dict[str, Any]) -> Optional[Tuple[str, str, int]]: + iface = str(event.get("iface") or "") + tcp_stream = event.get("tcp_stream") + if tcp_stream is not None: + return (iface, "tcp", int(tcp_stream)) + udp_stream = event.get("udp_stream") + if udp_stream is not None: + return (iface, "udp", int(udp_stream)) + return None + + def _packet_signature( iface: str, protocol: int, @@ -163,6 +245,8 @@ def _parse_line(line: str, fallback_iface: str) -> Optional[Dict[str, Any]]: 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"]) + tcp_stream = _safe_int(row["tcp.stream"]) + udp_stream = _safe_int(row["udp.stream"]) 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 @@ -182,6 +266,21 @@ def _parse_line(line: str, fallback_iface: str) -> Optional[Dict[str, Any]]: "dst_ip": dst_ip, "src_port": src_port, "dst_port": dst_port, + "tcp_stream": tcp_stream, + "udp_stream": udp_stream, + "tcp_flags": _safe_int(row["tcp.flags"]), + "tcp_seq_raw": _safe_int(row["tcp.seq_raw"]), + "tcp_ack_raw": _safe_int(row["tcp.ack_raw"]), + "tcp_len": _safe_int(row["tcp.len"]) or 0, + "tcp_retransmission": _safe_bool_flag(row["tcp.analysis.retransmission"]), + "tcp_fast_retransmission": _safe_bool_flag(row["tcp.analysis.fast_retransmission"]), + "tcp_spurious_retransmission": _safe_bool_flag(row["tcp.analysis.spurious_retransmission"]), + "tcp_keep_alive": _safe_bool_flag(row["tcp.analysis.keep_alive"]), + "tcp_keep_alive_ack": _safe_bool_flag(row["tcp.analysis.keep_alive_ack"]), + "tcp_duplicate_ack": _safe_bool_flag(row["tcp.analysis.duplicate_ack"]), + "arp_opcode": _safe_int(row["arp.opcode"]), + "icmp_type": _safe_int(row["icmp.type"]) or _safe_int(row["icmpv6.type"]), + "icmp_code": _safe_int(row["icmp.code"]) or _safe_int(row["icmpv6.code"]), "protocol_col": protocol_col, "info_col": info_col, "frame_protocols": _safe_text(row["frame.protocols"]), @@ -217,13 +316,12 @@ def _parse_line(line: str, fallback_iface: str) -> Optional[Dict[str, Any]]: } -def _build_enrichment(event: Dict[str, Any]) -> Dict[str, Any]: +def _derive_flow_context(event: Dict[str, Any]) -> Dict[str, Any]: + protocol_col = _safe_text(event.get("protocol_col")) + protocols = str(event.get("frame_protocols") or "") 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 @@ -245,9 +343,54 @@ def _build_enrichment(event: Dict[str, Any]) -> Dict[str, Any]: 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"}: + elif protocol_col: app_protocol = protocol_col + if app_protocol == "TCP": + app_category = app_category or "Transport" + elif app_protocol == "UDP": + app_category = app_category or "Transport" + elif app_protocol in {"ARP", "ICMP", "ICMPV6", "IP", "IPV6"}: + app_category = app_category or "Network" + elif app_protocol and app_category is None: + app_category = "Protocol" + + 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, + "last_seen_ms": int(event.get("observed_at_ms") or 0), + } + + +def _build_enrichment(event: Dict[str, Any], flow_context: Optional[Dict[str, Any]] = None) -> 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")) + flow_context = dict(flow_context or {}) + derived_context = _derive_flow_context(event) + + app_protocol = derived_context.get("app_protocol") or flow_context.get("app_protocol") + app_category = derived_context.get("app_category") or flow_context.get("app_category") + app_hostname = derived_context.get("app_hostname") or flow_context.get("app_hostname") + app_is_encrypted = ( + derived_context.get("app_is_encrypted") + if derived_context.get("app_is_encrypted") is not None + else flow_context.get("app_is_encrypted") + ) + app_confidence = derived_context.get("app_confidence") or flow_context.get("app_confidence") + + tcp_flags = _safe_int(event.get("tcp_flags")) + tcp_len = int(event.get("tcp_len") or 0) + tcp_packet_type = _tcp_flags_summary(tcp_flags, tcp_len) + tcp_flag_names = _tcp_flag_names(tcp_flags) + tshark_meta = _jsonable( { "observed_at_ms": event.get("observed_at_ms"), @@ -255,12 +398,36 @@ def _build_enrichment(event: Dict[str, Any]) -> Dict[str, Any]: "info": info_col, "frame_protocols": event.get("frame_protocols"), "length": event.get("length"), + "tcp_stream": event.get("tcp_stream"), + "udp_stream": event.get("udp_stream"), + "tcp_packet_type": tcp_packet_type, + "tcp_flag_names": tcp_flag_names, } ) dpi_metadata = _jsonable( { "tshark": tshark_meta, + "tcp": _jsonable( + { + "stream": event.get("tcp_stream"), + "flags": tcp_flags, + "packet_type": tcp_packet_type, + "flag_names": tcp_flag_names, + "seq_raw": event.get("tcp_seq_raw"), + "ack_raw": event.get("tcp_ack_raw"), + "payload_len": tcp_len, + "retransmission": event.get("tcp_retransmission"), + "fast_retransmission": event.get("tcp_fast_retransmission"), + "spurious_retransmission": event.get("tcp_spurious_retransmission"), + "keep_alive": event.get("tcp_keep_alive"), + "keep_alive_ack": event.get("tcp_keep_alive_ack"), + "duplicate_ack": event.get("tcp_duplicate_ack"), + } + ), + "udp": _jsonable({"stream": event.get("udp_stream")}), + "icmp": _jsonable({"type": event.get("icmp_type"), "code": event.get("icmp_code")}), + "arp": _jsonable({"opcode": event.get("arp_opcode")}), "http": http_meta, "tls": tls_meta, "dns": dns_meta, @@ -271,7 +438,7 @@ def _build_enrichment(event: Dict[str, Any]) -> Dict[str, Any]: "app_protocol": app_protocol, "app_master_protocol": app_protocol, "app_category": app_category, - "app_confidence": "high" if app_protocol else None, + "app_confidence": app_confidence, "app_hostname": app_hostname, "app_is_encrypted": app_is_encrypted, "dpi_metadata": dpi_metadata or None, @@ -287,11 +454,26 @@ def _has_useful_enrichment(enrichment: Dict[str, Any]) -> bool: 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) + for value in dpi_metadata.values(): if isinstance(value, dict) and value: return True - return False + return bool(dpi_metadata) + + +def _build_stream_enrichment(flow_context: Dict[str, Any]) -> Dict[str, Any]: + return { + "app_protocol": flow_context.get("app_protocol"), + "app_master_protocol": flow_context.get("app_master_protocol"), + "app_category": flow_context.get("app_category"), + "app_confidence": flow_context.get("app_confidence"), + "app_hostname": flow_context.get("app_hostname"), + "app_is_encrypted": flow_context.get("app_is_encrypted"), + "capture_sources": ["tshark"], + } + + +def _is_generic_stream_protocol(protocol_name: Optional[str]) -> bool: + return str(protocol_name or "").upper() in {"", "TCP", "UDP", "IP", "IPV6", "ETH"} class TsharkManager: @@ -301,6 +483,7 @@ class TsharkManager: 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._flow_context: Dict[Tuple[str, str, int], Dict[str, Any]] = {} self._last_error_by_iface: Dict[str, str] = {} self._stats: Dict[str, Any] = { "raw_lines_total": 0, @@ -336,6 +519,7 @@ class TsharkManager: for iface in list(self._workers.keys()): self._stop_worker_locked(iface) self._cache.clear() + self._flow_context.clear() def lookup_packet(self, pkt: Any, iface: Optional[str]) -> Dict[str, Any]: if not self._enabled or not iface: @@ -370,8 +554,10 @@ class TsharkManager: if not entries: self._cache.pop(signature, None) self._stats["lookup_hits"] += 1 + stream_key = _stream_key(event) + flow_context = dict(self._flow_context.get(stream_key) or {}) if stream_key is not None else None - return _build_enrichment(event) + return _build_enrichment(event, flow_context=flow_context) def get_debug_snapshot(self) -> Dict[str, Any]: with self._lock: @@ -387,6 +573,7 @@ class TsharkManager: for iface, worker in self._workers.items() }, "cache_entries": sum(len(entries) for entries in self._cache.values()), + "flow_context_entries": len(self._flow_context), "last_error_by_iface": dict(self._last_error_by_iface), "stats": dict(self._stats), } @@ -472,6 +659,9 @@ class TsharkManager: stale_keys = [key for key in self._cache.keys() if key[0] == iface] for key in stale_keys: self._cache.pop(key, None) + stale_flow_keys = [key for key in self._flow_context.keys() if key[0] == iface] + for key in stale_flow_keys: + self._flow_context.pop(key, None) def _read_loop(self, iface: str, process: subprocess.Popen[str]) -> None: stdout = process.stdout @@ -503,14 +693,28 @@ class TsharkManager: int(event["length"]), ) + stream_key = _stream_key(event) + flow_context: Optional[Dict[str, Any]] = None + with self._lock: + if stream_key is not None: + flow_context = dict(self._flow_context.get(stream_key) or {}) + derived_context = _derive_flow_context(event) + if stream_key is not None and derived_context.get("app_protocol") is not None: + merged_context = dict(flow_context or {}) + merged_context.update({k: v for k, v in derived_context.items() if v is not None}) + self._flow_context[stream_key] = merged_context + flow_context = merged_context 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"), + "tcp_stream": event.get("tcp_stream"), + "udp_stream": event.get("udp_stream"), "protocol_col": event.get("protocol_col"), "info_col": event.get("info_col"), + "tcp_packet_type": _tcp_flags_summary(_safe_int(event.get("tcp_flags")), int(event.get("tcp_len") or 0)), "protocol": event.get("protocol"), "src_ip": event.get("src_ip"), "dst_ip": event.get("dst_ip"), @@ -523,17 +727,23 @@ class TsharkManager: } self._purge_cache_locked(now_ms=int(event["observed_at_ms"])) - enrichment = _build_enrichment(event) + enrichment = _build_enrichment(event, flow_context=flow_context) if _has_useful_enrichment(enrichment): with self._lock: self._stats["events_useful"] += 1 - self._schedule_backfill(event, enrichment) + self._schedule_backfill(event, enrichment, flow_context=flow_context or derived_context) 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: + def _schedule_backfill( + self, + event: Dict[str, Any], + enrichment: Dict[str, Any], + *, + flow_context: Optional[Dict[str, Any]] = None, + ) -> 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: @@ -565,6 +775,46 @@ class TsharkManager: event["dst_port"], ) + stream_key = _stream_key(event) + if stream_key is None or not flow_context: + return + + stream_enrichment = _build_stream_enrichment(flow_context) + if _is_generic_stream_protocol(stream_enrichment.get("app_protocol")) and stream_enrichment.get("app_hostname") is None: + return + if not any( + stream_enrichment.get(key) is not None + for key in ( + "app_protocol", + "app_master_protocol", + "app_category", + "app_confidence", + "app_hostname", + "app_is_encrypted", + ) + ): + return + + stream_kind = str(stream_key[1]) + stream_id = int(stream_key[2]) + stream_updated = await web_db.backfill_stream_metadata( + iface=str(event["iface"]), + protocol=int(event["protocol"]), + stream_kind=stream_kind, + stream_id=stream_id, + observed_at_ms=int(event["observed_at_ms"]), + enrichment=stream_enrichment, + window_ms=max(settings.tshark_match_window_ms, 10_000), + ) + if stream_updated: + logger.debug( + "Backfilled tshark stream context for %s packets on %s %s.stream=%s", + stream_updated, + event["iface"], + stream_kind, + stream_id, + ) + try: future = asyncio.run_coroutine_threadsafe(_backfill(), web_loop) future.add_done_callback(self._log_backfill_result) @@ -598,5 +848,13 @@ class TsharkManager: for key in stale_keys: self._cache.pop(key, None) + stale_flow_keys = [] + for key, payload in self._flow_context.items(): + if now_ms - int(payload.get("last_seen_ms") or 0) > expiry_ms: + stale_flow_keys.append(key) + + for key in stale_flow_keys: + self._flow_context.pop(key, None) + tshark_manager = TsharkManager()