overload batching improvements
Some checks failed
Build and Deploy MITM Webserver / traffic_target (push) Successful in 0s
Build and Deploy MITM Webserver / build (push) Failing after 23m3s

This commit is contained in:
2026-05-03 20:11:42 +02:00
parent bf63c7c1a9
commit c40ddf4ded

View File

@@ -297,7 +297,7 @@ class DatabasePool:
self._publish_packet_row(pkt_info, row) self._publish_packet_row(pkt_info, row)
async def upsert_packets(self, packets: List[Dict[str, Any]]) -> None: async def upsert_packets(self, packets: List[Dict[str, Any]]) -> None:
"""Insert or update packet records using one database connection for the batch.""" """Insert or update packet records using one database round-trip optimized batch path."""
if not packets: if not packets:
return return
if self._pool is None: if self._pool is None:
@@ -305,21 +305,20 @@ class DatabasePool:
if self._pool is None: if self._pool is None:
return return
published_rows: List[tuple[Dict[str, Any], Optional[Dict[str, Any]]]] = []
try: try:
params = [self._packet_upsert_params(pkt_info) for pkt_info in packets]
async with self._pool.acquire() as conn: async with self._pool.acquire() as conn:
for pkt_info in packets: await conn.executemany(self._packet_upsert_sql(returning=False), params)
row = await self._upsert_packet_with_conn(conn, pkt_info)
published_rows.append((pkt_info, row))
except Exception: except Exception:
logger.exception("DB packet batch upsert failed") logger.exception("DB packet batch upsert failed")
raise raise
for pkt_info, row in published_rows:
self._publish_packet_row(pkt_info, row)
async def _upsert_packet_with_conn(self, conn: asyncpg.Connection, pkt_info: Dict[str, Any]) -> Optional[Dict[str, Any]]: async def _upsert_packet_with_conn(self, conn: asyncpg.Connection, pkt_info: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Insert or update one packet record using an already-acquired connection.""" """Insert or update one packet record using an already-acquired connection."""
row = await conn.fetchrow(self._packet_upsert_sql(returning=True), *self._packet_upsert_params(pkt_info))
return dict(row) if row else None
def _packet_upsert_params(self, pkt_info: Dict[str, Any]) -> tuple[Any, ...]:
_normalize_json_fields(pkt_info) _normalize_json_fields(pkt_info)
_attach_derived_fields(pkt_info) _attach_derived_fields(pkt_info)
@@ -327,8 +326,50 @@ class DatabasePool:
telemetry_metadata = pkt_info.get("telemetry_metadata") telemetry_metadata = pkt_info.get("telemetry_metadata")
capture_observations = pkt_info.get("capture_observations") capture_observations = pkt_info.get("capture_observations")
row = await conn.fetchrow( return (
""" pkt_info.get("timestamp"),
pkt_info["correlation_key"],
pkt_info.get("packet_id"),
pkt_info.get("packet_uid"),
pkt_info.get("flow_id"),
pkt_info.get("capture_session_id"),
pkt_info.get("correlation_source"),
pkt_info.get("skb_mark"),
pkt_info.get("capture_iface"),
pkt_info.get("ingress_if"),
pkt_info.get("egress_if"),
pkt_info.get("verdict"),
pkt_info.get("verdict_reason"),
pkt_info.get("verdict_confidence"),
pkt_info.get("ingress_seen_at"),
pkt_info.get("egress_seen_at"),
pkt_info.get("verdict_seen_at"),
pkt_info.get("src_mac"),
pkt_info.get("dst_mac"),
pkt_info.get("eth_type_raw"),
pkt_info.get("vlan_id"),
pkt_info.get("src_ip"),
pkt_info.get("dst_ip"),
pkt_info.get("protocol_raw"),
pkt_info.get("src_port"),
pkt_info.get("dst_port"),
pkt_info.get("length"),
pkt_info.get("capture_sources"),
pkt_info.get("app_protocol"),
pkt_info.get("app_category"),
pkt_info.get("app_confidence"),
pkt_info.get("app_hostname"),
pkt_info.get("app_is_encrypted"),
pkt_info.get("app_risk_score"),
json.dumps(dpi_metadata) if dpi_metadata is not None else None,
json.dumps(pkt_info.get("capture_metadata")) if pkt_info.get("capture_metadata") is not None else None,
json.dumps(telemetry_metadata) if telemetry_metadata is not None else None,
json.dumps(capture_observations) if capture_observations is not None else None,
pkt_info.get("raw"),
)
def _packet_upsert_sql(self, *, returning: bool) -> str:
sql = """
INSERT INTO packets ( INSERT INTO packets (
timestamp, timestamp,
correlation_key, correlation_key,
@@ -425,49 +466,10 @@ class DatabasePool:
telemetry_metadata = COALESCE(EXCLUDED.telemetry_metadata, packets.telemetry_metadata), telemetry_metadata = COALESCE(EXCLUDED.telemetry_metadata, packets.telemetry_metadata),
capture_observations = COALESCE(EXCLUDED.capture_observations, packets.capture_observations), capture_observations = COALESCE(EXCLUDED.capture_observations, packets.capture_observations),
raw = COALESCE(EXCLUDED.raw, packets.raw) raw = COALESCE(EXCLUDED.raw, packets.raw)
RETURNING * """
""", if returning:
pkt_info.get("timestamp"), sql += "\nRETURNING *"
pkt_info["correlation_key"], return sql
pkt_info.get("packet_id"),
pkt_info.get("packet_uid"),
pkt_info.get("flow_id"),
pkt_info.get("capture_session_id"),
pkt_info.get("correlation_source"),
pkt_info.get("skb_mark"),
pkt_info.get("capture_iface"),
pkt_info.get("ingress_if"),
pkt_info.get("egress_if"),
pkt_info.get("verdict"),
pkt_info.get("verdict_reason"),
pkt_info.get("verdict_confidence"),
pkt_info.get("ingress_seen_at"),
pkt_info.get("egress_seen_at"),
pkt_info.get("verdict_seen_at"),
pkt_info.get("src_mac"),
pkt_info.get("dst_mac"),
pkt_info.get("eth_type_raw"),
pkt_info.get("vlan_id"),
pkt_info.get("src_ip"),
pkt_info.get("dst_ip"),
pkt_info.get("protocol_raw"),
pkt_info.get("src_port"),
pkt_info.get("dst_port"),
pkt_info.get("length"),
pkt_info.get("capture_sources"),
pkt_info.get("app_protocol"),
pkt_info.get("app_category"),
pkt_info.get("app_confidence"),
pkt_info.get("app_hostname"),
pkt_info.get("app_is_encrypted"),
pkt_info.get("app_risk_score"),
json.dumps(dpi_metadata) if dpi_metadata is not None else None,
json.dumps(pkt_info.get("capture_metadata")) if pkt_info.get("capture_metadata") is not None else None,
json.dumps(telemetry_metadata) if telemetry_metadata is not None else None,
json.dumps(capture_observations) if capture_observations is not None else None,
pkt_info.get("raw"),
)
return dict(row) if row else None
def _publish_packet_row(self, pkt_info: Dict[str, Any], row: Optional[Dict[str, Any]]) -> None: def _publish_packet_row(self, pkt_info: Dict[str, Any], row: Optional[Dict[str, Any]]) -> None:
if row: if row: