Merge remote-tracking branch 'origin/master' into mintlify/cb114598

This commit is contained in:
mintlify[bot]
2026-07-08 12:32:39 +00:00
29 changed files with 1561 additions and 1315 deletions
+1 -1
View File
@@ -27,7 +27,7 @@ Please add a detailed description of how to review this PR.
- [ ] Review if the code is being covered by tests.
- [ ] Review if code is being documented following this specification https://github.com/google/styleguide/blob/gh-pages/pyguide.md#38-comments-and-docstrings
- [ ] Review if backport is needed.
- [ ] Review if is needed to change the [Readme.md](https://github.com/prowler-cloud/prowler/blob/master/README.md)
- [ ] Review if is needed to change the [README.md](https://github.com/prowler-cloud/prowler/blob/master/README.md)
- [ ] Ensure a changelog fragment is added under [prowler/changelog.d/](https://github.com/prowler-cloud/prowler/tree/master/prowler/changelog.d), if applicable.
#### SDK/CLI
@@ -187,7 +187,7 @@ def test_user_me_when_inviting_users(create_test_user, tenants_fixture, roles_fi
@pytest.mark.django_db
class TestTokenSwitchTenant:
def test_switch_tenant_with_valid_token(self, tenants_fixture, providers_fixture):
def test_switch_tenant_with_valid_token(self, tenants_fixture, aws_provider):
client = APIClient()
test_user = "test_email@prowler.com"
@@ -1396,7 +1396,7 @@ class TestAPIKeyMultiTenantWorkflows:
assert me_response2.json()["data"]["id"] == str(user.id)
def test_api_key_cannot_access_different_tenant_resources(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""API key from one tenant cannot access resources from another tenant.
+6 -6
View File
@@ -40,10 +40,10 @@ class TestSetTenantDecorator:
@pytest.mark.django_db
class TestHandleProviderDeletionDecorator:
def test_success_no_exception(self, tenants_fixture, providers_fixture):
def test_success_no_exception(self, tenants_fixture, aws_provider):
"""Decorated function runs normally when no exception is raised."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
@handle_provider_deletion
def task_func(**kwargs):
@@ -127,11 +127,11 @@ class TestHandleProviderDeletionDecorator:
@patch("api.decorators.rls_transaction")
@patch("api.decorators.Provider.objects.filter")
def test_provider_exists_reraises_original(
self, mock_filter, mock_rls, tenants_fixture, providers_fixture
self, mock_filter, mock_rls, tenants_fixture, aws_provider
):
"""Re-raises original exception when provider still exists."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
mock_rls.return_value.__enter__ = lambda s: None
mock_rls.return_value.__exit__ = lambda s, *args: None
@@ -187,11 +187,11 @@ class TestHandleProviderDeletionDecorator:
@patch("api.decorators.rls_transaction")
@patch("api.decorators.Provider.objects.filter")
def test_database_error_provider_exists_reraises(
self, mock_filter, mock_rls, tenants_fixture, providers_fixture
self, mock_filter, mock_rls, tenants_fixture, aws_provider
):
"""Re-raises original DatabaseError when provider still exists."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
mock_rls.return_value.__enter__ = lambda s: None
mock_rls.return_value.__exit__ = lambda s, *args: None
+10 -10
View File
@@ -19,8 +19,8 @@ from django.db import IntegrityError
@pytest.mark.django_db
class TestResourceModel:
def test_setting_tags(self, providers_fixture):
provider, *_ = providers_fixture
def test_setting_tags(self, aws_provider):
provider = aws_provider
tenant_id = provider.tenant_id
resource = Resource.objects.create(
@@ -111,9 +111,9 @@ class TestResourceModel:
# @pytest.mark.django_db
# class TestFindingModel:
# def test_add_finding_with_long_uid(
# self, providers_fixture, scans_fixture, resources_fixture
# self, aws_provider, scans_fixture, resources_fixture
# ):
# provider, *_ = providers_fixture
# provider = aws_provider
# tenant_id = provider.tenant_id
# long_uid = "1" * 500
@@ -372,8 +372,8 @@ class TestSAMLConfigurationModel:
@pytest.mark.django_db
class TestProviderComplianceScoreModel:
def test_create_provider_compliance_score(self, providers_fixture, scans_fixture):
provider = providers_fixture[0]
def test_create_provider_compliance_score(self, aws_provider, scans_fixture):
provider = aws_provider
scan = scans_fixture[0]
scan.completed_at = datetime.now(UTC)
scan.save()
@@ -393,9 +393,9 @@ class TestProviderComplianceScoreModel:
assert score.requirement_status == StatusChoices.PASS
def test_unique_constraint_per_provider_compliance_requirement(
self, providers_fixture, scans_fixture
self, aws_provider, scans_fixture
):
provider = providers_fixture[0]
provider = aws_provider
scan = scans_fixture[0]
scan.completed_at = datetime.now(UTC)
scan.save()
@@ -422,9 +422,9 @@ class TestProviderComplianceScoreModel:
)
def test_different_providers_same_requirement_allowed(
self, providers_fixture, scans_fixture
self, aws_provider_pair, scans_fixture
):
provider1, provider2, *_ = providers_fixture
provider1, provider2 = aws_provider_pair
scan1 = scans_fixture[0]
scan1.completed_at = datetime.now(UTC)
scan1.save()
+33 -41
View File
@@ -434,11 +434,11 @@ class TestUserViewSet:
@pytest.mark.django_db
class TestProviderViewSet:
def test_list_providers_with_all_permissions(
self, authenticated_client_rbac, providers_fixture
self, authenticated_client_rbac, aws_provider
):
response = authenticated_client_rbac.get(reverse("provider-list"))
assert response.status_code == status.HTTP_200_OK
assert len(response.json()["data"]) == len(providers_fixture)
assert len(response.json()["data"]) == 1
def test_list_providers_with_no_permissions(
self, authenticated_client_no_permissions_rbac
@@ -450,9 +450,9 @@ class TestProviderViewSet:
assert len(response.json()["data"]) == 0
def test_retrieve_provider_with_all_permissions(
self, authenticated_client_rbac, providers_fixture
self, authenticated_client_rbac, aws_provider
):
provider = providers_fixture[0]
provider = aws_provider
response = authenticated_client_rbac.get(
reverse("provider-detail", kwargs={"pk": provider.id})
)
@@ -460,9 +460,9 @@ class TestProviderViewSet:
assert response.json()["data"]["attributes"]["alias"] == provider.alias
def test_retrieve_provider_with_no_permissions(
self, authenticated_client_no_permissions_rbac, providers_fixture
self, authenticated_client_no_permissions_rbac, aws_provider
):
provider = providers_fixture[0]
provider = aws_provider
response = authenticated_client_no_permissions_rbac.get(
reverse("provider-detail", kwargs={"pk": provider.id})
)
@@ -486,9 +486,9 @@ class TestProviderViewSet:
assert response.status_code == status.HTTP_403_FORBIDDEN
def test_partial_update_provider_with_all_permissions(
self, authenticated_client_rbac, providers_fixture
self, authenticated_client_rbac, aws_provider
):
provider = providers_fixture[0]
provider = aws_provider
payload = {
"data": {
"type": "providers",
@@ -505,9 +505,9 @@ class TestProviderViewSet:
assert response.json()["data"]["attributes"]["alias"] == "updated_alias"
def test_partial_update_provider_with_no_permissions(
self, authenticated_client_no_permissions_rbac, providers_fixture
self, authenticated_client_no_permissions_rbac, aws_provider
):
provider = providers_fixture[0]
provider = aws_provider
update_payload = {
"data": {
"type": "providers",
@@ -528,7 +528,7 @@ class TestProviderViewSet:
mock_delete_task,
mock_task_get,
authenticated_client_rbac,
providers_fixture,
aws_provider,
tasks_fixture,
):
prowler_task = tasks_fixture[0]
@@ -537,7 +537,7 @@ class TestProviderViewSet:
mock_delete_task.return_value = task_mock
mock_task_get.return_value = prowler_task
provider1, *_ = providers_fixture
provider1 = aws_provider
response = authenticated_client_rbac.delete(
reverse("provider-detail", kwargs={"pk": provider1.id})
)
@@ -549,9 +549,9 @@ class TestProviderViewSet:
assert response.headers["Content-Location"] == f"/api/v1/tasks/{task_mock.id}"
def test_delete_provider_with_no_permissions(
self, authenticated_client_no_permissions_rbac, providers_fixture
self, authenticated_client_no_permissions_rbac, aws_provider
):
provider = providers_fixture[0]
provider = aws_provider
response = authenticated_client_no_permissions_rbac.delete(
reverse("provider-detail", kwargs={"pk": provider.id})
)
@@ -564,7 +564,7 @@ class TestProviderViewSet:
mock_provider_connection,
mock_task_get,
authenticated_client_rbac,
providers_fixture,
aws_provider,
tasks_fixture,
):
prowler_task = tasks_fixture[0]
@@ -574,7 +574,7 @@ class TestProviderViewSet:
mock_provider_connection.return_value = task_mock
mock_task_get.return_value = prowler_task
provider1, *_ = providers_fixture
provider1 = aws_provider
assert provider1.connected is None
assert provider1.connection_last_checked_at is None
@@ -589,9 +589,9 @@ class TestProviderViewSet:
assert response.headers["Content-Location"] == f"/api/v1/tasks/{task_mock.id}"
def test_connection_with_no_permissions(
self, authenticated_client_no_permissions_rbac, providers_fixture
self, authenticated_client_no_permissions_rbac, aws_provider
):
provider = providers_fixture[0]
provider = aws_provider
response = authenticated_client_no_permissions_rbac.post(
reverse("provider-connection", kwargs={"pk": provider.id})
)
@@ -604,10 +604,10 @@ class TestLimitedVisibility:
TEST_PASSWORD = "Thisisapassword123@"
@pytest.fixture
def limited_admin_user(self, django_db_blocker, tenants_fixture, providers_fixture):
def limited_admin_user(self, django_db_blocker, tenants_fixture, aws_provider):
with django_db_blocker.unblock():
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
user = User.objects.create_user(
name="testing",
email=self.TEST_EMAIL,
@@ -654,25 +654,17 @@ class TestLimitedVisibility:
@pytest.fixture
def authenticated_client_rbac_limited(
self, limited_admin_user, tenants_fixture, client
self,
limited_admin_user,
tenants_fixture,
authenticated_client_for_tenant_factory,
):
client.user = limited_admin_user
tenant_id = tenants_fixture[0].id
serializer = TokenSerializer(
data={
"type": "tokens",
"email": self.TEST_EMAIL,
"password": self.TEST_PASSWORD,
"tenant_id": tenant_id,
}
return authenticated_client_for_tenant_factory(
limited_admin_user, tenants_fixture[0]
)
serializer.is_valid(raise_exception=True)
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
def test_integrations(
self, authenticated_client_rbac_limited, integrations_fixture, providers_fixture
self, authenticated_client_rbac_limited, integrations_fixture
):
# Integration 2 is related to provider1 and provider 2
# This user cannot see provider 2
@@ -692,7 +684,7 @@ class TestLimitedVisibility:
def test_overviews_providers(
self,
authenticated_client_rbac_limited,
providers_fixture,
provider_factory,
):
# By default, the associated provider is the one which has the overview data
response = authenticated_client_rbac_limited.get(reverse("overview-providers"))
@@ -702,7 +694,7 @@ class TestLimitedVisibility:
# Changing the provider visibility, no data should be returned
# Only the associated provider to that group is changed
new_provider = providers_fixture[1]
new_provider = provider_factory()
ProviderGroupMembership.objects.all().update(provider=new_provider)
response = authenticated_client_rbac_limited.get(reverse("overview-providers"))
@@ -722,7 +714,7 @@ class TestLimitedVisibility:
self,
endpoint_name,
authenticated_client_rbac_limited,
providers_fixture,
provider_factory,
):
# By default, the associated provider is the one which has the overview data
response = authenticated_client_rbac_limited.get(
@@ -735,7 +727,7 @@ class TestLimitedVisibility:
# Changing the provider visibility, no data should be returned
# Only the associated provider to that group is changed
new_provider = providers_fixture[1]
new_provider = provider_factory()
ProviderGroupMembership.objects.all().update(provider=new_provider)
response = authenticated_client_rbac_limited.get(
@@ -750,7 +742,7 @@ class TestLimitedVisibility:
def test_overviews_services(
self,
authenticated_client_rbac_limited,
providers_fixture,
provider_factory,
):
# By default, the associated provider is the one which has the overview data
response = authenticated_client_rbac_limited.get(
@@ -762,7 +754,7 @@ class TestLimitedVisibility:
# Changing the provider visibility, no data should be returned
# Only the associated provider to that group is changed
new_provider = providers_fixture[1]
new_provider = provider_factory()
ProviderGroupMembership.objects.all().update(provider=new_provider)
response = authenticated_client_rbac_limited.get(
+25 -1
View File
@@ -1,9 +1,33 @@
import logging
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
from config.settings import sentry as sentry_settings
from config.settings.sentry import before_send
def test_initialize_sentry_skips_without_dsn():
with (
patch.object(sentry_settings.env, "str", return_value=""),
patch.object(sentry_settings.sentry_sdk, "init") as mock_init,
):
sentry_settings.initialize_sentry()
mock_init.assert_not_called()
def test_initialize_sentry_uses_configured_dsn():
sentry_dsn = "https://fake-public-key@sentry.example.invalid/1"
with (
patch.object(sentry_settings.env, "str", return_value=sentry_dsn),
patch.object(sentry_settings.sentry_sdk, "init") as mock_init,
):
sentry_settings.initialize_sentry()
assert mock_init.call_args.kwargs["dsn"] == sentry_dsn
assert mock_init.call_args.kwargs["before_send"] is sentry_settings.before_send
def _make_log_record(msg, level=logging.ERROR, name="test", args=None):
"""Build a real LogRecord so getMessage() works like in production."""
record = logging.LogRecord(
+2 -2
View File
@@ -188,10 +188,10 @@ class TestProwlerProviderConnectionTest:
@pytest.mark.django_db
@patch("api.utils.return_prowler_provider")
def test_prowler_provider_connection_test_without_secret(
self, mock_return_prowler_provider, providers_fixture
self, mock_return_prowler_provider, aws_provider
):
mock_return_prowler_provider.return_value = MagicMock()
connection = prowler_provider_connection_test(providers_fixture[0])
connection = prowler_provider_connection_test(aws_provider)
assert connection.is_connected is False
assert isinstance(connection.error, Provider.secret.RelatedObjectDoesNotExist)
File diff suppressed because it is too large Load Diff
+24 -16
View File
@@ -115,19 +115,27 @@ def before_send(event, hint):
return event
sentry_sdk.init(
dsn=env.str("DJANGO_SENTRY_DSN", ""),
# Add data like request headers and IP for users,
# see https://docs.sentry.io/platforms/python/data-management/data-collected/ for more info
before_send=before_send,
send_default_pii=True,
traces_sample_rate=env.float("DJANGO_SENTRY_TRACES_SAMPLE_RATE", default=0.02),
_experiments={
# Set continuous_profiling_auto_start to True
# to automatically start the profiler on when
# possible.
"continuous_profiling_auto_start": True,
},
attach_stacktrace=True,
ignore_errors=IGNORED_EXCEPTIONS,
)
def initialize_sentry():
sentry_dsn = env.str("DJANGO_SENTRY_DSN", "")
if not sentry_dsn:
return
sentry_sdk.init(
dsn=sentry_dsn,
# Add data like request headers and IP for users,
# see https://docs.sentry.io/platforms/python/data-management/data-collected/ for more info
before_send=before_send,
send_default_pii=True,
traces_sample_rate=env.float("DJANGO_SENTRY_TRACES_SAMPLE_RATE", default=0.02),
_experiments={
# Set continuous_profiling_auto_start to True
# to automatically start the profiler on when
# possible.
"continuous_profiling_auto_start": True,
},
attach_stacktrace=True,
ignore_errors=IGNORED_EXCEPTIONS,
)
initialize_sentry()
+254 -196
View File
@@ -2,6 +2,7 @@ import logging
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from allauth.socialaccount.models import SocialLogin
@@ -50,12 +51,14 @@ from api.v1.serializers import TokenSerializer
from django.conf import settings
from django.db import connection as django_connection
from django.db import connections as django_connections
from django.test import Client
from django.urls import reverse
from django_celery_results.models import TaskResult
from prowler.lib.check.models import Severity
from prowler.lib.outputs.finding import Status
from rest_framework import status
from rest_framework.test import APIClient
from rest_framework_simplejwt.tokens import AccessToken
from tasks.jobs.backfill import (
aggregate_scan_category_summaries,
aggregate_scan_resource_group_summaries,
@@ -358,22 +361,42 @@ def create_test_user_rbac_manage_account(django_db_setup, django_db_blocker):
return user
def first_membership_tenant(user):
return user.memberships.order_by("date_joined").first().tenant
def access_token_for_tenant(user, tenant):
access_token = AccessToken.for_user(user)
access_token["tenant_id"] = str(tenant.id)
access_token.payload["nbf"] = access_token["iat"]
return str(access_token)
def authenticate_client_for_tenant(client, user, tenant):
client.user = user
client.defaults["HTTP_AUTHORIZATION"] = (
f"Bearer {access_token_for_tenant(user, tenant)}"
)
return client
@pytest.fixture
def authenticated_client_for_tenant_factory():
def create_authenticated_client(user, tenant):
return authenticate_client_for_tenant(Client(), user, tenant)
return create_authenticated_client
@pytest.fixture
def authenticated_client_rbac_manage_account(
create_test_user_rbac_manage_account, tenants_fixture, client
create_test_user_rbac_manage_account, client
):
client.user = create_test_user_rbac_manage_account
serializer = TokenSerializer(
data={
"type": "tokens",
"email": "rbac_manage_account@rbac.com",
"password": TEST_PASSWORD,
}
return authenticate_client_for_tenant(
client,
create_test_user_rbac_manage_account,
first_membership_tenant(create_test_user_rbac_manage_account),
)
serializer.is_valid()
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
@pytest.fixture(scope="function")
@@ -410,86 +433,43 @@ def create_test_user_rbac_manage_users_only(django_db_setup, django_db_blocker):
def authenticated_client_rbac_manage_users_only(
create_test_user_rbac_manage_users_only, client
):
client.user = create_test_user_rbac_manage_users_only
serializer = TokenSerializer(
data={
"type": "tokens",
"email": "rbac_manage_users_only@rbac.com",
"password": TEST_PASSWORD,
}
return authenticate_client_for_tenant(
client,
create_test_user_rbac_manage_users_only,
first_membership_tenant(create_test_user_rbac_manage_users_only),
)
serializer.is_valid()
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
@pytest.fixture
def authenticated_client_rbac(create_test_user_rbac, tenants_fixture, client):
client.user = create_test_user_rbac
tenant_id = tenants_fixture[0].id
serializer = TokenSerializer(
data={
"type": "tokens",
"email": "rbac@rbac.com",
"password": TEST_PASSWORD,
"tenant_id": tenant_id,
}
return authenticate_client_for_tenant(
client, create_test_user_rbac, tenants_fixture[0]
)
serializer.is_valid(raise_exception=True)
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
@pytest.fixture
def authenticated_client_rbac_noroles(
create_test_user_rbac_no_roles, tenants_fixture, client
):
client.user = create_test_user_rbac_no_roles
serializer = TokenSerializer(
data={
"type": "tokens",
"email": "rbac_noroles@rbac.com",
"password": TEST_PASSWORD,
}
return authenticate_client_for_tenant(
client, create_test_user_rbac_no_roles, tenants_fixture[0]
)
serializer.is_valid()
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
@pytest.fixture
def authenticated_client_no_permissions_rbac(
create_test_user_rbac_limited, tenants_fixture, client
):
client.user = create_test_user_rbac_limited
serializer = TokenSerializer(
data={
"type": "tokens",
"email": "rbac_limited@rbac.com",
"password": TEST_PASSWORD,
}
return authenticate_client_for_tenant(
client, create_test_user_rbac_limited, tenants_fixture[0]
)
serializer.is_valid()
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
@pytest.fixture
def authenticated_client(
create_test_user, tenants_fixture, set_user_admin_roles_fixture, client
):
client.user = create_test_user
serializer = TokenSerializer(
data={"type": "tokens", "email": TEST_USER, "password": TEST_PASSWORD}
)
serializer.is_valid()
access_token = serializer.validated_data["access"]
client.defaults["HTTP_AUTHORIZATION"] = f"Bearer {access_token}"
return client
return authenticate_client_for_tenant(client, create_test_user, tenants_fixture[0])
@pytest.fixture
@@ -590,109 +570,191 @@ def users_fixture(django_user_model):
@pytest.fixture
def providers_fixture(tenants_fixture):
tenant, *_ = tenants_fixture
provider1 = Provider.objects.create(
provider="aws",
uid="123456789012",
alias="aws_testing_1",
tenant_id=tenant.id,
)
provider2 = Provider.objects.create(
provider="aws",
uid="123456789013",
alias="aws_testing_2",
tenant_id=tenant.id,
)
provider3 = Provider.objects.create(
provider="gcp",
uid="a12322-test321",
alias="gcp_testing",
tenant_id=tenant.id,
)
provider4 = Provider.objects.create(
provider="kubernetes",
uid="kubernetes-test-12345",
alias="k8s_testing",
tenant_id=tenant.id,
)
provider5 = Provider.objects.create(
provider="azure",
uid="37b065f8-26b0-4218-a665-0b23d07b27d9",
alias="azure_testing",
tenant_id=tenant.id,
scanner_args={"key1": "value1", "key2": {"key21": "value21"}},
)
provider6 = Provider.objects.create(
provider="m365",
uid="m365.test.com",
alias="m365_testing",
tenant_id=tenant.id,
)
provider7 = Provider.objects.create(
provider="oraclecloud",
uid="ocid1.tenancy.oc1..aaaaaaaa3dwoazoox4q7wrvriywpokp5grlhgnkwtyt6dmwyou7no6mdmzda",
alias="oci_testing",
tenant_id=tenant.id,
)
provider8 = Provider.objects.create(
provider="mongodbatlas",
uid="64b1d3c0e4b03b1234567890",
alias="mongodbatlas_testing",
tenant_id=tenant.id,
)
provider9 = Provider.objects.create(
provider="alibabacloud",
uid="1234567890123456",
alias="alibabacloud_testing",
tenant_id=tenant.id,
)
provider10 = Provider.objects.create(
provider="cloudflare",
uid="a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4",
alias="cloudflare_testing",
tenant_id=tenant.id,
)
provider11 = Provider.objects.create(
provider="openstack",
uid="a1b2c3d4-e5f6-7890-abcd-ef1234567890",
alias="openstack_testing",
tenant_id=tenant.id,
)
provider12 = Provider.objects.create(
provider="googleworkspace",
uid="C12345678",
alias="googleworkspace_testing",
tenant_id=tenant.id,
)
provider13 = Provider.objects.create(
provider="vercel",
uid="team_abcdef1234567890ab",
alias="vercel_testing",
tenant_id=tenant.id,
)
provider14 = Provider.objects.create(
provider="okta",
uid="acme.okta.com",
alias="okta_testing",
tenant_id=tenant.id,
def provider_factory(tenants_fixture):
tenant = tenants_fixture[0]
counters = {}
def next_counter(provider):
counters[provider] = counters.get(provider, 0) + 1
return counters[provider]
def defaults_for(provider, sequence):
return {
Provider.ProviderChoices.AWS.value: {
"uid": f"{123456789011 + sequence:012d}",
"alias": f"aws_testing_{sequence}",
},
Provider.ProviderChoices.AZURE.value: {
"uid": str(uuid4()),
"alias": f"azure_testing_{sequence}",
"scanner_args": {"key1": "value1", "key2": {"key21": "value21"}},
},
Provider.ProviderChoices.GCP.value: {
"uid": f"a12322-test{sequence:05d}",
"alias": f"gcp_testing_{sequence}",
},
Provider.ProviderChoices.KUBERNETES.value: {
"uid": f"kubernetes-test-{sequence}",
"alias": f"k8s_testing_{sequence}",
},
Provider.ProviderChoices.M365.value: {
"uid": f"m365-{sequence}.test.com",
"alias": f"m365_testing_{sequence}",
},
Provider.ProviderChoices.GITHUB.value: {
"uid": f"github-test-{sequence}",
"alias": f"github_testing_{sequence}",
},
Provider.ProviderChoices.MONGODBATLAS.value: {
"uid": f"64b1d3c0e4b03b{sequence:010x}",
"alias": f"mongodbatlas_testing_{sequence}",
},
Provider.ProviderChoices.IAC.value: {
"uid": f"https://github.com/prowler-cloud/test-{sequence}.git",
"alias": f"iac_testing_{sequence}",
},
Provider.ProviderChoices.ORACLECLOUD.value: {
"uid": f"ocid1.tenancy.oc1..aaaaaaaa{sequence:024d}",
"alias": f"oci_testing_{sequence}",
},
Provider.ProviderChoices.ALIBABACLOUD.value: {
"uid": f"{1234567890123455 + sequence:016d}",
"alias": f"alibabacloud_testing_{sequence}",
},
Provider.ProviderChoices.CLOUDFLARE.value: {
"uid": f"{0x1000000000000000000000000000000 + sequence:032x}",
"alias": f"cloudflare_testing_{sequence}",
},
Provider.ProviderChoices.OPENSTACK.value: {
"uid": f"openstack-project-{sequence}",
"alias": f"openstack_testing_{sequence}",
},
Provider.ProviderChoices.IMAGE.value: {
"uid": f"registry.example.com/prowler/test:{sequence}",
"alias": f"image_testing_{sequence}",
},
Provider.ProviderChoices.GOOGLEWORKSPACE.value: {
"uid": f"C{12345677 + sequence}",
"alias": f"googleworkspace_testing_{sequence}",
},
Provider.ProviderChoices.VERCEL.value: {
"uid": f"team_{sequence:016x}",
"alias": f"vercel_testing_{sequence}",
},
Provider.ProviderChoices.OKTA.value: {
"uid": f"acme-{sequence}.okta.com",
"alias": f"okta_testing_{sequence}",
},
}[provider]
def create_provider(provider=Provider.ProviderChoices.AWS.value, **overrides):
provider_value = getattr(provider, "value", provider)
selected_tenant = overrides.pop("tenant", tenant)
sequence = next_counter(provider_value)
attributes = {
"provider": provider_value,
"tenant_id": selected_tenant.id,
**defaults_for(provider_value, sequence),
}
attributes.update(overrides)
return Provider.objects.create(**attributes)
return create_provider
@pytest.fixture
def aws_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.AWS.value)
@pytest.fixture
def aws_provider_pair(aws_provider, provider_factory):
return (
aws_provider,
provider_factory(Provider.ProviderChoices.AWS.value),
)
return (
provider1,
provider2,
provider3,
provider4,
provider5,
provider6,
provider7,
provider8,
provider9,
provider10,
provider11,
provider12,
provider13,
provider14,
@pytest.fixture
def azure_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.AZURE.value)
@pytest.fixture
def gcp_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.GCP.value)
@pytest.fixture
def kubernetes_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.KUBERNETES.value)
@pytest.fixture
def m365_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.M365.value)
@pytest.fixture
def github_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.GITHUB.value)
@pytest.fixture
def mongodbatlas_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.MONGODBATLAS.value)
@pytest.fixture
def iac_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.IAC.value)
@pytest.fixture
def oraclecloud_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.ORACLECLOUD.value)
@pytest.fixture
def alibabacloud_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.ALIBABACLOUD.value)
@pytest.fixture
def cloudflare_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.CLOUDFLARE.value)
@pytest.fixture
def openstack_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.OPENSTACK.value)
@pytest.fixture
def image_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.IMAGE.value)
@pytest.fixture
def googleworkspace_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.GOOGLEWORKSPACE.value)
@pytest.fixture
def vercel_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.VERCEL.value)
@pytest.fixture
def okta_provider(provider_factory):
return provider_factory(Provider.ProviderChoices.OKTA.value)
@pytest.fixture
def all_provider_types_fixture(provider_factory):
return tuple(
provider_factory(provider_choice.value)
for provider_choice in Provider.ProviderChoices
)
@@ -797,7 +859,7 @@ def roles_fixture(tenants_fixture):
@pytest.fixture
def provider_secret_fixture(providers_fixture):
def provider_secret_fixture(all_provider_types_fixture):
return tuple(
ProviderSecret.objects.create(
tenant_id=provider.tenant_id,
@@ -806,14 +868,14 @@ def provider_secret_fixture(providers_fixture):
secret={"key": "value"},
name=provider.alias,
)
for provider in providers_fixture
for provider in all_provider_types_fixture
)
@pytest.fixture
def scans_fixture(tenants_fixture, providers_fixture):
def scans_fixture(tenants_fixture, aws_provider_pair):
tenant, *_ = tenants_fixture
provider, provider2, *_ = providers_fixture
provider, provider2 = aws_provider_pair
now = datetime.now(UTC)
@@ -876,8 +938,8 @@ def tasks_fixture(tenants_fixture):
@pytest.fixture
def resources_fixture(providers_fixture):
provider, *_ = providers_fixture
def resources_fixture(aws_provider_pair):
provider, provider2 = aws_provider_pair
tags = [
ResourceTag.objects.create(
@@ -918,8 +980,8 @@ def resources_fixture(providers_fixture):
resource2.upsert_or_delete_tags(tags)
resource3 = Resource.objects.create(
tenant_id=providers_fixture[1].tenant_id,
provider=providers_fixture[1],
tenant_id=provider2.tenant_id,
provider=provider2,
uid="arn:aws:ec2:us-east-1:123456789012:bucket/i-1234567890abcdef2",
name="My Bucket 3",
region="us-east-1",
@@ -1267,9 +1329,9 @@ def get_api_tokens(
@pytest.fixture
def scan_summaries_fixture(tenants_fixture, providers_fixture):
def scan_summaries_fixture(tenants_fixture, aws_provider):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
scan = Scan.objects.create(
name="overview scan",
provider=provider,
@@ -1346,8 +1408,8 @@ def scan_summaries_fixture(tenants_fixture, providers_fixture):
@pytest.fixture
def integrations_fixture(providers_fixture):
provider1, provider2, *_ = providers_fixture
def integrations_fixture(aws_provider_pair):
provider1, provider2 = aws_provider_pair
tenant_id = provider1.tenant_id
integration1 = Integration.objects.create(
tenant_id=tenant_id,
@@ -1408,9 +1470,9 @@ def lighthouse_config_fixture(authenticated_client, tenants_fixture):
@pytest.fixture(scope="function")
def latest_scan_finding(authenticated_client, providers_fixture, resources_fixture):
provider = providers_fixture[0]
tenant_id = str(providers_fixture[0].tenant_id)
def latest_scan_finding(authenticated_client, aws_provider, resources_fixture):
provider = aws_provider
tenant_id = str(aws_provider.tenant_id)
resource = resources_fixture[0]
scan = Scan.objects.create(
name="latest completed scan",
@@ -1521,10 +1583,10 @@ def findings_with_multiple_categories(scans_fixture, resources_fixture):
@pytest.fixture(scope="function")
def latest_scan_finding_with_categories(
authenticated_client, providers_fixture, resources_fixture
authenticated_client, aws_provider, resources_fixture
):
provider = providers_fixture[0]
tenant_id = str(providers_fixture[0].tenant_id)
provider = aws_provider
tenant_id = str(aws_provider.tenant_id)
resource = resources_fixture[0]
scan = Scan.objects.create(
name="latest completed scan with categories",
@@ -1558,9 +1620,9 @@ def latest_scan_finding_with_categories(
@pytest.fixture(scope="function")
def latest_scan_resource(authenticated_client, providers_fixture):
provider = providers_fixture[0]
tenant_id = str(providers_fixture[0].tenant_id)
def latest_scan_resource(authenticated_client, aws_provider):
provider = aws_provider
tenant_id = str(aws_provider.tenant_id)
scan = Scan.objects.create(
name="latest completed scan for resource",
provider=provider,
@@ -2025,11 +2087,11 @@ def get_authorization_header(access_token: str) -> dict:
@pytest.fixture
def provider_compliance_scores_fixture(
tenants_fixture, providers_fixture, scans_fixture
tenants_fixture, aws_provider_pair, scans_fixture
):
"""Create ProviderComplianceScore entries for compliance watchlist tests."""
tenant = tenants_fixture[0]
provider1, provider2, *_ = providers_fixture
provider1, provider2 = aws_provider_pair
scan1, _, scan3 = scans_fixture
scan1.completed_at = datetime.now(UTC) - timedelta(hours=1)
@@ -2126,9 +2188,7 @@ def tenant_compliance_summary_fixture(tenants_fixture):
@pytest.fixture
def finding_groups_fixture(
tenants_fixture, providers_fixture, scans_fixture, resources_fixture
):
def finding_groups_fixture(tenants_fixture, scans_fixture, resources_fixture):
"""
Create a comprehensive set of findings for testing Finding Groups aggregation.
@@ -2147,7 +2207,6 @@ def finding_groups_fixture(
- Finding counts (pass, fail, muted, new, changed)
"""
tenant = tenants_fixture[0]
provider1, provider2, *_ = providers_fixture
scan1, scan2, *_ = scans_fixture
resource1, resource2, *_ = resources_fixture
@@ -2398,7 +2457,7 @@ def finding_groups_fixture(
@pytest.fixture
def finding_groups_title_variants_fixture(
tenants_fixture, providers_fixture, scans_fixture, resources_fixture
tenants_fixture, scans_fixture, resources_fixture
):
"""
Two providers report the same check_id with different checktitle values.
@@ -2409,7 +2468,6 @@ def finding_groups_title_variants_fixture(
of which title variant matches the search term.
"""
tenant = tenants_fixture[0]
provider1, provider2, *_ = providers_fixture
scan1, scan2, *_ = scans_fixture
resource1, resource2, *_ = resources_fixture
@@ -0,0 +1,62 @@
import io
import struct
import zlib
from types import ModuleType, SimpleNamespace
from typing import Any
PNG_SIGNATURE = b"\x89PNG\r\n\x1a\n"
def _png_chunk(chunk_type: bytes, data: bytes) -> bytes:
checksum = zlib.crc32(chunk_type + data) & 0xFFFFFFFF
return (
struct.pack(">I", len(data)) + chunk_type + data + struct.pack(">I", checksum)
)
def _build_tiny_png() -> bytes:
ihdr = struct.pack(">IIBBBBB", 1, 1, 8, 2, 0, 0, 0)
# Filter byte 0 plus one white RGB pixel.
idat = zlib.compress(b"\x00\xff\xff\xff")
return (
PNG_SIGNATURE
+ _png_chunk(b"IHDR", ihdr)
+ _png_chunk(b"IDAT", idat)
+ _png_chunk(b"IEND", b"")
)
_TINY_PNG_BYTES = _build_tiny_png()
def fake_png_buffer() -> io.BytesIO:
return io.BytesIO(_TINY_PNG_BYTES)
def patch_chart_helpers(
monkeypatch: Any, module: ModuleType, names: tuple[str, ...]
) -> dict[str, list[dict[str, Any]]]:
calls: dict[str, list[dict[str, Any]]] = {name: [] for name in names}
def _build_fake_chart(name: str):
def _fake_chart(*args: Any, **kwargs: Any) -> io.BytesIO:
calls[name].append({"args": args, "kwargs": kwargs})
return fake_png_buffer()
return _fake_chart
for name in names:
monkeypatch.setattr(module, name, _build_fake_chart(name))
return calls
def patch_report_gc(monkeypatch: Any) -> None:
from tasks.jobs import report as report_module
from tasks.jobs.reports import base as base_report_module
from tasks.jobs.reports import threatscore as threatscore_report_module
gc_stub = SimpleNamespace(collect=lambda: 0)
monkeypatch.setattr(report_module, "gc", gc_stub)
monkeypatch.setattr(base_report_module, "gc", gc_stub)
monkeypatch.setattr(threatscore_report_module, "gc", gc_stub)
@@ -106,16 +106,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_drop_db,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -204,16 +200,12 @@ class TestAttackPathsRun:
self,
mock_graph_database_preflight,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -311,16 +303,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -410,18 +398,14 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
"""Failure during ingestion (before set_provider_graph_data_ready(False))
must NOT flip graph_data_ready to True for providers that never had data."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -513,16 +497,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -626,16 +606,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -739,16 +715,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -857,16 +829,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -981,16 +949,12 @@ class TestAttackPathsRun:
mock_event_loop,
mock_stringify,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1078,17 +1042,13 @@ class TestAttackPathsRun:
@pytest.mark.django_db
class TestFailAttackPathsScan:
def test_marks_executing_scan_as_failed(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import fail_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1120,17 +1080,13 @@ class TestFailAttackPathsScan:
}
def test_drops_temp_database_even_when_drop_fails(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import fail_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1156,17 +1112,13 @@ class TestFailAttackPathsScan:
assert attack_paths_scan.state == StateChoices.FAILED
def test_skips_already_failed_scan(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import fail_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1203,17 +1155,13 @@ class TestFailAttackPathsScan:
fail_attack_paths_scan(str(tenant.id), "nonexistent", "setup exploded")
def test_fail_recovers_graph_data_ready_when_data_exists(
self, tenants_fixture, providers_fixture, scans_fixture, sink_backend_stub
self, tenants_fixture, aws_provider, scans_fixture, sink_backend_stub
):
from tasks.jobs.attack_paths.db_utils import fail_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1243,17 +1191,13 @@ class TestFailAttackPathsScan:
mock_set_ready.assert_called_once_with(attack_paths_scan, True)
def test_fail_leaves_graph_data_ready_false_when_no_data(
self, tenants_fixture, providers_fixture, scans_fixture, sink_backend_stub
self, tenants_fixture, aws_provider, scans_fixture, sink_backend_stub
):
from tasks.jobs.attack_paths.db_utils import fail_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1279,17 +1223,13 @@ class TestFailAttackPathsScan:
mock_set_ready.assert_not_called()
def test_recover_graph_data_ready_never_raises(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import recover_graph_data_ready
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -1372,10 +1312,8 @@ class TestAttackPathsFindingsHelpers:
assert mock_run_write.call_count == len(FINDINGS_INDEX_STATEMENTS)
def test_load_findings_batches_requests(self, providers_fixture):
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
def test_load_findings_batches_requests(self, aws_provider):
provider = aws_provider
# Create a generator that yields two batches of dicts (pre-converted)
def findings_generator():
@@ -1423,12 +1361,10 @@ class TestAttackPathsFindingsHelpers:
def test_stream_findings_with_resources_returns_latest_scan_data(
self,
tenants_fixture,
providers_fixture,
aws_provider,
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
resource = Resource.objects.create(
tenant_id=tenant.id,
@@ -1527,13 +1463,11 @@ class TestAttackPathsFindingsHelpers:
def test_enrich_batch_with_resources_single_resource(
self,
tenants_fixture,
providers_fixture,
aws_provider,
):
"""One finding + one resource = one output dict"""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
resource = Resource.objects.create(
tenant_id=tenant.id,
@@ -1611,13 +1545,11 @@ class TestAttackPathsFindingsHelpers:
def test_enrich_batch_with_resources_multiple_resources(
self,
tenants_fixture,
providers_fixture,
aws_provider,
):
"""One finding + three resources = three output dicts"""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
resources = []
for i in range(3):
@@ -1703,13 +1635,11 @@ class TestAttackPathsFindingsHelpers:
def test_enrich_batch_with_resources_no_resources_skips(
self,
tenants_fixture,
providers_fixture,
aws_provider,
):
"""Finding without resources should be skipped"""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = Scan.objects.create(
name="Test Scan",
@@ -1768,11 +1698,9 @@ class TestAttackPathsFindingsHelpers:
assert len(result) == 0
mock_logger.warning.assert_not_called()
def test_generator_is_lazy(self, providers_fixture):
def test_generator_is_lazy(self, aws_provider):
"""Generator should not execute queries until iterated"""
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan_id = "some-scan-id"
with patch("tasks.jobs.attack_paths.findings.rls_transaction") as mock_rls:
@@ -1782,11 +1710,9 @@ class TestAttackPathsFindingsHelpers:
# Nothing should be called yet
mock_rls.assert_not_called()
def test_load_findings_empty_generator(self, providers_fixture):
def test_load_findings_empty_generator(self, aws_provider):
"""Empty generator should not call neo4j"""
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
mock_session = MagicMock()
config = SimpleNamespace(update_tag=12345)
@@ -2248,15 +2174,11 @@ class TestAttackPathsDbUtilsGraphDataReady:
"""Tests for db_utils functions related to graph_data_ready lifecycle."""
def test_database_defaults_allow_legacy_insert_without_cutover_columns(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan_id = uuid4()
now = datetime.now(tz=UTC)
@@ -2300,17 +2222,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
)
def test_create_attack_paths_scan_first_scan_defaults_to_false(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import create_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
with patch(
"tasks.jobs.attack_paths.db_utils.rls_transaction",
@@ -2326,17 +2244,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.sink_backend == "neo4j"
def test_create_attack_paths_scan_inherits_true_from_previous(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import create_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2371,18 +2285,14 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.sink_backend == "neptune"
def test_create_attack_paths_scan_prefers_active_sink_ready_scan(
self, tenants_fixture, providers_fixture, scans_fixture, settings
self, tenants_fixture, aws_provider, scans_fixture, settings
):
from tasks.jobs.attack_paths.db_utils import create_attack_paths_scan
settings.ATTACK_PATHS_SINK_DATABASE = "neo4j"
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2425,17 +2335,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.sink_backend == "neo4j"
def test_create_attack_paths_scan_inherits_is_migrated_false_from_legacy_ready(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import create_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
# Previous scan is ready but pre-cutover (legacy Neo4j graph shape)
AttackPathsScan.objects.create(
@@ -2471,17 +2377,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.sink_backend == "neo4j"
def test_create_attack_paths_scan_inherits_false_when_no_previous_ready(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import create_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2514,17 +2416,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.sink_backend == "neo4j"
def test_set_graph_data_ready_updates_field(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import set_graph_data_ready
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2553,17 +2451,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.graph_data_ready is True
def test_finish_attack_paths_scan_does_not_modify_graph_data_ready(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import finish_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2584,17 +2478,13 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.graph_data_ready is True
def test_finish_attack_paths_scan_preserves_graph_data_ready_on_failure(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import finish_attack_paths_scan
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
attack_paths_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2619,18 +2509,14 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert attack_paths_scan.graph_data_ready is True
def test_set_provider_graph_data_ready_updates_all_scans_for_provider_sink(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import set_provider_graph_data_ready
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan_a = scans_fixture[0]
scan_a.provider = provider
scan_a.save()
scan_b = Scan.objects.create(
name="Second Scan",
@@ -2669,18 +2555,14 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert new_ap_scan.graph_data_ready is False
def test_set_provider_graph_data_ready_preserves_other_sink_scans(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import set_provider_graph_data_ready
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
scan = scans_fixture[0]
scan.provider = provider
scan.save()
legacy_scan = AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2711,22 +2593,14 @@ class TestAttackPathsDbUtilsGraphDataReady:
assert neptune_scan.graph_data_ready is False
def test_set_provider_graph_data_ready_does_not_affect_other_providers(
self, tenants_fixture, providers_fixture, scans_fixture
self, tenants_fixture, aws_provider_pair, scans_fixture
):
from tasks.jobs.attack_paths.db_utils import set_provider_graph_data_ready
tenant = tenants_fixture[0]
provider_a = providers_fixture[0]
provider_a.provider = Provider.ProviderChoices.AWS
provider_a.save()
provider_b = providers_fixture[1]
provider_b.provider = Provider.ProviderChoices.AWS
provider_b.save()
provider_a, provider_b = aws_provider_pair
scan_a = scans_fixture[0]
scan_a.provider = provider_a
scan_a.save()
scan_b = Scan.objects.create(
name="Scan for provider B",
@@ -2808,15 +2682,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
# Recent scan — should still be cleaned up because worker is dead
ap_scan, task_result = self._create_executing_scan(
@@ -2857,15 +2729,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
old_start = datetime.now(tz=UTC) - timedelta(hours=49)
ap_scan, task_result = self._create_executing_scan(
@@ -2894,15 +2764,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
# Recent scan on live worker — should be skipped
self._create_executing_scan(tenant, provider, worker="live-worker@host")
@@ -2924,15 +2792,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
AttackPathsScan.objects.create(
tenant_id=tenant.id,
@@ -2966,15 +2832,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
self._create_executing_scan(tenant, provider, worker="dead-worker@host")
@@ -2996,15 +2860,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant1 = tenants_fixture[0]
tenant2 = tenants_fixture[1]
provider1 = providers_fixture[0]
provider1.provider = Provider.ProviderChoices.AWS
provider1.save()
provider1 = aws_provider
provider2 = Provider.objects.create(
provider="aws",
@@ -3043,15 +2905,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
ap_scan, _ = self._create_executing_scan(
tenant, provider, worker="dead-worker@host"
@@ -3074,15 +2934,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
# Old scan with no Task/TaskResult
old_start = datetime.now(tz=UTC) - timedelta(hours=49)
@@ -3113,15 +2971,13 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
# Two scans on the same dead worker
self._create_executing_scan(tenant, provider, worker="shared-worker@host")
@@ -3194,14 +3050,12 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
ap_scan, task_result = self._create_scheduled_scan(
tenant,
@@ -3249,14 +3103,12 @@ class TestCleanupStaleAttackPathsScans:
mock_drop_db,
mock_recover,
tenants_fixture,
providers_fixture,
aws_provider,
):
from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider.provider = Provider.ProviderChoices.AWS
provider.save()
provider = aws_provider
ap_scan, _ = self._create_scheduled_scan(
tenant,
+3 -3
View File
@@ -39,9 +39,9 @@ def resource_scan_summary_data(scans_fixture):
@pytest.fixture(scope="function")
def get_not_completed_scans(providers_fixture):
provider_id = providers_fixture[0].id
tenant_id = providers_fixture[0].tenant_id
def get_not_completed_scans(aws_provider):
provider_id = aws_provider.id
tenant_id = aws_provider.tenant_id
scan_1 = Scan.objects.create(
tenant_id=tenant_id,
trigger=Scan.TriggerChoices.MANUAL,
+6 -6
View File
@@ -10,8 +10,8 @@ from tasks.beat import schedule_provider_scan
@pytest.mark.django_db
class TestScheduleProviderScan:
def test_schedule_provider_scan_success(self, providers_fixture):
provider_instance, *_ = providers_fixture
def test_schedule_provider_scan_success(self, aws_provider):
provider_instance = aws_provider
with patch(
"tasks.tasks.perform_scheduled_scan_task.apply_async"
@@ -41,8 +41,8 @@ class TestScheduleProviderScan:
"provider_id": str(provider_instance.id),
}
def test_schedule_provider_scan_already_exists(self, providers_fixture):
provider_instance, *_ = providers_fixture
def test_schedule_provider_scan_already_exists(self, aws_provider):
provider_instance = aws_provider
# First, schedule the scan
with patch("tasks.tasks.perform_scheduled_scan_task.apply_async"):
@@ -56,8 +56,8 @@ class TestScheduleProviderScan:
exc_info.value
)
def test_remove_periodic_task(self, providers_fixture):
provider_instance = providers_fixture[0]
def test_remove_periodic_task(self, aws_provider):
provider_instance = aws_provider
assert Scan.objects.count() == 0
with patch("tasks.tasks.perform_scheduled_scan_task.apply_async"):
+11 -11
View File
@@ -9,7 +9,7 @@ from tasks.jobs.deletion import delete_provider, delete_tenant
@pytest.mark.django_db
class TestDeleteProvider:
def test_delete_provider_success(self, providers_fixture):
def test_delete_provider_success(self, aws_provider):
with (
patch(
"tasks.jobs.deletion.graph_database.get_database_name",
@@ -19,7 +19,7 @@ class TestDeleteProvider:
"tasks.jobs.deletion.graph_database.drop_subgraph"
) as mock_drop_subgraph,
):
instance = providers_fixture[0]
instance = aws_provider
tenant_id = str(instance.tenant_id)
result = delete_provider(tenant_id, instance.id)
@@ -53,9 +53,9 @@ class TestDeleteProvider:
mock_drop_subgraph.assert_not_called()
def test_delete_provider_drops_temp_attack_paths_databases(
self, providers_fixture, create_attack_paths_scan
self, aws_provider, create_attack_paths_scan
):
instance = providers_fixture[0]
instance = aws_provider
tenant_id = str(instance.tenant_id)
aps1 = create_attack_paths_scan(instance)
@@ -84,9 +84,9 @@ class TestDeleteProvider:
mock_drop_database.assert_has_calls(expected_tmp_calls, any_order=True)
def test_delete_provider_drops_graph_data_from_all_recorded_sinks(
self, providers_fixture, create_attack_paths_scan
self, aws_provider, create_attack_paths_scan
):
instance = providers_fixture[0]
instance = aws_provider
tenant_id = str(instance.tenant_id)
create_attack_paths_scan(instance, sink_backend="neo4j")
create_attack_paths_scan(instance, sink_backend="neptune")
@@ -124,9 +124,9 @@ class TestDeleteProvider:
)
def test_delete_provider_continues_when_temp_db_drop_fails(
self, providers_fixture, create_attack_paths_scan
self, aws_provider, create_attack_paths_scan
):
instance = providers_fixture[0]
instance = aws_provider
tenant_id = str(instance.tenant_id)
create_attack_paths_scan(instance)
@@ -151,10 +151,10 @@ class TestDeleteProvider:
def test_delete_provider_recalculates_tenant_compliance_summary(
self,
providers_fixture,
aws_provider_pair,
provider_compliance_scores_fixture,
):
instance = providers_fixture[0]
instance = aws_provider_pair[0]
tenant_id = instance.tenant_id
TenantComplianceSummary.objects.create(
@@ -199,7 +199,7 @@ class TestDeleteProvider:
@pytest.mark.django_db
class TestDeleteTenant:
def test_delete_tenant_success(self, tenants_fixture, providers_fixture):
def test_delete_tenant_success(self, tenants_fixture, aws_provider):
"""
Test successful deletion of a tenant and its related data.
"""
+52 -15
View File
@@ -45,14 +45,48 @@ from tasks.jobs.reports import (
get_color_for_risk_level,
get_color_for_weight,
)
from tasks.jobs.reports import cis as cis_report_module
from tasks.jobs.reports import csa as csa_report_module
from tasks.jobs.reports import ens as ens_report_module
from tasks.jobs.reports import nis2 as nis2_report_module
from tasks.jobs.reports import threatscore as threatscore_report_module
from tasks.jobs.threatscore_utils import (
_aggregate_requirement_statistics_from_database,
_load_findings_for_requirement_checks,
)
from tasks.tests.report_test_helpers import patch_chart_helpers, patch_report_gc
matplotlib.use("Agg") # Use non-interactive backend for tests
@pytest.fixture
def patch_report_rendering(monkeypatch):
patch_report_gc(monkeypatch)
patch_chart_helpers(
monkeypatch,
cis_report_module,
(
"create_pie_chart",
"create_horizontal_bar_chart",
"create_stacked_bar_chart",
),
)
patch_chart_helpers(
monkeypatch, csa_report_module, ("create_horizontal_bar_chart",)
)
patch_chart_helpers(
monkeypatch,
ens_report_module,
("create_horizontal_bar_chart", "create_radar_chart"),
)
patch_chart_helpers(
monkeypatch, nis2_report_module, ("create_horizontal_bar_chart",)
)
patch_chart_helpers(
monkeypatch, threatscore_report_module, ("create_vertical_bar_chart",)
)
@pytest.mark.django_db
class TestAggregateRequirementStatistics:
"""Test suite for _aggregate_requirement_statistics_from_database function."""
@@ -355,7 +389,7 @@ class TestPDFStylesCreation:
class TestLoadFindingsForChecks:
"""Test suite for _load_findings_for_requirement_checks function."""
def test_empty_check_ids_returns_empty(self, tenants_fixture, providers_fixture):
def test_empty_check_ids_returns_empty(self, tenants_fixture):
"""Test that empty check_ids list returns empty dict."""
tenant = tenants_fixture[0]
@@ -1041,6 +1075,7 @@ class TestStaleCleanupProtectionHelpers:
@pytest.mark.django_db
@pytest.mark.usefixtures("patch_report_rendering")
class TestGenerateThreatscoreReportFunction:
"""Test suite for generate_threatscore_report function."""
@@ -1050,12 +1085,12 @@ class TestGenerateThreatscoreReportFunction:
mock_build_provider_metadata,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test that exceptions during report generation are properly handled."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
mock_build_provider_metadata.side_effect = Exception("Test exception")
@@ -1072,6 +1107,7 @@ class TestGenerateThreatscoreReportFunction:
@pytest.mark.django_db
@pytest.mark.usefixtures("patch_report_rendering")
class TestGenerateComplianceReportsOptimized:
"""Test suite for generate_compliance_reports function."""
@@ -1087,12 +1123,12 @@ class TestGenerateComplianceReportsOptimized:
mock_upload,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test that function returns early when scan has no findings."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
result = generate_compliance_reports(
tenant_id=str(tenant.id),
@@ -1144,14 +1180,14 @@ class TestGenerateComplianceReportsOptimized:
mock_upload,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Scan with no findings and ``generate_cis=True`` must yield a flat
``{"upload": False, "path": ""}`` entry, consistent with the other
frameworks (no nested dict, no sentinel keys)."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
result = generate_compliance_reports(
tenant_id=str(tenant.id),
@@ -1439,6 +1475,7 @@ class TestGenerateComplianceReportsOptimized:
@pytest.mark.django_db
@pytest.mark.usefixtures("patch_report_rendering")
class TestGenerateComplianceReportsCIS:
"""Test suite covering the CIS branch of generate_compliance_reports."""
@@ -1468,7 +1505,7 @@ class TestGenerateComplianceReportsCIS:
monkeypatch,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""CIS branch should generate a single PDF for the highest version.
@@ -1478,7 +1515,7 @@ class TestGenerateComplianceReportsCIS:
"""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._force_scan_has_findings(monkeypatch)
@@ -1527,12 +1564,12 @@ class TestGenerateComplianceReportsCIS:
monkeypatch,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""A failure in the latest CIS variant must be surfaced in the flat results entry."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._force_scan_has_findings(monkeypatch)
@@ -1574,14 +1611,14 @@ class TestGenerateComplianceReportsCIS:
monkeypatch,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""When ``Compliance.get_bulk`` returns no CIS entry the CIS branch
must skip cleanly and record a flat ``{"upload": False, "path": ""}``
entry — no hard-coded provider whitelist is consulted."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._force_scan_has_findings(monkeypatch)
mock_stats.return_value = {}
@@ -1613,12 +1650,12 @@ class TestGenerateComplianceReportsCIS:
monkeypatch,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""CIS output dir errors must be captured in results (not raised)."""
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._force_scan_has_findings(monkeypatch)
mock_stats.return_value = {}
+39 -225
View File
@@ -43,6 +43,7 @@ from tasks.jobs.reports import ( # Configuration; Colors; Components; Charts; B
get_framework_config,
get_status_color,
)
from tasks.tests.report_test_helpers import PNG_SIGNATURE, fake_png_buffer
# =============================================================================
# Configuration Tests
@@ -452,174 +453,47 @@ class TestSectionHeader:
# =============================================================================
class TestChartCreation:
"""Tests for chart creation functions."""
class TestChartRenderingSmoke:
"""Small real-render coverage for the chart helpers."""
def test_create_vertical_bar_chart(self):
"""Test vertical bar chart creation."""
buffer = create_vertical_bar_chart(
labels=["A", "B", "C"],
values=[80, 60, 40],
)
assert isinstance(buffer, io.BytesIO)
assert buffer.getvalue() # Not empty
@pytest.mark.parametrize(
("chart_helper", "kwargs"),
[
(
create_vertical_bar_chart,
{"labels": ["Section 1", "Section 2"], "values": [90, 70]},
),
(
create_horizontal_bar_chart,
{"labels": ["Category 1", "Category 2"], "values": [85, 65]},
),
(
create_radar_chart,
{"labels": ["A", "B", "C"], "values": [50, 60, 70]},
),
(
create_pie_chart,
{"labels": ["Pass", "Fail"], "values": [80, 20]},
),
(
create_stacked_bar_chart,
{
"labels": ["Section 1", "Section 2"],
"data_series": {"Pass": [8, 6], "Fail": [2, 4]},
},
),
],
)
def test_chart_helper_renders_valid_png(self, chart_helper, kwargs):
buffer = chart_helper(**kwargs)
image_bytes = buffer.getvalue()
def test_create_vertical_bar_chart_with_options(self):
"""Test vertical bar chart with custom options."""
buffer = create_vertical_bar_chart(
labels=["Section 1", "Section 2"],
values=[90, 70],
ylabel="Compliance",
title="Test Chart",
figsize=(8, 6),
)
assert isinstance(buffer, io.BytesIO)
assert image_bytes
assert image_bytes.startswith(PNG_SIGNATURE)
def test_create_horizontal_bar_chart(self):
"""Test horizontal bar chart creation."""
buffer = create_horizontal_bar_chart(
labels=["Category 1", "Category 2", "Category 3"],
values=[85, 65, 45],
)
assert isinstance(buffer, io.BytesIO)
assert buffer.getvalue()
def test_create_horizontal_bar_chart_with_options(self):
"""Test horizontal bar chart with custom options."""
buffer = create_horizontal_bar_chart(
labels=["A", "B"],
values=[100, 50],
xlabel="Percentage",
title="Custom Chart",
)
assert isinstance(buffer, io.BytesIO)
def test_create_radar_chart(self):
"""Test radar chart creation."""
buffer = create_radar_chart(
labels=["Dim 1", "Dim 2", "Dim 3", "Dim 4", "Dim 5"],
values=[80, 70, 60, 90, 75],
)
assert isinstance(buffer, io.BytesIO)
assert buffer.getvalue()
def test_create_radar_chart_with_options(self):
"""Test radar chart with custom options."""
buffer = create_radar_chart(
labels=["A", "B", "C"],
values=[50, 60, 70],
color="#FF0000",
fill_alpha=0.5,
title="Custom Radar",
)
assert isinstance(buffer, io.BytesIO)
def test_create_pie_chart(self):
"""Test pie chart creation."""
buffer = create_pie_chart(
labels=["Pass", "Fail"],
values=[80, 20],
)
assert isinstance(buffer, io.BytesIO)
assert buffer.getvalue()
def test_create_pie_chart_with_options(self):
"""Test pie chart with custom options."""
buffer = create_pie_chart(
labels=["Pass", "Fail", "Manual"],
values=[60, 30, 10],
colors=["#4CAF50", "#F44336", "#9E9E9E"],
title="Status Distribution",
autopct="%1.0f%%",
)
assert isinstance(buffer, io.BytesIO)
def test_create_stacked_bar_chart(self):
"""Test stacked bar chart creation."""
buffer = create_stacked_bar_chart(
labels=["Section 1", "Section 2", "Section 3"],
data_series={
"Pass": [8, 6, 4],
"Fail": [2, 4, 6],
},
)
assert isinstance(buffer, io.BytesIO)
assert buffer.getvalue()
def test_create_stacked_bar_chart_with_options(self):
"""Test stacked bar chart with custom options."""
buffer = create_stacked_bar_chart(
labels=["A", "B"],
data_series={
"Pass": [10, 5],
"Fail": [2, 3],
"Manual": [1, 2],
},
colors={
"Pass": "#4CAF50",
"Fail": "#F44336",
"Manual": "#9E9E9E",
},
xlabel="Categories",
ylabel="Requirements",
title="Requirements by Status",
)
assert isinstance(buffer, io.BytesIO)
def test_create_stacked_bar_chart_without_legend(self):
"""Test stacked bar chart without legend."""
buffer = create_stacked_bar_chart(
labels=["X", "Y"],
data_series={"A": [1, 2]},
show_legend=False,
)
assert isinstance(buffer, io.BytesIO)
def test_create_vertical_bar_chart_without_labels(self):
"""Test vertical bar chart without value labels."""
buffer = create_vertical_bar_chart(
labels=["A", "B"],
values=[50, 75],
show_labels=False,
)
assert isinstance(buffer, io.BytesIO)
def test_create_vertical_bar_chart_with_explicit_colors(self):
"""Test vertical bar chart with explicit color list."""
buffer = create_vertical_bar_chart(
labels=["Pass", "Fail"],
values=[80, 20],
colors=["#4CAF50", "#F44336"],
)
assert isinstance(buffer, io.BytesIO)
def test_create_horizontal_bar_chart_auto_figsize(self):
"""Test horizontal bar chart auto-calculates figure size for many items."""
labels = [f"Item {i}" for i in range(20)]
values = [50 + i * 2 for i in range(20)]
buffer = create_horizontal_bar_chart(
labels=labels,
values=values,
)
assert isinstance(buffer, io.BytesIO)
def test_create_horizontal_bar_chart_with_explicit_colors(self):
"""Test horizontal bar chart with explicit colors."""
buffer = create_horizontal_bar_chart(
labels=["A", "B", "C"],
values=[80, 60, 40],
colors=["#4CAF50", "#FFEB3B", "#F44336"],
)
assert isinstance(buffer, io.BytesIO)
def test_create_radar_chart_with_custom_ticks(self):
"""Test radar chart with custom y-axis ticks."""
buffer = create_radar_chart(
labels=["A", "B", "C", "D"],
values=[25, 50, 75, 100],
y_ticks=[0, 25, 50, 75, 100],
)
assert isinstance(buffer, io.BytesIO)
buffer.seek(0)
assert Image(buffer, width=1 * inch, height=1 * inch)
# =============================================================================
@@ -1056,10 +930,7 @@ class TestExampleReportGenerator:
]
def create_charts_section(self, data):
chart_buffer = create_vertical_bar_chart(
labels=["Pass", "Fail"],
values=[80, 20],
)
chart_buffer = fake_png_buffer()
return [Image(chart_buffer, width=6 * inch, height=4 * inch)]
def create_requirements_index(self, data):
@@ -1150,63 +1021,6 @@ class TestExampleReportGenerator:
# =============================================================================
class TestChartEdgeCases:
"""Tests for chart edge cases."""
def test_vertical_bar_chart_empty_data(self):
"""Test vertical bar chart with empty data."""
buffer = create_vertical_bar_chart(labels=[], values=[])
assert isinstance(buffer, io.BytesIO)
def test_vertical_bar_chart_single_item(self):
"""Test vertical bar chart with single item."""
buffer = create_vertical_bar_chart(labels=["Single"], values=[75.0])
assert isinstance(buffer, io.BytesIO)
def test_horizontal_bar_chart_empty_data(self):
"""Test horizontal bar chart with empty data."""
buffer = create_horizontal_bar_chart(labels=[], values=[])
assert isinstance(buffer, io.BytesIO)
def test_horizontal_bar_chart_single_item(self):
"""Test horizontal bar chart with single item."""
buffer = create_horizontal_bar_chart(labels=["Single"], values=[50.0])
assert isinstance(buffer, io.BytesIO)
def test_radar_chart_minimum_points(self):
"""Test radar chart with minimum number of points (3)."""
buffer = create_radar_chart(
labels=["A", "B", "C"],
values=[30.0, 60.0, 90.0],
)
assert isinstance(buffer, io.BytesIO)
def test_pie_chart_single_slice(self):
"""Test pie chart with single slice."""
buffer = create_pie_chart(labels=["Only"], values=[100.0])
assert isinstance(buffer, io.BytesIO)
def test_pie_chart_many_slices(self):
"""Test pie chart with many slices."""
labels = [f"Item {i}" for i in range(10)]
values = [10.0] * 10
buffer = create_pie_chart(labels=labels, values=values)
assert isinstance(buffer, io.BytesIO)
def test_stacked_bar_chart_single_series(self):
"""Test stacked bar chart with single series."""
buffer = create_stacked_bar_chart(
labels=["A", "B"],
data_series={"Only": [10.0, 20.0]},
)
assert isinstance(buffer, io.BytesIO)
def test_stacked_bar_chart_empty_data(self):
"""Test stacked bar chart with empty data."""
buffer = create_stacked_bar_chart(labels=[], data_series={})
assert isinstance(buffer, io.BytesIO)
class TestComponentEdgeCases:
"""Tests for component edge cases."""
@@ -4,11 +4,13 @@ import pytest
from api.models import StatusChoices
from reportlab.platypus import Image, LongTable, Paragraph, Table
from tasks.jobs.reports import FRAMEWORK_REGISTRY, ComplianceData, RequirementData
from tasks.jobs.reports import cis as cis_report_module
from tasks.jobs.reports.cis import (
CISReportGenerator,
_normalize_profile,
_profile_badge_text,
)
from tasks.tests.report_test_helpers import patch_chart_helpers
# =============================================================================
# Fixtures
@@ -399,18 +401,69 @@ class TestCISExecutiveSummary:
class TestCISChartsSection:
def test_charts_rendered(self, cis_generator, populated_cis_compliance_data):
elements = cis_generator.create_charts_section(populated_cis_compliance_data)
# At least 1 image for the pie + 1 for section bar + 1 for stacked
images = [e for e in elements if isinstance(e, Image)]
assert len(images) >= 1
def test_charts_rendered(
self, monkeypatch, cis_generator, populated_cis_compliance_data
):
chart_calls = patch_chart_helpers(
monkeypatch,
cis_report_module,
(
"create_pie_chart",
"create_horizontal_bar_chart",
"create_stacked_bar_chart",
),
)
def test_charts_no_data_no_crash(self, cis_generator, basic_cis_compliance_data):
elements = cis_generator.create_charts_section(populated_cis_compliance_data)
images = [e for e in elements if isinstance(e, Image)]
assert len(images) == 3
pie_kwargs = chart_calls["create_pie_chart"][0]["kwargs"]
assert pie_kwargs["labels"] == ["Pass (2)", "Fail (2)", "Manual (1)"]
assert pie_kwargs["values"] == [2, 2, 1]
assert pie_kwargs["colors"]
bar_kwargs = chart_calls["create_horizontal_bar_chart"][0]["kwargs"]
assert set(bar_kwargs["labels"]) == {
"1 Identity and Access Management",
"2 Storage",
}
assert bar_kwargs["values"] == [50.0, 50.0]
assert bar_kwargs["xlabel"] == "Compliance (%)"
assert bar_kwargs["color_func"]
assert bar_kwargs["label_fontsize"] == 9
stacked_kwargs = chart_calls["create_stacked_bar_chart"][0]["kwargs"]
assert stacked_kwargs["labels"] == ["Level 1", "Level 2"]
assert stacked_kwargs["data_series"] == {
"Pass": [1, 1],
"Fail": [2, 0],
"Manual": [0, 1],
}
assert stacked_kwargs["xlabel"] == "Profile"
assert stacked_kwargs["ylabel"] == "Requirements"
def test_charts_no_data_no_crash(
self, monkeypatch, cis_generator, basic_cis_compliance_data
):
chart_calls = patch_chart_helpers(
monkeypatch,
cis_report_module,
(
"create_pie_chart",
"create_horizontal_bar_chart",
"create_stacked_bar_chart",
),
)
basic_cis_compliance_data.requirements = []
basic_cis_compliance_data.attributes_by_requirement_id = {}
elements = cis_generator.create_charts_section(basic_cis_compliance_data)
# Must not raise; may or may not have any Image
assert isinstance(elements, list)
assert chart_calls["create_pie_chart"] == []
assert chart_calls["create_horizontal_bar_chart"] == []
assert chart_calls["create_stacked_bar_chart"] == []
# =============================================================================
@@ -2,9 +2,11 @@ import io
from unittest.mock import Mock
import pytest
from reportlab.platypus import PageBreak, Paragraph, Table
from reportlab.platypus import Image, PageBreak, Paragraph, Table
from tasks.jobs.reports import FRAMEWORK_REGISTRY, ComplianceData, RequirementData
from tasks.jobs.reports import csa as csa_report_module
from tasks.jobs.reports.csa import CSAReportGenerator
from tasks.tests.report_test_helpers import patch_chart_helpers
# Use string status values directly to avoid Django DB initialization
@@ -29,6 +31,13 @@ def csa_generator():
return CSAReportGenerator(config)
@pytest.fixture
def patched_csa_charts(monkeypatch):
return patch_chart_helpers(
monkeypatch, csa_report_module, ("create_horizontal_bar_chart",)
)
@pytest.fixture
def mock_csa_requirement_attribute_iam():
"""Create a mock CSA CCM requirement attribute for Identity & Access Management."""
@@ -320,7 +329,7 @@ class TestCSAChartsSection:
"""Test suite for CSA charts section generation."""
def test_charts_section_has_section_chart_title(
self, csa_generator, basic_csa_compliance_data
self, csa_generator, basic_csa_compliance_data, patched_csa_charts
):
"""Test that charts section has section compliance title."""
basic_csa_compliance_data.requirements = []
@@ -331,9 +340,14 @@ class TestCSAChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "Section" in content or "Compliance" in content
assert any(isinstance(e, Image) for e in elements)
chart_kwargs = patched_csa_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == []
assert chart_kwargs["values"] == []
assert chart_kwargs["xlabel"] == "Compliance (%)"
def test_charts_section_has_page_break(
self, csa_generator, basic_csa_compliance_data
self, csa_generator, basic_csa_compliance_data, patched_csa_charts
):
"""Test that charts section has page breaks."""
basic_csa_compliance_data.requirements = []
@@ -343,12 +357,14 @@ class TestCSAChartsSection:
page_breaks = [e for e in elements if isinstance(e, PageBreak)]
assert len(page_breaks) >= 1
assert len(patched_csa_charts["create_horizontal_bar_chart"]) == 1
def test_charts_section_has_section_breakdown(
self,
csa_generator,
basic_csa_compliance_data,
mock_csa_requirement_attribute_iam,
patched_csa_charts,
):
"""Test that charts section includes section breakdown table."""
basic_csa_compliance_data.requirements = [
@@ -372,6 +388,11 @@ class TestCSAChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "Section" in content or "Breakdown" in content
assert any(isinstance(e, Image) for e in elements)
chart_kwargs = patched_csa_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["Identity & Access Management"]
assert chart_kwargs["values"] == [100.0]
assert chart_kwargs["color_func"]
# =============================================================================
@@ -387,6 +408,7 @@ class TestCSASectionChart:
csa_generator,
basic_csa_compliance_data,
mock_csa_requirement_attribute_iam,
patched_csa_charts,
):
"""Test that section chart is created successfully."""
basic_csa_compliance_data.requirements = [
@@ -409,12 +431,17 @@ class TestCSASectionChart:
assert isinstance(chart_buffer, io.BytesIO)
assert chart_buffer.getvalue() # Not empty
chart_kwargs = patched_csa_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["Identity & Access Management"]
assert chart_kwargs["values"] == [100.0]
assert chart_kwargs["xlabel"] == "Compliance (%)"
def test_section_chart_excludes_manual(
self,
csa_generator,
basic_csa_compliance_data,
mock_csa_requirement_attribute_iam,
patched_csa_charts,
):
"""Test that manual requirements are excluded from section chart."""
basic_csa_compliance_data.requirements = [
@@ -447,6 +474,9 @@ class TestCSASectionChart:
# Should not raise any errors
chart_buffer = csa_generator._create_section_chart(basic_csa_compliance_data)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_csa_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["Identity & Access Management"]
assert chart_kwargs["values"] == [100.0]
def test_section_chart_multiple_sections(
self,
@@ -455,6 +485,7 @@ class TestCSASectionChart:
mock_csa_requirement_attribute_iam,
mock_csa_requirement_attribute_logging,
mock_csa_requirement_attribute_crypto,
patched_csa_charts,
):
"""Test section chart with multiple sections."""
basic_csa_compliance_data.requirements = [
@@ -501,6 +532,13 @@ class TestCSASectionChart:
chart_buffer = csa_generator._create_section_chart(basic_csa_compliance_data)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_csa_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == [
"Cryptography & Encryption",
"Identity & Access Management",
"Logging and Monitoring",
]
assert chart_kwargs["values"] == [100.0, 100.0, 0.0]
# =============================================================================
+90 -14
View File
@@ -2,9 +2,11 @@ import io
from unittest.mock import Mock, patch
import pytest
from reportlab.platypus import PageBreak, Paragraph, Table
from reportlab.platypus import Image, PageBreak, Paragraph, Table
from tasks.jobs.reports import FRAMEWORK_REGISTRY, ComplianceData, RequirementData
from tasks.jobs.reports import ens as ens_report_module
from tasks.jobs.reports.ens import ENSReportGenerator
from tasks.tests.report_test_helpers import patch_chart_helpers
# Use string status values directly to avoid Django DB initialization
@@ -29,6 +31,15 @@ def ens_generator():
return ENSReportGenerator(config)
@pytest.fixture
def patched_ens_charts(monkeypatch):
return patch_chart_helpers(
monkeypatch,
ens_report_module,
("create_horizontal_bar_chart", "create_radar_chart"),
)
@pytest.fixture
def mock_ens_requirement_attribute():
"""Create a mock ENS requirement attribute with all fields."""
@@ -355,7 +366,7 @@ class TestENSChartsSection:
"""Test suite for ENS charts section generation."""
def test_charts_section_has_page_breaks(
self, ens_generator, basic_ens_compliance_data
self, ens_generator, basic_ens_compliance_data, patched_ens_charts
):
"""Test that charts section has page breaks between charts."""
basic_ens_compliance_data.requirements = []
@@ -365,9 +376,25 @@ class TestENSChartsSection:
page_breaks = [e for e in elements if isinstance(e, PageBreak)]
assert len(page_breaks) >= 2 # At least 2 page breaks for different charts
assert any(isinstance(e, Image) for e in elements)
assert len(patched_ens_charts["create_horizontal_bar_chart"]) == 1
assert len(patched_ens_charts["create_radar_chart"]) == 1
marco_kwargs = patched_ens_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert marco_kwargs["labels"] == []
assert marco_kwargs["values"] == []
radar_kwargs = patched_ens_charts["create_radar_chart"][0]["kwargs"]
assert radar_kwargs["labels"] == ens_report_module.DIMENSION_NAMES
assert radar_kwargs["values"] == [100, 100, 100, 100, 100]
assert radar_kwargs["color"] == "#2196F3"
def test_charts_section_has_marco_category_chart(
self, ens_generator, basic_ens_compliance_data, mock_ens_requirement_attribute
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute,
patched_ens_charts,
):
"""Test that charts section contains Marco/Categoría chart."""
basic_ens_compliance_data.requirements = [
@@ -391,9 +418,18 @@ class TestENSChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "Marco" in content or "Categoría" in content
assert any(isinstance(e, Image) for e in elements)
chart_kwargs = patched_ens_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["Operacional - Gestión de incidentes"]
assert chart_kwargs["values"] == [100.0]
assert chart_kwargs["xlabel"] == "Porcentaje de Cumplimiento (%)"
def test_charts_section_has_dimensions_radar(
self, ens_generator, basic_ens_compliance_data, mock_ens_requirement_attribute
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute,
patched_ens_charts,
):
"""Test that charts section contains dimensions radar chart."""
basic_ens_compliance_data.requirements = [
@@ -417,9 +453,17 @@ class TestENSChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "Dimensiones" in content or "dimensiones" in content.lower()
radar_kwargs = patched_ens_charts["create_radar_chart"][0]["kwargs"]
assert radar_kwargs["labels"] == ens_report_module.DIMENSION_NAMES
assert radar_kwargs["values"] == [100, 100, 100.0, 100.0, 100]
assert radar_kwargs["color"] == "#2196F3"
def test_charts_section_has_tipo_distribution(
self, ens_generator, basic_ens_compliance_data, mock_ens_requirement_attribute
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute,
patched_ens_charts,
):
"""Test that charts section contains tipo distribution."""
basic_ens_compliance_data.requirements = [
@@ -443,6 +487,8 @@ class TestENSChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "Tipo" in content or "tipo" in content.lower()
assert len(patched_ens_charts["create_horizontal_bar_chart"]) == 1
assert len(patched_ens_charts["create_radar_chart"]) == 1
# =============================================================================
@@ -829,7 +875,11 @@ class TestENSDimensionHandling:
"""Test suite for ENS security dimension handling."""
def test_dimensions_as_list(
self, ens_generator, basic_ens_compliance_data, mock_ens_requirement_attribute
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute,
patched_ens_charts,
):
"""Test handling dimensions as a list."""
# mock_ens_requirement_attribute has Dimensiones as list
@@ -837,9 +887,9 @@ class TestENSDimensionHandling:
RequirementData(
id="REQ-001",
description="Test requirement",
status=StatusChoices.PASS,
passed_findings=10,
failed_findings=0,
status=StatusChoices.FAIL,
passed_findings=0,
failed_findings=10,
total_findings=10,
),
]
@@ -854,12 +904,16 @@ class TestENSDimensionHandling:
basic_ens_compliance_data
)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_ens_charts["create_radar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ens_report_module.DIMENSION_NAMES
assert chart_kwargs["values"] == [100, 100, 0.0, 0.0, 100]
def test_dimensions_as_string(
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute_medio,
patched_ens_charts,
):
"""Test handling dimensions as comma-separated string."""
# mock_ens_requirement_attribute_medio has Dimensiones as string
@@ -867,9 +921,9 @@ class TestENSDimensionHandling:
RequirementData(
id="REQ-001",
description="Test requirement",
status=StatusChoices.PASS,
passed_findings=10,
failed_findings=0,
status=StatusChoices.FAIL,
passed_findings=0,
failed_findings=10,
total_findings=10,
),
]
@@ -884,12 +938,16 @@ class TestENSDimensionHandling:
basic_ens_compliance_data
)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_ens_charts["create_radar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ens_report_module.DIMENSION_NAMES
assert chart_kwargs["values"] == [0.0, 0.0, 100, 100, 100]
def test_dimensions_empty(
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute_opcional,
patched_ens_charts,
):
"""Test handling empty dimensions."""
# mock_ens_requirement_attribute_opcional has empty Dimensiones
@@ -916,6 +974,9 @@ class TestENSDimensionHandling:
basic_ens_compliance_data
)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_ens_charts["create_radar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ens_report_module.DIMENSION_NAMES
assert chart_kwargs["values"] == [100, 100, 100, 100, 100]
# =============================================================================
@@ -1061,7 +1122,11 @@ class TestENSMarcoCategoryChart:
"""Test suite for ENS Marco/Categoría chart."""
def test_marco_category_chart_creation(
self, ens_generator, basic_ens_compliance_data, mock_ens_requirement_attribute
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute,
patched_ens_charts,
):
"""Test that Marco/Categoría chart is created successfully."""
basic_ens_compliance_data.requirements = [
@@ -1086,9 +1151,17 @@ class TestENSMarcoCategoryChart:
assert isinstance(chart_buffer, io.BytesIO)
assert chart_buffer.getvalue() # Not empty
chart_kwargs = patched_ens_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["Operacional - Gestión de incidentes"]
assert chart_kwargs["values"] == [100.0]
assert chart_kwargs["xlabel"] == "Porcentaje de Cumplimiento (%)"
def test_marco_category_chart_excludes_manual(
self, ens_generator, basic_ens_compliance_data, mock_ens_requirement_attribute
self,
ens_generator,
basic_ens_compliance_data,
mock_ens_requirement_attribute,
patched_ens_charts,
):
"""Test that manual requirements are excluded from chart."""
basic_ens_compliance_data.requirements = [
@@ -1123,6 +1196,9 @@ class TestENSMarcoCategoryChart:
basic_ens_compliance_data
)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_ens_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["Operacional - Gestión de incidentes"]
assert chart_kwargs["values"] == [100.0]
# =============================================================================
@@ -2,9 +2,11 @@ import io
from unittest.mock import Mock, patch
import pytest
from reportlab.platypus import PageBreak, Paragraph, Table
from reportlab.platypus import Image, PageBreak, Paragraph, Table
from tasks.jobs.reports import FRAMEWORK_REGISTRY, ComplianceData, RequirementData
from tasks.jobs.reports import nis2 as nis2_report_module
from tasks.jobs.reports.nis2 import NIS2ReportGenerator, _extract_section_number
from tasks.tests.report_test_helpers import patch_chart_helpers
# Use string status values directly to avoid Django DB initialization
@@ -29,6 +31,13 @@ def nis2_generator():
return NIS2ReportGenerator(config)
@pytest.fixture
def patched_nis2_charts(monkeypatch):
return patch_chart_helpers(
monkeypatch, nis2_report_module, ("create_horizontal_bar_chart",)
)
@pytest.fixture
def mock_nis2_requirement_attribute_section1():
"""Create a mock NIS2 requirement attribute for Section 1."""
@@ -380,7 +389,7 @@ class TestNIS2ChartsSection:
"""Test suite for NIS2 charts section generation."""
def test_charts_section_has_section_chart_title(
self, nis2_generator, basic_nis2_compliance_data
self, nis2_generator, basic_nis2_compliance_data, patched_nis2_charts
):
"""Test that charts section has section compliance title."""
basic_nis2_compliance_data.requirements = []
@@ -391,9 +400,14 @@ class TestNIS2ChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "Section" in content or "Compliance" in content
assert any(isinstance(e, Image) for e in elements)
chart_kwargs = patched_nis2_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == []
assert chart_kwargs["values"] == []
assert chart_kwargs["xlabel"] == "Compliance (%)"
def test_charts_section_has_page_break(
self, nis2_generator, basic_nis2_compliance_data
self, nis2_generator, basic_nis2_compliance_data, patched_nis2_charts
):
"""Test that charts section has page breaks."""
basic_nis2_compliance_data.requirements = []
@@ -403,12 +417,14 @@ class TestNIS2ChartsSection:
page_breaks = [e for e in elements if isinstance(e, PageBreak)]
assert len(page_breaks) >= 1
assert len(patched_nis2_charts["create_horizontal_bar_chart"]) == 1
def test_charts_section_has_subsection_breakdown(
self,
nis2_generator,
basic_nis2_compliance_data,
mock_nis2_requirement_attribute_section1,
patched_nis2_charts,
):
"""Test that charts section includes subsection breakdown table."""
basic_nis2_compliance_data.requirements = [
@@ -434,6 +450,11 @@ class TestNIS2ChartsSection:
paragraphs = [e for e in elements if isinstance(e, Paragraph)]
content = " ".join(str(p.text) for p in paragraphs)
assert "SubSection" in content or "Breakdown" in content
assert any(isinstance(e, Image) for e in elements)
chart_kwargs = patched_nis2_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["1. Policy on Security"]
assert chart_kwargs["values"] == [100.0]
assert chart_kwargs["color_func"]
# =============================================================================
@@ -449,6 +470,7 @@ class TestNIS2SectionChart:
nis2_generator,
basic_nis2_compliance_data,
mock_nis2_requirement_attribute_section1,
patched_nis2_charts,
):
"""Test that section chart is created successfully."""
basic_nis2_compliance_data.requirements = [
@@ -473,12 +495,17 @@ class TestNIS2SectionChart:
assert isinstance(chart_buffer, io.BytesIO)
assert chart_buffer.getvalue() # Not empty
chart_kwargs = patched_nis2_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["1. Policy on Security"]
assert chart_kwargs["values"] == [100.0]
assert chart_kwargs["xlabel"] == "Compliance (%)"
def test_section_chart_excludes_manual(
self,
nis2_generator,
basic_nis2_compliance_data,
mock_nis2_requirement_attribute_section1,
patched_nis2_charts,
):
"""Test that manual requirements are excluded from section chart."""
basic_nis2_compliance_data.requirements = [
@@ -515,6 +542,9 @@ class TestNIS2SectionChart:
# Should not raise any errors
chart_buffer = nis2_generator._create_section_chart(basic_nis2_compliance_data)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_nis2_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == ["1. Policy on Security"]
assert chart_kwargs["values"] == [100.0]
def test_section_chart_multiple_sections(
self,
@@ -523,6 +553,7 @@ class TestNIS2SectionChart:
mock_nis2_requirement_attribute_section1,
mock_nis2_requirement_attribute_section2,
mock_nis2_requirement_attribute_section11,
patched_nis2_charts,
):
"""Test section chart with multiple sections."""
basic_nis2_compliance_data.requirements = [
@@ -571,6 +602,13 @@ class TestNIS2SectionChart:
chart_buffer = nis2_generator._create_section_chart(basic_nis2_compliance_data)
assert isinstance(chart_buffer, io.BytesIO)
chart_kwargs = patched_nis2_charts["create_horizontal_bar_chart"][0]["kwargs"]
assert chart_kwargs["labels"] == [
"1. Policy on Security",
"2. Risk Management",
"11. Access Control",
]
assert chart_kwargs["values"] == [100.0, 0.0, 100.0]
# =============================================================================
@@ -10,6 +10,8 @@ from tasks.jobs.reports import (
RequirementData,
ThreatScoreReportGenerator,
)
from tasks.jobs.reports import threatscore as threatscore_report_module
from tasks.tests.report_test_helpers import patch_chart_helpers
# =============================================================================
# Fixtures
@@ -23,6 +25,13 @@ def threatscore_generator():
return ThreatScoreReportGenerator(config)
@pytest.fixture
def patched_threatscore_charts(monkeypatch):
return patch_chart_helpers(
monkeypatch, threatscore_report_module, ("create_vertical_bar_chart",)
)
@pytest.fixture
def mock_requirement_attribute():
"""Create a mock requirement attribute with numeric values."""
@@ -677,7 +686,7 @@ class TestSectionScoreChart:
"""Test suite for section score chart generation."""
def test_create_section_chart_empty_data(
self, threatscore_generator, basic_compliance_data
self, threatscore_generator, basic_compliance_data, patched_threatscore_charts
):
"""Test chart creation with no requirements."""
basic_compliance_data.requirements = []
@@ -689,9 +698,22 @@ class TestSectionScoreChart:
assert isinstance(result, io.BytesIO)
assert result.getvalue() # Should have content
chart_kwargs = patched_threatscore_charts["create_vertical_bar_chart"][0][
"kwargs"
]
assert chart_kwargs["labels"] == []
assert chart_kwargs["values"] == []
assert chart_kwargs["ylabel"] == "Compliance Score (%)"
assert chart_kwargs["xlabel"] == ""
assert chart_kwargs["color_func"]
assert chart_kwargs["rotation"] == 0
def test_create_section_chart_single_section(
self, threatscore_generator, basic_compliance_data, mock_requirement_attribute
self,
threatscore_generator,
basic_compliance_data,
mock_requirement_attribute,
patched_threatscore_charts,
):
"""Test chart creation with a single section."""
basic_compliance_data.requirements = [
@@ -713,9 +735,14 @@ class TestSectionScoreChart:
)
assert isinstance(result, io.BytesIO)
chart_kwargs = patched_threatscore_charts["create_vertical_bar_chart"][0][
"kwargs"
]
assert chart_kwargs["labels"] == ["1. IAM"]
assert chart_kwargs["values"] == [100.0]
def test_create_section_chart_multiple_sections(
self, threatscore_generator, basic_compliance_data
self, threatscore_generator, basic_compliance_data, patched_threatscore_charts
):
"""Test chart creation with multiple sections."""
mock_attr_1 = Mock()
@@ -756,9 +783,14 @@ class TestSectionScoreChart:
)
assert isinstance(result, io.BytesIO)
chart_kwargs = patched_threatscore_charts["create_vertical_bar_chart"][0][
"kwargs"
]
assert chart_kwargs["labels"] == ["1. IAM", "2. Attack Surface"]
assert chart_kwargs["values"] == [100.0, 50.0]
def test_create_section_chart_no_findings_section_gets_100(
self, threatscore_generator, basic_compliance_data
self, threatscore_generator, basic_compliance_data, patched_threatscore_charts
):
"""Test that sections without findings get 100% score."""
mock_attr = Mock()
@@ -786,6 +818,11 @@ class TestSectionScoreChart:
)
assert isinstance(result, io.BytesIO)
chart_kwargs = patched_threatscore_charts["create_vertical_bar_chart"][0][
"kwargs"
]
assert chart_kwargs["labels"] == ["1. IAM"]
assert chart_kwargs["values"] == [100.0]
# =============================================================================
@@ -797,7 +834,11 @@ class TestExecutiveSummary:
"""Test suite for executive summary generation."""
def test_executive_summary_contains_chart(
self, threatscore_generator, basic_compliance_data, mock_requirement_attribute
self,
threatscore_generator,
basic_compliance_data,
mock_requirement_attribute,
patched_threatscore_charts,
):
"""Test that executive summary contains a chart."""
basic_compliance_data.requirements = [
@@ -818,9 +859,18 @@ class TestExecutiveSummary:
assert len(elements) > 0
assert any(isinstance(e, Image) for e in elements)
chart_kwargs = patched_threatscore_charts["create_vertical_bar_chart"][0][
"kwargs"
]
assert chart_kwargs["labels"] == ["1. IAM"]
assert chart_kwargs["values"] == [100.0]
def test_executive_summary_contains_score_table(
self, threatscore_generator, basic_compliance_data, mock_requirement_attribute
self,
threatscore_generator,
basic_compliance_data,
mock_requirement_attribute,
patched_threatscore_charts,
):
"""Test that executive summary contains a score table."""
basic_compliance_data.requirements = [
+44 -44
View File
@@ -76,7 +76,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
with (
patch("api.db_utils.rls_transaction"),
@@ -134,7 +134,7 @@ class TestPerformScan:
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
# Ensure the provider type is 'aws' to match our mocks
provider.provider = Provider.ProviderChoices.AWS
@@ -243,11 +243,11 @@ class TestPerformScan:
mock_prowler_scan_class,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
tenant_id = str(tenant.id)
scan_id = str(scan.id)
@@ -268,11 +268,11 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
tenant_id = str(tenant.id)
scan_id = str(scan.id)
@@ -304,11 +304,11 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
tenant_id = str(tenant.id)
scan_id = str(scan.id)
@@ -511,7 +511,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test that failed findings increment the failed_findings_count"""
with (
@@ -532,7 +532,7 @@ class TestPerformScan:
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
# Ensure the provider type is 'aws'
provider.provider = Provider.ProviderChoices.AWS
@@ -589,7 +589,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test that multiple FAIL findings on the same resource increment the counter correctly"""
with (
@@ -606,7 +606,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -706,7 +706,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test that muted FAIL findings do not increment the failed_findings_count"""
with (
@@ -723,7 +723,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -777,13 +777,13 @@ class TestPerformScan:
def test_perform_prowler_scan_reset_failed_findings_count(
self,
tenants_fixture,
providers_fixture,
aws_provider,
resources_fixture,
):
"""Test that failed_findings_count is reset to 0 at the beginning of each scan"""
# Use existing resource from fixture and set initial failed_findings_count
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
resource = resources_fixture[0]
# Set a non-zero failed_findings_count initially
@@ -956,7 +956,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test active MuteRule mutes findings with correct reason"""
with (
@@ -973,7 +973,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -1073,7 +1073,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test inactive MuteRule does not mute findings"""
with (
@@ -1090,7 +1090,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -1159,7 +1159,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test mutelist processor takes precedence over MuteRule"""
with (
@@ -1176,7 +1176,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -1245,7 +1245,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test MuteRule with multiple finding UIDs mutes all findings"""
with (
@@ -1262,7 +1262,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -1344,7 +1344,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test scan continues when MuteRule loading fails"""
with (
@@ -1362,7 +1362,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -1427,7 +1427,7 @@ class TestPerformScan:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
):
"""Test muted_at timestamp is set correctly for muted findings"""
with (
@@ -1444,7 +1444,7 @@ class TestPerformScan:
):
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.AWS
provider.save()
@@ -2031,7 +2031,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
resources_fixture,
):
@@ -2082,7 +2082,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
with patch(
@@ -2120,7 +2120,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
"""Re-running compliance materialization must not raise nor duplicate rows.
@@ -2175,7 +2175,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
with patch(
@@ -2183,7 +2183,7 @@ class TestCreateComplianceRequirements:
) as mock_compliance_template:
tenant = tenants_fixture[0]
scan = scans_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
provider.provider = Provider.ProviderChoices.KUBERNETES
provider.save()
@@ -2221,7 +2221,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
with patch(
@@ -2240,7 +2240,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
with patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider:
@@ -2324,7 +2324,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
with patch(
@@ -2362,7 +2362,7 @@ class TestCreateComplianceRequirements:
self,
tenants_fixture,
scans_fixture,
providers_fixture,
aws_provider,
findings_fixture,
):
with patch(
@@ -4689,7 +4689,7 @@ class TestUpdateProviderComplianceScores:
self,
mock_psycopg_connection,
tenants_fixture,
providers_fixture,
aws_provider,
scans_fixture,
settings,
):
@@ -4797,7 +4797,7 @@ class TestResetEphemeralResourceFindingsCount:
)
def test_resets_only_resources_missing_from_full_scope_scan(
self, tenants_fixture, scans_fixture, providers_fixture, resources_fixture
self, tenants_fixture, scans_fixture, aws_provider, resources_fixture
):
tenant, *_ = tenants_fixture
scan1, scan2, *_ = scans_fixture
@@ -4877,7 +4877,7 @@ class TestResetEphemeralResourceFindingsCount:
assert result["reason"] == "scan not found"
def test_skips_when_newer_scan_completed_for_same_provider(
self, tenants_fixture, scans_fixture, providers_fixture, resources_fixture
self, tenants_fixture, scans_fixture, aws_provider, resources_fixture
):
# If a newer completed scan exists for the same provider, our
# ResourceScanSummary set is stale relative to the resources' current
@@ -4886,7 +4886,7 @@ class TestResetEphemeralResourceFindingsCount:
tenant, *_ = tenants_fixture
scan1, *_ = scans_fixture
provider, *_ = providers_fixture
provider = aws_provider
_, resource2, _ = resources_fixture
Resource.objects.filter(id=resource2.id).update(failed_findings_count=5)
@@ -4916,7 +4916,7 @@ class TestResetEphemeralResourceFindingsCount:
assert resource2.failed_findings_count == 5
def test_does_not_touch_other_providers_resources(
self, tenants_fixture, scans_fixture, providers_fixture, resources_fixture
self, tenants_fixture, scans_fixture, aws_provider, resources_fixture
):
tenant, *_ = tenants_fixture
scan1, *_ = scans_fixture
@@ -4982,14 +4982,14 @@ class TestResetEphemeralResourceFindingsCount:
assert resource2.failed_findings_count == 5
def test_ignores_sibling_scan_with_null_completed_at(
self, tenants_fixture, scans_fixture, providers_fixture, resources_fixture
self, tenants_fixture, scans_fixture, aws_provider, resources_fixture
):
# Postgres orders NULL first in DESC; a sibling COMPLETED scan with a
# missing completed_at must not be treated as the latest scan and
# cause us to incorrectly skip the reset.
tenant, *_ = tenants_fixture
scan1, *_ = scans_fixture
provider, *_ = providers_fixture
provider = aws_provider
resource1, resource2, _ = resources_fixture
Resource.objects.filter(id=resource2.id).update(failed_findings_count=5)
+34 -37
View File
@@ -1964,11 +1964,11 @@ class TestCleanupOrphanScheduledScans:
)
def test_cleanup_deletes_orphan_when_both_available_and_scheduled_exist(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Test that AVAILABLE scan is deleted when SCHEDULED also exists."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
# Create orphan AVAILABLE scan
@@ -2004,11 +2004,11 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=scheduled_scan.id).exists()
def test_cleanup_does_not_delete_when_only_available_exists(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Test that AVAILABLE scan is NOT deleted when no SCHEDULED exists."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
# Create only AVAILABLE scan (normal first scan scenario)
@@ -2033,11 +2033,11 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=available_scan.id).exists()
def test_cleanup_does_not_delete_when_only_scheduled_exists(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Test that nothing is deleted when only SCHEDULED exists."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
# Create only SCHEDULED scan (normal subsequent scan scenario)
@@ -2062,11 +2062,11 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=scheduled_scan.id).exists()
def test_cleanup_returns_zero_when_no_scans_exist(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Test that cleanup returns 0 when no scans exist."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
# Execute cleanup with no scans
@@ -2079,11 +2079,11 @@ class TestCleanupOrphanScheduledScans:
assert deleted_count == 0
def test_cleanup_deletes_multiple_orphan_available_scans(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Test that multiple AVAILABLE orphan scans are all deleted."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
# Create multiple orphan AVAILABLE scans
@@ -2128,12 +2128,11 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=scheduled_scan.id).exists()
def test_cleanup_does_not_affect_different_provider(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider_pair
):
"""Test that cleanup only affects scans for the specified provider."""
tenant = tenants_fixture[0]
provider1 = providers_fixture[0]
provider2 = providers_fixture[1]
provider1, provider2 = aws_provider_pair
periodic_task1 = self._create_periodic_task(provider1.id, tenant.id)
periodic_task2 = self._create_periodic_task(provider2.id, tenant.id)
@@ -2178,12 +2177,10 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=scheduled_scan_p1.id).exists()
assert Scan.objects.filter(id=available_scan_p2.id).exists()
def test_cleanup_does_not_affect_manual_scans(
self, tenants_fixture, providers_fixture
):
def test_cleanup_does_not_affect_manual_scans(self, tenants_fixture, aws_provider):
"""Test that cleanup only affects SCHEDULED trigger scans, not MANUAL."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
# Create orphan AVAILABLE scheduled scan
@@ -2229,11 +2226,11 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=manual_scan.id).exists()
def test_cleanup_does_not_affect_different_scheduler_task(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Test that cleanup only affects scans with the specified scheduler_task_id."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task1 = self._create_periodic_task(provider.id, tenant.id)
# Create another periodic task
@@ -2288,11 +2285,11 @@ class TestCleanupOrphanScheduledScans:
assert Scan.objects.filter(id=available_scan_other_task.id).exists()
def test_cleanup_keeps_db_queued_scheduled_scans(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""DB-queued scheduled scans have a task and must not be deleted as orphans."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
task_result = TaskResult.objects.create(
task_id=str(uuid.uuid4()),
@@ -2381,11 +2378,11 @@ class TestPerformScheduledScanTask:
return task_result
def test_queues_scheduled_scan_when_scheduled_scan_is_executing(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Queue a scheduled run when another scheduled scan is executing."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
task_id = str(uuid.uuid4())
self._create_task_result(tenant.id, task_id)
@@ -2431,11 +2428,11 @@ class TestPerformScheduledScanTask:
)
def test_queues_scheduled_scan_when_manual_scan_is_pending(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Queue one scheduled run when a manual scan is already dispatched."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._create_periodic_task(provider.id, tenant.id)
task_id = str(uuid.uuid4())
self._create_task_result(tenant.id, task_id)
@@ -2488,11 +2485,11 @@ class TestPerformScheduledScanTask:
assert scheduled_scan.scheduled_at > datetime.now(UTC)
def test_coalesces_scheduled_scan_when_one_is_already_queued(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Reuse the existing queued scheduled scan instead of adding another."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
task_id = str(uuid.uuid4())
self._create_task_result(tenant.id, task_id)
@@ -2557,11 +2554,11 @@ class TestPerformScheduledScanTask:
)
def test_creates_next_scheduled_scan_after_completion(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Create a next scheduled scan after a successful run completes."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._create_periodic_task(provider.id, tenant.id)
task_id = str(uuid.uuid4())
self._create_task_result(tenant.id, task_id)
@@ -2616,11 +2613,11 @@ class TestPerformScheduledScanTask:
)
def test_next_scheduled_scan_failure_does_not_mask_completed_scan(
self, tenants_fixture, providers_fixture, caplog
self, tenants_fixture, aws_provider, caplog
):
"""Keep scheduled scan success when next-run creation fails."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
self._create_periodic_task(provider.id, tenant.id)
task_id = str(uuid.uuid4())
self._create_task_result(tenant.id, task_id)
@@ -2651,11 +2648,11 @@ class TestPerformScheduledScanTask:
assert "Failed to ensure next scheduled scan" in caplog.text
def test_dedupes_multiple_scheduled_scans_before_run(
self, tenants_fixture, providers_fixture
self, tenants_fixture, aws_provider
):
"""Ensure duplicated scheduled scans are removed before executing."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
periodic_task = self._create_periodic_task(provider.id, tenant.id)
task_id = str(uuid.uuid4())
self._create_task_result(tenant.id, task_id)
@@ -2767,12 +2764,12 @@ class TestPerformScanTask:
def test_dispatches_next_queued_scan_after_completion(
self,
tenants_fixture,
providers_fixture,
aws_provider,
django_capture_on_commit_callbacks,
):
"""Dispatch the next queued scan for the provider after completion."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
current_scan = Scan.objects.create(
tenant_id=tenant.id,
provider=provider,
@@ -2830,11 +2827,11 @@ class TestPerformScanTask:
)
def test_dispatch_failure_does_not_mask_completed_scan(
self, tenants_fixture, providers_fixture, caplog
self, tenants_fixture, aws_provider, caplog
):
"""Keep scan success when queued dispatch fails after completion."""
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
current_scan = Scan.objects.create(
tenant_id=tenant.id,
provider=provider,
+3 -6
View File
@@ -3212,21 +3212,18 @@ class YourProviderAPITestCase(APITestCase):
#### 2.6.1. Add your mocked provider to the tests
If needed, add your mocked provider to the tests config file so you can use it on the tests.
If needed, add a named provider fixture or extend the provider factory defaults so tests can request only the provider they need.
**File:** `api/src/backend/conftest.py`
```python
@pytest.fixture
def providers_fixture(tenants_fixture):
tenant, *_ = tenants_fixture
providerX = Provider.objects.create(
def your_provider(provider_factory):
return provider_factory(
provider="your_provider",
uid="your_uid",
alias="your_alias",
tenant_id=tenant.id,
)
return provider1, provider2, provider3, ... providerX
```
### 2.7. Compliance and Output Support
+15 -3
View File
@@ -31,7 +31,7 @@ allowed-tools: Read, Edit, Write, Glob, Grep, Bash, WebFetch, WebSearch, Task
```text
create_test_user (session) ─► tenants_fixture (function) ─► authenticated_client
│
└─► providers_fixture ─► scans_fixture ─► findings_fixture
└─► aws_provider ─► scans_fixture ─► findings_fixture
```
### Key Fixtures
@@ -40,8 +40,12 @@ create_test_user (session) ─► tenants_fixture (function) ─► authenticate
|---------|-------------|
| `create_test_user` | Session user (`dev@prowler.com`) |
| `tenants_fixture` | 3 tenants: [0],[1] have membership, [2] isolated |
| `authenticated_client` | JWT client for tenant[0] |
| `providers_fixture` | 9 providers in tenant[0] |
| `authenticated_client` | Django test client with JWT for tenant[0] |
| `authenticated_client_for_tenant_factory` | Creates a Django test client with JWT for a specific user and tenant |
| `provider_factory` | Creates one validated provider with provider-specific defaults |
| `aws_provider` | 1 AWS provider in tenant[0] |
| `aws_provider_pair` | 2 AWS providers in tenant[0] |
| `all_provider_types_fixture` | 1 provider for every supported provider type |
| `tasks_fixture` | 2 Celery tasks with TaskResult |
### RBAC Fixtures
@@ -52,6 +56,14 @@ create_test_user (session) ─► tenants_fixture (function) ─► authenticate
| `authenticated_client_rbac_noroles` | Membership but NO roles |
| `authenticated_client_no_permissions_rbac` | All permissions = False |
Use `authenticated_client` for normal view behavior tests. It uses a cheap JWT
and still runs the real request authentication path. Use serializer-generated
JWTs or API-key clients only when the test is specifically about token
obtain/refresh, invalid tokens, expired tokens, tenant switching by token, API
keys, or unauthenticated 401 behavior. Use
`authenticated_client_for_tenant_factory` when a test needs a cheap JWT client
for a different user or tenant.
---
## 2. JSON:API Requests
+24 -24
View File
@@ -22,12 +22,12 @@ from api.rls import Tenant
class TestProviderViewSet:
"""Example API tests for Provider endpoints."""
def test_list_providers(self, authenticated_client, providers_fixture):
def test_list_providers(self, authenticated_client, aws_provider):
"""GET list returns all providers for authenticated tenant."""
response = authenticated_client.get(reverse("provider-list"))
assert response.status_code == status.HTTP_200_OK
assert len(response.json()["data"]) == len(providers_fixture)
assert len(response.json()["data"]) == 1
def test_create_provider(self, authenticated_client):
"""POST with JSON:API format creates provider."""
@@ -49,9 +49,9 @@ class TestProviderViewSet:
assert response.status_code == status.HTTP_201_CREATED
assert response.json()["data"]["attributes"]["uid"] == "123456789012"
def test_update_provider(self, authenticated_client, providers_fixture):
def test_update_provider(self, authenticated_client, aws_provider):
"""PATCH with JSON:API format updates provider."""
provider = providers_fixture[0]
provider = aws_provider
payload = {
"data": {
@@ -95,7 +95,7 @@ class TestRLSIsolation:
assert response.status_code == status.HTTP_404_NOT_FOUND
def test_list_excludes_other_tenants(
self, authenticated_client, providers_fixture, tenants_fixture
self, authenticated_client, aws_provider, tenants_fixture
):
"""List endpoints only return resources from user's tenants."""
# Create provider in isolated tenant
@@ -109,8 +109,8 @@ class TestRLSIsolation:
response = authenticated_client.get(reverse("provider-list"))
assert response.status_code == status.HTTP_200_OK
# Should only see providers_fixture (9 providers in tenant[0])
assert len(response.json()["data"]) == len(providers_fixture)
# Should only see the AWS provider in tenant[0]
assert len(response.json()["data"]) == 1
@pytest.mark.django_db
@@ -136,7 +136,7 @@ class TestRBACPermissions:
response = authenticated_client_rbac_noroles.get(reverse("user-list"))
assert response.status_code == status.HTTP_403_FORBIDDEN
def test_admin_sees_all(self, authenticated_client_rbac, providers_fixture):
def test_admin_sees_all(self, authenticated_client_rbac, aws_provider):
"""Admin with unlimited_visibility=True sees all providers."""
response = authenticated_client_rbac.get(reverse("provider-list"))
assert response.status_code == status.HTTP_200_OK
@@ -153,11 +153,11 @@ class TestAsyncOperations:
mock_delete_task,
mock_task_get,
authenticated_client,
providers_fixture,
aws_provider,
tasks_fixture,
):
"""DELETE returns 202 Accepted with Content-Location header."""
provider = providers_fixture[0]
provider = aws_provider
prowler_task = tasks_fixture[0]
# Mock the Celery task
@@ -184,11 +184,11 @@ class TestAsyncOperations:
mock_scan_task,
mock_task_get,
authenticated_client,
providers_fixture,
aws_provider,
tasks_fixture,
):
"""POST to scan trigger returns 202 with task location."""
provider = providers_fixture[0]
provider = aws_provider
prowler_task = tasks_fixture[0]
task_mock = Mock()
@@ -208,9 +208,9 @@ class TestAsyncOperations:
class TestJSONAPIResponses:
"""Example JSON:API response handling."""
def test_read_single_resource(self, authenticated_client, providers_fixture):
def test_read_single_resource(self, authenticated_client, aws_provider):
"""Read data from single resource response."""
provider = providers_fixture[0]
provider = aws_provider
response = authenticated_client.get(
reverse("provider-detail", kwargs={"pk": provider.id})
)
@@ -222,12 +222,12 @@ class TestJSONAPIResponses:
assert resource_id == str(provider.id)
assert attrs["provider"] == provider.provider
def test_read_list_response(self, authenticated_client, providers_fixture):
def test_read_list_response(self, authenticated_client, aws_provider):
"""Read data from list response."""
response = authenticated_client.get(reverse("provider-list"))
items = response.json()["data"]
assert len(items) == len(providers_fixture)
assert len(items) == 1
def test_read_relationships(self, authenticated_client, scans_fixture):
"""Read relationship data."""
@@ -262,9 +262,9 @@ class TestJSONAPIResponses:
class TestSoftDelete:
"""Example soft-delete manager tests."""
def test_objects_excludes_soft_deleted(self, providers_fixture):
def test_objects_excludes_soft_deleted(self, aws_provider):
"""Default manager excludes soft-deleted records."""
provider = providers_fixture[0]
provider = aws_provider
provider.is_deleted = True
provider.save()
@@ -284,12 +284,12 @@ class TestSoftDelete:
class TestCeleryTaskLogic:
"""Example: Testing Celery task logic directly with apply()."""
def test_task_logic_directly(self, tenants_fixture, providers_fixture):
def test_task_logic_directly(self, tenants_fixture, aws_provider):
"""Use apply() for synchronous execution without Celery worker."""
from tasks.tasks import check_provider_connection_task
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
# Execute task synchronously (no broker needed)
result = check_provider_connection_task.apply(
@@ -328,12 +328,12 @@ class TestSetTenantDecorator:
"""Example: Testing @set_tenant decorator behavior."""
@patch("api.decorators.connection")
def test_sets_rls_context(self, mock_conn, tenants_fixture, providers_fixture):
def test_sets_rls_context(self, mock_conn, tenants_fixture, aws_provider):
"""Verify @set_tenant sets RLS context via SET_CONFIG_QUERY."""
from tasks.tasks import check_provider_connection_task
tenant = tenants_fixture[0]
provider = providers_fixture[0]
provider = aws_provider
# Call task with tenant_id - decorator sets RLS and pops it
check_provider_connection_task.apply(
@@ -349,13 +349,13 @@ class TestBeatScheduling:
"""Example: Testing Beat scheduled task creation."""
@patch("tasks.beat.perform_scheduled_scan_task.apply_async")
def test_schedule_provider_scan(self, mock_apply, providers_fixture):
def test_schedule_provider_scan(self, mock_apply, aws_provider):
"""Verify periodic task is created with correct settings."""
from django_celery_beat.models import PeriodicTask
from tasks.beat import schedule_provider_scan
provider = providers_fixture[0]
provider = aws_provider
mock_apply.return_value = Mock(id="task-123")
schedule_provider_scan(provider)
@@ -24,7 +24,7 @@ create_test_user (session)
│ └─► authenticated_client
│ └─► (most API tests use this)
│
├─► providers_fixture
├─► aws_provider
│ └─► scans_fixture
│ └─► findings_fixture
│
@@ -102,12 +102,20 @@ Authentication tests:
```python
@pytest.mark.django_db
class TestProviderViewSet:
def test_list(self, authenticated_client, providers_fixture):
# authenticated_client has JWT for tenant[0]
# providers_fixture has 9 providers in tenant[0]
def test_list(self, authenticated_client, aws_provider):
# authenticated_client is a Django test client with JWT for tenant[0]
# aws_provider creates one validated AWS provider in tenant[0]
...
```
Use serializer-generated JWTs or API-key clients for authentication behavior
tests only: token obtain/refresh, invalid or expired tokens, token-scoped tenant
switching, API keys, and unauthenticated 401 responses. Regular view tests
should use `authenticated_client` so they still exercise `request.user`,
`request.auth["tenant_id"]`, RLS, and RBAC without paying token serializer cost.
Use `authenticated_client_for_tenant_factory` when a test needs the same cheap
JWT path for a different user or tenant.
### RBAC Tests
```python
+2 -2
View File
@@ -164,9 +164,9 @@ class Test_ec2_ami_public:
```python
@pytest.mark.django_db
class TestResourceModel:
def test_create_resource_with_tags(self, providers_fixture):
def test_create_resource_with_tags(self, aws_provider):
# Given
provider, *_ = providers_fixture
provider = aws_provider
tenant_id = provider.tenant_id
# When