overload batching improvements
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user