Files
routeros-suricata-tzsp/app/flow_tracker.py
T
2026-08-16 11:46:35 +02:00

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 ""