fix(api): avoid field-named annotation in attack surface aggregation (#12889)

This commit is contained in:
Rubén De la Torre Vico
2026-09-28 11:39:36 +02:00
committed by GitHub
parent c114aa304b
commit 453c953f37
2 changed files with 76 additions and 7 deletions
+4 -2
View File
@@ -2113,7 +2113,9 @@ def aggregate_attack_surface(tenant_id: str, scan_id: str):
.annotate(
total=Count("id"),
failed=Count("id", filter=Q(status="FAIL", muted=False)),
muted=Count("id", filter=Q(status="FAIL", muted=True)),
# Not `muted`: an annotation named after a model field comes
# back as `muted_new` from the psqlextra queryset.
muted_count=Count("id", filter=Q(status="FAIL", muted=True)),
)
)
@@ -2124,7 +2126,7 @@ def aggregate_attack_surface(tenant_id: str, scan_id: str):
aggregated_counts[attack_surface_type]["total"] += stats["total"] or 0
aggregated_counts[attack_surface_type]["failed"] += stats["failed"] or 0
aggregated_counts[attack_surface_type]["muted"] += stats["muted"] or 0
aggregated_counts[attack_surface_type]["muted"] += stats["muted_count"] or 0
overview_objects = []
for attack_surface_type, counts in aggregated_counts.items():
+72 -5
View File
@@ -12,6 +12,7 @@ from api.db_router import MainRouter
from api.db_utils import rls_transaction
from api.exceptions import ProviderConnectionError, ProviderDeletedException
from api.models import (
AttackSurfaceOverview,
Finding,
MuteRule,
Provider,
@@ -5327,8 +5328,13 @@ class TestAggregateAttackSurface:
mock_queryset = MagicMock()
mock_queryset.values.return_value = mock_queryset
mock_queryset.annotate.return_value = [
{"check_id": "check_internet_1", "total": 10, "failed": 3, "muted": 1},
{"check_id": "check_secrets_1", "total": 5, "failed": 2, "muted": 0},
{
"check_id": "check_internet_1",
"total": 10,
"failed": 3,
"muted_count": 1,
},
{"check_id": "check_secrets_1", "total": 5, "failed": 2, "muted_count": 0},
]
ctx = MagicMock()
@@ -5377,7 +5383,7 @@ class TestAggregateAttackSurface:
mock_queryset = MagicMock()
mock_queryset.values.return_value = mock_queryset
mock_queryset.annotate.return_value = [
{"check_id": "check_internet_1", "total": 5, "failed": 1, "muted": 0},
{"check_id": "check_internet_1", "total": 5, "failed": 1, "muted_count": 0},
]
ctx = MagicMock()
@@ -5460,8 +5466,13 @@ class TestAggregateAttackSurface:
mock_queryset = MagicMock()
mock_queryset.values.return_value = mock_queryset
mock_queryset.annotate.return_value = [
{"check_id": "check_internet_1", "total": 10, "failed": 3, "muted": 1},
{"check_id": "check_internet_2", "total": 5, "failed": 2, "muted": 0},
{
"check_id": "check_internet_1",
"total": 10,
"failed": 3,
"muted_count": 1,
},
{"check_id": "check_internet_2", "total": 5, "failed": 2, "muted_count": 0},
]
ctx = MagicMock()
@@ -5482,6 +5493,62 @@ class TestAggregateAttackSurface:
assert overview.failed_findings == 5 # 3 + 2
assert overview.muted_failed_findings == 1 # 1 + 0
@patch("tasks.jobs.scan._get_attack_surface_mapping_from_provider")
def test_aggregate_attack_surface_counts_real_findings(
self, mock_get_mapping, tenants_fixture, scans_fixture
):
"""Run the aggregation query against real Finding rows.
The other tests mock the queryset, so they never execute the real
`annotate`. This one guards the row keys the query returns."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
mock_get_mapping.return_value = {
"privilege-escalation": {"check_privesc_1"},
"secrets": {"check_secrets_1"},
}
def create_finding(uid, check_id, status, muted):
Finding.objects.create(
tenant_id=tenant.id,
uid=uid,
scan=scan,
status=status,
status_extended="status extended",
impact=Severity.high,
severity=Severity.high,
raw_result={"status": status},
check_id=check_id,
check_metadata={"CheckId": check_id},
muted=muted,
first_seen_at="2024-01-02T00:00:00Z",
)
create_finding("privesc_fail", "check_privesc_1", Status.FAIL, False)
create_finding("privesc_fail_2", "check_privesc_1", Status.FAIL, False)
create_finding("privesc_fail_muted", "check_privesc_1", Status.FAIL, True)
create_finding("privesc_pass", "check_privesc_1", Status.PASS, False)
create_finding("secrets_pass_muted", "check_secrets_1", Status.PASS, True)
create_finding("unmapped_fail", "check_unmapped", Status.FAIL, False)
aggregate_attack_surface(str(tenant.id), str(scan.id))
overviews = {
overview.attack_surface_type: overview
for overview in AttackSurfaceOverview.objects.filter(
tenant_id=tenant.id, scan_id=scan.id
)
}
assert set(overviews) == {"privilege-escalation", "secrets"}
assert overviews["privilege-escalation"].total_findings == 4
assert overviews["privilege-escalation"].failed_findings == 2
assert overviews["privilege-escalation"].muted_failed_findings == 1
assert overviews["secrets"].total_findings == 1
assert overviews["secrets"].failed_findings == 0
assert overviews["secrets"].muted_failed_findings == 0
@patch("tasks.jobs.scan.Scan.all_objects.select_related")
@patch("tasks.jobs.scan.rls_transaction")
def test_aggregate_attack_surface_uses_select_related(