mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-24 04:51:51 +00:00
feat(overviews): Compliance watchlist endpoint (#9596)
Co-authored-by: Adrián Jesús Peña Rodríguez <adrianjpr@gmail.com>
This commit is contained in:
committed by
GitHub
parent
6c01151d78
commit
5f2cb614ad
@@ -5,6 +5,7 @@ All notable changes to the **Prowler API** are documented in this file.
|
||||
## [1.18.0] (Prowler UNRELEASED)
|
||||
|
||||
### Added
|
||||
- `/api/v1/overviews/compliance-watchlist` to retrieve the compliance watchlist [(#9596)](https://github.com/prowler-cloud/prowler/pull/9596)
|
||||
- Support AlibabaCloud provider [(#9485)](https://github.com/prowler-cloud/prowler/pull/9485)
|
||||
|
||||
---
|
||||
|
||||
@@ -37,6 +37,7 @@ from api.models import (
|
||||
PermissionChoices,
|
||||
Processor,
|
||||
Provider,
|
||||
ProviderComplianceScore,
|
||||
ProviderGroup,
|
||||
ProviderSecret,
|
||||
Resource,
|
||||
@@ -92,6 +93,54 @@ class ChoiceInFilter(BaseInFilter, ChoiceFilter):
|
||||
pass
|
||||
|
||||
|
||||
class BaseProviderFilter(FilterSet):
|
||||
"""
|
||||
Abstract base filter for models with direct FK to Provider.
|
||||
|
||||
Provides standard provider_id and provider_type filters.
|
||||
Subclasses must define Meta.model.
|
||||
"""
|
||||
|
||||
provider_id = UUIDFilter(field_name="provider__id", lookup_expr="exact")
|
||||
provider_id__in = UUIDInFilter(field_name="provider__id", lookup_expr="in")
|
||||
provider_type = ChoiceFilter(
|
||||
field_name="provider__provider", choices=Provider.ProviderChoices.choices
|
||||
)
|
||||
provider_type__in = ChoiceInFilter(
|
||||
field_name="provider__provider",
|
||||
choices=Provider.ProviderChoices.choices,
|
||||
lookup_expr="in",
|
||||
)
|
||||
|
||||
class Meta:
|
||||
abstract = True
|
||||
fields = {}
|
||||
|
||||
|
||||
class BaseScanProviderFilter(FilterSet):
|
||||
"""
|
||||
Abstract base filter for models with FK to Scan (and Scan has FK to Provider).
|
||||
|
||||
Provides standard provider_id and provider_type filters via scan relationship.
|
||||
Subclasses must define Meta.model.
|
||||
"""
|
||||
|
||||
provider_id = UUIDFilter(field_name="scan__provider__id", lookup_expr="exact")
|
||||
provider_id__in = UUIDInFilter(field_name="scan__provider__id", lookup_expr="in")
|
||||
provider_type = ChoiceFilter(
|
||||
field_name="scan__provider__provider", choices=Provider.ProviderChoices.choices
|
||||
)
|
||||
provider_type__in = ChoiceInFilter(
|
||||
field_name="scan__provider__provider",
|
||||
choices=Provider.ProviderChoices.choices,
|
||||
lookup_expr="in",
|
||||
)
|
||||
|
||||
class Meta:
|
||||
abstract = True
|
||||
fields = {}
|
||||
|
||||
|
||||
class CommonFindingFilters(FilterSet):
|
||||
# We filter providers from the scan in findings
|
||||
provider = UUIDFilter(field_name="scan__provider__id", lookup_expr="exact")
|
||||
@@ -1086,39 +1135,25 @@ class ThreatScoreSnapshotFilter(FilterSet):
|
||||
}
|
||||
|
||||
|
||||
class AttackSurfaceOverviewFilter(FilterSet):
|
||||
class AttackSurfaceOverviewFilter(BaseScanProviderFilter):
|
||||
"""Filter for attack surface overview aggregations by provider."""
|
||||
|
||||
provider_id = UUIDFilter(field_name="scan__provider__id", lookup_expr="exact")
|
||||
provider_id__in = UUIDInFilter(field_name="scan__provider__id", lookup_expr="in")
|
||||
provider_type = ChoiceFilter(
|
||||
field_name="scan__provider__provider", choices=Provider.ProviderChoices.choices
|
||||
)
|
||||
provider_type__in = ChoiceInFilter(
|
||||
field_name="scan__provider__provider",
|
||||
choices=Provider.ProviderChoices.choices,
|
||||
lookup_expr="in",
|
||||
)
|
||||
|
||||
class Meta:
|
||||
class Meta(BaseScanProviderFilter.Meta):
|
||||
model = AttackSurfaceOverview
|
||||
fields = {}
|
||||
|
||||
|
||||
class CategoryOverviewFilter(FilterSet):
|
||||
provider_id = UUIDFilter(field_name="scan__provider__id", lookup_expr="exact")
|
||||
provider_id__in = UUIDInFilter(field_name="scan__provider__id", lookup_expr="in")
|
||||
provider_type = ChoiceFilter(
|
||||
field_name="scan__provider__provider", choices=Provider.ProviderChoices.choices
|
||||
)
|
||||
provider_type__in = ChoiceInFilter(
|
||||
field_name="scan__provider__provider",
|
||||
choices=Provider.ProviderChoices.choices,
|
||||
lookup_expr="in",
|
||||
)
|
||||
class CategoryOverviewFilter(BaseScanProviderFilter):
|
||||
"""Filter for category overview aggregations by provider."""
|
||||
|
||||
category = CharFilter(field_name="category", lookup_expr="exact")
|
||||
category__in = CharInFilter(field_name="category", lookup_expr="in")
|
||||
|
||||
class Meta:
|
||||
class Meta(BaseScanProviderFilter.Meta):
|
||||
model = ScanCategorySummary
|
||||
fields = {}
|
||||
|
||||
|
||||
class ComplianceWatchlistFilter(BaseProviderFilter):
|
||||
"""Filter for compliance watchlist overview by provider."""
|
||||
|
||||
class Meta(BaseProviderFilter.Meta):
|
||||
model = ProviderComplianceScore
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
import uuid
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
import api.db_utils
|
||||
import api.rls
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [
|
||||
("api", "0065_alibabacloud_provider"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name="ProviderComplianceScore",
|
||||
fields=[
|
||||
(
|
||||
"id",
|
||||
models.UUIDField(
|
||||
default=uuid.uuid4,
|
||||
editable=False,
|
||||
primary_key=True,
|
||||
serialize=False,
|
||||
),
|
||||
),
|
||||
("compliance_id", models.TextField()),
|
||||
("requirement_id", models.TextField()),
|
||||
(
|
||||
"requirement_status",
|
||||
api.db_utils.StatusEnumField(
|
||||
choices=[
|
||||
("FAIL", "Fail"),
|
||||
("PASS", "Pass"),
|
||||
("MANUAL", "Manual"),
|
||||
]
|
||||
),
|
||||
),
|
||||
("scan_completed_at", models.DateTimeField()),
|
||||
(
|
||||
"provider",
|
||||
models.ForeignKey(
|
||||
on_delete=django.db.models.deletion.CASCADE,
|
||||
related_name="compliance_scores",
|
||||
related_query_name="compliance_score",
|
||||
to="api.provider",
|
||||
),
|
||||
),
|
||||
(
|
||||
"scan",
|
||||
models.ForeignKey(
|
||||
on_delete=django.db.models.deletion.CASCADE,
|
||||
related_name="compliance_scores",
|
||||
related_query_name="compliance_score",
|
||||
to="api.scan",
|
||||
),
|
||||
),
|
||||
(
|
||||
"tenant",
|
||||
models.ForeignKey(
|
||||
on_delete=django.db.models.deletion.CASCADE,
|
||||
to="api.tenant",
|
||||
),
|
||||
),
|
||||
],
|
||||
options={
|
||||
"db_table": "provider_compliance_scores",
|
||||
"abstract": False,
|
||||
},
|
||||
),
|
||||
migrations.AddConstraint(
|
||||
model_name="providercompliancescore",
|
||||
constraint=models.UniqueConstraint(
|
||||
fields=("tenant_id", "provider_id", "compliance_id", "requirement_id"),
|
||||
name="unique_provider_compliance_req",
|
||||
),
|
||||
),
|
||||
migrations.AddConstraint(
|
||||
model_name="providercompliancescore",
|
||||
constraint=api.rls.RowLevelSecurityConstraint(
|
||||
"tenant_id",
|
||||
name="rls_on_providercompliancescore",
|
||||
statements=["SELECT", "INSERT", "UPDATE", "DELETE"],
|
||||
),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name="providercompliancescore",
|
||||
index=models.Index(
|
||||
fields=["tenant_id", "provider_id", "compliance_id"],
|
||||
name="pcs_tenant_prov_comp_idx",
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,61 @@
|
||||
import uuid
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.db import migrations, models
|
||||
|
||||
import api.rls
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [
|
||||
("api", "0066_provider_compliance_score"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name="TenantComplianceSummary",
|
||||
fields=[
|
||||
(
|
||||
"id",
|
||||
models.UUIDField(
|
||||
default=uuid.uuid4,
|
||||
editable=False,
|
||||
primary_key=True,
|
||||
serialize=False,
|
||||
),
|
||||
),
|
||||
("compliance_id", models.TextField()),
|
||||
("requirements_passed", models.IntegerField(default=0)),
|
||||
("requirements_failed", models.IntegerField(default=0)),
|
||||
("requirements_manual", models.IntegerField(default=0)),
|
||||
("total_requirements", models.IntegerField(default=0)),
|
||||
("updated_at", models.DateTimeField(auto_now=True)),
|
||||
(
|
||||
"tenant",
|
||||
models.ForeignKey(
|
||||
on_delete=django.db.models.deletion.CASCADE,
|
||||
to="api.tenant",
|
||||
),
|
||||
),
|
||||
],
|
||||
options={
|
||||
"db_table": "tenant_compliance_summaries",
|
||||
"abstract": False,
|
||||
},
|
||||
),
|
||||
migrations.AddConstraint(
|
||||
model_name="tenantcompliancesummary",
|
||||
constraint=models.UniqueConstraint(
|
||||
fields=("tenant_id", "compliance_id"),
|
||||
name="unique_tenant_compliance_summary",
|
||||
),
|
||||
),
|
||||
migrations.AddConstraint(
|
||||
model_name="tenantcompliancesummary",
|
||||
constraint=api.rls.RowLevelSecurityConstraint(
|
||||
"tenant_id",
|
||||
name="rls_on_tenantcompliancesummary",
|
||||
statements=["SELECT", "INSERT", "UPDATE", "DELETE"],
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -2605,3 +2605,92 @@ class AttackSurfaceOverview(RowLevelSecurityProtectedModel):
|
||||
|
||||
class JSONAPIMeta:
|
||||
resource_name = "attack-surface-overviews"
|
||||
|
||||
|
||||
class ProviderComplianceScore(RowLevelSecurityProtectedModel):
|
||||
"""
|
||||
Compliance requirement status from latest completed scan per provider.
|
||||
|
||||
Used for efficient compliance watchlist queries with FAIL-dominant aggregation
|
||||
across multiple providers. Updated via atomic upsert after each scan completion.
|
||||
"""
|
||||
|
||||
id = models.UUIDField(primary_key=True, default=uuid4, editable=False)
|
||||
|
||||
scan = models.ForeignKey(
|
||||
Scan,
|
||||
on_delete=models.CASCADE,
|
||||
related_name="compliance_scores",
|
||||
related_query_name="compliance_score",
|
||||
)
|
||||
|
||||
provider = models.ForeignKey(
|
||||
Provider,
|
||||
on_delete=models.CASCADE,
|
||||
related_name="compliance_scores",
|
||||
related_query_name="compliance_score",
|
||||
)
|
||||
|
||||
compliance_id = models.TextField()
|
||||
requirement_id = models.TextField()
|
||||
requirement_status = StatusEnumField(choices=StatusChoices)
|
||||
|
||||
scan_completed_at = models.DateTimeField()
|
||||
|
||||
class Meta(RowLevelSecurityProtectedModel.Meta):
|
||||
db_table = "provider_compliance_scores"
|
||||
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=("tenant_id", "provider_id", "compliance_id", "requirement_id"),
|
||||
name="unique_provider_compliance_req",
|
||||
),
|
||||
RowLevelSecurityConstraint(
|
||||
field="tenant_id",
|
||||
name="rls_on_%(class)s",
|
||||
statements=["SELECT", "INSERT", "UPDATE", "DELETE"],
|
||||
),
|
||||
]
|
||||
|
||||
indexes = [
|
||||
models.Index(
|
||||
fields=["tenant_id", "provider_id", "compliance_id"],
|
||||
name="pcs_tenant_prov_comp_idx",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class TenantComplianceSummary(RowLevelSecurityProtectedModel):
|
||||
"""
|
||||
Pre-aggregated compliance counts per tenant with FAIL-dominant logic applied.
|
||||
|
||||
One row per (tenant, compliance_id). Used for fast watchlist queries when
|
||||
no provider filter is applied. Recalculated after each scan by aggregating
|
||||
across all providers with FAIL-dominant logic at requirement level.
|
||||
"""
|
||||
|
||||
id = models.UUIDField(primary_key=True, default=uuid4, editable=False)
|
||||
|
||||
compliance_id = models.TextField()
|
||||
|
||||
requirements_passed = models.IntegerField(default=0)
|
||||
requirements_failed = models.IntegerField(default=0)
|
||||
requirements_manual = models.IntegerField(default=0)
|
||||
total_requirements = models.IntegerField(default=0)
|
||||
|
||||
updated_at = models.DateTimeField(auto_now=True)
|
||||
|
||||
class Meta(RowLevelSecurityProtectedModel.Meta):
|
||||
db_table = "tenant_compliance_summaries"
|
||||
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=("tenant_id", "compliance_id"),
|
||||
name="unique_tenant_compliance_summary",
|
||||
),
|
||||
RowLevelSecurityConstraint(
|
||||
field="tenant_id",
|
||||
name="rls_on_%(class)s",
|
||||
statements=["SELECT", "INSERT", "UPDATE", "DELETE"],
|
||||
),
|
||||
]
|
||||
|
||||
@@ -4883,6 +4883,42 @@ paths:
|
||||
schema:
|
||||
$ref: '#/components/schemas/PaginatedCategoryOverviewList'
|
||||
description: ''
|
||||
/api/v1/overviews/compliance-watchlist:
|
||||
get:
|
||||
operationId: overviews_compliance_watchlist_retrieve
|
||||
description: |-
|
||||
Get compliance watchlist overview with FAIL-dominant aggregation.
|
||||
|
||||
Without filters: uses pre-aggregated TenantComplianceSummary (~70 rows).
|
||||
With provider filters: queries ProviderComplianceScore with FAIL-dominant logic.
|
||||
parameters:
|
||||
- in: query
|
||||
name: fields[compliance-watchlist-overviews]
|
||||
schema:
|
||||
type: array
|
||||
items:
|
||||
type: string
|
||||
enum:
|
||||
- id
|
||||
- compliance_id
|
||||
- requirements_passed
|
||||
- requirements_failed
|
||||
- requirements_manual
|
||||
- total_requirements
|
||||
description: endpoint return only specific fields in the response on a per-type
|
||||
basis by including a fields[TYPE] query parameter.
|
||||
explode: false
|
||||
tags:
|
||||
- Overview
|
||||
security:
|
||||
- JWT or API Key: []
|
||||
responses:
|
||||
'200':
|
||||
content:
|
||||
application/vnd.api+json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ComplianceWatchlistOverviewResponse'
|
||||
description: ''
|
||||
/api/v1/overviews/findings:
|
||||
get:
|
||||
operationId: overviews_findings_retrieve
|
||||
@@ -11417,6 +11453,50 @@ components:
|
||||
type: string
|
||||
required:
|
||||
- regions
|
||||
ComplianceWatchlistOverview:
|
||||
type: object
|
||||
required:
|
||||
- type
|
||||
- id
|
||||
additionalProperties: false
|
||||
properties:
|
||||
type:
|
||||
type: string
|
||||
description: The [type](https://jsonapi.org/format/#document-resource-object-identification)
|
||||
member is used to describe resource objects that share common attributes
|
||||
and relationships.
|
||||
enum:
|
||||
- compliance-watchlist-overviews
|
||||
id: {}
|
||||
attributes:
|
||||
type: object
|
||||
properties:
|
||||
id:
|
||||
type: string
|
||||
compliance_id:
|
||||
type: string
|
||||
requirements_passed:
|
||||
type: integer
|
||||
requirements_failed:
|
||||
type: integer
|
||||
requirements_manual:
|
||||
type: integer
|
||||
total_requirements:
|
||||
type: integer
|
||||
required:
|
||||
- id
|
||||
- compliance_id
|
||||
- requirements_passed
|
||||
- requirements_failed
|
||||
- requirements_manual
|
||||
- total_requirements
|
||||
ComplianceWatchlistOverviewResponse:
|
||||
type: object
|
||||
properties:
|
||||
data:
|
||||
$ref: '#/components/schemas/ComplianceWatchlistOverview'
|
||||
required:
|
||||
- data
|
||||
Finding:
|
||||
type: object
|
||||
required:
|
||||
|
||||
@@ -1,9 +1,21 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from allauth.socialaccount.models import SocialApp
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.db import IntegrityError
|
||||
|
||||
from api.db_router import MainRouter
|
||||
from api.models import Resource, ResourceTag, SAMLConfiguration, SAMLDomainIndex
|
||||
from api.models import (
|
||||
ProviderComplianceScore,
|
||||
Resource,
|
||||
ResourceTag,
|
||||
SAMLConfiguration,
|
||||
SAMLDomainIndex,
|
||||
StateChoices,
|
||||
StatusChoices,
|
||||
TenantComplianceSummary,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
@@ -324,3 +336,159 @@ class TestSAMLConfigurationModel:
|
||||
errors = exc_info.value.message_dict
|
||||
assert "metadata_xml" in errors
|
||||
assert "There is a problem with your metadata." in errors["metadata_xml"][0]
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestProviderComplianceScoreModel:
|
||||
def test_create_provider_compliance_score(self, providers_fixture, scans_fixture):
|
||||
provider = providers_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
scan.completed_at = datetime.now(timezone.utc)
|
||||
scan.save()
|
||||
|
||||
score = ProviderComplianceScore.objects.create(
|
||||
tenant_id=provider.tenant_id,
|
||||
provider=provider,
|
||||
scan=scan,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.PASS,
|
||||
scan_completed_at=scan.completed_at,
|
||||
)
|
||||
|
||||
assert score.compliance_id == "aws_cis_2.0"
|
||||
assert score.requirement_id == "req_1"
|
||||
assert score.requirement_status == StatusChoices.PASS
|
||||
|
||||
def test_unique_constraint_per_provider_compliance_requirement(
|
||||
self, providers_fixture, scans_fixture
|
||||
):
|
||||
provider = providers_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
scan.completed_at = datetime.now(timezone.utc)
|
||||
scan.save()
|
||||
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=provider.tenant_id,
|
||||
provider=provider,
|
||||
scan=scan,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.PASS,
|
||||
scan_completed_at=scan.completed_at,
|
||||
)
|
||||
|
||||
with pytest.raises(IntegrityError):
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=provider.tenant_id,
|
||||
provider=provider,
|
||||
scan=scan,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.FAIL,
|
||||
scan_completed_at=scan.completed_at,
|
||||
)
|
||||
|
||||
def test_different_providers_same_requirement_allowed(
|
||||
self, providers_fixture, scans_fixture
|
||||
):
|
||||
provider1, provider2, *_ = providers_fixture
|
||||
scan1 = scans_fixture[0]
|
||||
scan1.completed_at = datetime.now(timezone.utc)
|
||||
scan1.save()
|
||||
|
||||
scan2 = scans_fixture[2]
|
||||
scan2.state = StateChoices.COMPLETED
|
||||
scan2.completed_at = datetime.now(timezone.utc)
|
||||
scan2.save()
|
||||
|
||||
score1 = ProviderComplianceScore.objects.create(
|
||||
tenant_id=provider1.tenant_id,
|
||||
provider=provider1,
|
||||
scan=scan1,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.PASS,
|
||||
scan_completed_at=scan1.completed_at,
|
||||
)
|
||||
|
||||
score2 = ProviderComplianceScore.objects.create(
|
||||
tenant_id=provider2.tenant_id,
|
||||
provider=provider2,
|
||||
scan=scan2,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.FAIL,
|
||||
scan_completed_at=scan2.completed_at,
|
||||
)
|
||||
|
||||
assert score1.id != score2.id
|
||||
assert score1.requirement_status != score2.requirement_status
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestTenantComplianceSummaryModel:
|
||||
def test_create_tenant_compliance_summary(self, tenants_fixture):
|
||||
tenant = tenants_fixture[0]
|
||||
|
||||
summary = TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirements_passed=5,
|
||||
requirements_failed=2,
|
||||
requirements_manual=1,
|
||||
total_requirements=8,
|
||||
)
|
||||
|
||||
assert summary.compliance_id == "aws_cis_2.0"
|
||||
assert summary.requirements_passed == 5
|
||||
assert summary.requirements_failed == 2
|
||||
assert summary.requirements_manual == 1
|
||||
assert summary.total_requirements == 8
|
||||
assert summary.updated_at is not None
|
||||
|
||||
def test_unique_constraint_per_tenant_compliance(self, tenants_fixture):
|
||||
tenant = tenants_fixture[0]
|
||||
|
||||
TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirements_passed=5,
|
||||
requirements_failed=2,
|
||||
requirements_manual=1,
|
||||
total_requirements=8,
|
||||
)
|
||||
|
||||
with pytest.raises(IntegrityError):
|
||||
TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirements_passed=3,
|
||||
requirements_failed=4,
|
||||
requirements_manual=1,
|
||||
total_requirements=8,
|
||||
)
|
||||
|
||||
def test_different_tenants_same_compliance_allowed(self, tenants_fixture):
|
||||
tenant1, tenant2, *_ = tenants_fixture
|
||||
|
||||
summary1 = TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant1.id,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirements_passed=5,
|
||||
requirements_failed=2,
|
||||
requirements_manual=1,
|
||||
total_requirements=8,
|
||||
)
|
||||
|
||||
summary2 = TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant2.id,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirements_passed=3,
|
||||
requirements_failed=4,
|
||||
requirements_manual=1,
|
||||
total_requirements=8,
|
||||
)
|
||||
|
||||
assert summary1.id != summary2.id
|
||||
assert summary1.requirements_passed != summary2.requirements_passed
|
||||
|
||||
@@ -7956,6 +7956,97 @@ class TestOverviewViewSet:
|
||||
assert data[0]["attributes"]["failed_findings"] == 13
|
||||
assert data[0]["attributes"]["new_failed_findings"] == 5
|
||||
|
||||
def test_compliance_watchlist_no_filters_uses_tenant_summary(
|
||||
self, authenticated_client, tenant_compliance_summary_fixture
|
||||
):
|
||||
response = authenticated_client.get(reverse("overview-compliance-watchlist"))
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
data = response.json()["data"]
|
||||
|
||||
assert len(data) == 2
|
||||
|
||||
by_id = {item["id"]: item["attributes"] for item in data}
|
||||
assert "aws_cis_2.0" in by_id
|
||||
assert by_id["aws_cis_2.0"]["requirements_passed"] == 1
|
||||
assert by_id["aws_cis_2.0"]["requirements_failed"] == 2
|
||||
assert by_id["aws_cis_2.0"]["requirements_manual"] == 1
|
||||
assert by_id["aws_cis_2.0"]["total_requirements"] == 4
|
||||
|
||||
assert "gdpr_aws" in by_id
|
||||
assert by_id["gdpr_aws"]["requirements_passed"] == 5
|
||||
assert by_id["gdpr_aws"]["requirements_failed"] == 0
|
||||
assert by_id["gdpr_aws"]["total_requirements"] == 7
|
||||
|
||||
def test_compliance_watchlist_with_provider_filter_uses_provider_scores(
|
||||
self,
|
||||
authenticated_client,
|
||||
provider_compliance_scores_fixture,
|
||||
providers_fixture,
|
||||
):
|
||||
provider1 = providers_fixture[0]
|
||||
url = f"{reverse('overview-compliance-watchlist')}?filter[provider_id]={provider1.id}"
|
||||
response = authenticated_client.get(url)
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
data = response.json()["data"]
|
||||
|
||||
assert len(data) == 2
|
||||
by_id = {item["id"]: item["attributes"] for item in data}
|
||||
|
||||
assert by_id["aws_cis_2.0"]["requirements_passed"] == 1
|
||||
assert by_id["aws_cis_2.0"]["requirements_failed"] == 1
|
||||
assert by_id["aws_cis_2.0"]["requirements_manual"] == 1
|
||||
assert by_id["aws_cis_2.0"]["total_requirements"] == 3
|
||||
|
||||
def test_compliance_watchlist_fail_dominant_logic(
|
||||
self, authenticated_client, provider_compliance_scores_fixture
|
||||
):
|
||||
response = authenticated_client.get(
|
||||
f"{reverse('overview-compliance-watchlist')}?filter[provider_type]=aws"
|
||||
)
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
data = response.json()["data"]
|
||||
|
||||
by_id = {item["id"]: item["attributes"] for item in data}
|
||||
aws_cis = by_id["aws_cis_2.0"]
|
||||
|
||||
assert aws_cis["requirements_failed"] == 2
|
||||
assert aws_cis["requirements_passed"] == 0
|
||||
assert aws_cis["requirements_manual"] == 1
|
||||
assert aws_cis["total_requirements"] == 3
|
||||
|
||||
def test_compliance_watchlist_provider_id_in_filter(
|
||||
self,
|
||||
authenticated_client,
|
||||
provider_compliance_scores_fixture,
|
||||
providers_fixture,
|
||||
):
|
||||
provider1, provider2, *_ = providers_fixture
|
||||
url = (
|
||||
f"{reverse('overview-compliance-watchlist')}"
|
||||
f"?filter[provider_id__in]={provider1.id},{provider2.id}"
|
||||
)
|
||||
response = authenticated_client.get(url)
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
data = response.json()["data"]
|
||||
assert len(data) >= 1
|
||||
|
||||
def test_compliance_watchlist_empty_result(self, authenticated_client):
|
||||
response = authenticated_client.get(reverse("overview-compliance-watchlist"))
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
data = response.json()["data"]
|
||||
assert data == []
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"invalid_provider_type",
|
||||
["invalid", "not_a_provider", "AWS", "awss"],
|
||||
)
|
||||
def test_compliance_watchlist_invalid_provider_type_filter(
|
||||
self, authenticated_client, invalid_provider_type
|
||||
):
|
||||
url = f"{reverse('overview-compliance-watchlist')}?filter[provider_type]={invalid_provider_type}"
|
||||
response = authenticated_client.get(url)
|
||||
assert response.status_code == status.HTTP_400_BAD_REQUEST
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestScheduleViewSet:
|
||||
|
||||
@@ -2303,6 +2303,20 @@ class CategoryOverviewSerializer(BaseSerializerV1):
|
||||
resource_name = "category-overviews"
|
||||
|
||||
|
||||
class ComplianceWatchlistOverviewSerializer(BaseSerializerV1):
|
||||
"""Serializer for compliance watchlist overview with FAIL-dominant aggregation."""
|
||||
|
||||
id = serializers.CharField(source="compliance_id")
|
||||
compliance_id = serializers.CharField()
|
||||
requirements_passed = serializers.IntegerField()
|
||||
requirements_failed = serializers.IntegerField()
|
||||
requirements_manual = serializers.IntegerField()
|
||||
total_requirements = serializers.IntegerField()
|
||||
|
||||
class JSONAPIMeta:
|
||||
resource_name = "compliance-watchlist-overviews"
|
||||
|
||||
|
||||
class OverviewRegionSerializer(serializers.Serializer):
|
||||
id = serializers.SerializerMethodField()
|
||||
provider_type = serializers.CharField()
|
||||
|
||||
@@ -101,6 +101,7 @@ from api.filters import (
|
||||
AttackSurfaceOverviewFilter,
|
||||
CategoryOverviewFilter,
|
||||
ComplianceOverviewFilter,
|
||||
ComplianceWatchlistFilter,
|
||||
CustomDjangoFilterBackend,
|
||||
DailySeveritySummaryFilter,
|
||||
FindingFilter,
|
||||
@@ -144,6 +145,7 @@ from api.models import (
|
||||
MuteRule,
|
||||
Processor,
|
||||
Provider,
|
||||
ProviderComplianceScore,
|
||||
ProviderGroup,
|
||||
ProviderGroupMembership,
|
||||
ProviderSecret,
|
||||
@@ -163,6 +165,7 @@ from api.models import (
|
||||
StateChoices,
|
||||
Task,
|
||||
TenantAPIKey,
|
||||
TenantComplianceSummary,
|
||||
ThreatScoreSnapshot,
|
||||
User,
|
||||
UserRoleRelationship,
|
||||
@@ -185,6 +188,7 @@ from api.v1.serializers import (
|
||||
ComplianceOverviewDetailThreatscoreSerializer,
|
||||
ComplianceOverviewMetadataSerializer,
|
||||
ComplianceOverviewSerializer,
|
||||
ComplianceWatchlistOverviewSerializer,
|
||||
FindingDynamicFilterSerializer,
|
||||
FindingMetadataSerializer,
|
||||
FindingSerializer,
|
||||
@@ -4142,6 +4146,8 @@ class OverviewViewSet(BaseRLSViewSet):
|
||||
return AttackSurfaceOverviewSerializer
|
||||
elif self.action == "categories":
|
||||
return CategoryOverviewSerializer
|
||||
elif self.action == "compliance_watchlist":
|
||||
return ComplianceWatchlistOverviewSerializer
|
||||
return super().get_serializer_class()
|
||||
|
||||
def get_filterset_class(self):
|
||||
@@ -4157,6 +4163,8 @@ class OverviewViewSet(BaseRLSViewSet):
|
||||
return CategoryOverviewFilter
|
||||
elif self.action == "attack_surface":
|
||||
return AttackSurfaceOverviewFilter
|
||||
elif self.action == "compliance_watchlist":
|
||||
return ComplianceWatchlistFilter
|
||||
return None
|
||||
|
||||
def filter_queryset(self, queryset):
|
||||
@@ -4240,6 +4248,8 @@ class OverviewViewSet(BaseRLSViewSet):
|
||||
self.request.query_params, exclude_keys=set(exclude_keys or [])
|
||||
)
|
||||
filterset = filterset_class(normalized_params, queryset=queryset)
|
||||
if not filterset.is_valid():
|
||||
raise ValidationError(filterset.errors)
|
||||
return filterset.qs
|
||||
|
||||
def _latest_scan_ids_for_allowed_providers(self, tenant_id, provider_filters=None):
|
||||
@@ -4256,9 +4266,10 @@ class OverviewViewSet(BaseRLSViewSet):
|
||||
)
|
||||
|
||||
def _extract_provider_filters_from_params(self):
|
||||
"""Extract provider filters from query params to apply on Scan queryset."""
|
||||
"""Extract and validate provider filters from query params."""
|
||||
params = self.request.query_params
|
||||
filters = {}
|
||||
valid_provider_types = {c[0] for c in Provider.ProviderChoices.choices}
|
||||
|
||||
provider_id = params.get("filter[provider_id]")
|
||||
if provider_id:
|
||||
@@ -4270,11 +4281,21 @@ class OverviewViewSet(BaseRLSViewSet):
|
||||
|
||||
provider_type = params.get("filter[provider_type]")
|
||||
if provider_type:
|
||||
if provider_type not in valid_provider_types:
|
||||
raise ValidationError(
|
||||
{"provider_type": f"Invalid choice: {provider_type}"}
|
||||
)
|
||||
filters["provider__provider"] = provider_type
|
||||
|
||||
provider_type_in = params.get("filter[provider_type__in]")
|
||||
if provider_type_in:
|
||||
filters["provider__provider__in"] = provider_type_in.split(",")
|
||||
types = provider_type_in.split(",")
|
||||
invalid = [t for t in types if t not in valid_provider_types]
|
||||
if invalid:
|
||||
raise ValidationError(
|
||||
{"provider_type__in": f"Invalid choices: {', '.join(invalid)}"}
|
||||
)
|
||||
filters["provider__provider__in"] = types
|
||||
|
||||
return filters
|
||||
|
||||
@@ -4984,6 +5005,92 @@ class OverviewViewSet(BaseRLSViewSet):
|
||||
status=status.HTTP_200_OK,
|
||||
)
|
||||
|
||||
@action(
|
||||
detail=False,
|
||||
methods=["get"],
|
||||
url_name="compliance-watchlist",
|
||||
url_path="compliance-watchlist",
|
||||
)
|
||||
def compliance_watchlist(self, request):
|
||||
"""
|
||||
Get compliance watchlist overview with FAIL-dominant aggregation.
|
||||
|
||||
Without filters: uses pre-aggregated TenantComplianceSummary (~70 rows).
|
||||
With provider filters: queries ProviderComplianceScore with FAIL-dominant logic.
|
||||
"""
|
||||
tenant_id = request.tenant_id
|
||||
rbac_filter = self._get_provider_filter()
|
||||
query_params = request.query_params
|
||||
|
||||
has_provider_filter = any(
|
||||
key.startswith("filter[provider") for key in query_params.keys()
|
||||
)
|
||||
has_rbac_restriction = bool(rbac_filter)
|
||||
|
||||
if not has_provider_filter and not has_rbac_restriction:
|
||||
response_data = list(
|
||||
TenantComplianceSummary.objects.filter(tenant_id=tenant_id)
|
||||
.values(
|
||||
"compliance_id",
|
||||
"requirements_passed",
|
||||
"requirements_failed",
|
||||
"requirements_manual",
|
||||
"total_requirements",
|
||||
)
|
||||
.order_by("compliance_id")
|
||||
)
|
||||
else:
|
||||
base_queryset = ProviderComplianceScore.objects.filter(
|
||||
tenant_id=tenant_id, **rbac_filter
|
||||
)
|
||||
|
||||
filtered_queryset = self._apply_filterset(
|
||||
base_queryset, ComplianceWatchlistFilter
|
||||
)
|
||||
|
||||
aggregation = (
|
||||
filtered_queryset.values("compliance_id", "requirement_id")
|
||||
.annotate(
|
||||
has_fail=Sum(
|
||||
Case(When(requirement_status="FAIL", then=1), default=0)
|
||||
),
|
||||
has_manual=Sum(
|
||||
Case(When(requirement_status="MANUAL", then=1), default=0)
|
||||
),
|
||||
)
|
||||
.values("compliance_id", "requirement_id", "has_fail", "has_manual")
|
||||
)
|
||||
|
||||
compliance_data = defaultdict(
|
||||
lambda: {
|
||||
"requirements_passed": 0,
|
||||
"requirements_failed": 0,
|
||||
"requirements_manual": 0,
|
||||
"total_requirements": 0,
|
||||
}
|
||||
)
|
||||
|
||||
for row in aggregation:
|
||||
cid = row["compliance_id"]
|
||||
compliance_data[cid]["total_requirements"] += 1
|
||||
|
||||
if row["has_fail"] and row["has_fail"] > 0:
|
||||
compliance_data[cid]["requirements_failed"] += 1
|
||||
elif row["has_manual"] and row["has_manual"] > 0:
|
||||
compliance_data[cid]["requirements_manual"] += 1
|
||||
else:
|
||||
compliance_data[cid]["requirements_passed"] += 1
|
||||
|
||||
response_data = [
|
||||
{"compliance_id": cid, **data}
|
||||
for cid, data in sorted(compliance_data.items())
|
||||
]
|
||||
|
||||
return Response(
|
||||
self.get_serializer(response_data, many=True).data,
|
||||
status=status.HTTP_200_OK,
|
||||
)
|
||||
|
||||
|
||||
@extend_schema(tags=["Schedule"])
|
||||
@extend_schema_view(
|
||||
|
||||
@@ -30,6 +30,7 @@ from api.models import (
|
||||
MuteRule,
|
||||
Processor,
|
||||
Provider,
|
||||
ProviderComplianceScore,
|
||||
ProviderGroup,
|
||||
ProviderSecret,
|
||||
Resource,
|
||||
@@ -45,6 +46,7 @@ from api.models import (
|
||||
StatusChoices,
|
||||
Task,
|
||||
TenantAPIKey,
|
||||
TenantComplianceSummary,
|
||||
User,
|
||||
UserRoleRelationship,
|
||||
)
|
||||
@@ -1631,6 +1633,108 @@ def get_authorization_header(access_token: str) -> dict:
|
||||
return {"Authorization": f"Bearer {access_token}"}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def provider_compliance_scores_fixture(
|
||||
tenants_fixture, providers_fixture, scans_fixture
|
||||
):
|
||||
"""Create ProviderComplianceScore entries for compliance watchlist tests."""
|
||||
tenant = tenants_fixture[0]
|
||||
provider1, provider2, *_ = providers_fixture
|
||||
scan1, _, scan3 = scans_fixture
|
||||
|
||||
scan1.completed_at = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
scan1.save()
|
||||
scan3.state = StateChoices.COMPLETED
|
||||
scan3.completed_at = datetime.now(timezone.utc)
|
||||
scan3.save()
|
||||
|
||||
scores = [
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider1,
|
||||
scan=scan1,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.PASS,
|
||||
scan_completed_at=scan1.completed_at,
|
||||
),
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider1,
|
||||
scan=scan1,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_2",
|
||||
requirement_status=StatusChoices.FAIL,
|
||||
scan_completed_at=scan1.completed_at,
|
||||
),
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider1,
|
||||
scan=scan1,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_3",
|
||||
requirement_status=StatusChoices.MANUAL,
|
||||
scan_completed_at=scan1.completed_at,
|
||||
),
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider2,
|
||||
scan=scan3,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_1",
|
||||
requirement_status=StatusChoices.FAIL,
|
||||
scan_completed_at=scan3.completed_at,
|
||||
),
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider2,
|
||||
scan=scan3,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirement_id="req_2",
|
||||
requirement_status=StatusChoices.PASS,
|
||||
scan_completed_at=scan3.completed_at,
|
||||
),
|
||||
ProviderComplianceScore.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider1,
|
||||
scan=scan1,
|
||||
compliance_id="gdpr_aws",
|
||||
requirement_id="gdpr_req_1",
|
||||
requirement_status=StatusChoices.PASS,
|
||||
scan_completed_at=scan1.completed_at,
|
||||
),
|
||||
]
|
||||
|
||||
return scores
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tenant_compliance_summary_fixture(tenants_fixture):
|
||||
"""Create TenantComplianceSummary entries for compliance watchlist tests."""
|
||||
tenant = tenants_fixture[0]
|
||||
|
||||
summaries = [
|
||||
TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
compliance_id="aws_cis_2.0",
|
||||
requirements_passed=1,
|
||||
requirements_failed=2,
|
||||
requirements_manual=1,
|
||||
total_requirements=4,
|
||||
),
|
||||
TenantComplianceSummary.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
compliance_id="gdpr_aws",
|
||||
requirements_passed=5,
|
||||
requirements_failed=0,
|
||||
requirements_manual=2,
|
||||
total_requirements=7,
|
||||
),
|
||||
]
|
||||
|
||||
return summaries
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(items):
|
||||
"""Ensure test_rbac.py is executed first."""
|
||||
items.sort(key=lambda item: 0 if "test_rbac.py" in item.nodeid else 1)
|
||||
|
||||
@@ -1,17 +1,28 @@
|
||||
from collections import defaultdict
|
||||
from datetime import timedelta
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from django.db.models import Sum
|
||||
from django.utils import timezone
|
||||
from tasks.jobs.queries import (
|
||||
COMPLIANCE_UPSERT_PROVIDER_SCORE_SQL,
|
||||
COMPLIANCE_UPSERT_TENANT_SUMMARY_ALL_SQL,
|
||||
)
|
||||
from tasks.jobs.scan import aggregate_category_counts
|
||||
|
||||
from api.db_router import READ_REPLICA_ALIAS
|
||||
from api.db_utils import rls_transaction
|
||||
from api.db_router import READ_REPLICA_ALIAS, MainRouter
|
||||
from api.db_utils import (
|
||||
POSTGRES_TENANT_VAR,
|
||||
SET_CONFIG_QUERY,
|
||||
psycopg_connection,
|
||||
rls_transaction,
|
||||
)
|
||||
from api.models import (
|
||||
ComplianceOverviewSummary,
|
||||
ComplianceRequirementOverview,
|
||||
DailySeveritySummary,
|
||||
Finding,
|
||||
ProviderComplianceScore,
|
||||
Resource,
|
||||
ResourceFindingMapping,
|
||||
ResourceScanSummary,
|
||||
@@ -21,6 +32,8 @@ from api.models import (
|
||||
StateChoices,
|
||||
)
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
def backfill_resource_scan_summaries(tenant_id: str, scan_id: str):
|
||||
with rls_transaction(tenant_id, using=READ_REPLICA_ALIAS):
|
||||
@@ -341,3 +354,114 @@ def backfill_scan_category_summaries(tenant_id: str, scan_id: str):
|
||||
)
|
||||
|
||||
return {"status": "backfilled", "categories_count": len(category_counts)}
|
||||
|
||||
|
||||
def backfill_provider_compliance_scores(tenant_id: str) -> dict:
|
||||
"""
|
||||
Backfill ProviderComplianceScore from latest completed scan per provider.
|
||||
|
||||
For each provider with completed scans, finds the most recent scan and
|
||||
upserts compliance requirement statuses with FAIL-dominant aggregation.
|
||||
|
||||
Args:
|
||||
tenant_id: Target tenant UUID
|
||||
|
||||
Returns:
|
||||
dict: Statistics about the backfill operation
|
||||
"""
|
||||
with rls_transaction(tenant_id, using=READ_REPLICA_ALIAS):
|
||||
completed_scans = Scan.all_objects.filter(
|
||||
tenant_id=tenant_id,
|
||||
state=StateChoices.COMPLETED,
|
||||
completed_at__isnull=False,
|
||||
)
|
||||
if not completed_scans.exists():
|
||||
return {"status": "no completed scans"}
|
||||
|
||||
existing_providers = set(
|
||||
ProviderComplianceScore.objects.filter(tenant_id=tenant_id)
|
||||
.values_list("provider_id", flat=True)
|
||||
.distinct()
|
||||
)
|
||||
|
||||
if existing_providers:
|
||||
completed_scans = completed_scans.exclude(
|
||||
provider_id__in=existing_providers
|
||||
)
|
||||
|
||||
scan_info = list(
|
||||
completed_scans.order_by("provider_id", "-completed_at")
|
||||
.distinct("provider_id")
|
||||
.values("id", "provider_id", "completed_at")
|
||||
)
|
||||
|
||||
if not scan_info:
|
||||
return {"status": "no scans to process"}
|
||||
|
||||
total_upserted = 0
|
||||
providers_processed = 0
|
||||
providers_skipped = 0
|
||||
|
||||
for scan in scan_info:
|
||||
provider_id = scan["provider_id"]
|
||||
|
||||
scan_id = scan["id"]
|
||||
|
||||
try:
|
||||
with psycopg_connection(MainRouter.default_db) as connection:
|
||||
connection.autocommit = False
|
||||
try:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(
|
||||
SET_CONFIG_QUERY, [POSTGRES_TENANT_VAR, tenant_id]
|
||||
)
|
||||
cursor.execute(
|
||||
COMPLIANCE_UPSERT_PROVIDER_SCORE_SQL,
|
||||
[tenant_id, str(scan_id)],
|
||||
)
|
||||
upserted = cursor.rowcount
|
||||
connection.commit()
|
||||
total_upserted += upserted
|
||||
providers_processed += 1
|
||||
except Exception:
|
||||
connection.rollback()
|
||||
raise
|
||||
except Exception as e:
|
||||
providers_skipped += 1
|
||||
logger.exception(
|
||||
"Error backfilling provider %s for tenant %s: %s",
|
||||
provider_id,
|
||||
tenant_id,
|
||||
e,
|
||||
)
|
||||
|
||||
# Recalculate tenant summary after all providers are backfilled
|
||||
if providers_processed > 0:
|
||||
with psycopg_connection(MainRouter.default_db) as connection:
|
||||
connection.autocommit = False
|
||||
try:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(SET_CONFIG_QUERY, [POSTGRES_TENANT_VAR, tenant_id])
|
||||
# Advisory lock to prevent race conditions
|
||||
cursor.execute(
|
||||
"SELECT pg_advisory_xact_lock(hashtext(%s))", [tenant_id]
|
||||
)
|
||||
cursor.execute(
|
||||
COMPLIANCE_UPSERT_TENANT_SUMMARY_ALL_SQL,
|
||||
[tenant_id, tenant_id],
|
||||
)
|
||||
tenant_summary_count = cursor.rowcount
|
||||
connection.commit()
|
||||
except Exception:
|
||||
connection.rollback()
|
||||
raise
|
||||
else:
|
||||
tenant_summary_count = 0
|
||||
|
||||
return {
|
||||
"status": "backfilled",
|
||||
"providers_processed": providers_processed,
|
||||
"providers_skipped": providers_skipped,
|
||||
"total_upserted": total_upserted,
|
||||
"tenant_summary_count": tenant_summary_count,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
"""
|
||||
Shared SQL queries for tasks.
|
||||
|
||||
This module centralizes raw SQL queries used across multiple task modules
|
||||
to ensure consistency and maintainability.
|
||||
"""
|
||||
|
||||
# =============================================================================
|
||||
# COMPLIANCE SCORE QUERIES
|
||||
# =============================================================================
|
||||
|
||||
# Upsert provider compliance scores from a scan's compliance requirements.
|
||||
# Uses FAIL-dominant aggregation: FAIL > MANUAL > PASS
|
||||
# Parameters: [tenant_id, scan_id]
|
||||
COMPLIANCE_UPSERT_PROVIDER_SCORE_SQL = """
|
||||
INSERT INTO provider_compliance_scores
|
||||
(id, tenant_id, provider_id, scan_id, compliance_id, requirement_id,
|
||||
requirement_status, scan_completed_at)
|
||||
SELECT
|
||||
gen_random_uuid(),
|
||||
agg.tenant_id,
|
||||
agg.provider_id,
|
||||
agg.scan_id,
|
||||
agg.compliance_id,
|
||||
agg.requirement_id,
|
||||
agg.requirement_status,
|
||||
agg.completed_at
|
||||
FROM (
|
||||
SELECT DISTINCT ON (cro.compliance_id, cro.requirement_id)
|
||||
cro.tenant_id,
|
||||
s.provider_id,
|
||||
cro.scan_id,
|
||||
cro.compliance_id,
|
||||
cro.requirement_id,
|
||||
(CASE
|
||||
WHEN bool_or(cro.requirement_status = 'FAIL')
|
||||
OVER (PARTITION BY cro.compliance_id, cro.requirement_id) THEN 'FAIL'
|
||||
WHEN bool_or(cro.requirement_status = 'MANUAL')
|
||||
OVER (PARTITION BY cro.compliance_id, cro.requirement_id) THEN 'MANUAL'
|
||||
ELSE 'PASS'
|
||||
END)::status as requirement_status,
|
||||
s.completed_at
|
||||
FROM compliance_requirements_overviews cro
|
||||
JOIN scans s ON s.id = cro.scan_id
|
||||
WHERE cro.tenant_id = %s AND cro.scan_id = %s
|
||||
ORDER BY cro.compliance_id, cro.requirement_id
|
||||
) agg
|
||||
ON CONFLICT (tenant_id, provider_id, compliance_id, requirement_id)
|
||||
DO UPDATE SET
|
||||
requirement_status = EXCLUDED.requirement_status,
|
||||
scan_id = EXCLUDED.scan_id,
|
||||
scan_completed_at = EXCLUDED.scan_completed_at
|
||||
WHERE EXCLUDED.scan_completed_at > provider_compliance_scores.scan_completed_at
|
||||
"""
|
||||
|
||||
# Upsert tenant compliance summary for specific compliance IDs.
|
||||
# Aggregates across all providers with FAIL-dominant logic at requirement level.
|
||||
# Parameters: [tenant_id, tenant_id, compliance_ids_array]
|
||||
COMPLIANCE_UPSERT_TENANT_SUMMARY_SQL = """
|
||||
INSERT INTO tenant_compliance_summaries
|
||||
(id, tenant_id, compliance_id,
|
||||
requirements_passed, requirements_failed, requirements_manual,
|
||||
total_requirements, updated_at)
|
||||
SELECT
|
||||
gen_random_uuid(),
|
||||
%s as tenant_id,
|
||||
compliance_id,
|
||||
COUNT(*) FILTER (WHERE req_status = 'PASS') as requirements_passed,
|
||||
COUNT(*) FILTER (WHERE req_status = 'FAIL') as requirements_failed,
|
||||
COUNT(*) FILTER (WHERE req_status = 'MANUAL') as requirements_manual,
|
||||
COUNT(*) as total_requirements,
|
||||
NOW() as updated_at
|
||||
FROM (
|
||||
SELECT
|
||||
compliance_id,
|
||||
requirement_id,
|
||||
CASE
|
||||
WHEN bool_or(requirement_status = 'FAIL') THEN 'FAIL'
|
||||
WHEN bool_or(requirement_status = 'MANUAL') THEN 'MANUAL'
|
||||
ELSE 'PASS'
|
||||
END as req_status
|
||||
FROM provider_compliance_scores
|
||||
WHERE tenant_id = %s AND compliance_id = ANY(%s)
|
||||
GROUP BY compliance_id, requirement_id
|
||||
) req_agg
|
||||
GROUP BY compliance_id
|
||||
ON CONFLICT (tenant_id, compliance_id)
|
||||
DO UPDATE SET
|
||||
requirements_passed = EXCLUDED.requirements_passed,
|
||||
requirements_failed = EXCLUDED.requirements_failed,
|
||||
requirements_manual = EXCLUDED.requirements_manual,
|
||||
total_requirements = EXCLUDED.total_requirements,
|
||||
updated_at = NOW()
|
||||
"""
|
||||
|
||||
# Upsert tenant compliance summary for ALL compliance IDs in tenant.
|
||||
# Used by backfill when recalculating entire tenant summary.
|
||||
# Parameters: [tenant_id, tenant_id]
|
||||
COMPLIANCE_UPSERT_TENANT_SUMMARY_ALL_SQL = """
|
||||
INSERT INTO tenant_compliance_summaries
|
||||
(id, tenant_id, compliance_id,
|
||||
requirements_passed, requirements_failed, requirements_manual,
|
||||
total_requirements, updated_at)
|
||||
SELECT
|
||||
gen_random_uuid(),
|
||||
%s as tenant_id,
|
||||
compliance_id,
|
||||
COUNT(*) FILTER (WHERE req_status = 'PASS') as requirements_passed,
|
||||
COUNT(*) FILTER (WHERE req_status = 'FAIL') as requirements_failed,
|
||||
COUNT(*) FILTER (WHERE req_status = 'MANUAL') as requirements_manual,
|
||||
COUNT(*) as total_requirements,
|
||||
NOW() as updated_at
|
||||
FROM (
|
||||
SELECT
|
||||
compliance_id,
|
||||
requirement_id,
|
||||
CASE
|
||||
WHEN bool_or(requirement_status = 'FAIL') THEN 'FAIL'
|
||||
WHEN bool_or(requirement_status = 'MANUAL') THEN 'MANUAL'
|
||||
ELSE 'PASS'
|
||||
END as req_status
|
||||
FROM provider_compliance_scores
|
||||
WHERE tenant_id = %s
|
||||
GROUP BY compliance_id, requirement_id
|
||||
) req_agg
|
||||
GROUP BY compliance_id
|
||||
ON CONFLICT (tenant_id, compliance_id)
|
||||
DO UPDATE SET
|
||||
requirements_passed = EXCLUDED.requirements_passed,
|
||||
requirements_failed = EXCLUDED.requirements_failed,
|
||||
requirements_manual = EXCLUDED.requirements_manual,
|
||||
total_requirements = EXCLUDED.total_requirements,
|
||||
updated_at = NOW()
|
||||
"""
|
||||
@@ -14,6 +14,10 @@ from config.env import env
|
||||
from config.settings.celery import CELERY_DEADLOCK_ATTEMPTS
|
||||
from django.db import IntegrityError, OperationalError
|
||||
from django.db.models import Case, Count, IntegerField, Prefetch, Q, Sum, When
|
||||
from tasks.jobs.queries import (
|
||||
COMPLIANCE_UPSERT_PROVIDER_SCORE_SQL,
|
||||
COMPLIANCE_UPSERT_TENANT_SUMMARY_SQL,
|
||||
)
|
||||
from tasks.utils import CustomEncoder
|
||||
|
||||
from api.compliance import PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE
|
||||
@@ -1489,3 +1493,140 @@ def aggregate_daily_severity(tenant_id: str, scan_id: str):
|
||||
"date": str(scan_date),
|
||||
"severity_data": severity_data,
|
||||
}
|
||||
|
||||
|
||||
def update_provider_compliance_scores(tenant_id: str, scan_id: str):
|
||||
"""
|
||||
Update ProviderComplianceScore with requirement statuses from a completed scan.
|
||||
|
||||
Uses atomic SQL upsert with ON CONFLICT for concurrency safety. Only updates
|
||||
if the new scan is more recent than existing data. Also cleans up stale
|
||||
requirements that no longer exist in the new scan.
|
||||
|
||||
Reads from primary DB (not replica) to avoid replication lag issues since
|
||||
this runs immediately after create_compliance_requirements_task.
|
||||
|
||||
Args:
|
||||
tenant_id: Tenant that owns the scan.
|
||||
scan_id: Scan UUID whose compliance data should be materialized.
|
||||
|
||||
Returns:
|
||||
dict: Statistics about the upsert operation.
|
||||
"""
|
||||
with rls_transaction(tenant_id):
|
||||
scan = (
|
||||
Scan.all_objects.filter(
|
||||
tenant_id=tenant_id,
|
||||
id=scan_id,
|
||||
state=StateChoices.COMPLETED,
|
||||
)
|
||||
.select_related("provider")
|
||||
.first()
|
||||
)
|
||||
|
||||
if not scan:
|
||||
logger.warning(
|
||||
f"Scan {scan_id} not found or not completed for compliance score update"
|
||||
)
|
||||
return {"status": "skipped", "reason": "scan not completed"}
|
||||
|
||||
if not scan.completed_at:
|
||||
logger.warning(f"Scan {scan_id} has no completed_at timestamp")
|
||||
return {"status": "skipped", "reason": "no completed_at"}
|
||||
|
||||
provider_id = str(scan.provider_id)
|
||||
scan_completed_at = scan.completed_at
|
||||
|
||||
delete_stale_sql = """
|
||||
DELETE FROM provider_compliance_scores pcs
|
||||
WHERE pcs.tenant_id = %s
|
||||
AND pcs.provider_id = %s
|
||||
AND pcs.scan_completed_at < %s
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM compliance_requirements_overviews cro
|
||||
WHERE cro.tenant_id = pcs.tenant_id
|
||||
AND cro.scan_id = %s
|
||||
AND cro.compliance_id = pcs.compliance_id
|
||||
AND cro.requirement_id = pcs.requirement_id
|
||||
)
|
||||
RETURNING compliance_id
|
||||
"""
|
||||
|
||||
compliance_ids_sql = """
|
||||
SELECT DISTINCT compliance_id
|
||||
FROM compliance_requirements_overviews
|
||||
WHERE tenant_id = %s AND scan_id = %s
|
||||
"""
|
||||
|
||||
try:
|
||||
with psycopg_connection(MainRouter.default_db) as connection:
|
||||
connection.autocommit = False
|
||||
try:
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(SET_CONFIG_QUERY, [POSTGRES_TENANT_VAR, tenant_id])
|
||||
|
||||
# Update requirement-level scores per provider
|
||||
cursor.execute(
|
||||
COMPLIANCE_UPSERT_PROVIDER_SCORE_SQL, [tenant_id, scan_id]
|
||||
)
|
||||
upserted_count = cursor.rowcount
|
||||
|
||||
cursor.execute(compliance_ids_sql, [tenant_id, scan_id])
|
||||
scan_rows = cursor.fetchall()
|
||||
if not isinstance(scan_rows, (list, tuple)):
|
||||
scan_rows = []
|
||||
scan_compliance_ids = {row[0] for row in scan_rows}
|
||||
|
||||
cursor.execute(
|
||||
delete_stale_sql,
|
||||
[tenant_id, provider_id, scan_completed_at, scan_id],
|
||||
)
|
||||
deleted_rows = cursor.fetchall()
|
||||
if not isinstance(deleted_rows, (list, tuple)):
|
||||
deleted_rows = []
|
||||
deleted_ids = {row[0] for row in deleted_rows}
|
||||
stale_deleted = len(deleted_ids)
|
||||
|
||||
impacted_compliance_ids = sorted(scan_compliance_ids | deleted_ids)
|
||||
|
||||
if impacted_compliance_ids:
|
||||
# Advisory lock on tenant to prevent race conditions when
|
||||
# multiple scans complete simultaneously for the same tenant
|
||||
cursor.execute(
|
||||
"SELECT pg_advisory_xact_lock(hashtext(%s))", [tenant_id]
|
||||
)
|
||||
|
||||
# Recalculate tenant-level summary (FAIL-dominant across all providers)
|
||||
cursor.execute(
|
||||
COMPLIANCE_UPSERT_TENANT_SUMMARY_SQL,
|
||||
[tenant_id, tenant_id, impacted_compliance_ids],
|
||||
)
|
||||
tenant_summary_count = cursor.rowcount
|
||||
else:
|
||||
tenant_summary_count = 0
|
||||
|
||||
connection.commit()
|
||||
except Exception:
|
||||
connection.rollback()
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
f"Provider compliance scores updated for scan {scan_id}: "
|
||||
f"{upserted_count} upserted, {stale_deleted} stale deleted, "
|
||||
f"{tenant_summary_count} tenant summaries upserted"
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"scan_id": str(scan_id),
|
||||
"provider_id": provider_id,
|
||||
"upserted": upserted_count,
|
||||
"stale_deleted": stale_deleted,
|
||||
"tenant_summary_count": tenant_summary_count,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error updating provider compliance scores for scan {scan_id}: {e}"
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -11,6 +11,7 @@ from django_celery_beat.models import PeriodicTask
|
||||
from tasks.jobs.backfill import (
|
||||
backfill_compliance_summaries,
|
||||
backfill_daily_severity_summaries,
|
||||
backfill_provider_compliance_scores,
|
||||
backfill_resource_scan_summaries,
|
||||
backfill_scan_category_summaries,
|
||||
)
|
||||
@@ -44,6 +45,7 @@ from tasks.jobs.scan import (
|
||||
aggregate_findings,
|
||||
create_compliance_requirements,
|
||||
perform_prowler_scan,
|
||||
update_provider_compliance_scores,
|
||||
)
|
||||
from tasks.utils import batched, get_next_execution_datetime
|
||||
|
||||
@@ -122,9 +124,10 @@ def _perform_scan_complete_tasks(tenant_id: str, scan_id: str, provider_id: str)
|
||||
scan_id (str): The ID of the scan that was performed.
|
||||
provider_id (str): The primary key of the Provider instance that was scanned.
|
||||
"""
|
||||
create_compliance_requirements_task.apply_async(
|
||||
kwargs={"tenant_id": tenant_id, "scan_id": scan_id}
|
||||
)
|
||||
chain(
|
||||
create_compliance_requirements_task.si(tenant_id=tenant_id, scan_id=scan_id),
|
||||
update_provider_compliance_scores_task.si(tenant_id=tenant_id, scan_id=scan_id),
|
||||
).apply_async()
|
||||
aggregate_attack_surface_task.apply_async(
|
||||
kwargs={"tenant_id": tenant_id, "scan_id": scan_id}
|
||||
)
|
||||
@@ -610,6 +613,20 @@ def backfill_scan_category_summaries_task(tenant_id: str, scan_id: str):
|
||||
return backfill_scan_category_summaries(tenant_id=tenant_id, scan_id=scan_id)
|
||||
|
||||
|
||||
@shared_task(name="backfill-provider-compliance-scores", queue="backfill")
|
||||
def backfill_provider_compliance_scores_task(tenant_id: str):
|
||||
"""
|
||||
Backfill ProviderComplianceScore from latest completed scan per provider.
|
||||
|
||||
Used to populate the compliance watchlist materialized table for tenants
|
||||
that had scans before the feature was deployed.
|
||||
|
||||
Args:
|
||||
tenant_id: Target tenant UUID.
|
||||
"""
|
||||
return backfill_provider_compliance_scores(tenant_id=tenant_id)
|
||||
|
||||
|
||||
@shared_task(base=RLSTask, name="scan-compliance-overviews", queue="compliance")
|
||||
@handle_provider_deletion
|
||||
def create_compliance_requirements_task(tenant_id: str, scan_id: str):
|
||||
@@ -643,6 +660,21 @@ def aggregate_attack_surface_task(tenant_id: str, scan_id: str):
|
||||
return aggregate_attack_surface(tenant_id=tenant_id, scan_id=scan_id)
|
||||
|
||||
|
||||
@shared_task(name="scan-provider-compliance-scores", queue="compliance")
|
||||
def update_provider_compliance_scores_task(tenant_id: str, scan_id: str):
|
||||
"""
|
||||
Update provider compliance scores from a completed scan.
|
||||
|
||||
This task materializes compliance requirement statuses into ProviderComplianceScore
|
||||
for efficient watchlist queries. Uses atomic upsert with concurrency protection.
|
||||
|
||||
Args:
|
||||
tenant_id (str): The tenant ID for which to update scores.
|
||||
scan_id (str): The ID of the scan whose data should be materialized.
|
||||
"""
|
||||
return update_provider_compliance_scores(tenant_id=tenant_id, scan_id=scan_id)
|
||||
|
||||
|
||||
@shared_task(name="scan-daily-severity", queue="overview")
|
||||
@handle_provider_deletion
|
||||
def aggregate_daily_severity_task(tenant_id: str, scan_id: str):
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from tasks.jobs.backfill import (
|
||||
backfill_compliance_summaries,
|
||||
backfill_provider_compliance_scores,
|
||||
backfill_resource_scan_summaries,
|
||||
backfill_scan_category_summaries,
|
||||
)
|
||||
@@ -260,3 +263,62 @@ class TestBackfillScanCategorySummaries:
|
||||
assert summary.total_findings == 1
|
||||
assert summary.failed_findings == 1
|
||||
assert summary.new_failed_findings == 1
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestBackfillProviderComplianceScores:
|
||||
def test_no_completed_scans(self, tenants_fixture):
|
||||
tenant = tenants_fixture[2]
|
||||
result = backfill_provider_compliance_scores(str(tenant.id))
|
||||
assert result == {"status": "no completed scans"}
|
||||
|
||||
def test_no_scans_to_process(self, tenants_fixture, scans_fixture):
|
||||
tenant = tenants_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
scan.completed_at = None
|
||||
scan.save()
|
||||
|
||||
result = backfill_provider_compliance_scores(str(tenant.id))
|
||||
assert result == {"status": "no completed scans"}
|
||||
|
||||
@patch("tasks.jobs.backfill.psycopg_connection")
|
||||
def test_successful_backfill_executes_sql_queries(
|
||||
self,
|
||||
mock_psycopg_connection,
|
||||
tenants_fixture,
|
||||
scans_fixture,
|
||||
settings,
|
||||
):
|
||||
"""Test successful backfill executes SQL queries and returns correct stats."""
|
||||
settings.DATABASES.setdefault("admin", settings.DATABASES["default"])
|
||||
tenant = tenants_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
|
||||
# Set completed_at to make the scan eligible for backfill
|
||||
scan.completed_at = datetime.now(timezone.utc)
|
||||
scan.save()
|
||||
|
||||
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
|
||||
connection.autocommit = True
|
||||
|
||||
context_manager = MagicMock()
|
||||
context_manager.__enter__.return_value = connection
|
||||
context_manager.__exit__.return_value = False
|
||||
mock_psycopg_connection.return_value = context_manager
|
||||
|
||||
cursor.rowcount = 5
|
||||
|
||||
result = backfill_provider_compliance_scores(str(tenant.id))
|
||||
|
||||
assert result["status"] == "backfilled"
|
||||
assert result["providers_processed"] == 1
|
||||
assert result["providers_skipped"] == 0
|
||||
assert result["total_upserted"] == 5
|
||||
assert result["tenant_summary_count"] == 5
|
||||
|
||||
@@ -24,6 +24,7 @@ from tasks.jobs.scan import (
|
||||
aggregate_findings,
|
||||
create_compliance_requirements,
|
||||
perform_prowler_scan,
|
||||
update_provider_compliance_scores,
|
||||
)
|
||||
from tasks.utils import CustomEncoder
|
||||
|
||||
@@ -4022,3 +4023,123 @@ class TestAggregateCategoryCounts:
|
||||
assert len(cache) == 3
|
||||
for cat in ["security", "compliance", "data-protection"]:
|
||||
assert cache[(cat, "low")] == {"total": 1, "failed": 1, "new_failed": 1}
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestUpdateProviderComplianceScores:
|
||||
@patch("tasks.jobs.scan.psycopg_connection")
|
||||
def test_update_provider_compliance_scores_basic(
|
||||
self,
|
||||
mock_psycopg_connection,
|
||||
tenants_fixture,
|
||||
scans_fixture,
|
||||
settings,
|
||||
):
|
||||
settings.DATABASES.setdefault("admin", settings.DATABASES["default"])
|
||||
tenant = tenants_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
tenant_id = str(tenant.id)
|
||||
scan_id = str(scan.id)
|
||||
|
||||
scan.state = StateChoices.COMPLETED
|
||||
scan.completed_at = datetime.now(timezone.utc)
|
||||
scan.save()
|
||||
|
||||
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
|
||||
connection.autocommit = True
|
||||
|
||||
context_manager = MagicMock()
|
||||
context_manager.__enter__.return_value = connection
|
||||
context_manager.__exit__.return_value = False
|
||||
mock_psycopg_connection.return_value = context_manager
|
||||
|
||||
cursor.rowcount = 2
|
||||
|
||||
result = update_provider_compliance_scores(tenant_id, scan_id)
|
||||
|
||||
assert result["status"] == "completed"
|
||||
assert result["upserted"] == 2
|
||||
assert cursor.execute.call_count >= 3
|
||||
connection.commit.assert_called_once()
|
||||
|
||||
def test_update_provider_compliance_scores_skips_incomplete_scan(
|
||||
self, tenants_fixture, scans_fixture
|
||||
):
|
||||
tenant = tenants_fixture[0]
|
||||
scan = scans_fixture[1]
|
||||
tenant_id = str(tenant.id)
|
||||
scan_id = str(scan.id)
|
||||
|
||||
result = update_provider_compliance_scores(tenant_id, scan_id)
|
||||
|
||||
assert result["status"] == "skipped"
|
||||
assert result["reason"] == "scan not completed"
|
||||
|
||||
def test_update_provider_compliance_scores_skips_no_completed_at(
|
||||
self, tenants_fixture, scans_fixture
|
||||
):
|
||||
tenant = tenants_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
tenant_id = str(tenant.id)
|
||||
scan_id = str(scan.id)
|
||||
|
||||
scan.state = StateChoices.COMPLETED
|
||||
scan.completed_at = None
|
||||
scan.save()
|
||||
|
||||
result = update_provider_compliance_scores(tenant_id, scan_id)
|
||||
|
||||
assert result["status"] == "skipped"
|
||||
assert result["reason"] == "no completed_at"
|
||||
|
||||
@patch("tasks.jobs.scan.psycopg_connection")
|
||||
def test_update_provider_compliance_scores_executes_sql_queries(
|
||||
self,
|
||||
mock_psycopg_connection,
|
||||
tenants_fixture,
|
||||
providers_fixture,
|
||||
scans_fixture,
|
||||
settings,
|
||||
):
|
||||
settings.DATABASES.setdefault("admin", settings.DATABASES["default"])
|
||||
tenant = tenants_fixture[0]
|
||||
scan = scans_fixture[0]
|
||||
tenant_id = str(tenant.id)
|
||||
scan_id = str(scan.id)
|
||||
|
||||
scan.state = StateChoices.COMPLETED
|
||||
scan.completed_at = datetime.now(timezone.utc)
|
||||
scan.save()
|
||||
|
||||
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
|
||||
|
||||
cursor.rowcount = 1
|
||||
cursor.fetchall.side_effect = [[("aws_cis_2.0",)], []]
|
||||
|
||||
result = update_provider_compliance_scores(tenant_id, scan_id)
|
||||
|
||||
assert result["status"] == "completed"
|
||||
|
||||
calls = [str(c) for c in cursor.execute.call_args_list]
|
||||
assert any("provider_compliance_scores" in c for c in calls)
|
||||
assert any("tenant_compliance_summaries" in c for c in calls)
|
||||
assert any("pg_advisory_xact_lock" in c for c in calls)
|
||||
|
||||
@@ -730,7 +730,9 @@ class TestGenerateOutputs:
|
||||
|
||||
class TestScanCompleteTasks:
|
||||
@patch("tasks.tasks.aggregate_attack_surface_task.apply_async")
|
||||
@patch("tasks.tasks.create_compliance_requirements_task.apply_async")
|
||||
@patch("tasks.tasks.chain")
|
||||
@patch("tasks.tasks.create_compliance_requirements_task.si")
|
||||
@patch("tasks.tasks.update_provider_compliance_scores_task.si")
|
||||
@patch("tasks.tasks.perform_scan_summary_task.si")
|
||||
@patch("tasks.tasks.generate_outputs_task.si")
|
||||
@patch("tasks.tasks.generate_compliance_reports_task.si")
|
||||
@@ -741,15 +743,22 @@ class TestScanCompleteTasks:
|
||||
mock_compliance_reports_task,
|
||||
mock_outputs_task,
|
||||
mock_scan_summary_task,
|
||||
mock_update_compliance_scores_task,
|
||||
mock_compliance_requirements_task,
|
||||
mock_chain,
|
||||
mock_attack_surface_task,
|
||||
):
|
||||
"""Test that scan complete tasks are properly orchestrated with optimized reports."""
|
||||
_perform_scan_complete_tasks("tenant-id", "scan-id", "provider-id")
|
||||
|
||||
# Verify compliance requirements task is called
|
||||
# Verify compliance requirements task is called via chain
|
||||
mock_compliance_requirements_task.assert_called_once_with(
|
||||
kwargs={"tenant_id": "tenant-id", "scan_id": "scan-id"},
|
||||
tenant_id="tenant-id", scan_id="scan-id"
|
||||
)
|
||||
|
||||
# Verify update provider compliance scores task is called via chain
|
||||
mock_update_compliance_scores_task.assert_called_once_with(
|
||||
tenant_id="tenant-id", scan_id="scan-id"
|
||||
)
|
||||
|
||||
# Verify attack surface task is called
|
||||
|
||||
Reference in New Issue
Block a user