338 lines
12 KiB
Python
338 lines
12 KiB
Python
from __future__ import annotations
|
|
|
|
import collections
|
|
import hashlib
|
|
import socket
|
|
import struct
|
|
import time
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timezone
|
|
from typing import Any
|
|
|
|
from .live import LiveEventPipeline, TrafficNormalizer
|
|
|
|
_ETH_IPV4 = 0x0800
|
|
_ETH_IPV6 = 0x86DD
|
|
_VLAN_TYPES = {0x8100, 0x88A8, 0x9100}
|
|
_IP_PROTO_NAMES = {1: "ICMP", 6: "TCP", 17: "UDP", 58: "ICMPV6"}
|
|
_IPV6_EXTENSIONS = {0, 43, 44, 51, 60}
|
|
|
|
|
|
@dataclass
|
|
class _FlowState:
|
|
stable_id: str
|
|
src_ip: str
|
|
src_port: int
|
|
dest_ip: str
|
|
dest_port: int
|
|
proto: str
|
|
app_proto: str
|
|
first_seen: float
|
|
last_seen: float
|
|
last_published: float
|
|
bytes_to_server: int = 0
|
|
bytes_to_client: int = 0
|
|
packets_to_server: int = 0
|
|
packets_to_client: int = 0
|
|
|
|
|
|
class FlowTracker:
|
|
"""Bounded L3/L4 session tracker used only for immediate dashboard updates.
|
|
|
|
Suricata remains the source of durable EVE history. This tracker emits
|
|
non-persistent updates from TZSP frames so long-lived sessions are visible
|
|
before Suricata closes and writes the final flow event.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
normalizer: TrafficNormalizer,
|
|
pipeline: LiveEventPipeline,
|
|
update_interval_seconds: float = 1.0,
|
|
idle_seconds: float = 120.0,
|
|
max_flows: int = 20000,
|
|
) -> None:
|
|
self.normalizer = normalizer
|
|
self.pipeline = pipeline
|
|
self.update_interval = max(0.25, float(update_interval_seconds))
|
|
self.idle_seconds = max(10.0, float(idle_seconds))
|
|
self.max_flows = max(1000, int(max_flows))
|
|
self._flows: collections.OrderedDict[tuple[Any, ...], _FlowState] = collections.OrderedDict()
|
|
self._last_cleanup = time.monotonic()
|
|
self._published = 0
|
|
self._evicted = 0
|
|
self._parse_errors = 0
|
|
self._throughput_samples = 0
|
|
self._rate_started = time.monotonic()
|
|
self._rate_counters = {
|
|
"bytes_total": 0, "bytes_in": 0, "bytes_out": 0,
|
|
"bytes_internal": 0, "bytes_external": 0,
|
|
"packets_total": 0, "packets_in": 0, "packets_out": 0,
|
|
"packets_internal": 0, "packets_external": 0,
|
|
}
|
|
# Monotonic counters are kept separately from the one-second sampling
|
|
# bucket above. Prometheus can safely apply rate()/increase() to these
|
|
# even when its scrape interval differs from the UI sampling interval.
|
|
self._traffic_counters = dict.fromkeys(self._rate_counters, 0)
|
|
|
|
def observe(self, frame: bytes) -> None:
|
|
parsed = _parse_frame(frame)
|
|
if parsed is None:
|
|
self._parse_errors += 1
|
|
return
|
|
src_ip, src_port, dest_ip, dest_port, proto = parsed
|
|
now = time.monotonic()
|
|
self._record_throughput(src_ip, dest_ip, len(frame), now)
|
|
|
|
# Building per-flow state is only needed for the optional Live Sessions
|
|
# stream. The overview throughput counters above stay active at all times,
|
|
# but when no browser requested live streaming we avoid OrderedDict churn,
|
|
# hashing and periodic synthetic flow updates for every captured packet.
|
|
live_needed = getattr(self.pipeline, "has_live_subscribers", None)
|
|
if callable(live_needed) and not live_needed():
|
|
if self._flows and now - self._last_cleanup >= 10.0:
|
|
self._flows.clear()
|
|
self._last_cleanup = now
|
|
return
|
|
|
|
key = _canonical_key(src_ip, src_port, dest_ip, dest_port, proto)
|
|
state = self._flows.get(key)
|
|
if state is None:
|
|
stable_id = "live-" + hashlib.blake2s(repr(key).encode("utf-8"), digest_size=10).hexdigest()
|
|
state = _FlowState(
|
|
stable_id=stable_id,
|
|
src_ip=src_ip,
|
|
src_port=src_port,
|
|
dest_ip=dest_ip,
|
|
dest_port=dest_port,
|
|
proto=proto,
|
|
app_proto=_guess_app(proto, src_port, dest_port),
|
|
first_seen=now,
|
|
last_seen=now,
|
|
last_published=0.0,
|
|
)
|
|
self._flows[key] = state
|
|
else:
|
|
state.last_seen = now
|
|
self._flows.move_to_end(key)
|
|
|
|
frame_bytes = len(frame)
|
|
if src_ip == state.src_ip and src_port == state.src_port:
|
|
state.bytes_to_server += frame_bytes
|
|
state.packets_to_server += 1
|
|
else:
|
|
state.bytes_to_client += frame_bytes
|
|
state.packets_to_client += 1
|
|
|
|
if state.last_published == 0.0 or now - state.last_published >= self.update_interval:
|
|
self._publish(state, now)
|
|
|
|
if len(self._flows) > self.max_flows:
|
|
while len(self._flows) > self.max_flows:
|
|
self._flows.popitem(last=False)
|
|
self._evicted += 1
|
|
if now - self._last_cleanup >= 10.0:
|
|
self._cleanup(now)
|
|
|
|
def status(self) -> dict[str, Any]:
|
|
return {
|
|
"active_flows": len(self._flows),
|
|
"max_flows": self.max_flows,
|
|
"published_updates": self._published,
|
|
"evicted_flows": self._evicted,
|
|
"parse_errors": self._parse_errors,
|
|
"throughput_samples": self._throughput_samples,
|
|
"update_interval_seconds": self.update_interval,
|
|
"traffic_counters": dict(self._traffic_counters),
|
|
}
|
|
|
|
def _record_throughput(self, src_ip: str, dest_ip: str, frame_bytes: int, now: float) -> None:
|
|
direction = self.normalizer._direction(src_ip, dest_ip)
|
|
counters = self._rate_counters
|
|
totals = self._traffic_counters
|
|
|
|
for target in (counters, totals):
|
|
target["bytes_total"] += frame_bytes
|
|
target["packets_total"] += 1
|
|
if direction == "inbound":
|
|
for target in (counters, totals):
|
|
target["bytes_in"] += frame_bytes
|
|
target["packets_in"] += 1
|
|
elif direction == "outbound":
|
|
for target in (counters, totals):
|
|
target["bytes_out"] += frame_bytes
|
|
target["packets_out"] += 1
|
|
elif direction == "internal":
|
|
for target in (counters, totals):
|
|
target["bytes_internal"] += frame_bytes
|
|
target["packets_internal"] += 1
|
|
else:
|
|
for target in (counters, totals):
|
|
target["bytes_external"] += frame_bytes
|
|
target["packets_external"] += 1
|
|
|
|
elapsed = now - self._rate_started
|
|
if elapsed < 1.0:
|
|
return
|
|
sample = dict(counters)
|
|
sample["ts_ms"] = int(time.time() * 1000)
|
|
sample["interval_ms"] = max(1, round(elapsed * 1000))
|
|
self.pipeline.publish_throughput(sample)
|
|
self._throughput_samples += 1
|
|
for key in counters:
|
|
counters[key] = 0
|
|
self._rate_started = now
|
|
|
|
def _publish(self, state: _FlowState, now: float) -> None:
|
|
state.last_published = now
|
|
event = {
|
|
"timestamp": datetime.now(timezone.utc).isoformat(),
|
|
"event_type": "flow",
|
|
"flow_id": state.stable_id,
|
|
"src_ip": state.src_ip,
|
|
"src_port": state.src_port or None,
|
|
"dest_ip": state.dest_ip,
|
|
"dest_port": state.dest_port or None,
|
|
"proto": state.proto,
|
|
"app_proto": state.app_proto,
|
|
"flow": {
|
|
"bytes_toserver": state.bytes_to_server,
|
|
"bytes_toclient": state.bytes_to_client,
|
|
"pkts_toserver": state.packets_to_server,
|
|
"pkts_toclient": state.packets_to_client,
|
|
"state": "live",
|
|
"reason": "tzsp",
|
|
},
|
|
}
|
|
record = self.normalizer.normalize(
|
|
event,
|
|
id=state.stable_id,
|
|
live=True,
|
|
source="tzsp",
|
|
age_seconds=round(now - state.first_seen, 3),
|
|
)
|
|
if record is not None:
|
|
self.pipeline.publish(record, persist=False)
|
|
self._published += 1
|
|
|
|
def _cleanup(self, now: float) -> None:
|
|
cutoff = now - self.idle_seconds
|
|
while self._flows:
|
|
_key, state = next(iter(self._flows.items()))
|
|
if state.last_seen >= cutoff:
|
|
break
|
|
self._flows.popitem(last=False)
|
|
self._last_cleanup = now
|
|
|
|
|
|
def _canonical_key(src: str, src_port: int, dst: str, dst_port: int, proto: str) -> tuple[Any, ...]:
|
|
left = (src, src_port)
|
|
right = (dst, dst_port)
|
|
if left <= right:
|
|
return proto, left, right
|
|
return proto, right, left
|
|
|
|
|
|
def _parse_frame(frame: bytes) -> tuple[str, int, str, int, str] | None:
|
|
if len(frame) < 14:
|
|
return None
|
|
offset = 14
|
|
ethertype = struct.unpack_from("!H", frame, 12)[0]
|
|
for _ in range(2):
|
|
if ethertype not in _VLAN_TYPES or len(frame) < offset + 4:
|
|
break
|
|
ethertype = struct.unpack_from("!H", frame, offset + 2)[0]
|
|
offset += 4
|
|
|
|
if ethertype == _ETH_IPV4:
|
|
return _parse_ipv4(frame, offset)
|
|
if ethertype == _ETH_IPV6:
|
|
return _parse_ipv6(frame, offset)
|
|
return None
|
|
|
|
|
|
def _parse_ipv4(frame: bytes, offset: int) -> tuple[str, int, str, int, str] | None:
|
|
if len(frame) < offset + 20:
|
|
return None
|
|
version_ihl = frame[offset]
|
|
if version_ihl >> 4 != 4:
|
|
return None
|
|
header_len = (version_ihl & 0x0F) * 4
|
|
if header_len < 20 or len(frame) < offset + header_len:
|
|
return None
|
|
protocol = frame[offset + 9]
|
|
src = socket.inet_ntop(socket.AF_INET, frame[offset + 12 : offset + 16])
|
|
dst = socket.inet_ntop(socket.AF_INET, frame[offset + 16 : offset + 20])
|
|
frag = struct.unpack_from("!H", frame, offset + 6)[0] & 0x1FFF
|
|
l4_offset = offset + header_len
|
|
src_port, dst_port = _ports(frame, l4_offset, protocol) if frag == 0 else (0, 0)
|
|
return src, src_port, dst, dst_port, _IP_PROTO_NAMES.get(protocol, f"IP{protocol}")
|
|
|
|
|
|
def _parse_ipv6(frame: bytes, offset: int) -> tuple[str, int, str, int, str] | None:
|
|
if len(frame) < offset + 40 or frame[offset] >> 4 != 6:
|
|
return None
|
|
next_header = frame[offset + 6]
|
|
src = socket.inet_ntop(socket.AF_INET6, frame[offset + 8 : offset + 24])
|
|
dst = socket.inet_ntop(socket.AF_INET6, frame[offset + 24 : offset + 40])
|
|
l4_offset = offset + 40
|
|
fragmented_nonzero = False
|
|
|
|
for _ in range(6):
|
|
if next_header not in _IPV6_EXTENSIONS:
|
|
break
|
|
if next_header == 44: # Fragment header: fixed 8 bytes.
|
|
if len(frame) < l4_offset + 8:
|
|
return src, 0, dst, 0, "IPV6"
|
|
fragment_bits = struct.unpack_from("!H", frame, l4_offset + 2)[0]
|
|
fragmented_nonzero = (fragment_bits >> 3) != 0
|
|
next_header = frame[l4_offset]
|
|
l4_offset += 8
|
|
continue
|
|
if next_header == 51: # Authentication Header length is in 32-bit words minus 2.
|
|
if len(frame) < l4_offset + 2:
|
|
return src, 0, dst, 0, "IPV6"
|
|
following = frame[l4_offset]
|
|
header_len = (frame[l4_offset + 1] + 2) * 4
|
|
else:
|
|
if len(frame) < l4_offset + 2:
|
|
return src, 0, dst, 0, "IPV6"
|
|
following = frame[l4_offset]
|
|
header_len = (frame[l4_offset + 1] + 1) * 8
|
|
if header_len <= 0 or len(frame) < l4_offset + header_len:
|
|
return src, 0, dst, 0, "IPV6"
|
|
next_header = following
|
|
l4_offset += header_len
|
|
|
|
src_port, dst_port = (0, 0) if fragmented_nonzero else _ports(frame, l4_offset, next_header)
|
|
return src, src_port, dst, dst_port, _IP_PROTO_NAMES.get(next_header, f"IP{next_header}")
|
|
|
|
|
|
def _ports(frame: bytes, offset: int, protocol: int) -> tuple[int, int]:
|
|
if protocol not in {6, 17} or len(frame) < offset + 4:
|
|
return 0, 0
|
|
return struct.unpack_from("!HH", frame, offset)
|
|
|
|
|
|
def _guess_app(proto: str, src_port: int, dst_port: int) -> str:
|
|
ports = {src_port, dst_port}
|
|
if 53 in ports:
|
|
return "dns"
|
|
if proto == "UDP" and 443 in ports:
|
|
return "quic"
|
|
if 443 in ports:
|
|
return "tls"
|
|
if 80 in ports or 8080 in ports:
|
|
return "http"
|
|
if 22 in ports:
|
|
return "ssh"
|
|
if 3389 in ports:
|
|
return "rdp"
|
|
if 445 in ports:
|
|
return "smb"
|
|
if 8291 in ports:
|
|
return "winbox"
|
|
if 123 in ports:
|
|
return "ntp"
|
|
return ""
|