perf(api): ingest compliance overviews in a single transaction (#11875)

This commit is contained in:
Pedro Martín
2026-07-20 15:15:39 +02:00
committed by GitHub
parent e035e0ff62
commit 4e22289a19
3 changed files with 384 additions and 113 deletions
@@ -0,0 +1 @@
Compliance overview ingest now runs in a single transaction per scan with a configurable `COPY` batch size (`DJANGO_COMPLIANCE_COPY_BATCH_SIZE`, default 2000), reducing write pressure on the database
+177 -83
View File
@@ -6,7 +6,7 @@ import re
import time
import uuid
from collections import defaultdict
from collections.abc import Iterable
from collections.abc import Callable, Iterable
from datetime import UTC, datetime
from typing import Any
@@ -49,6 +49,7 @@ from celery.utils.log import get_task_logger
from config.django.base import DJANGO_FINDINGS_BATCH_SIZE
from config.env import env
from config.settings.celery import CELERY_DEADLOCK_ATTEMPTS
from django.core.exceptions import ImproperlyConfigured
from django.db import DatabaseError, IntegrityError, OperationalError, transaction
from django.db.models import (
Case,
@@ -99,6 +100,16 @@ COMPLIANCE_REQUIREMENT_COPY_COLUMNS = (
FINDINGS_MICRO_BATCH_SIZE = env.int("DJANGO_FINDINGS_MICRO_BATCH_SIZE", default=3000)
# Controls how many rows each ORM bulk_create/bulk_update call sends to Postgres.
SCAN_DB_BATCH_SIZE = env.int("DJANGO_SCAN_DB_BATCH_SIZE", default=1000)
# Rows per COPY statement when ingesting compliance requirement overviews. All
# batches of a scan share one transaction/commit; the batch size only bounds the
# client-side CSV buffer and how long each individual COPY statement runs on the
# writer (memory footprint, lock time and slow-statement logging under load).
COMPLIANCE_COPY_BATCH_SIZE = env.int("DJANGO_COMPLIANCE_COPY_BATCH_SIZE", default=2000)
if COMPLIANCE_COPY_BATCH_SIZE < 1:
raise ImproperlyConfigured(
"DJANGO_COMPLIANCE_COPY_BATCH_SIZE must be a positive integer, got "
f"{COMPLIANCE_COPY_BATCH_SIZE}"
)
# Throttle scan progress persistence: minimum progress delta (fraction 0-1)
# between two persisted progress updates.
PROGRESS_THROTTLE_DELTA = env.float("DJANGO_SCAN_PROGRESS_THROTTLE_DELTA", default=0.01)
@@ -356,30 +367,36 @@ def _bulk_update_resource_failed_findings_counts(
raise
def _copy_compliance_requirement_rows(
tenant_id: str, rows: list[dict[str, Any]]
) -> None:
"""Stream compliance requirement rows into Postgres using COPY.
class ComplianceRowScopeError(ValueError):
"""A compliance requirement row does not belong to the scan being ingested."""
We leverage the admin connection (when available) to bypass the COPY + RLS
restriction, writing only the fields required by
``ComplianceRequirementOverview``.
Args:
tenant_id: Target tenant UUID.
rows: List of row dictionaries prepared by
:func:`create_compliance_requirements`.
def _compliance_requirement_rows_to_csv(
rows: list[dict[str, Any]], tenant_id: str, scan_id: str
) -> io.StringIO:
"""Serialize compliance requirement rows into a CSV buffer for COPY.
COPY runs on the admin connection, which bypasses RLS, so every row is
checked against the expected tenant/scan before it is written: a mismatched
row would otherwise be inserted verbatim into another tenant's data.
"""
csv_buffer = io.StringIO()
writer = csv.writer(csv_buffer)
datetime_now = datetime.now(tz=UTC)
for row in rows:
row_tenant_id = str(row.get("tenant_id"))
row_scan_id = str(row.get("scan_id"))
if row_tenant_id != tenant_id or row_scan_id != scan_id:
raise ComplianceRowScopeError(
"Compliance requirement row does not belong to the scan being "
f"ingested (expected tenant {tenant_id} / scan {scan_id}, got "
f"tenant {row_tenant_id} / scan {row_scan_id})"
)
writer.writerow(
[
str(row.get("id")),
str(row.get("tenant_id")),
row_tenant_id,
(row.get("inserted_at") or datetime_now).isoformat(),
row.get("compliance_id") or "",
row.get("framework") or "",
@@ -393,65 +410,100 @@ def _copy_compliance_requirement_rows(
row.get("total_checks", 0),
row.get("passed_findings", 0),
row.get("total_findings", 0),
str(row.get("scan_id")),
row_scan_id,
]
)
csv_buffer.seek(0)
return csv_buffer
def _copy_compliance_requirement_rows(
tenant_id: str, scan_id: str, rows: Iterable[dict[str, Any]], batch_size: int
) -> int:
"""Replace a scan's compliance requirement rows using batched COPY.
We leverage the admin connection (when available) to bypass the COPY + RLS
restriction. The scan's DELETE and every COPY batch run on one connection
inside a single transaction with a single commit, so the writer takes one
fsync per scan instead of one per batch, and a failed ingest rolls back
without committing a partial delete/insert (which a retry would otherwise
delete again, feeding dead rows to autovacuum).
Args:
tenant_id: Target tenant UUID.
scan_id: Scan whose previous rows are replaced.
rows: Iterable of row dictionaries, consumed lazily batch by batch.
batch_size: Number of rows per COPY statement.
Returns:
int: total number of rows staged and committed.
Raises:
ComplianceRowScopeError: A row belongs to another tenant or scan.
"""
# Normalized once so the per-row scope check compares like with like even if
# the caller passes UUID instances instead of strings.
tenant_id = str(tenant_id)
scan_id = str(scan_id)
total_rows = 0
batch_num = 0
copy_sql = (
"COPY compliance_requirements_overviews ("
+ ", ".join(COMPLIANCE_REQUIREMENT_COPY_COLUMNS)
+ ") FROM STDIN WITH (FORMAT CSV, DELIMITER ',', QUOTE '\"', ESCAPE '\"', NULL '\\N')"
)
try:
with psycopg_connection(MainRouter.admin_db) as connection:
connection.autocommit = False
try:
with connection.cursor() as cursor:
cursor.execute(SET_CONFIG_QUERY, [POSTGRES_TENANT_VAR, tenant_id])
cursor.copy_expert(copy_sql, csv_buffer)
connection.commit()
except Exception:
connection.rollback()
raise
finally:
csv_buffer.close()
with psycopg_connection(MainRouter.admin_db) as connection:
connection.autocommit = False
try:
with connection.cursor() as cursor:
cursor.execute(SET_CONFIG_QUERY, [POSTGRES_TENANT_VAR, tenant_id])
# Idempotent re-run: clearing this scan's rows inside the same
# transaction keeps delete + reinsert atomic.
cursor.execute(
"DELETE FROM compliance_requirements_overviews "
"WHERE tenant_id = %s AND scan_id = %s",
[tenant_id, scan_id],
)
for batch, _is_last in batched(rows, batch_size):
if not batch:
continue
batch_num += 1
csv_buffer = _compliance_requirement_rows_to_csv(
batch, tenant_id, scan_id
)
try:
cursor.copy_expert(copy_sql, csv_buffer)
finally:
csv_buffer.close()
total_rows += len(batch)
logger.info(
f"Compliance COPY batch {batch_num}: staged {len(batch)} rows "
f"({total_rows} total)"
)
connection.commit()
except Exception:
connection.rollback()
raise
return total_rows
def _persist_compliance_requirement_rows(
tenant_id: str, rows: Iterable[dict[str, Any]], batch_size: int = 10000
def _bulk_create_compliance_requirement_rows(
tenant_id: str, scan_id: str, rows: Iterable[dict[str, Any]], batch_size: int
) -> int:
"""Persist compliance requirement rows using batched COPY with ORM fallback.
"""Replace a scan's compliance requirement rows via the ORM.
``rows`` is consumed lazily in batches, so peak memory stays at ~``batch_size``
rows instead of the full set. A batch that fails COPY falls back to an ORM
``bulk_create`` of just that batch.
Args:
tenant_id: Target tenant UUID.
rows: Iterable of row dictionaries reflecting the compliance overview
state for a scan.
batch_size: Number of rows per COPY batch (default: 10000).
Returns:
int: total number of rows persisted.
Fallback for when COPY is unavailable; the delete and every ``bulk_create``
share one RLS transaction so the replacement stays atomic.
"""
total_rows = 0
batch_num = 0
for batch, _is_last in batched(rows, batch_size):
if not batch:
continue
batch_num += 1
try:
_copy_compliance_requirement_rows(tenant_id, batch)
except Exception as error:
logger.exception(
f"COPY bulk insert for compliance requirements batch {batch_num} "
"failed; falling back to ORM bulk_create for this batch",
exc_info=error,
)
with rls_transaction(tenant_id):
ComplianceRequirementOverview.objects.filter(scan_id=scan_id).delete()
for batch, _is_last in batched(rows, batch_size):
if not batch:
continue
fallback_objects = [
ComplianceRequirementOverview(
id=row["id"],
@@ -473,20 +525,58 @@ def _persist_compliance_requirement_rows(
)
for row in batch
]
with rls_transaction(tenant_id):
ComplianceRequirementOverview.objects.bulk_create(
fallback_objects, batch_size=500
)
total_rows += len(batch)
logger.info(
f"Compliance COPY batch {batch_num}: inserted {len(batch)} rows "
f"({total_rows} total)"
)
ComplianceRequirementOverview.objects.bulk_create(
fallback_objects, batch_size=500
)
total_rows += len(batch)
return total_rows
def _persist_compliance_requirement_rows(
tenant_id: str,
scan_id: str,
rows_factory: Callable[[], Iterable[dict[str, Any]]],
batch_size: int | None = None,
) -> int:
"""Persist a scan's compliance requirement rows, replacing any previous ones.
``rows_factory`` must return a fresh row iterator on every call: the COPY
path consumes it lazily in batches (peak memory ~``batch_size`` rows), and
if COPY fails the whole ingest falls back to a single ORM transaction that
re-iterates the rows.
Args:
tenant_id: Target tenant UUID.
scan_id: Scan whose compliance overview rows are being replaced.
rows_factory: Callable returning an iterable of row dictionaries.
batch_size: Rows per COPY/bulk_create batch (default:
``COMPLIANCE_COPY_BATCH_SIZE``).
Returns:
int: total number of rows persisted.
"""
if batch_size is None:
batch_size = COMPLIANCE_COPY_BATCH_SIZE
try:
return _copy_compliance_requirement_rows(
tenant_id, scan_id, rows_factory(), batch_size
)
except ComplianceRowScopeError:
# Cross-tenant/scan rows are a bug in the caller, not a COPY failure:
# retrying through the ORM would persist the very rows we rejected.
raise
except Exception as error:
logger.exception(
"COPY bulk insert for compliance requirements failed; "
"falling back to ORM bulk_create",
exc_info=error,
)
return _bulk_create_compliance_requirement_rows(
tenant_id, scan_id, rows_factory(), batch_size
)
def _create_compliance_summaries(
tenant_id: str, scan_id: str, requirement_statuses: dict
) -> None:
@@ -885,15 +975,19 @@ def _process_finding_micro_batch(
# Denormalized resource arrays populated directly on insert
# (was previously a separate bulk_update; saves a CASE WHEN
# over thousands of rows per micro-batch).
resource_regions=[resource_instance.region]
if resource_instance.region
else [],
resource_services=[resource_instance.service]
if resource_instance.service
else [],
resource_types=[resource_instance.type]
if resource_instance.type
else [],
resource_regions=(
[resource_instance.region]
if resource_instance.region
else []
),
resource_services=(
[resource_instance.service]
if resource_instance.service
else []
),
resource_types=(
[resource_instance.type] if resource_instance.type else []
),
)
findings_to_create.append(finding_instance)
resource_denormalized_data.append(
@@ -1708,8 +1802,10 @@ def create_compliance_requirements(tenant_id: str, scan_id: str):
)
# Yield rows lazily (consumed batch-by-batch by COPY) so peak memory
# stays bounded; tally requirement_statuses in the same pass.
# stays bounded; tally requirement_statuses in the same pass. The
# ORM fallback re-iterates from scratch, so the tally resets first.
def _iter_compliance_requirement_rows():
requirement_statuses.clear()
for region in regions:
region_stats = region_requirement_stats.get(region, {})
region_findings = findings_count_by_compliance.get(region, {})
@@ -1773,12 +1869,10 @@ def create_compliance_requirements(tenant_id: str, scan_id: str):
"total_findings": total_findings,
}
# Idempotent re-run: clear this scan's rows before re-inserting.
with rls_transaction(tenant_id):
ComplianceRequirementOverview.objects.filter(scan_id=scan_id).delete()
# The delete of the scan's previous rows happens inside the same
# transaction as the inserts (see _copy_compliance_requirement_rows).
requirements_created = _persist_compliance_requirement_rows(
tenant_id, _iter_compliance_requirement_rows()
tenant_id_str, scan_id_str, _iter_compliance_requirement_rows
)
# Create pre-aggregated summaries for fast compliance overview lookups
+206 -30
View File
@@ -26,6 +26,7 @@ from prowler.lib.check.models import Severity
from prowler.lib.outputs.finding import Status
from tasks.jobs.scan import (
_ATTACK_SURFACE_MAPPING_CACHE,
ComplianceRowScopeError,
_aggregate_findings_by_region,
_bulk_update_resource_failed_findings_counts,
_copy_compliance_requirement_rows,
@@ -2314,9 +2315,9 @@ class TestCreateComplianceRequirements:
create_compliance_requirements(tenant_id, scan_id)
mock_persist.assert_called_once()
persisted_rows = mock_persist.call_args[0][1]
rows_factory = mock_persist.call_args[0][2]
requirement_row = next(
row for row in persisted_rows if row["requirement_id"] == "1.1"
row for row in rows_factory() if row["requirement_id"] == "1.1"
)
assert requirement_row["requirement_status"] == "FAIL"
@@ -2454,18 +2455,26 @@ class TestComplianceRequirementCopy:
}
with patch.object(MainRouter, "admin_db", "admin"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
mock_psycopg_connection.assert_called_once_with("admin")
connection.cursor.assert_called_once()
cursor.execute.assert_called_once()
# One execute for set_config plus one for the scan's DELETE.
assert cursor.execute.call_count == 2
delete_sql, delete_params = cursor.execute.call_args_list[1][0]
assert "DELETE FROM compliance_requirements_overviews" in delete_sql
assert delete_params == [str(row["tenant_id"]), str(row["scan_id"])]
cursor.copy_expert.assert_called_once()
connection.commit.assert_called_once()
csv_rows = list(csv.reader(StringIO(captured["data"])))
assert csv_rows[0][0] == str(row["id"])
assert csv_rows[0][5] == ""
assert csv_rows[0][-1] == str(row["scan_id"])
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.filter")
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.bulk_create")
@patch("tasks.jobs.scan.rls_transaction")
@patch(
@@ -2473,7 +2482,7 @@ class TestComplianceRequirementCopy:
side_effect=Exception("copy failed"),
)
def test_persist_compliance_requirement_rows_fallback(
self, mock_copy, mock_rls_transaction, mock_bulk_create
self, mock_copy, mock_rls_transaction, mock_bulk_create, mock_filter
):
inserted_at = datetime.now(UTC)
row = {
@@ -2494,16 +2503,22 @@ class TestComplianceRequirementCopy:
}
tenant_id = row["tenant_id"]
scan_id = str(row["scan_id"])
ctx = MagicMock()
ctx.__enter__.return_value = None
ctx.__exit__.return_value = False
mock_rls_transaction.return_value = ctx
_persist_compliance_requirement_rows(tenant_id, [row])
_persist_compliance_requirement_rows(tenant_id, scan_id, lambda: [row])
mock_copy.assert_called_once_with(tenant_id, [row])
mock_copy.assert_called_once()
assert mock_copy.call_args[0][0] == tenant_id
assert mock_copy.call_args[0][1] == scan_id
mock_rls_transaction.assert_called_once_with(tenant_id)
# The fallback replaces the scan's rows: delete + insert atomically.
mock_filter.assert_called_once_with(scan_id=scan_id)
mock_filter.return_value.delete.assert_called_once()
mock_bulk_create.assert_called_once()
args, kwargs = mock_bulk_create.call_args
@@ -2515,13 +2530,18 @@ class TestComplianceRequirementCopy:
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.bulk_create")
@patch("tasks.jobs.scan.rls_transaction")
@patch("tasks.jobs.scan._copy_compliance_requirement_rows")
@patch("tasks.jobs.scan._copy_compliance_requirement_rows", return_value=0)
def test_persist_compliance_requirement_rows_no_rows(
self, mock_copy, mock_rls_transaction, mock_bulk_create
):
_persist_compliance_requirement_rows(str(uuid.uuid4()), [])
# Even with no rows the COPY path runs: it must clear the scan's
# previous rows so a re-run with fewer findings drops stale data.
total = _persist_compliance_requirement_rows(
str(uuid.uuid4()), str(uuid.uuid4()), lambda: []
)
mock_copy.assert_not_called()
assert total == 0
mock_copy.assert_called_once()
mock_rls_transaction.assert_not_called()
mock_bulk_create.assert_not_called()
@@ -2610,11 +2630,12 @@ class TestComplianceRequirementCopy:
]
with patch.object(MainRouter, "admin_db", "admin"):
_copy_compliance_requirement_rows(tenant_id, rows)
_copy_compliance_requirement_rows(tenant_id, str(scan_id), rows, 2000)
mock_psycopg_connection.assert_called_once_with("admin")
connection.cursor.assert_called_once()
cursor.execute.assert_called_once()
# set_config + DELETE of the scan's previous rows.
assert cursor.execute.call_count == 2
cursor.copy_expert.assert_called_once()
csv_rows = list(csv.reader(StringIO(captured["data"])))
@@ -2644,6 +2665,60 @@ class TestComplianceRequirementCopy:
assert csv_rows[2][5] == "2.0"
assert csv_rows[2][9] == "MANUAL"
@patch("tasks.jobs.scan.psycopg_connection")
def test_copy_compliance_requirement_rows_batches_share_one_transaction(
self, mock_psycopg_connection, settings
):
"""Every COPY batch runs on the same connection with a single commit."""
settings.DATABASES.setdefault("admin", settings.DATABASES["default"])
connection = MagicMock()
cursor = MagicMock()
cursor_context = MagicMock()
cursor_context.__enter__.return_value = cursor
cursor_context.__exit__.return_value = False
connection.cursor.return_value = cursor_context
connection.__enter__.return_value = connection
connection.__exit__.return_value = False
context_manager = MagicMock()
context_manager.__enter__.return_value = connection
context_manager.__exit__.return_value = False
mock_psycopg_connection.return_value = context_manager
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
inserted_at = datetime.now(UTC)
rows = [
{
"id": uuid.uuid4(),
"tenant_id": tenant_id,
"inserted_at": inserted_at,
"compliance_id": "cisa_aws",
"framework": "CISA",
"version": "1.0",
"description": f"Requirement {index}",
"region": "us-east-1",
"requirement_id": f"req-{index}",
"requirement_status": "PASS",
"passed_checks": 1,
"failed_checks": 0,
"total_checks": 1,
"scan_id": scan_id,
}
for index in range(3)
]
with patch.object(MainRouter, "admin_db", "admin"):
total = _copy_compliance_requirement_rows(tenant_id, scan_id, rows, 1)
assert total == 3
# One connection, three COPY statements, one commit for the whole scan.
mock_psycopg_connection.assert_called_once_with("admin")
assert cursor.copy_expert.call_count == 3
connection.commit.assert_called_once()
connection.rollback.assert_not_called()
@patch("tasks.jobs.scan.psycopg_connection")
def test_copy_compliance_requirement_rows_null_values(
self, mock_psycopg_connection, settings
@@ -2691,7 +2766,9 @@ class TestComplianceRequirementCopy:
}
with patch.object(MainRouter, "admin_db", "admin"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
csv_rows = list(csv.reader(StringIO(captured["data"])))
assert len(csv_rows) == 1
@@ -2747,7 +2824,9 @@ class TestComplianceRequirementCopy:
}
with patch.object(MainRouter, "admin_db", "admin"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
# Verify CSV was generated (csv module handles escaping automatically)
csv_rows = list(csv.reader(StringIO(captured["data"])))
@@ -2808,7 +2887,9 @@ class TestComplianceRequirementCopy:
before_call = datetime.now(UTC)
with patch.object(MainRouter, "admin_db", "admin"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
after_call = datetime.now(UTC)
csv_rows = list(csv.reader(StringIO(captured["data"])))
@@ -2861,12 +2942,84 @@ class TestComplianceRequirementCopy:
with patch.object(MainRouter, "admin_db", "admin"):
with pytest.raises(Exception, match="COPY command failed"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
# Verify rollback was called
connection.rollback.assert_called_once()
connection.commit.assert_not_called()
@pytest.mark.parametrize("mismatched_field", ["tenant_id", "scan_id"])
@patch("tasks.jobs.scan.psycopg_connection")
def test_copy_compliance_requirement_rows_rejects_out_of_scope_rows(
self, mock_psycopg_connection, mismatched_field, settings
):
"""COPY bypasses RLS, so rows from another tenant/scan must be rejected."""
settings.DATABASES.setdefault("admin", settings.DATABASES["default"])
connection = MagicMock()
cursor = MagicMock()
cursor_context = MagicMock()
cursor_context.__enter__.return_value = cursor
cursor_context.__exit__.return_value = False
connection.cursor.return_value = cursor_context
connection.__enter__.return_value = connection
connection.__exit__.return_value = False
context_manager = MagicMock()
context_manager.__enter__.return_value = connection
context_manager.__exit__.return_value = False
mock_psycopg_connection.return_value = context_manager
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
row = {
"id": uuid.uuid4(),
"tenant_id": tenant_id,
"compliance_id": "test",
"framework": "Test",
"version": "1.0",
"description": "desc",
"region": "us-east-1",
"requirement_id": "req-1",
"requirement_status": "PASS",
"passed_checks": 1,
"failed_checks": 0,
"total_checks": 1,
"scan_id": scan_id,
}
row[mismatched_field] = str(uuid.uuid4())
with patch.object(MainRouter, "admin_db", "admin"):
with pytest.raises(ComplianceRowScopeError):
_copy_compliance_requirement_rows(tenant_id, scan_id, [row], 2000)
cursor.copy_expert.assert_not_called()
connection.rollback.assert_called_once()
connection.commit.assert_not_called()
@patch("tasks.jobs.scan.ComplianceRequirementOverview")
@patch("tasks.jobs.scan.rls_transaction")
@patch(
"tasks.jobs.scan._copy_compliance_requirement_rows",
side_effect=ComplianceRowScopeError("out of scope"),
)
def test_persist_compliance_requirement_rows_does_not_fall_back_on_scope_error(
self, mock_copy, mock_rls_transaction, mock_model
):
"""A scope violation is a caller bug: the ORM fallback must not persist it."""
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
with pytest.raises(ComplianceRowScopeError):
_persist_compliance_requirement_rows(tenant_id, scan_id, lambda: [])
mock_copy.assert_called_once()
mock_rls_transaction.assert_not_called()
mock_model.objects.filter.assert_not_called()
mock_model.objects.bulk_create.assert_not_called()
@patch("tasks.jobs.scan.psycopg_connection")
def test_copy_compliance_requirement_rows_transaction_rollback_on_set_config_error(
self, mock_psycopg_connection, settings
@@ -2909,7 +3062,9 @@ class TestComplianceRequirementCopy:
with patch.object(MainRouter, "admin_db", "admin"):
with pytest.raises(Exception, match="SET prowler.tenant_id failed"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
# Verify rollback was called
connection.rollback.assert_called_once()
@@ -2955,7 +3110,9 @@ class TestComplianceRequirementCopy:
}
with patch.object(MainRouter, "admin_db", "admin"):
_copy_compliance_requirement_rows(str(row["tenant_id"]), [row])
_copy_compliance_requirement_rows(
str(row["tenant_id"]), str(row["scan_id"]), [row], 2000
)
# Verify commit was called and rollback was not
connection.commit.assert_called_once()
@@ -2966,9 +3123,10 @@ class TestComplianceRequirementCopy:
@patch("tasks.jobs.scan._copy_compliance_requirement_rows")
def test_persist_compliance_requirement_rows_success(self, mock_copy):
"""Test successful COPY path without fallback to ORM."""
mock_copy.return_value = None # Success, no exception
mock_copy.return_value = 1 # Success, no exception
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
rows = [
{
"id": uuid.uuid4(),
@@ -2984,16 +3142,21 @@ class TestComplianceRequirementCopy:
"passed_checks": 1,
"failed_checks": 0,
"total_checks": 1,
"scan_id": uuid.uuid4(),
"scan_id": scan_id,
}
]
_persist_compliance_requirement_rows(tenant_id, rows)
total = _persist_compliance_requirement_rows(tenant_id, scan_id, lambda: rows)
# Verify COPY was called
mock_copy.assert_called_once_with(tenant_id, rows)
assert total == 1
mock_copy.assert_called_once()
copy_args = mock_copy.call_args[0]
assert copy_args[0] == tenant_id
assert copy_args[1] == scan_id
assert list(copy_args[2]) == rows
@patch("tasks.jobs.scan.logger")
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.filter")
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.bulk_create")
@patch("tasks.jobs.scan.rls_transaction")
@patch(
@@ -3001,7 +3164,12 @@ class TestComplianceRequirementCopy:
side_effect=Exception("COPY failed"),
)
def test_persist_compliance_requirement_rows_fallback_logging(
self, mock_copy, mock_rls_transaction, mock_bulk_create, mock_logger
self,
mock_copy,
mock_rls_transaction,
mock_bulk_create,
mock_filter,
mock_logger,
):
"""Test logger.exception is called when COPY fails and fallback occurs."""
tenant_id = str(uuid.uuid4())
@@ -3027,7 +3195,9 @@ class TestComplianceRequirementCopy:
ctx.__exit__.return_value = False
mock_rls_transaction.return_value = ctx
_persist_compliance_requirement_rows(tenant_id, [row])
_persist_compliance_requirement_rows(
tenant_id, str(row["scan_id"]), lambda: [row]
)
# Verify logger.exception was called
mock_logger.exception.assert_called_once()
@@ -3036,6 +3206,7 @@ class TestComplianceRequirementCopy:
assert "falling back to ORM" in args[0]
assert kwargs.get("exc_info") is not None
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.filter")
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.bulk_create")
@patch("tasks.jobs.scan.rls_transaction")
@patch(
@@ -3043,7 +3214,7 @@ class TestComplianceRequirementCopy:
side_effect=Exception("copy failed"),
)
def test_persist_compliance_requirement_rows_fallback_multiple_rows(
self, mock_copy, mock_rls_transaction, mock_bulk_create
self, mock_copy, mock_rls_transaction, mock_bulk_create, mock_filter
):
"""Test ORM fallback with multiple rows."""
tenant_id = str(uuid.uuid4())
@@ -3090,10 +3261,14 @@ class TestComplianceRequirementCopy:
ctx.__exit__.return_value = False
mock_rls_transaction.return_value = ctx
_persist_compliance_requirement_rows(tenant_id, rows)
total = _persist_compliance_requirement_rows(
tenant_id, str(scan_id), lambda: rows
)
mock_copy.assert_called_once_with(tenant_id, rows)
assert total == 2
mock_copy.assert_called_once()
mock_rls_transaction.assert_called_once_with(tenant_id)
mock_filter.assert_called_once_with(scan_id=str(scan_id))
mock_bulk_create.assert_called_once()
args, kwargs = mock_bulk_create.call_args
@@ -3117,6 +3292,7 @@ class TestComplianceRequirementCopy:
assert objects[1].passed_checks == 2
assert objects[1].failed_checks == 3
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.filter")
@patch("tasks.jobs.scan.ComplianceRequirementOverview.objects.bulk_create")
@patch("tasks.jobs.scan.rls_transaction")
@patch(
@@ -3124,7 +3300,7 @@ class TestComplianceRequirementCopy:
side_effect=Exception("copy failed"),
)
def test_persist_compliance_requirement_rows_fallback_all_fields(
self, mock_copy, mock_rls_transaction, mock_bulk_create
self, mock_copy, mock_rls_transaction, mock_bulk_create, mock_filter
):
"""Test ORM fallback correctly maps all fields from row dict to model."""
tenant_id = str(uuid.uuid4())
@@ -3154,7 +3330,7 @@ class TestComplianceRequirementCopy:
ctx.__exit__.return_value = False
mock_rls_transaction.return_value = ctx
_persist_compliance_requirement_rows(tenant_id, [row])
_persist_compliance_requirement_rows(tenant_id, str(scan_id), lambda: [row])
args, kwargs = mock_bulk_create.call_args
objects = args[0]