test ebpf
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 1m40s
All checks were successful
Build and Deploy MITM Webserver / build (push) Successful in 1m40s
This commit is contained in:
@@ -31,7 +31,10 @@ from src.utilities.interface_bridge_helpers import (
|
||||
check_interface_up,
|
||||
get_bridge_ports_once,
|
||||
)
|
||||
from src.utilities.bridge_telemetry import bridge_telemetry_manager
|
||||
from src.utilities.ndpi_classifier import ndpi_classifier
|
||||
from src.utilities.packet_identity import build_packet_uid
|
||||
from src.utilities.packet_tracker import packet_tracker
|
||||
from src.Models.etherType import EtherTypeEnum, ethertype_from_int
|
||||
from src.Models.ip_protocol import IPProtocolEnum, protocol_from_number
|
||||
|
||||
@@ -55,13 +58,12 @@ sessions: Dict[str, Dict[str, Any]] = {}
|
||||
# PacketInfo typing
|
||||
# -------------------------
|
||||
class PacketInfo(TypedDict, total=False):
|
||||
"""
|
||||
TypedDict for the parsed packet info produced by parse_packet.
|
||||
Fields marked optional (total=False) for flexibility across contexts.
|
||||
"""
|
||||
"""TypedDict for parsed packet data used by persistence and telemetry."""
|
||||
|
||||
packet_uid: str
|
||||
iface: str
|
||||
length: int
|
||||
raw: bytes # original raw bytes (kept for buffering; DB helper may convert to base64)
|
||||
raw: bytes
|
||||
src_mac: Optional[str]
|
||||
dst_mac: Optional[str]
|
||||
eth_type_raw: Optional[int]
|
||||
@@ -82,6 +84,18 @@ class PacketInfo(TypedDict, total=False):
|
||||
app_is_encrypted: Optional[bool]
|
||||
app_risk_score: Optional[int]
|
||||
dpi_metadata: Optional[Dict[str, Any]]
|
||||
ip_id: Optional[int]
|
||||
icmp_type: Optional[int]
|
||||
icmp_code: Optional[int]
|
||||
arp_op: Optional[int]
|
||||
tcp_seq: Optional[int]
|
||||
tcp_ack: Optional[int]
|
||||
tcp_flags: Optional[int]
|
||||
icmp_embedded_src_ip: Optional[str]
|
||||
icmp_embedded_dst_ip: Optional[str]
|
||||
icmp_embedded_protocol: Optional[int]
|
||||
icmp_embedded_src_port: Optional[int]
|
||||
icmp_embedded_dst_port: Optional[int]
|
||||
|
||||
|
||||
# small bounded buffer for packets produced before shared_objects is ready
|
||||
@@ -105,24 +119,19 @@ threading.Thread(target=_start_async_loop, args=(async_loop,), daemon=True).star
|
||||
# Helpers for buffer draining
|
||||
# -------------------------
|
||||
def drain_buffer_to_shared_db() -> None:
|
||||
"""
|
||||
Attempt to schedule buffered packets for insertion on shared_objects.web_loop.
|
||||
Call this from main.py after shared_objects.db and shared_objects.web_loop are initialized.
|
||||
"""
|
||||
"""Replay buffered packets once the shared DB loop becomes available."""
|
||||
try:
|
||||
web_loop = getattr(shared_objects, "web_loop", None)
|
||||
web_db = getattr(shared_objects, "db", None)
|
||||
if web_db is None or web_loop is None:
|
||||
return
|
||||
|
||||
# schedule draining on the web loop to avoid blocking this thread
|
||||
def _drain() -> None:
|
||||
while _PACKET_BUFFER:
|
||||
pkt = _PACKET_BUFFER.pop(0)
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(web_db.insert_packet(pkt), web_loop)
|
||||
packet_tracker.observe_packet(pkt)
|
||||
except Exception:
|
||||
# re-buffer first element and stop to avoid busy loop
|
||||
_PACKET_BUFFER.insert(0, pkt)
|
||||
break
|
||||
|
||||
@@ -180,6 +189,18 @@ def parse_packet(pkt, bridge_label: str) -> None:
|
||||
"app_is_encrypted": None,
|
||||
"app_risk_score": None,
|
||||
"dpi_metadata": None,
|
||||
"ip_id": None,
|
||||
"icmp_type": None,
|
||||
"icmp_code": None,
|
||||
"arp_op": None,
|
||||
"tcp_seq": None,
|
||||
"tcp_ack": None,
|
||||
"tcp_flags": None,
|
||||
"icmp_embedded_src_ip": None,
|
||||
"icmp_embedded_dst_ip": None,
|
||||
"icmp_embedded_protocol": None,
|
||||
"icmp_embedded_src_port": None,
|
||||
"icmp_embedded_dst_port": None,
|
||||
}
|
||||
|
||||
# Ethernet layer
|
||||
@@ -220,6 +241,10 @@ def parse_packet(pkt, bridge_label: str) -> None:
|
||||
pkt_info["protocol_name"] = "ARP"
|
||||
pkt_info["src_ip"] = _safe_get_attr(arp, "psrc")
|
||||
pkt_info["dst_ip"] = _safe_get_attr(arp, "pdst")
|
||||
try:
|
||||
pkt_info["arp_op"] = int(_safe_get_attr(arp, "op"))
|
||||
except Exception:
|
||||
pkt_info["arp_op"] = None
|
||||
pkt_info["src_port"] = None
|
||||
pkt_info["dst_port"] = None
|
||||
|
||||
@@ -228,6 +253,10 @@ def parse_packet(pkt, bridge_label: str) -> None:
|
||||
ip = pkt[IP]
|
||||
pkt_info["src_ip"] = _safe_get_attr(ip, "src")
|
||||
pkt_info["dst_ip"] = _safe_get_attr(ip, "dst")
|
||||
try:
|
||||
pkt_info["ip_id"] = int(_safe_get_attr(ip, "id"))
|
||||
except Exception:
|
||||
pkt_info["ip_id"] = None
|
||||
|
||||
try:
|
||||
proto_num = int(_safe_get_attr(ip, "proto"))
|
||||
@@ -245,12 +274,23 @@ def parse_packet(pkt, bridge_label: str) -> None:
|
||||
pkt_info["protocol_name"] = "TCP"
|
||||
pkt_info["src_port"] = _safe_get_attr(pkt[TCP], "sport")
|
||||
pkt_info["dst_port"] = _safe_get_attr(pkt[TCP], "dport")
|
||||
try:
|
||||
pkt_info["tcp_seq"] = int(_safe_get_attr(pkt[TCP], "seq"))
|
||||
pkt_info["tcp_ack"] = int(_safe_get_attr(pkt[TCP], "ack"))
|
||||
pkt_info["tcp_flags"] = int(_safe_get_attr(pkt[TCP], "flags"))
|
||||
except Exception:
|
||||
pass
|
||||
elif proto_num == 17 and UDP in pkt:
|
||||
pkt_info["protocol_name"] = "UDP"
|
||||
pkt_info["src_port"] = _safe_get_attr(pkt[UDP], "sport")
|
||||
pkt_info["dst_port"] = _safe_get_attr(pkt[UDP], "dport")
|
||||
elif proto_num == 1 and ICMP in pkt:
|
||||
pkt_info["protocol_name"] = "ICMP"
|
||||
try:
|
||||
pkt_info["icmp_type"] = int(_safe_get_attr(pkt[ICMP], "type"))
|
||||
pkt_info["icmp_code"] = int(_safe_get_attr(pkt[ICMP], "code"))
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
if pkt_info.get("protocol_name") is None:
|
||||
pkt_info["protocol_name"] = f"IP_PROTO_{proto_num}" if proto_num is not None else None
|
||||
@@ -277,12 +317,23 @@ def parse_packet(pkt, bridge_label: str) -> None:
|
||||
pkt_info["protocol_name"] = "TCP"
|
||||
pkt_info["src_port"] = _safe_get_attr(pkt[TCP], "sport")
|
||||
pkt_info["dst_port"] = _safe_get_attr(pkt[TCP], "dport")
|
||||
try:
|
||||
pkt_info["tcp_seq"] = int(_safe_get_attr(pkt[TCP], "seq"))
|
||||
pkt_info["tcp_ack"] = int(_safe_get_attr(pkt[TCP], "ack"))
|
||||
pkt_info["tcp_flags"] = int(_safe_get_attr(pkt[TCP], "flags"))
|
||||
except Exception:
|
||||
pass
|
||||
elif nh == 17 and UDP in pkt:
|
||||
pkt_info["protocol_name"] = "UDP"
|
||||
pkt_info["src_port"] = _safe_get_attr(pkt[UDP], "sport")
|
||||
pkt_info["dst_port"] = _safe_get_attr(pkt[UDP], "dport")
|
||||
elif ICMPv6Unknown in pkt:
|
||||
pkt_info["protocol_name"] = "ICMPv6"
|
||||
try:
|
||||
pkt_info["icmp_type"] = int(_safe_get_attr(pkt[ICMPv6Unknown], "type"))
|
||||
pkt_info["icmp_code"] = int(_safe_get_attr(pkt[ICMPv6Unknown], "code"))
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
if pkt_info.get("protocol_name") is None:
|
||||
pkt_info["protocol_name"] = f"IPV6_PROTO_{nh}" if nh is not None else None
|
||||
@@ -299,21 +350,35 @@ def parse_packet(pkt, bridge_label: str) -> None:
|
||||
except Exception:
|
||||
logger.exception("nDPI enrichment failed")
|
||||
|
||||
# Submit DB insert to shared web loop if available, otherwise buffer
|
||||
if ICMP in pkt:
|
||||
inner = pkt[ICMP].payload
|
||||
if inner and IP in inner:
|
||||
inner_ip = inner[IP]
|
||||
pkt_info["icmp_embedded_src_ip"] = _safe_get_attr(inner_ip, "src")
|
||||
pkt_info["icmp_embedded_dst_ip"] = _safe_get_attr(inner_ip, "dst")
|
||||
try:
|
||||
pkt_info["icmp_embedded_protocol"] = int(_safe_get_attr(inner_ip, "proto"))
|
||||
except Exception:
|
||||
pkt_info["icmp_embedded_protocol"] = None
|
||||
|
||||
if TCP in inner:
|
||||
pkt_info["icmp_embedded_src_port"] = _safe_get_attr(inner[TCP], "sport")
|
||||
pkt_info["icmp_embedded_dst_port"] = _safe_get_attr(inner[TCP], "dport")
|
||||
elif UDP in inner:
|
||||
pkt_info["icmp_embedded_src_port"] = _safe_get_attr(inner[UDP], "sport")
|
||||
pkt_info["icmp_embedded_dst_port"] = _safe_get_attr(inner[UDP], "dport")
|
||||
|
||||
pkt_info["packet_uid"] = build_packet_uid(pkt_info)
|
||||
|
||||
try:
|
||||
web_loop = getattr(shared_objects, "web_loop", None)
|
||||
web_db = getattr(shared_objects, "db", None)
|
||||
if web_db is not None and web_loop is not None:
|
||||
try:
|
||||
fut = asyncio.run_coroutine_threadsafe(web_db.insert_packet(pkt_info), web_loop)
|
||||
# best-effort non-blocking check for immediate errors
|
||||
try:
|
||||
fut.result(timeout=0.005)
|
||||
except Exception:
|
||||
pass
|
||||
logger.debug("Scheduled insert for packet on %s (len=%d)", pkt_info.get("iface"), pkt_info.get("length"))
|
||||
packet_tracker.observe_packet(pkt_info)
|
||||
logger.debug("Tracked packet on %s (len=%d)", pkt_info.get("iface"), pkt_info.get("length"))
|
||||
except Exception as e:
|
||||
logger.exception("Failed to schedule insert_packet for %s — buffering: %s", pkt_info.get("iface"), e)
|
||||
logger.exception("Failed to track packet for %s — buffering: %s", pkt_info.get("iface"), e)
|
||||
_PACKET_BUFFER.append(pkt_info)
|
||||
if len(_PACKET_BUFFER) > _BUFFER_CAPACITY:
|
||||
_PACKET_BUFFER.pop(0)
|
||||
@@ -399,6 +464,14 @@ def _ensure_socket_for_session(sockets: Dict[str, socket.socket], iface: str, br
|
||||
logger.warning("Failed to create AF_PACKET socket for %s (label=%s)", iface, bridge_label)
|
||||
|
||||
|
||||
def _sync_bridge_telemetry() -> None:
|
||||
interfaces = sorted({iface for session in sessions.values() for iface in session.get("ports", [])})
|
||||
try:
|
||||
bridge_telemetry_manager.update_interfaces(interfaces)
|
||||
except Exception:
|
||||
logger.exception("Failed to update bridge telemetry collector")
|
||||
|
||||
|
||||
# -------------------------
|
||||
# Per-session reader loop
|
||||
# -------------------------
|
||||
@@ -545,6 +618,7 @@ def start_afpacket_sniffer(target: str, target_is_interface: bool = False) -> st
|
||||
t = threading.Thread(target=_session_reader_loop, args=(session_id,), daemon=True)
|
||||
session["thread"] = t
|
||||
t.start()
|
||||
_sync_bridge_telemetry()
|
||||
logger.info("Started sniffer session %s label=%s ports=%s", session_id, target, ports)
|
||||
return session_id
|
||||
|
||||
@@ -566,6 +640,7 @@ def stop_afpacket_sniffer(session_id: Optional[str] = None, target: Optional[str
|
||||
t = s.get("thread")
|
||||
if t and isinstance(t, threading.Thread):
|
||||
t.join(timeout=2)
|
||||
_sync_bridge_telemetry()
|
||||
logger.info("Stopped session %s", session_id)
|
||||
return
|
||||
|
||||
@@ -586,6 +661,7 @@ def stop_afpacket_sniffer(session_id: Optional[str] = None, target: Optional[str
|
||||
except Exception:
|
||||
pass
|
||||
logger.info("Removed target %s from session %s", target, sid)
|
||||
_sync_bridge_telemetry()
|
||||
return
|
||||
|
||||
# Global stop: stop all sessions
|
||||
@@ -602,6 +678,10 @@ def stop_afpacket_sniffer(session_id: Optional[str] = None, target: Optional[str
|
||||
logger.exception("Failed to schedule DB pool close")
|
||||
|
||||
logger.info("All sniffer sessions stopped")
|
||||
try:
|
||||
bridge_telemetry_manager.stop()
|
||||
except Exception:
|
||||
logger.exception("Failed to stop bridge telemetry collector")
|
||||
|
||||
|
||||
def get_sniffer_status() -> Dict[str, Dict[str, object]]:
|
||||
@@ -643,4 +723,5 @@ def get_internal_debug_state() -> dict:
|
||||
for sid, s in sessions.items()
|
||||
},
|
||||
"buffer_len": len(_PACKET_BUFFER),
|
||||
"telemetry_ports": sorted({iface for session in sessions.values() for iface in session.get("ports", [])}),
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user