From c40ddf4ded7ad16ea248b038facd373d9f13315d Mon Sep 17 00:00:00 2001 From: malmert Date: Sun, 3 May 2026 20:11:42 +0200 Subject: [PATCH] overload batching improvements --- backend/src/utilities/database.py | 108 +++++++++++++++--------------- 1 file changed, 55 insertions(+), 53 deletions(-) diff --git a/backend/src/utilities/database.py b/backend/src/utilities/database.py index e38eb16..1032fc2 100644 --- a/backend/src/utilities/database.py +++ b/backend/src/utilities/database.py @@ -297,7 +297,7 @@ class DatabasePool: self._publish_packet_row(pkt_info, row) 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: return if self._pool is None: @@ -305,21 +305,20 @@ class DatabasePool: if self._pool is None: return - published_rows: List[tuple[Dict[str, Any], Optional[Dict[str, Any]]]] = [] try: + params = [self._packet_upsert_params(pkt_info) for pkt_info in packets] async with self._pool.acquire() as conn: - for pkt_info in packets: - row = await self._upsert_packet_with_conn(conn, pkt_info) - published_rows.append((pkt_info, row)) + await conn.executemany(self._packet_upsert_sql(returning=False), params) except Exception: logger.exception("DB packet batch upsert failed") 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]]: """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) _attach_derived_fields(pkt_info) @@ -327,8 +326,50 @@ class DatabasePool: telemetry_metadata = pkt_info.get("telemetry_metadata") 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 ( timestamp, correlation_key, @@ -425,49 +466,10 @@ class DatabasePool: telemetry_metadata = COALESCE(EXCLUDED.telemetry_metadata, packets.telemetry_metadata), capture_observations = COALESCE(EXCLUDED.capture_observations, packets.capture_observations), raw = COALESCE(EXCLUDED.raw, packets.raw) - RETURNING * - """, - 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"), - ) - return dict(row) if row else None + """ + if returning: + sql += "\nRETURNING *" + return sql def _publish_packet_row(self, pkt_info: Dict[str, Any], row: Optional[Dict[str, Any]]) -> None: if row: