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