fix(api): recover missing scan resources (#12089)

Co-authored-by: Hugo Pereira Brito <101209179+HugoPBrito@users.noreply.github.com>
Co-authored-by: Hugo P.Brito <hugopbrit@gmail.com>
This commit is contained in:
Prowler Bot
2026-07-22 14:10:36 +02:00
committed by GitHub
parent 0772169b5c
commit b313a6dde1
3 changed files with 797 additions and 42 deletions
@@ -0,0 +1 @@
Scan findings now recover resources missing from the in-memory cache after resource pre-resolution, preventing valid findings from being skipped
+150 -20
View File
@@ -1,3 +1,4 @@
import copy
import csv
import io
import json
@@ -605,6 +606,45 @@ def _process_finding_micro_batch(
scan_resource_groups_cache: Dict tracking resource group counts {(resource_group, severity): {"total", "failed", "new_failed"}}.
group_resources_cache: Dict tracking unique resources per group {resource_group: set(resource_uids)}.
"""
def build_resource_defaults_from_finding(finding: ProwlerFinding) -> dict[str, Any]:
check_metadata = finding.get_metadata()
group = check_metadata.get("resourcegroup") or None
return {
"tenant_id": tenant_id,
"provider": provider_instance,
"uid": finding.resource_uid,
"region": finding.region,
"service": finding.service_name,
"type": finding.resource_type,
"name": finding.resource_name,
"groups": [group] if group else None,
}
def recover_resource_after_cache_miss(finding: ProwlerFinding) -> Resource:
resource_uid = finding.resource_uid
resource_instance = Resource.objects.filter(
tenant_id=tenant_id,
provider_id=provider_instance.id,
uid=resource_uid,
).first()
if resource_instance is None:
try:
with transaction.atomic():
resource_instance = Resource.objects.create(
**build_resource_defaults_from_finding(finding)
)
except IntegrityError:
resource_instance = Resource.objects.filter(
tenant_id=tenant_id,
provider_id=provider_instance.id,
uid=resource_uid,
).first()
if resource_instance is None:
raise
return cache_resource(resource_uid, resource_instance)
# Accumulate objects for bulk operations
findings_to_create = []
dirty_resources = {}
@@ -643,7 +683,103 @@ def _process_finding_micro_batch(
# All DB writes for this micro-batch run inside ONE rls_transaction,
# with deadlock-retry at micro-batch granularity instead of per-finding.
missing_cache_value = object()
for attempt in range(CELERY_DEADLOCK_ATTEMPTS):
resource_cache_originals: dict[str, Resource | object] = {}
failed_count_originals: dict[str, int | None] = {}
resource_field_originals: dict[str, dict[str, Any]] = {}
tag_cache_original = dict(tag_cache)
scan_resource_cache_original = set(scan_resource_cache)
scan_categories_cache_original = {
key: value.copy() for key, value in scan_categories_cache.items()
}
scan_resource_groups_cache_original = {
key: value.copy() for key, value in scan_resource_groups_cache.items()
}
group_resources_cache_original = {
key: set(value) for key, value in group_resources_cache.items()
}
def cache_resource(resource_uid: str, resource_instance: Resource) -> Resource:
if resource_uid not in resource_cache_originals:
resource_cache_originals[resource_uid] = resource_cache.get(
resource_uid, missing_cache_value
)
resource_cache[resource_uid] = resource_instance
if resource_uid not in resource_failed_findings_cache:
failed_count_originals[resource_uid] = None
resource_failed_findings_cache[resource_uid] = 0
return resource_instance
def snapshot_failed_count(resource_uid: str) -> None:
if resource_uid not in failed_count_originals:
failed_count_originals[resource_uid] = (
resource_failed_findings_cache.get(resource_uid)
)
def snapshot_resource_fields(
resource_uid: str, resource_instance: Resource
) -> None:
if resource_uid in resource_field_originals:
return
resource_field_originals[resource_uid] = {
field: copy.deepcopy(getattr(resource_instance, field))
for field in (
"name",
"metadata",
"details",
"partition",
"region",
"service",
"type",
"groups",
"updated_at",
)
}
def restore_attempt_caches() -> None:
for resource_uid, original_fields in resource_field_originals.items():
resource_instance = resource_cache.get(resource_uid)
if resource_instance is None:
continue
for field, value in original_fields.items():
setattr(resource_instance, field, value)
for resource_uid, original_resource in resource_cache_originals.items():
if original_resource is missing_cache_value:
resource_cache.pop(resource_uid, None)
else:
resource_cache[resource_uid] = original_resource
for resource_uid, original_count in failed_count_originals.items():
if original_count is None:
resource_failed_findings_cache.pop(resource_uid, None)
else:
resource_failed_findings_cache[resource_uid] = original_count
tag_cache.clear()
tag_cache.update(tag_cache_original)
scan_resource_cache.clear()
scan_resource_cache.update(scan_resource_cache_original)
scan_categories_cache.clear()
scan_categories_cache.update(
{
key: value.copy()
for key, value in scan_categories_cache_original.items()
}
)
scan_resource_groups_cache.clear()
scan_resource_groups_cache.update(
{
key: value.copy()
for key, value in scan_resource_groups_cache_original.items()
}
)
group_resources_cache.clear()
group_resources_cache.update(
{
key: set(value)
for key, value in group_resources_cache_original.items()
}
)
try:
with rls_transaction(tenant_id):
# 1) Pre-resolve Resources in bulk
@@ -678,19 +814,8 @@ def _process_finding_micro_batch(
resources_to_create = []
for uid in missing_uids:
f = first_finding_per_uid[uid]
check_metadata = f.get_metadata()
group = check_metadata.get("resourcegroup") or None
resources_to_create.append(
Resource(
tenant_id=tenant_id,
provider=provider_instance,
uid=uid,
region=f.region,
service=f.service_name,
type=f.resource_type,
name=f.resource_name,
groups=[group] if group else None,
)
Resource(**build_resource_defaults_from_finding(f))
)
Resource.objects.bulk_create(
resources_to_create,
@@ -711,8 +836,7 @@ def _process_finding_micro_batch(
}
)
for uid, r in existing_resources.items():
resource_cache[uid] = r
resource_failed_findings_cache.setdefault(uid, 0)
cache_resource(uid, r)
# 2) Pre-resolve ResourceTags in bulk
batch_tag_kv: set[tuple[str, str]] = set()
@@ -758,47 +882,50 @@ def _process_finding_micro_batch(
resource_uid = finding.resource_uid
resource_instance = resource_cache.get(resource_uid)
if resource_instance is None:
# Should be unreachable after the pre-resolve step. Defensive log.
logger.error(
f"Resource {resource_uid} missing from cache after pre-resolve "
f"on scan {scan_instance.id}; skipping finding."
)
continue
resource_instance = recover_resource_after_cache_miss(finding)
# Detect resource field changes (defer save until end-of-batch bulk_update).
check_metadata = finding.get_metadata()
group = check_metadata.get("resourcegroup") or None
updated = False
if finding.region and resource_instance.region != finding.region:
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.region = finding.region
updated = True
if (
finding.resource_name
and resource_instance.name != finding.resource_name
):
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.name = finding.resource_name
updated = True
if resource_instance.service != finding.service_name:
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.service = finding.service_name
updated = True
if resource_instance.type != finding.resource_type:
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.type = finding.resource_type
updated = True
if resource_instance.metadata != finding.resource_metadata:
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.metadata = json.dumps(
finding.resource_metadata, cls=CustomEncoder
)
updated = True
if resource_instance.details != finding.resource_details:
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.details = finding.resource_details
updated = True
if resource_instance.partition != finding.partition:
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.partition = finding.partition
updated = True
if group and (
not resource_instance.groups
or group not in resource_instance.groups
):
snapshot_resource_fields(resource_uid, resource_instance)
resource_instance.groups = (resource_instance.groups or []) + [
group
]
@@ -860,6 +987,7 @@ def _process_finding_micro_batch(
muted_reason = mute_rules_cache[finding_uid]
if status == FindingStatus.FAIL and not is_muted:
snapshot_failed_count(resource_uid)
resource_failed_findings_cache[resource_uid] += 1
check_metadata["compliance"] = finding.compliance
@@ -1013,6 +1141,7 @@ def _process_finding_micro_batch(
if r is None:
continue
# Manually bump updated_at since bulk_update bypasses auto_now.
snapshot_resource_fields(uid, r)
r.updated_at = now_utc
resources_to_bulk_update.append(r)
if resources_to_bulk_update:
@@ -1034,6 +1163,7 @@ def _process_finding_micro_batch(
# Successful execution: leave deadlock retry loop.
break
except (OperationalError, IntegrityError) as db_err:
restore_attempt_caches()
if attempt < CELERY_DEADLOCK_ATTEMPTS - 1:
logger.warning(
f"{'Deadlock error' if isinstance(db_err, OperationalError) else 'Integrity error'} "
+646 -22
View File
@@ -2,6 +2,7 @@ import csv
import json
import re
import uuid
from collections.abc import MutableMapping
from contextlib import contextmanager
from datetime import UTC, datetime
from io import StringIO
@@ -15,13 +16,16 @@ from api.models import (
MuteRule,
Provider,
Resource,
ResourceFindingMapping,
ResourceScanSummary,
ResourceTag,
ResourceTagMapping,
Scan,
ScanSummary,
StateChoices,
StatusChoices,
)
from django.db import IntegrityError, OperationalError
from django.db import IntegrityError, OperationalError, transaction
from prowler.lib.check.models import Severity
from prowler.lib.outputs.finding import Status
from tasks.jobs.scan import (
@@ -52,6 +56,12 @@ def noop_rls_transaction(*args, **kwargs):
yield
@contextmanager
def atomic_rls_transaction(*args, **kwargs):
with transaction.atomic():
yield
class FakeFinding:
def __init__(self, **attrs):
self.metadata = attrs.pop("metadata", {})
@@ -70,6 +80,32 @@ class FakeFinding:
return self.metadata
class CacheMissAfterPreResolve(MutableMapping):
def __init__(self, missing_uid):
self._cache = {}
self.missing_uid = missing_uid
def __contains__(self, key):
if key == self.missing_uid:
return True
return key in self._cache
def __getitem__(self, key):
return self._cache[key]
def __setitem__(self, key, value):
self._cache[key] = value
def __delitem__(self, key):
del self._cache[key]
def __iter__(self):
return iter(self._cache)
def __len__(self):
return len(self._cache)
@pytest.mark.django_db
class TestPerformScan:
def test_perform_prowler_scan_success(
@@ -1054,8 +1090,12 @@ class TestPerformScan:
perform_prowler_scan(tenant_id, scan_id, provider_id, [])
# Verify findings are muted with correct reason
fail_finding_db = Finding.objects.get(uid=finding_uid_1)
pass_finding_db = Finding.objects.get(uid=finding_uid_2)
fail_finding_db = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding_uid_1
)
pass_finding_db = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding_uid_2
)
assert fail_finding_db.muted
assert fail_finding_db.muted_reason == mute_rule_reason
@@ -1066,7 +1106,9 @@ class TestPerformScan:
assert pass_finding_db.muted_at is not None
# Verify failed_findings_count is 0 for muted FAIL finding
resource_1 = Resource.objects.get(uid="resource_uid_1")
resource_1 = Resource.objects.get(
tenant_id=tenant.id, provider_id=provider.id, uid="resource_uid_1"
)
assert resource_1.failed_findings_count == 0
def test_perform_prowler_scan_with_inactive_mute_rules(
@@ -1146,13 +1188,17 @@ class TestPerformScan:
perform_prowler_scan(tenant_id, scan_id, provider_id, [])
# Verify finding is NOT muted
finding_db = Finding.objects.get(uid=finding_uid)
finding_db = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding_uid
)
assert not finding_db.muted
assert finding_db.muted_reason is None
assert finding_db.muted_at is None
# Verify failed_findings_count increments for FAIL finding
resource = Resource.objects.get(uid="resource_uid_inactive")
resource = Resource.objects.get(
tenant_id=tenant.id, provider_id=provider.id, uid="resource_uid_inactive"
)
assert resource.failed_findings_count == 1
def test_perform_prowler_scan_mutelist_overrides_mute_rules(
@@ -1232,13 +1278,17 @@ class TestPerformScan:
perform_prowler_scan(tenant_id, scan_id, provider_id, [])
# Verify mutelist reason takes precedence
finding_db = Finding.objects.get(uid=finding_uid)
finding_db = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding_uid
)
assert finding_db.muted
assert finding_db.muted_reason == "Muted by mutelist"
assert finding_db.muted_at is not None
# Verify failed_findings_count is 0
resource = Resource.objects.get(uid="resource_both")
resource = Resource.objects.get(
tenant_id=tenant.id, provider_id=provider.id, uid="resource_both"
)
assert resource.failed_findings_count == 0
def test_perform_prowler_scan_mute_rules_multiple_findings(
@@ -1330,14 +1380,20 @@ class TestPerformScan:
# Verify all findings are muted with same reason
for uid in finding_uids:
finding_db = Finding.objects.get(uid=uid)
finding_db = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=uid
)
assert finding_db.muted
assert finding_db.muted_reason == mute_rule_reason
assert finding_db.muted_at is not None
# Verify all resources have failed_findings_count = 0
for i in range(len(finding_uids)):
resource = Resource.objects.get(uid=f"resource_bulk_{i}")
resource = Resource.objects.get(
tenant_id=tenant.id,
provider_id=provider.id,
uid=f"resource_bulk_{i}",
)
assert resource.failed_findings_count == 0
def test_perform_prowler_scan_mute_rules_error_handling(
@@ -1415,12 +1471,18 @@ class TestPerformScan:
assert scan.state == StateChoices.COMPLETED
# Verify finding is not muted (mute_rules_cache was empty dict)
finding_db = Finding.objects.get(uid="finding_error_handling")
finding_db = Finding.objects.get(
tenant_id=tenant.id,
scan_id=scan.id,
uid="finding_error_handling",
)
assert not finding_db.muted
assert finding_db.muted_reason is None
# Verify failed_findings_count increments
resource = Resource.objects.get(uid="resource_error")
resource = Resource.objects.get(
tenant_id=tenant.id, provider_id=provider.id, uid="resource_error"
)
assert resource.failed_findings_count == 1
def test_perform_prowler_scan_muted_at_timestamp(
@@ -1502,7 +1564,9 @@ class TestPerformScan:
after_scan = datetime.now(UTC)
# Verify muted_at is within the scan time window
finding_db = Finding.objects.get(uid=finding_uid)
finding_db = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding_uid
)
assert finding_db.muted
assert finding_db.muted_at is not None
assert before_scan <= finding_db.muted_at <= after_scan
@@ -1513,6 +1577,548 @@ class TestPerformScan:
@pytest.mark.django_db
class TestProcessFindingMicroBatch:
def _process_one_finding_micro_batch(
self,
tenant,
scan,
provider,
finding,
resource_cache=None,
resource_failed_findings_cache=None,
):
resource_cache = resource_cache if resource_cache is not None else {}
resource_failed_findings_cache = (
resource_failed_findings_cache
if resource_failed_findings_cache is not None
else {}
)
caches = {
"resource_cache": resource_cache,
"tag_cache": {},
"last_status_cache": {},
"resource_failed_findings_cache": resource_failed_findings_cache,
"unique_resources": set(),
"scan_resource_cache": set(),
"mute_rules_cache": {},
"scan_categories_cache": {},
"scan_resource_groups_cache": {},
"group_resources_cache": {},
}
with (
patch("tasks.jobs.scan.rls_transaction", new=noop_rls_transaction),
patch("api.db_utils.rls_transaction", new=noop_rls_transaction),
):
_process_finding_micro_batch(
str(tenant.id),
[finding],
scan,
provider,
caches["resource_cache"],
caches["tag_cache"],
caches["last_status_cache"],
caches["resource_failed_findings_cache"],
caches["unique_resources"],
caches["scan_resource_cache"],
caches["mute_rules_cache"],
caches["scan_categories_cache"],
caches["scan_resource_groups_cache"],
caches["group_resources_cache"],
)
return caches
def test_process_finding_micro_batch_fallback_creates_resource_after_cache_miss(
self, tenants_fixture, scans_fixture
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = scan.provider
resource_uid = "arn:aws:accessanalyzer:us-east-1:123456789012:analyzer/unknown"
finding = FakeFinding(
uid="finding-cache-miss-create",
status=StatusChoices.FAIL,
status_extended="missing analyzer",
severity=Severity.medium,
check_id="accessanalyzer_enabled",
resource_uid=resource_uid,
resource_name="analyzer/unknown",
region="us-east-1",
service_name="accessanalyzer",
resource_type="analyzer",
resource_tags={},
resource_metadata={},
resource_details={},
partition="aws",
raw={},
compliance={},
metadata={"resourcegroup": "identity"},
muted=False,
)
caches = self._process_one_finding_micro_batch(
tenant,
scan,
provider,
finding,
resource_cache=CacheMissAfterPreResolve(resource_uid),
)
resource = Resource.objects.get(
tenant_id=tenant.id, provider_id=provider.id, uid=resource_uid
)
created_finding = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
)
assert created_finding.scan_id == scan.id
assert resource.provider_id == provider.id
assert resource.region == finding.region
assert resource.service == finding.service_name
assert resource.type == finding.resource_type
assert resource.name == finding.resource_name
assert resource.groups == ["identity"]
assert resource.findings.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
assert caches["resource_cache"][resource_uid].id == resource.id
assert caches["resource_failed_findings_cache"][resource_uid] == 1
def test_process_finding_micro_batch_fallback_recovers_existing_resource_after_cache_miss(
self, tenants_fixture, scans_fixture
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = scan.provider
resource_uid = "arn:aws:guardduty:us-east-1:123456789012:detector/unknown"
existing_resource = Resource.objects.create(
tenant_id=tenant.id,
provider=provider,
uid=resource_uid,
name="detector/unknown",
region="us-east-1",
service="guardduty",
type="detector",
)
finding = FakeFinding(
uid="finding-cache-miss-existing",
status=StatusChoices.FAIL,
status_extended="missing detector",
severity=Severity.high,
check_id="guardduty_enabled",
resource_uid=resource_uid,
resource_name=existing_resource.name,
region=existing_resource.region,
service_name=existing_resource.service,
resource_type=existing_resource.type,
resource_tags={},
resource_metadata={},
resource_details={},
partition="aws",
raw={},
compliance={},
metadata={},
muted=False,
)
caches = self._process_one_finding_micro_batch(
tenant,
scan,
provider,
finding,
resource_cache=CacheMissAfterPreResolve(resource_uid),
)
assert (
Resource.objects.filter(
tenant_id=tenant.id, provider_id=provider.id, uid=resource_uid
).count()
== 1
)
created_finding = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
)
existing_resource.refresh_from_db()
assert created_finding.scan_id == scan.id
assert existing_resource.findings.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
assert caches["resource_cache"][resource_uid].id == existing_resource.id
assert caches["resource_failed_findings_cache"][resource_uid] == 1
def test_process_finding_micro_batch_fallback_recovers_after_create_race(
self, tenants_fixture, scans_fixture
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = scan.provider
resource_uid = "arn:aws:securityhub:us-east-1:123456789012:hub/unknown"
raced_resource = Resource.objects.create(
tenant_id=tenant.id,
provider=provider,
uid=resource_uid,
name="hub/unknown",
region="us-east-1",
service="securityhub",
type="hub",
)
finding = FakeFinding(
uid="finding-cache-miss-failure",
status=StatusChoices.FAIL,
status_extended="missing hub",
severity=Severity.high,
check_id="securityhub_enabled",
resource_uid=resource_uid,
resource_name="hub/unknown",
region="us-east-1",
service_name="securityhub",
resource_type="hub",
resource_tags={},
resource_metadata={},
resource_details={},
partition="aws",
raw={},
compliance={},
metadata={},
muted=False,
)
resource_filter_result = MagicMock()
resource_filter_result.first.side_effect = [None, raced_resource]
with (
patch.object(
Resource.objects,
"filter",
return_value=resource_filter_result,
),
patch.object(
Resource.objects,
"create",
side_effect=IntegrityError("duplicate resource"),
),
):
caches = self._process_one_finding_micro_batch(
tenant,
scan,
provider,
finding,
resource_cache=CacheMissAfterPreResolve(resource_uid),
)
assert (
Resource.objects.filter(
tenant_id=tenant.id, provider_id=provider.id, uid=resource_uid
).count()
== 1
)
created_finding = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
)
raced_resource.refresh_from_db()
assert created_finding.scan_id == scan.id
assert raced_resource.findings.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
assert caches["resource_cache"][resource_uid].id == raced_resource.id
assert caches["resource_failed_findings_cache"][resource_uid] == 1
def test_process_finding_micro_batch_cache_miss_retry_drops_rolled_back_resource(
self, tenants_fixture, scans_fixture
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = scan.provider
resource_uid = "generic-resource-cache-miss-retry"
cached_resource = Resource.objects.create(
tenant_id=tenant.id,
provider=provider,
uid="generic-cached-resource-retry",
name="old-cached-resource",
region="us-west-2",
service="old-service",
type="old-type",
)
finding = FakeFinding(
uid="finding-cache-miss-retry-clean-resource-cache",
status=StatusChoices.FAIL,
status_extended="missing resource",
severity=Severity.high,
check_id="generic_resource_check",
resource_uid=resource_uid,
resource_name="generic-resource",
region="us-east-1",
service_name="generic-service",
resource_type="generic-type",
resource_tags={"team": "platform"},
resource_metadata={"owner": "security"},
resource_details={"id": "generic-resource"},
partition="aws",
raw={},
compliance={},
metadata={"categories": ["security"], "resourcegroup": "identity"},
muted=False,
)
cached_resource_finding = FakeFinding(
uid="finding-cache-miss-retry-restores-dirty-resource",
status=StatusChoices.FAIL,
status_extended="cached resource changed",
severity=Severity.high,
check_id="generic_cached_resource_check",
resource_uid=cached_resource.uid,
resource_name="new-cached-resource",
region="eu-west-1",
service_name="new-service",
resource_type="new-type",
resource_tags={},
resource_metadata={"owner": "platform"},
resource_details={"id": "cached-resource"},
partition="aws",
raw={},
compliance={},
metadata={"categories": ["security"], "resourcegroup": "identity"},
muted=False,
)
resource_cache = CacheMissAfterPreResolve(resource_uid)
resource_cache[cached_resource.uid] = cached_resource
tag_cache = {}
resource_failed_findings_cache = {cached_resource.uid: 0}
scan_resource_cache: set[tuple[str, str, str, str]] = set()
scan_categories_cache: dict[tuple[str, str], dict[str, int]] = {}
scan_resource_groups_cache: dict[tuple[str, str], dict[str, int]] = {}
group_resources_cache: dict[str, set] = {}
original_bulk_create = ResourceFindingMapping.objects.bulk_create
original_tag_mapping_bulk_create = ResourceTagMapping.objects.bulk_create
mapping_bulk_create_calls = []
tag_mapping_bulk_create_calls = []
def fail_once_then_bulk_create(objects, *args, **kwargs):
mapping_bulk_create_calls.append([str(obj.resource_id) for obj in objects])
if len(mapping_bulk_create_calls) == 1:
raise IntegrityError("rollback after fallback resource creation")
return original_bulk_create(objects, *args, **kwargs)
def track_tag_mappings_bulk_create(objects, *args, **kwargs):
tag_mapping_bulk_create_calls.append([str(obj.tag_id) for obj in objects])
return original_tag_mapping_bulk_create(objects, *args, **kwargs)
with (
patch("tasks.jobs.scan.CELERY_DEADLOCK_ATTEMPTS", 2),
patch("tasks.jobs.scan.rls_transaction", new=atomic_rls_transaction),
patch("api.db_utils.rls_transaction", new=atomic_rls_transaction),
patch.object(
ResourceTagMapping.objects,
"bulk_create",
side_effect=track_tag_mappings_bulk_create,
),
patch.object(
ResourceFindingMapping.objects,
"bulk_create",
side_effect=fail_once_then_bulk_create,
),
):
_process_finding_micro_batch(
str(tenant.id),
[finding, cached_resource_finding],
scan,
provider,
resource_cache,
tag_cache,
{},
resource_failed_findings_cache,
set(),
scan_resource_cache,
{},
scan_categories_cache,
scan_resource_groups_cache,
group_resources_cache,
)
resource = Resource.objects.get(
tenant_id=tenant.id,
provider_id=provider.id,
uid=resource_uid,
)
created_finding = Finding.objects.get(
tenant_id=tenant.id,
scan_id=scan.id,
uid=finding.uid,
)
cached_resource.refresh_from_db()
assert len(mapping_bulk_create_calls) == 2
assert mapping_bulk_create_calls[0] != mapping_bulk_create_calls[1]
assert len(tag_mapping_bulk_create_calls) == 2
assert tag_mapping_bulk_create_calls[0] != tag_mapping_bulk_create_calls[1]
assert created_finding.scan_id == scan.id
assert resource.findings.filter(
tenant_id=tenant.id,
scan_id=scan.id,
uid=finding.uid,
).exists()
assert cached_resource.findings.filter(
tenant_id=tenant.id,
scan_id=scan.id,
uid=cached_resource_finding.uid,
).exists()
assert cached_resource.name == cached_resource_finding.resource_name
assert cached_resource.region == cached_resource_finding.region
assert cached_resource.service == cached_resource_finding.service_name
assert cached_resource.type == cached_resource_finding.resource_type
assert resource_cache[resource_uid].id == resource.id
assert resource_failed_findings_cache[resource_uid] == 1
assert resource_failed_findings_cache[cached_resource.uid] == 1
assert scan_resource_cache == {
(
str(resource.id),
finding.service_name,
finding.region,
finding.resource_type,
),
(
str(cached_resource.id),
cached_resource_finding.service_name,
cached_resource_finding.region,
cached_resource_finding.resource_type,
),
}
assert (
tag_cache[("team", "platform")].id
== ResourceTag.objects.get(
tenant_id=tenant.id,
key="team",
value="platform",
).id
)
assert scan_categories_cache == {
("security", "high"): {"total": 2, "failed": 2, "new_failed": 2}
}
assert scan_resource_groups_cache == {
("identity", "high"): {"total": 2, "failed": 2, "new_failed": 2}
}
assert group_resources_cache == {
"identity": {resource_uid, cached_resource.uid}
}
def test_process_finding_micro_batch_propagates_retryable_cache_miss_db_errors(
self, tenants_fixture, scans_fixture
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = scan.provider
resource_uid = "arn:aws:securityhub:us-east-1:123456789012:hub/retryable"
finding = FakeFinding(
uid="finding-cache-miss-retryable-error",
status=StatusChoices.FAIL,
status_extended="missing hub",
severity=Severity.high,
check_id="securityhub_enabled",
resource_uid=resource_uid,
resource_name="hub/retryable",
region="us-east-1",
service_name="securityhub",
resource_type="hub",
resource_tags={},
resource_metadata={},
resource_details={},
partition="aws",
raw={},
compliance={},
metadata={},
muted=False,
)
with (
patch("tasks.jobs.scan.CELERY_DEADLOCK_ATTEMPTS", 1),
patch.object(
Resource.objects,
"create",
side_effect=OperationalError("deadlock detected"),
),
):
with pytest.raises(OperationalError, match="deadlock detected"):
self._process_one_finding_micro_batch(
tenant,
scan,
provider,
finding,
resource_cache=CacheMissAfterPreResolve(resource_uid),
)
assert not Finding.objects.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
def test_process_finding_micro_batch_propagates_unrecovered_cache_miss_integrity_error(
self, tenants_fixture, scans_fixture
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = scan.provider
resource_uid = "arn:aws:securityhub:us-east-1:123456789012:hub/unrecovered"
finding = FakeFinding(
uid="finding-cache-miss-unrecovered-integrity-error",
status=StatusChoices.FAIL,
status_extended="missing hub",
severity=Severity.high,
check_id="securityhub_enabled",
resource_uid=resource_uid,
resource_name="hub/unrecovered",
region="us-east-1",
service_name="securityhub",
resource_type="hub",
resource_tags={},
resource_metadata={},
resource_details={},
partition="aws",
raw={},
compliance={},
metadata={},
muted=False,
)
original_resource_filter = Resource.objects.filter
resource_filter_result = MagicMock()
resource_filter_result.first.side_effect = [None, None]
def resource_filter_side_effect(*args, **kwargs):
if kwargs.get("uid") == resource_uid:
return resource_filter_result
return original_resource_filter(*args, **kwargs)
with (
patch("tasks.jobs.scan.CELERY_DEADLOCK_ATTEMPTS", 1),
patch.object(
Resource.objects,
"filter",
side_effect=resource_filter_side_effect,
),
patch.object(
Resource.objects,
"create",
side_effect=IntegrityError("constraint violation"),
),
):
with pytest.raises(IntegrityError, match="constraint violation"):
self._process_one_finding_micro_batch(
tenant,
scan,
provider,
finding,
resource_cache=CacheMissAfterPreResolve(resource_uid),
)
assert not Finding.objects.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
def test_process_finding_micro_batch_creates_records_and_updates_caches(
self, tenants_fixture, scans_fixture
):
@@ -1573,8 +2179,12 @@ class TestProcessFindingMicroBatch:
group_resources_cache,
)
created_finding = Finding.objects.get(uid=finding.uid)
resource = Resource.objects.get(uid=finding.resource_uid)
created_finding = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
)
resource = Resource.objects.get(
tenant_id=tenant.id, provider_id=provider.id, uid=finding.resource_uid
)
assert created_finding.scan_id == scan.id
assert created_finding.status == StatusChoices.PASS
@@ -1602,7 +2212,9 @@ class TestProcessFindingMicroBatch:
assert set(resource.tags.values_list("key", "value")) == set(
finding.resource_tags.items()
)
assert resource.findings.filter(uid=finding.uid).exists()
assert resource.findings.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
assert resource_cache[finding.resource_uid].id == resource.id
assert resource_failed_findings_cache[finding.resource_uid] == 0
@@ -1691,7 +2303,9 @@ class TestProcessFindingMicroBatch:
)
existing_resource.refresh_from_db()
created_finding = Finding.objects.get(uid=finding.uid)
created_finding = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
)
assert created_finding.delta == Finding.DeltaChoices.CHANGED
assert created_finding.status == StatusChoices.FAIL
@@ -1725,7 +2339,9 @@ class TestProcessFindingMicroBatch:
assert set(existing_resource.tags.values_list("key", "value")) == {
("team", "devsec")
}
assert existing_resource.findings.filter(uid=finding.uid).exists()
assert existing_resource.findings.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=finding.uid
).exists()
assert resource_cache[finding.resource_uid].region == finding.region
assert resource_cache[finding.resource_uid].service == finding.service_name
@@ -1891,10 +2507,14 @@ class TestProcessFindingMicroBatch:
)
# Verify the long UID finding was NOT created
assert not Finding.objects.filter(uid=long_uid).exists()
assert not Finding.objects.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=long_uid
).exists()
# Verify the normal finding WAS created
assert Finding.objects.filter(uid=normal_finding.uid).exists()
assert Finding.objects.filter(
tenant_id=tenant.id, scan_id=scan.id, uid=normal_finding.uid
).exists()
# Verify logging was called for skipped finding
assert mock_logger.warning.called
@@ -2019,8 +2639,12 @@ class TestProcessFindingMicroBatch:
"new_failed": 1,
}
created_finding1 = Finding.objects.get(uid="finding-cat-1")
created_finding2 = Finding.objects.get(uid="finding-cat-2")
created_finding1 = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid="finding-cat-1"
)
created_finding2 = Finding.objects.get(
tenant_id=tenant.id, scan_id=scan.id, uid="finding-cat-2"
)
assert set(created_finding1.categories) == {"gen-ai", "security"}
assert set(created_finding2.categories) == {"security", "iam"}