Compare commits

...
19 Commits
Author SHA1 Message Date
Prowler BotandPepe Fagoaga 0d00e29608 chore(sdk): changelog for v5.30.3 (#11652)
Co-authored-by: Pepe Fagoaga <pepe@prowler.com>
2026-06-19 15:26:23 +02:00
Adrián PeñaandPepe Fagoaga f5ff30ad17 fix(saml): cross-tenant account takeover via SAML domain claiming (#11650)
Co-authored-by: Pepe Fagoaga <pepe@prowler.com>
2026-06-19 14:27:06 +02:00
Prowler BotandPedro Martín f6679fadf4 fix(compliance): multi-section undercount & leaked provider tab (#11635)
Co-authored-by: Pedro Martín <pedromarting3@gmail.com>
2026-06-18 10:40:20 +02:00
Prowler Botandprowler-bot dcc9401957 chore(release): Bump versions to v5.30.3 (#11627)
Co-authored-by: prowler-bot <179230569+prowler-bot@users.noreply.github.com>
2026-06-17 10:47:31 +02:00
Prowler BotandPepe Fagoaga 88f6848913 chore(changelog): v5.30.2 (#11625)
Co-authored-by: Pepe Fagoaga <pepe@prowler.com>
2026-06-17 09:42:05 +02:00
Prowler Botandlydiavilchez 2de298fb7b fix(cli): prevent unrelated built-in provider failures from aborting the CLI (#11620)
Co-authored-by: lydiavilchez <114735608+lydiavilchez@users.noreply.github.com>
2026-06-16 14:35:15 +02:00
11f0845a91 fix(gcp): surface organization-scan failures instead of silently scanning the home project (#11619)
Co-authored-by: Rubén De la Torre Vico <ruben@prowler.com>
Co-authored-by: Daniel Barranquero <danielbo2001@gmail.com>
2026-06-16 14:12:57 +02:00
Prowler BotandPedro Martín 42d99a17a6 perf(api): optimize scan-compliance-overviews task (#11613)
Co-authored-by: Pedro Martín <pedromarting3@gmail.com>
2026-06-16 12:18:55 +02:00
César Arroba 832f10b7f6 ci: narrow osv-scanner gate to CRITICAL on v5.30 (backport #11580) (#11616) 2026-06-16 11:35:30 +02:00
César Arroba d133ad18a4 ci: always run container and dependency vulnerability scans on PRs (v5.30 backport) (#11614) 2026-06-16 11:16:38 +02:00
3539940a26 fix(gcp): credit audit-filtered aggregated sinks in metric-filter checks (#11607)
Co-authored-by: Aline Almeida <aline@tuplita.ai>
Co-authored-by: Hugo Pereira Brito <101209179+HugoPBrito@users.noreply.github.com>
Co-authored-by: Hugo P.Brito <hugopbrit@gmail.com>
2026-06-16 10:34:23 +02:00
Prowler Botandprowler-bot 1192d94648 chore(release): Bump versions to v5.30.2 (#11571)
Co-authored-by: prowler-bot <179230569+prowler-bot@users.noreply.github.com>
2026-06-12 13:48:44 +02:00
Prowler BotandJosema Camacho a578f4af34 chore: prepare API and UI changelogs for 5.30.1 release (#11566)
Co-authored-by: Josema Camacho <josema@prowler.com>
2026-06-12 12:16:15 +02:00
Prowler BotandAlejandro Bailo d6528b674e fix(ui): show threat map data for okta and google workspace accounts (#11563)
Co-authored-by: Alejandro Bailo <59607668+alejandrobailo@users.noreply.github.com>
2026-06-12 10:18:43 +02:00
Prowler BotandJosema Camacho 75decbbedf fix(api): drop_subgraph deletes relationships then nodes to cut Neo4j memory (#11561)
Co-authored-by: Josema Camacho <josema@prowler.com>
2026-06-12 09:47:41 +02:00
Prowler BotandPedro Martín 4a14559a5f fix(compliance): resolve provider from scan in attributes endp (#11560)
Co-authored-by: Pedro Martín <pedromarting3@gmail.com>
2026-06-12 09:18:11 +02:00
Prowler BotandHugo Pereira Brito c6f8620a0d fix(api): normalize OCI scan region credentials (#11559)
Co-authored-by: Hugo Pereira Brito <101209179+HugoPBrito@users.noreply.github.com>
2026-06-11 17:55:26 +02:00
Prowler Botandprowler-bot ca4889b43e chore(release): Bump versions to v5.30.1 (#11547)
Co-authored-by: prowler-bot <179230569+prowler-bot@users.noreply.github.com>
2026-06-11 15:28:54 +02:00
Prowler Botandprowler-bot 057d061c7e chore(api): Update prowler dependency to v5.30 for release 5.30.0 (#11543)
Co-authored-by: prowler-bot <179230569+prowler-bot@users.noreply.github.com>
2026-06-11 11:15:18 +02:00
70 changed files with 3941 additions and 465 deletions
+1 -1
View File
@@ -145,7 +145,7 @@ SENTRY_RELEASE=local
NEXT_PUBLIC_SENTRY_ENVIRONMENT=${SENTRY_ENVIRONMENT}
#### Prowler release version ####
NEXT_PUBLIC_PROWLER_RELEASE_VERSION=v5.30.0
NEXT_PUBLIC_PROWLER_RELEASE_VERSION=v5.30.3
# Social login credentials
SOCIAL_GOOGLE_OAUTH_CALLBACK_URL="${AUTH_URL}/api/auth/callback/google"
+3 -3
View File
@@ -1,5 +1,5 @@
name: 'OSV-Scanner'
description: 'Install osv-scanner and scan a lockfile, failing on HIGH/CRITICAL/UNKNOWN severity findings. Posts/updates a PR comment with findings on pull_request events (requires pull-requests: write).'
description: 'Install osv-scanner and scan a lockfile, failing on CRITICAL severity findings. Posts/updates a PR comment with findings on pull_request events (requires pull-requests: write).'
author: 'Prowler'
inputs:
@@ -7,9 +7,9 @@ inputs:
description: 'Path to the lockfile to scan, relative to the repository root (e.g. uv.lock, api/uv.lock, ui/pnpm-lock.yaml).'
required: true
severity-levels:
description: 'Comma-separated severity levels that fail the scan. Default: HIGH,CRITICAL,UNKNOWN.'
description: 'Comma-separated severity levels that fail the scan. Default: CRITICAL.'
required: false
default: 'HIGH,CRITICAL,UNKNOWN'
default: 'CRITICAL'
version:
description: 'osv-scanner release tag to install. When overriding, you MUST also override binary-sha256.'
required: false
@@ -12,9 +12,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'api/**'
- '.github/workflows/api-container-checks.yml'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
-7
View File
@@ -16,13 +16,6 @@ on:
branches:
- "master"
- "v5.*"
paths:
- 'api/**'
- '.github/workflows/api-tests.yml'
- '.github/workflows/api-security.yml'
- '.github/actions/setup-python-uv/**'
- '.github/actions/osv-scanner/**'
- '.github/scripts/osv-scan.sh'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -12,9 +12,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'mcp_server/**'
- '.github/workflows/mcp-container-checks.yml'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
-6
View File
@@ -15,12 +15,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'mcp_server/pyproject.toml'
- 'mcp_server/uv.lock'
- '.github/workflows/mcp-security.yml'
- '.github/actions/osv-scanner/**'
- '.github/scripts/osv-scan.sh'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -15,12 +15,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'prowler/**'
- 'Dockerfile*'
- 'pyproject.toml'
- 'uv.lock'
- '.github/workflows/sdk-container-checks.yml'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
-10
View File
@@ -19,16 +19,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'prowler/**'
- 'tests/**'
- 'pyproject.toml'
- 'uv.lock'
- '.github/workflows/sdk-tests.yml'
- '.github/workflows/sdk-security.yml'
- '.github/actions/setup-python-uv/**'
- '.github/actions/osv-scanner/**'
- '.github/scripts/osv-scan.sh'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -12,9 +12,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'ui/**'
- '.github/workflows/ui-container-checks.yml'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
-6
View File
@@ -15,12 +15,6 @@ on:
branches:
- 'master'
- 'v5.*'
paths:
- 'ui/package.json'
- 'ui/pnpm-lock.yaml'
- '.github/workflows/ui-security.yml'
- '.github/actions/osv-scanner/**'
- '.github/scripts/osv-scan.sh'
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
+27 -1
View File
@@ -2,6 +2,32 @@
All notable changes to the **Prowler API** are documented in this file.
## [1.31.3] (Prowler v5.30.3)
### 🔐 Security
- SAML logins now link to an existing account only when the asserted email domain matches the ACS endpoint and the user is already a member of that domain's tenant, fixing a cross-tenant account takeover [(GHSA-h8m9-jgf8-vwvp)](https://github.com/prowler-cloud/prowler/security/advisories/GHSA-h8m9-jgf8-vwvp) [bf3b5c2ba713e533014927141b64948c82c8f32e](https://github.com/prowler-cloud/prowler/commit/bf3b5c2ba713e533014927141b64948c82c8f32e)
---
## [1.31.2] (Prowler v5.30.2)
### 🔄 Changed
- `scan-compliance-overviews` task now streams the findings aggregation and the requirement-row writes so it runs faster and its peak memory no longer grows with the number of regions and frameworks [(#11591)](https://github.com/prowler-cloud/prowler/pull/11591)
---
## [1.31.1] (Prowler v5.30.1)
### 🐞 Fixed
- `compliance-overviews/attributes` now resolves the provider from the scan, so multi-provider universal frameworks (e.g. CSA CCM) return the check IDs of the scan's provider and Azure/GCP requirement details show their findings instead of appearing empty [(#11546)](https://github.com/prowler-cloud/prowler/pull/11546)
- Attack Paths: `drop_subgraph` now deletes relationships first and then nodes in batches, using less memory on Neo4j when clearing a dense provider graph [(#11557)](https://github.com/prowler-cloud/prowler/pull/11557)
- OCI scans now use API key credentials with the configured region instead of falling back to `/home/prowler/.oci/config` [(#11558)](https://github.com/prowler-cloud/prowler/pull/11558)
---
## [1.31.0] (Prowler v5.30.0)
### 🚀 Added
@@ -19,7 +45,7 @@ All notable changes to the **Prowler API** are documented in this file.
- Workers now shut down gracefully on deploy or restart, finishing or re-queueing in-flight tasks instead of being force-killed and leaving them stuck [(#11416)](https://github.com/prowler-cloud/prowler/pull/11416)
- Resource `name` is now stored and refreshed on every scan, so resources no longer keep an empty name [(#11476)](https://github.com/prowler-cloud/prowler/pull/11476)
- Compliance catalog now warms in background during startup. `compliance-overviews/attributes` returns `503` while warming, so the first request after a deploy no longer trips the API timeout [(#4554)](https://github.com/prowler-cloud/prowler-cloud/pull/4554)
- Compliance catalog now warms in background during startup. `compliance-overviews/attributes` returns `503` while warming, so the first request after a deploy no longer trips the API timeout [(#11530)](https://github.com/prowler-cloud/prowler/pull/11530)
### 🔐 Security
+2 -2
View File
@@ -43,7 +43,7 @@ dependencies = [
"defusedxml==0.7.1",
"gunicorn==23.0.0",
"lxml==6.1.0",
"prowler @ git+https://github.com/prowler-cloud/prowler.git@master",
"prowler @ git+https://github.com/prowler-cloud/prowler.git@v5.30",
"psycopg2-binary==2.9.9",
"pytest-celery[redis] (==1.3.0)",
"sentry-sdk[django] (==2.56.0)",
@@ -68,7 +68,7 @@ name = "prowler-api"
package-mode = false
# Needed for the SDK compatibility
requires-python = ">=3.11,<3.13"
version = "1.31.0"
version = "1.31.3"
[tool.uv]
# Transitive pins matching master to avoid silent drift; bump deliberately.
+41 -1
View File
@@ -3,7 +3,14 @@ from django.db import transaction
from api.db_router import MainRouter
from api.db_utils import rls_transaction
from api.models import Membership, Role, Tenant, User, UserRoleRelationship
from api.models import (
Membership,
Role,
SAMLConfiguration,
Tenant,
User,
UserRoleRelationship,
)
class ProwlerSocialAccountAdapter(DefaultSocialAccountAdapter):
@@ -18,7 +25,40 @@ class ProwlerSocialAccountAdapter(DefaultSocialAccountAdapter):
# Link existing accounts with the same email address
email = sociallogin.account.extra_data.get("email")
if sociallogin.provider.id == "saml":
# For SAML, the asserted NameID email cannot be trusted on its own.
# Prevent cross-tenant account takeover (GHSA-h8m9-jgf8-vwvp) by
# linking only when the email domain matches the ACS endpoint and the
# existing user is already a member of that tenant.
email = sociallogin.user.email
if not email:
return
domain = email.rsplit("@", 1)[-1].lower()
resolver_match = getattr(request, "resolver_match", None)
organization_slug = (
(resolver_match.kwargs or {}).get("organization_slug", "")
if resolver_match
else ""
).lower()
# The ACS endpoint is scoped per email domain; reject mismatches so an
# attacker cannot replay an assertion through another tenant's endpoint.
if organization_slug != domain:
return
try:
saml_config = SAMLConfiguration.objects.using(MainRouter.admin_db).get(
email_domain=domain
)
except SAMLConfiguration.DoesNotExist:
return
existing_user = self.get_user_by_email(email)
if existing_user and existing_user.is_member_of_tenant(
str(saml_config.tenant_id)
):
sociallogin.connect(request, existing_user)
return
if email:
existing_user = self.get_user_by_email(email)
if existing_user:
+18 -2
View File
@@ -175,7 +175,8 @@ def drop_subgraph(database: str, provider_id: str) -> int:
"""
Delete all nodes for a provider from the tenant database.
Uses batched deletion to avoid memory issues with large graphs.
Deletes relationships then nodes in batches (not `DETACH DELETE`) so a dense
provider's graph cannot exceed Neo4j's transaction memory limit.
Silently returns 0 if the database doesn't exist.
"""
provider_label = get_provider_label(provider_id)
@@ -183,13 +184,28 @@ def drop_subgraph(database: str, provider_id: str) -> int:
try:
with get_session(database) as session:
# Phase 1: delete relationships incident to provider nodes in batches.
deleted_count = 1
while deleted_count > 0:
result = session.run(
f"""
MATCH (:`{provider_label}`)-[r]-()
WITH DISTINCT r LIMIT $batch_size
DELETE r
RETURN COUNT(r) AS deleted_rels_count
""",
{"batch_size": BATCH_SIZE},
)
deleted_count = result.single().get("deleted_rels_count", 0)
# Phase 2: delete the now relationship-free nodes in batches.
deleted_count = 1
while deleted_count > 0:
result = session.run(
f"""
MATCH (n:{PROVIDER_RESOURCE_LABEL}:`{provider_label}`)
WITH n LIMIT $batch_size
DETACH DELETE n
DELETE n
RETURN COUNT(n) AS deleted_nodes_count
""",
{"batch_size": BATCH_SIZE},
+1 -1
View File
@@ -1,7 +1,7 @@
openapi: 3.0.3
info:
title: Prowler API
version: 1.31.0
version: 1.31.3
description: |-
Prowler API specification.
+148 -14
View File
@@ -1,3 +1,4 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
@@ -5,9 +6,48 @@ from allauth.socialaccount.models import SocialLogin
from django.contrib.auth import get_user_model
from api.adapters import ProwlerSocialAccountAdapter
from api.db_router import MainRouter
from api.models import SAMLConfiguration
User = get_user_model()
# Minimal, well-formed IdP metadata accepted by SAMLConfiguration._parse_metadata.
VALID_METADATA = """<?xml version='1.0' encoding='UTF-8'?>
<md:EntityDescriptor entityID='TEST' xmlns:md='urn:oasis:names:tc:SAML:2.0:metadata'>
<md:IDPSSODescriptor WantAuthnRequestsSigned='false' protocolSupportEnumeration='urn:oasis:names:tc:SAML:2.0:protocol'>
<md:KeyDescriptor use='signing'>
<ds:KeyInfo xmlns:ds='http://www.w3.org/2000/09/xmldsig#'>
<ds:X509Data>
<ds:X509Certificate>FAKECERTDATA</ds:X509Certificate>
</ds:X509Data>
</ds:KeyInfo>
</md:KeyDescriptor>
<md:SingleSignOnService Binding='urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST' Location='https://idp.test/sso'/>
</md:IDPSSODescriptor>
</md:EntityDescriptor>
"""
def _saml_request(rf, organization_slug):
"""Build an ACS request whose resolver_match carries the organization slug,
mirroring how Django populates it after routing the SAML ACS URL."""
request = rf.post(f"/api/v1/accounts/saml/{organization_slug}/acs/finish/")
request.resolver_match = SimpleNamespace(
kwargs={"organization_slug": organization_slug}
)
return request
def _saml_sociallogin(user):
sociallogin = MagicMock(spec=SocialLogin)
sociallogin.account = MagicMock()
sociallogin.provider = MagicMock()
sociallogin.provider.id = "saml"
sociallogin.account.extra_data = {}
sociallogin.user = user
sociallogin.connect = MagicMock()
return sociallogin
@pytest.mark.django_db
class TestProwlerSocialAccountAdapter:
@@ -20,26 +60,99 @@ class TestProwlerSocialAccountAdapter:
adapter = ProwlerSocialAccountAdapter()
assert adapter.get_user_by_email("notfound@example.com") is None
def test_pre_social_login_links_existing_user(self, create_test_user, rf):
def test_pre_social_login_links_member_of_saml_tenant(
self, create_test_user, tenants_fixture, rf
):
"""A SAML login links to an existing account only when that user is
already a member of the tenant that owns the asserted email domain."""
adapter = ProwlerSocialAccountAdapter()
# create_test_user (dev@prowler.com) is a member of tenant1.
domain = create_test_user.email.rsplit("@", 1)[-1]
SAMLConfiguration.objects.using(MainRouter.admin_db).create(
email_domain=domain,
metadata_xml=VALID_METADATA,
tenant=tenants_fixture[0],
)
sociallogin = MagicMock(spec=SocialLogin)
sociallogin.account = MagicMock()
sociallogin.provider = MagicMock()
sociallogin.provider.id = "saml"
sociallogin.account.extra_data = {}
sociallogin.user = create_test_user
sociallogin.connect = MagicMock()
adapter.pre_social_login(rf.get("/"), sociallogin)
sociallogin = _saml_sociallogin(create_test_user)
adapter.pre_social_login(_saml_request(rf, domain), sociallogin)
call_args = sociallogin.connect.call_args
assert call_args is not None
called_request, called_user = call_args[0]
assert called_request.path == "/"
_, called_user = call_args[0]
assert called_user.email == create_test_user.email
def test_pre_social_login_blocks_cross_tenant_takeover(
self, create_test_user, tenants_fixture, rf
):
"""GHSA-h8m9-jgf8-vwvp: an attacker tenant that claims the victim's
email domain must NOT be able to link to the victim's account, because
the victim is not a member of the attacker's tenant."""
adapter = ProwlerSocialAccountAdapter()
domain = create_test_user.email.rsplit("@", 1)[-1]
# tenant3 is the attacker tenant; create_test_user is NOT a member of it.
attacker_tenant = tenants_fixture[2]
assert not create_test_user.is_member_of_tenant(str(attacker_tenant.id))
SAMLConfiguration.objects.using(MainRouter.admin_db).create(
email_domain=domain,
metadata_xml=VALID_METADATA,
tenant=attacker_tenant,
)
sociallogin = _saml_sociallogin(create_test_user)
adapter.pre_social_login(_saml_request(rf, domain), sociallogin)
sociallogin.connect.assert_not_called()
def test_pre_social_login_blocks_domain_slug_mismatch(
self, create_test_user, tenants_fixture, rf
):
"""The asserted email domain must match the ACS endpoint's slug, so an
assertion cannot be replayed through a different tenant's endpoint."""
adapter = ProwlerSocialAccountAdapter()
domain = create_test_user.email.rsplit("@", 1)[-1]
SAMLConfiguration.objects.using(MainRouter.admin_db).create(
email_domain=domain,
metadata_xml=VALID_METADATA,
tenant=tenants_fixture[0],
)
sociallogin = _saml_sociallogin(create_test_user)
# Slug points at a different domain than the asserted email.
adapter.pre_social_login(_saml_request(rf, "attacker.com"), sociallogin)
sociallogin.connect.assert_not_called()
def test_pre_social_login_blocks_when_no_saml_config(
self, create_test_user, tenants_fixture, rf
):
"""No SAML configuration for the domain means nothing to link against."""
adapter = ProwlerSocialAccountAdapter()
domain = create_test_user.email.rsplit("@", 1)[-1]
sociallogin = _saml_sociallogin(create_test_user)
adapter.pre_social_login(_saml_request(rf, domain), sociallogin)
sociallogin.connect.assert_not_called()
def test_pre_social_login_blocks_without_resolver_match(
self, create_test_user, tenants_fixture, rf
):
"""Fail closed: if the request has no resolver_match we cannot bind the
assertion to a tenant, so no linking happens."""
adapter = ProwlerSocialAccountAdapter()
domain = create_test_user.email.rsplit("@", 1)[-1]
SAMLConfiguration.objects.using(MainRouter.admin_db).create(
email_domain=domain,
metadata_xml=VALID_METADATA,
tenant=tenants_fixture[0],
)
sociallogin = _saml_sociallogin(create_test_user)
adapter.pre_social_login(rf.post("/"), sociallogin)
sociallogin.connect.assert_not_called()
def test_pre_social_login_no_link_if_email_missing(self, rf):
adapter = ProwlerSocialAccountAdapter()
@@ -47,14 +160,35 @@ class TestProwlerSocialAccountAdapter:
sociallogin.account = MagicMock()
sociallogin.provider = MagicMock()
sociallogin.user = MagicMock()
sociallogin.user.email = ""
sociallogin.provider.id = "saml"
sociallogin.account.extra_data = {}
sociallogin.connect = MagicMock()
adapter.pre_social_login(rf.get("/"), sociallogin)
adapter.pre_social_login(_saml_request(rf, "prowler.com"), sociallogin)
sociallogin.connect.assert_not_called()
def test_pre_social_login_non_saml_links_by_email(self, create_test_user, rf):
"""Non-SAML providers (e.g. Google/GitHub) still link to an existing
local account by email; the tenant binding only applies to SAML."""
adapter = ProwlerSocialAccountAdapter()
sociallogin = MagicMock(spec=SocialLogin)
sociallogin.account = MagicMock()
sociallogin.provider = MagicMock()
sociallogin.provider.id = "google"
sociallogin.account.extra_data = {"email": create_test_user.email}
sociallogin.user = create_test_user
sociallogin.connect = MagicMock()
adapter.pre_social_login(rf.get("/"), sociallogin)
call_args = sociallogin.connect.call_args
assert call_args is not None
_, called_user = call_args[0]
assert called_user.email == create_test_user.email
def test_save_user_saml_sets_session_flag(self, rf):
adapter = ProwlerSocialAccountAdapter()
request = rf.get("/")
@@ -542,3 +542,84 @@ class TestHasProviderData:
):
with pytest.raises(db_module.GraphDatabaseQueryException):
db_module.has_provider_data("db-tenant-abc", "provider-123")
class TestDropSubgraph:
"""Test drop_subgraph two-phase batched deletion of a provider's graph."""
@staticmethod
def _result(count):
result = MagicMock()
result.single.return_value.get.return_value = count
return result
@staticmethod
def _session_ctx(session):
ctx = MagicMock()
ctx.__enter__.return_value = session
ctx.__exit__.return_value = False
return ctx
def test_deletes_relationships_then_nodes_in_batches(self):
session = MagicMock()
# Phase 1 (relationships): one full batch then empty.
# Phase 2 (nodes): one full batch then empty.
session.run.side_effect = [
self._result(1000),
self._result(0),
self._result(1000),
self._result(0),
]
with patch(
"api.attack_paths.database.get_session",
return_value=self._session_ctx(session),
):
deleted = db_module.drop_subgraph("db-tenant-abc", "provider-123")
# Only phase-2 node counts contribute to the return value.
assert deleted == 1000
assert session.run.call_count == 4
queries = [call.args[0] for call in session.run.call_args_list]
# Regression guard: the memory blow-up was caused by DETACH DELETE.
assert all("DETACH DELETE" not in query for query in queries)
rel_queries = [query for query in queries if "DELETE r" in query]
node_queries = [query for query in queries if "DELETE n" in query]
assert rel_queries and node_queries
# DISTINCT avoids double-counting relationships matched from both ends.
assert all("DISTINCT r" in query for query in rel_queries)
# Relationships must be fully drained before nodes are deleted.
first_node = next(i for i, q in enumerate(queries) if "DELETE n" in q)
last_rel = max(i for i, q in enumerate(queries) if "DELETE r" in q)
assert last_rel < first_node
def test_returns_zero_when_database_not_found(self):
session_ctx = MagicMock()
session_ctx.__enter__.side_effect = db_module.GraphDatabaseQueryException(
message="Database does not exist",
code="Neo.ClientError.Database.DatabaseNotFound",
)
with patch(
"api.attack_paths.database.get_session",
return_value=session_ctx,
):
assert db_module.drop_subgraph("db-tenant-gone", "provider-123") == 0
def test_raises_on_other_errors(self):
session_ctx = MagicMock()
session_ctx.__enter__.side_effect = db_module.GraphDatabaseQueryException(
message="Connection refused",
code="Neo.TransientError.General.UnknownError",
)
with patch(
"api.attack_paths.database.get_session",
return_value=session_ctx,
):
with pytest.raises(db_module.GraphDatabaseQueryException):
db_module.drop_subgraph("db-tenant-abc", "provider-123")
+24
View File
@@ -357,6 +357,30 @@ class TestGetProwlerProviderKwargs:
expected_result = {**secret_dict, **expected_extra_kwargs}
assert result == expected_result
def test_get_prowler_provider_kwargs_oraclecloud_converts_region_string_to_set(
self,
):
secret_dict = {
"user": "ocid1.user.oc1..fake",
"fingerprint": "00:11:22:33:44:55:66:77",
"key_content": "-----BEGIN PRIVATE KEY-----\nfake\n-----END PRIVATE KEY-----",
"tenancy": "ocid1.tenancy.oc1..fake",
"region": "us-ashburn-1",
"pass_phrase": "fake-passphrase",
}
secret_mock = MagicMock()
secret_mock.secret = secret_dict
provider = MagicMock()
provider.provider = Provider.ProviderChoices.ORACLECLOUD.value
provider.secret = secret_mock
provider.uid = "ocid1.tenancy.oc1..fake"
result = get_prowler_provider_kwargs(provider)
expected_result = {**secret_dict, "region": {"us-ashburn-1"}}
assert result == expected_result
def test_get_prowler_provider_kwargs_with_mutelist(self):
provider_uid = "provider_uid"
secret_dict = {"key": "value"}
+312 -20
View File
@@ -9570,6 +9570,188 @@ class TestComplianceOverviewViewSet:
assert "Category" in first_attr
assert "AWSService" in first_attr
def test_compliance_overview_attributes_resolves_provider_from_scan(
self, authenticated_client, tenants_fixture, providers_fixture
):
# csa_ccm_4.0 is a multi-provider universal framework: a single
# compliance_id whose requirements expose different checks per provider.
# Passing a scan must return the check IDs for that scan's provider,
# otherwise the endpoint defaults to the first provider that declares the
# framework and azure/gcp requirements end up with check IDs that match
# no findings.
tenant = tenants_fixture[0]
gcp_provider = providers_fixture[2]
azure_provider = providers_fixture[4]
assert gcp_provider.provider == Provider.ProviderChoices.GCP.value
assert azure_provider.provider == Provider.ProviderChoices.AZURE.value
now = datetime.now(timezone.utc)
gcp_scan = Scan.objects.create(
name="gcp scan",
provider=gcp_provider,
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
tenant_id=tenant.id,
started_at=now,
completed_at=now,
)
azure_scan = Scan.objects.create(
name="azure scan",
provider=azure_provider,
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
tenant_id=tenant.id,
started_at=now,
completed_at=now,
)
def request_attributes(scan_id=None):
params = {"filter[compliance_id]": "csa_ccm_4.0"}
if scan_id is not None:
params["filter[scan_id]"] = str(scan_id)
return authenticated_client.get(
reverse("complianceoverview-attributes"), params
)
def collect_check_ids(scan_id=None):
response = request_attributes(scan_id)
assert response.status_code == status.HTTP_200_OK
check_ids = set()
for item in response.json()["data"]:
check_ids.update(item["attributes"]["attributes"]["check_ids"])
return check_ids
gcp_check_ids = collect_check_ids(gcp_scan.id)
azure_check_ids = collect_check_ids(azure_scan.id)
# Each scan resolves to its own provider's checks, and they differ.
assert gcp_check_ids
assert azure_check_ids
assert gcp_check_ids != azure_check_ids
# The returned check IDs belong to the SDK's per-provider definition.
from api.compliance import get_prowler_provider_compliance
def expected_check_ids(provider_type):
framework = get_prowler_provider_compliance(provider_type)["csa_ccm_4.0"]
expected = set()
for requirement in framework.requirements:
expected.update(requirement.checks.get(provider_type, []))
return expected
assert gcp_check_ids <= expected_check_ids(Provider.ProviderChoices.GCP.value)
assert azure_check_ids <= expected_check_ids(
Provider.ProviderChoices.AZURE.value
)
# An explicit scan_id is authoritative: a non-existent scan must fail
# closed with 404 instead of silently falling back to another provider.
missing_response = request_attributes("00000000-0000-0000-0000-000000000000")
assert missing_response.status_code == status.HTTP_404_NOT_FOUND
# A malformed scan_id is rejected with 404 as well.
malformed_response = request_attributes("not-a-uuid")
assert malformed_response.status_code == status.HTTP_404_NOT_FOUND
# An empty value (filter[scan_id]=) must not fall back to the legacy
# provider picker: the explicit (if blank) selector fails closed.
empty_response = request_attributes("")
assert empty_response.status_code == status.HTTP_404_NOT_FOUND
# A scan belonging to another tenant is not visible (RLS), so it must
# return 404 rather than leaking the fallback provider's check IDs.
other_tenant = Tenant.objects.create(name="Other Compliance Tenant")
foreign_provider = Provider.objects.create(
provider="gcp",
uid="foreign-gcp-test",
alias="foreign_gcp",
tenant_id=other_tenant.id,
)
foreign_scan = Scan.objects.create(
name="foreign scan",
provider=foreign_provider,
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
tenant_id=other_tenant.id,
started_at=now,
completed_at=now,
)
foreign_response = request_attributes(foreign_scan.id)
assert foreign_response.status_code == status.HTTP_404_NOT_FOUND
def test_compliance_overview_attributes_scan_scoped_by_provider_group(
self,
authenticated_client_no_permissions_rbac,
providers_fixture,
):
# A user with limited visibility (no UNLIMITED_VISIBILITY) must only be
# able to resolve scans for providers in its provider groups. Tenant RLS
# alone is not enough here: both scans belong to the same tenant, so the
# endpoint has to scope the scan lookup by provider group, otherwise a
# restricted user could read another provider's compliance metadata.
client = authenticated_client_no_permissions_rbac
limited_user = client.user
membership = Membership.objects.filter(user=limited_user).first()
tenant = membership.tenant
allowed_provider = providers_fixture[2]
denied_provider = providers_fixture[4]
assert allowed_provider.provider == Provider.ProviderChoices.GCP.value
assert denied_provider.provider == Provider.ProviderChoices.AZURE.value
provider_group = ProviderGroup.objects.create(
name="limited-compliance-group",
tenant_id=tenant.id,
)
ProviderGroupMembership.objects.create(
tenant_id=tenant.id,
provider_group=provider_group,
provider=allowed_provider,
)
RoleProviderGroupRelationship.objects.create(
tenant_id=tenant.id,
role=limited_user.roles.first(),
provider_group=provider_group,
)
now = datetime.now(timezone.utc)
allowed_scan = Scan.objects.create(
name="allowed scan",
provider=allowed_provider,
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
tenant_id=tenant.id,
started_at=now,
completed_at=now,
)
denied_scan = Scan.objects.create(
name="denied scan",
provider=denied_provider,
trigger=Scan.TriggerChoices.MANUAL,
state=StateChoices.COMPLETED,
tenant_id=tenant.id,
started_at=now,
completed_at=now,
)
def request_attributes(scan_id):
return client.get(
reverse("complianceoverview-attributes"),
{
"filter[compliance_id]": "csa_ccm_4.0",
"filter[scan_id]": str(scan_id),
},
)
# The scan in the user's provider group resolves normally.
assert request_attributes(allowed_scan.id).status_code == status.HTTP_200_OK
# The scan outside the user's provider group is invisible, so it fails
# closed with 404 instead of leaking the other provider's check IDs.
assert (
request_attributes(denied_scan.id).status_code == status.HTTP_404_NOT_FOUND
)
def test_compliance_overview_attributes_missing_compliance_id(
self, authenticated_client
):
@@ -12663,7 +12845,9 @@ class TestTenantFinishACSView:
)
request = RequestFactory().get(
reverse("saml_finish_acs", kwargs={"organization_slug": "testtenant"})
reverse(
"saml_finish_acs", kwargs={"organization_slug": saml_setup["domain"]}
)
)
request.user = user
request.session = {}
@@ -12683,18 +12867,23 @@ class TestTenantFinishACSView:
patch("api.models.User.objects.get") as mock_user_get,
):
mock_get_app_or_404.return_value = MagicMock(
provider="saml", client_id="testtenant", name="Test App", settings={}
provider="saml",
client_id=saml_setup["domain"],
name="Test App",
settings={},
)
mock_sa_get.return_value = social_account
mock_socialapp_get.return_value = MagicMock(provider_id="saml")
mock_saml_domain_get.return_value = SimpleNamespace(
tenant_id=tenants_fixture[0].id
)
mock_saml_config_get.return_value = MagicMock()
mock_saml_config_get.return_value = SimpleNamespace(
email_domain=saml_setup["domain"], tenant=tenants_fixture[0]
)
mock_user_get.return_value = user
view = TenantFinishACSView.as_view()
response = view(request, organization_slug="testtenant")
response = view(request, organization_slug=saml_setup["domain"])
assert response.status_code == 302
@@ -12733,6 +12922,81 @@ class TestTenantFinishACSView:
user.company_name = original_company
user.save()
def test_dispatch_rejects_assertion_email_domain_that_differs_from_slug(
self, tenants_fixture, saml_setup, monkeypatch
):
monkeypatch.setenv("AUTH_URL", "http://localhost")
monkeypatch.setenv("SAML_SSO_CALLBACK_URL", "http://localhost/sso-complete")
victim_tenant = tenants_fixture[0]
attacker_tenant = tenants_fixture[1]
attacker_domain = "attacker.com"
SAMLConfiguration.objects.using(MainRouter.admin_db).create(
email_domain=attacker_domain,
metadata_xml="""<?xml version='1.0' encoding='UTF-8'?>
<md:EntityDescriptor entityID='ATTACKER' xmlns:md='urn:oasis:names:tc:SAML:2.0:metadata'>
<md:IDPSSODescriptor WantAuthnRequestsSigned='false' protocolSupportEnumeration='urn:oasis:names:tc:SAML:2.0:protocol'>
<md:KeyDescriptor use='signing'>
<ds:KeyInfo xmlns:ds='http://www.w3.org/2000/09/xmldsig#'>
<ds:X509Data>
<ds:X509Certificate>TEST</ds:X509Certificate>
</ds:X509Data>
</ds:KeyInfo>
</md:KeyDescriptor>
<md:SingleSignOnService Binding='urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST' Location='https://ATTACKER/sso/saml'/>
</md:IDPSSODescriptor>
</md:EntityDescriptor>
""",
tenant=attacker_tenant,
)
user = User.objects.using(MainRouter.admin_db).create(
email=f"intruder@{saml_setup['domain']}", name="Intruder"
)
social_account = SocialAccount(
user=user,
provider="ATTACKER",
extra_data={
"firstName": ["Mallory"],
"lastName": ["Example"],
},
)
request = RequestFactory().get(
reverse("saml_finish_acs", kwargs={"organization_slug": attacker_domain})
)
request.user = user
request.session = {}
with (
patch(
"allauth.socialaccount.providers.saml.views.get_app_or_404"
) as mock_get_app_or_404,
patch(
"allauth.socialaccount.models.SocialAccount.objects.get"
) as mock_sa_get,
):
mock_get_app_or_404.return_value = MagicMock(
provider="saml",
provider_id="ATTACKER",
client_id=attacker_domain,
name="Attacker App",
settings={},
)
mock_sa_get.return_value = social_account
view = TenantFinishACSView.as_view()
response = view(request, organization_slug=attacker_domain)
assert response.status_code == 302
assert "sso_saml_failed=true" in response.url
assert not (
Membership.objects.using(MainRouter.admin_db)
.filter(user=user, tenant=victim_tenant)
.exists()
)
assert (
not SAMLToken.objects.using(MainRouter.admin_db).filter(user=user).exists()
)
def test_rollback_saml_user_when_error_occurs(self, users_fixture, monkeypatch):
"""Test that a user is properly deleted when created during SAML flow and an error occurs"""
monkeypatch.setenv("AUTH_URL", "http://localhost")
@@ -12802,7 +13066,9 @@ class TestTenantFinishACSView:
)
request = RequestFactory().get(
reverse("saml_finish_acs", kwargs={"organization_slug": "testtenant"})
reverse(
"saml_finish_acs", kwargs={"organization_slug": saml_setup["domain"]}
)
)
request.user = user
request.session = {}
@@ -12822,16 +13088,21 @@ class TestTenantFinishACSView:
patch("api.models.User.objects.get") as mock_user_get,
):
mock_get_app_or_404.return_value = MagicMock(
provider="saml", client_id="testtenant", name="Test App", settings={}
provider="saml",
client_id=saml_setup["domain"],
name="Test App",
settings={},
)
mock_sa_get.return_value = social_account
mock_socialapp_get.return_value = MagicMock(provider_id="saml")
mock_saml_domain_get.return_value = SimpleNamespace(tenant_id=tenant.id)
mock_saml_config_get.return_value = MagicMock()
mock_saml_config_get.return_value = SimpleNamespace(
email_domain=saml_setup["domain"], tenant=tenant
)
mock_user_get.return_value = user
view = TenantFinishACSView.as_view()
response = view(request, organization_slug="testtenant")
response = view(request, organization_slug=saml_setup["domain"])
assert response.status_code == 302
@@ -12889,7 +13160,9 @@ class TestTenantFinishACSView:
)
request = RequestFactory().get(
reverse("saml_finish_acs", kwargs={"organization_slug": "testtenant"})
reverse(
"saml_finish_acs", kwargs={"organization_slug": saml_setup["domain"]}
)
)
request.user = user
request.session = {}
@@ -12909,16 +13182,21 @@ class TestTenantFinishACSView:
patch("api.models.User.objects.get") as mock_user_get,
):
mock_get_app_or_404.return_value = MagicMock(
provider="saml", client_id="testtenant", name="Test App", settings={}
provider="saml",
client_id=saml_setup["domain"],
name="Test App",
settings={},
)
mock_sa_get.return_value = social_account
mock_socialapp_get.return_value = MagicMock(provider_id="saml")
mock_saml_domain_get.return_value = SimpleNamespace(tenant_id=tenant.id)
mock_saml_config_get.return_value = MagicMock()
mock_saml_config_get.return_value = SimpleNamespace(
email_domain=saml_setup["domain"], tenant=tenant
)
mock_user_get.return_value = user
view = TenantFinishACSView.as_view()
response = view(request, organization_slug="testtenant")
response = view(request, organization_slug=saml_setup["domain"])
assert response.status_code == 302
@@ -12973,7 +13251,9 @@ class TestTenantFinishACSView:
)
request = RequestFactory().get(
reverse("saml_finish_acs", kwargs={"organization_slug": "testtenant"})
reverse(
"saml_finish_acs", kwargs={"organization_slug": saml_setup["domain"]}
)
)
request.user = user
request.session = {}
@@ -12993,16 +13273,21 @@ class TestTenantFinishACSView:
patch("api.models.User.objects.get") as mock_user_get,
):
mock_get_app_or_404.return_value = MagicMock(
provider="saml", client_id="testtenant", name="Test App", settings={}
provider="saml",
client_id=saml_setup["domain"],
name="Test App",
settings={},
)
mock_sa_get.return_value = social_account
mock_socialapp_get.return_value = MagicMock(provider_id="saml")
mock_saml_domain_get.return_value = SimpleNamespace(tenant_id=tenant.id)
mock_saml_config_get.return_value = MagicMock()
mock_saml_config_get.return_value = SimpleNamespace(
email_domain=saml_setup["domain"], tenant=tenant
)
mock_user_get.return_value = user
view = TenantFinishACSView.as_view()
response = view(request, organization_slug="testtenant")
response = view(request, organization_slug=saml_setup["domain"])
assert response.status_code == 302
@@ -13056,7 +13341,9 @@ class TestTenantFinishACSView:
)
request = RequestFactory().get(
reverse("saml_finish_acs", kwargs={"organization_slug": "testtenant"})
reverse(
"saml_finish_acs", kwargs={"organization_slug": saml_setup["domain"]}
)
)
request.user = non_admin_user
request.session = {}
@@ -13076,16 +13363,21 @@ class TestTenantFinishACSView:
patch("api.models.User.objects.get") as mock_user_get,
):
mock_get_app_or_404.return_value = MagicMock(
provider="saml", client_id="testtenant", name="Test App", settings={}
provider="saml",
client_id=saml_setup["domain"],
name="Test App",
settings={},
)
mock_sa_get.return_value = social_account
mock_socialapp_get.return_value = MagicMock(provider_id="saml")
mock_saml_domain_get.return_value = SimpleNamespace(tenant_id=tenant.id)
mock_saml_config_get.return_value = MagicMock()
mock_saml_config_get.return_value = SimpleNamespace(
email_domain=saml_setup["domain"], tenant=tenant
)
mock_user_get.return_value = non_admin_user
view = TenantFinishACSView.as_view()
response = view(request, organization_slug="testtenant")
response = view(request, organization_slug=saml_setup["domain"])
assert response.status_code == 302
+6
View File
@@ -243,6 +243,12 @@ def get_prowler_provider_kwargs(
**prowler_provider_kwargs,
"filter_accounts": [provider.uid],
}
elif provider.provider == Provider.ProviderChoices.ORACLECLOUD.value:
if isinstance(prowler_provider_kwargs.get("region"), str):
prowler_provider_kwargs = {
**prowler_provider_kwargs,
"region": {prowler_provider_kwargs["region"]},
}
elif provider.provider == Provider.ProviderChoices.OPENSTACK.value:
# clouds_yaml_content, clouds_yaml_cloud and provider_id are validated
# in the provider itself, so it's not needed here.
+69 -9
View File
@@ -30,6 +30,7 @@ from dj_rest_auth.registration.views import SocialLoginView
from django.conf import settings as django_settings
from django.contrib.postgres.aggregates import ArrayAgg, BoolAnd, StringAgg
from django.contrib.postgres.search import SearchQuery
from django.core.exceptions import ValidationError as DjangoValidationError
from django.db import transaction
from django.db.models import (
BooleanField,
@@ -760,7 +761,10 @@ class TenantFinishACSView(FinishACSView):
try:
check = SAMLDomainIndex.objects.get(email_domain=organization_slug)
with rls_transaction(str(check.tenant_id)):
SAMLConfiguration.objects.get(tenant_id=str(check.tenant_id))
saml_config = SAMLConfiguration.objects.select_related("tenant").get(
tenant_id=str(check.tenant_id)
)
tenant = saml_config.tenant
social_app = SocialApp.objects.get(
provider="saml", client_id=organization_slug
)
@@ -780,6 +784,15 @@ class TenantFinishACSView(FinishACSView):
callback_url = env.str("AUTH_URL")
return redirect(f"{callback_url}?sso_saml_failed=true")
requested_domain = organization_slug.lower()
configured_domain = saml_config.email_domain.lower()
email_domain = user.email.rsplit("@", 1)[-1].lower()
if configured_domain != requested_domain or email_domain != configured_domain:
logger.error("SAML email domain does not match requested organization")
self._rollback_saml_user(request)
callback_url = env.str("AUTH_URL")
return redirect(f"{callback_url}?sso_saml_failed=true")
extra = social_account.extra_data
user.first_name = (
extra.get("firstName", [""])[0] if extra.get("firstName") else ""
@@ -793,13 +806,6 @@ class TenantFinishACSView(FinishACSView):
user.name = "N/A"
user.save()
email_domain = user.email.split("@")[-1]
tenant = (
SAMLConfiguration.objects.using(MainRouter.admin_db)
.get(email_domain=email_domain)
.tenant
)
role_name = (
extra.get("userType", ["no_permissions"])[0].strip()
if extra.get("userType")
@@ -4644,6 +4650,16 @@ class RoleProviderGroupRelationshipView(RelationshipView, BaseRLSViewSet):
location=OpenApiParameter.QUERY,
description="Compliance framework ID to get attributes for.",
),
OpenApiParameter(
name="filter[scan_id]",
required=False,
type=OpenApiTypes.UUID,
location=OpenApiParameter.QUERY,
description="Scan ID used to resolve the provider for "
"multi-provider universal frameworks (e.g. CSA CCM), so "
"the returned check IDs match the scan's provider. When omitted, "
"the first provider that declares the framework is used.",
),
],
responses={
200: OpenApiResponse(
@@ -5084,7 +5100,51 @@ class ComplianceOverviewViewSet(BaseRLSViewSet, TaskManagementMixin):
provider_type = None
# If we couldn't determine from database, try each provider type
# When a scan is provided, resolve the provider from it. Multi-provider
# universal frameworks (e.g. CSA CCM) share a single compliance_id
# across providers but expose different checks per provider, so the
# metadata (and therefore the check IDs the UI uses to fetch findings)
# must be returned for the scan's provider. Without this, the endpoint
# falls back to the first provider that declares the framework and
# returns its check IDs, leaving azure/gcp/... requirements with no
# matching findings.
scan_id = request.query_params.get("filter[scan_id]")
if "filter[scan_id]" in request.query_params:
# An explicit scan_id is authoritative: fail closed instead of
# falling back to another provider. Otherwise an invalid, empty
# (filter[scan_id]=) or inaccessible scan would silently return the
# first provider's check IDs, recreating the multi-provider mismatch
# this endpoint fixes.
if not scan_id:
raise NotFound(detail=f"Scan '{scan_id}' not found.")
# Tenant isolation is already enforced by Postgres RLS on the
# connection (see BaseRLSViewSet). Scope the lookup by provider
# group as well so a user with limited visibility can't resolve
# another provider's scan and read its compliance metadata, mirroring
# the RBAC scoping get_queryset() applies to the rest of the ViewSet.
role = get_role(request.user, request.tenant_id)
if getattr(role, Permissions.UNLIMITED_VISIBILITY.value, False):
scan_queryset = Scan.objects.filter(tenant_id=request.tenant_id)
else:
scan_queryset = Scan.objects.filter(provider__in=get_providers(role))
try:
scan = scan_queryset.select_related("provider").get(id=scan_id)
except (Scan.DoesNotExist, DjangoValidationError, ValueError):
raise NotFound(detail=f"Scan '{scan_id}' not found.")
provider_type = scan.provider.provider
if compliance_id not in get_compliance_frameworks(provider_type):
raise NotFound(
detail=(
f"Compliance framework '{compliance_id}' is not "
f"available for scan '{scan_id}'."
)
)
# Fall back to the first provider that declares the framework. Keeps the
# endpoint working for provider-agnostic callers that omit the scan.
if not provider_type:
for pt in Provider.ProviderChoices.values:
if compliance_id in get_compliance_frameworks(pt):
+161 -137
View File
@@ -5,6 +5,7 @@ import re
import time
import uuid
from collections import defaultdict
from collections.abc import Iterable
from datetime import datetime, timezone
from typing import Any
@@ -22,7 +23,6 @@ from django.db.models import (
Max,
Min,
OuterRef,
Prefetch,
Q,
Sum,
When,
@@ -357,68 +357,71 @@ def _copy_compliance_requirement_rows(
def _persist_compliance_requirement_rows(
tenant_id: str, rows: list[dict[str, Any]], batch_size: int = 10000
) -> None:
tenant_id: str, rows: Iterable[dict[str, Any]], batch_size: int = 10000
) -> int:
"""Persist compliance requirement rows using batched COPY with ORM fallback.
Splits large row sets into batches to reduce lock duration and improve concurrency.
``rows`` is consumed lazily in batches, so peak memory stays at ~``batch_size``
rows instead of the full set. A batch that fails COPY falls back to an ORM
``bulk_create`` of just that batch.
Args:
tenant_id: Target tenant UUID.
rows: Precomputed row dictionaries that reflect the compliance
overview state for a scan.
rows: Iterable of row dictionaries reflecting the compliance overview
state for a scan.
batch_size: Number of rows per COPY batch (default: 10000).
Returns:
int: total number of rows persisted.
"""
if not rows:
return
total_rows = len(rows)
total_batches = (total_rows + batch_size - 1) // batch_size
try:
# Process rows in batches to reduce lock duration
for batch_num in range(total_batches):
start_idx = batch_num * batch_size
end_idx = min(start_idx + batch_size, total_rows)
batch = rows[start_idx:end_idx]
total_rows = 0
batch_num = 0
for batch, _is_last in batched(rows, batch_size):
if not batch:
continue
batch_num += 1
try:
_copy_compliance_requirement_rows(tenant_id, batch)
except Exception as error:
logger.exception(
f"COPY bulk insert for compliance requirements batch {batch_num} "
"failed; falling back to ORM bulk_create for this batch",
exc_info=error,
)
fallback_objects = [
ComplianceRequirementOverview(
id=row["id"],
tenant_id=row["tenant_id"],
inserted_at=row["inserted_at"],
compliance_id=row["compliance_id"],
framework=row["framework"],
version=row["version"],
description=row["description"],
region=row["region"],
requirement_id=row["requirement_id"],
requirement_status=row["requirement_status"],
passed_checks=row["passed_checks"],
failed_checks=row["failed_checks"],
total_checks=row["total_checks"],
passed_findings=row.get("passed_findings", 0),
total_findings=row.get("total_findings", 0),
scan_id=row["scan_id"],
)
for row in batch
]
with rls_transaction(tenant_id):
ComplianceRequirementOverview.objects.bulk_create(
fallback_objects, batch_size=500
)
logger.info(
f"Compliance COPY batch {batch_num + 1}/{total_batches}: "
f"inserted {len(batch)} rows ({start_idx + len(batch)}/{total_rows} total)"
)
except Exception as error:
logger.exception(
"COPY bulk insert for compliance requirements failed; falling back to ORM bulk_create",
exc_info=error,
total_rows += len(batch)
logger.info(
f"Compliance COPY batch {batch_num}: inserted {len(batch)} rows "
f"({total_rows} total)"
)
# Fallback: use ORM bulk_create for all remaining rows
fallback_objects = [
ComplianceRequirementOverview(
id=row["id"],
tenant_id=row["tenant_id"],
inserted_at=row["inserted_at"],
compliance_id=row["compliance_id"],
framework=row["framework"],
version=row["version"],
description=row["description"],
region=row["region"],
requirement_id=row["requirement_id"],
requirement_status=row["requirement_status"],
passed_checks=row["passed_checks"],
failed_checks=row["failed_checks"],
total_checks=row["total_checks"],
passed_findings=row.get("passed_findings", 0),
total_findings=row.get("total_findings", 0),
scan_id=row["scan_id"],
)
for row in rows
]
with rls_transaction(tenant_id):
ComplianceRequirementOverview.objects.bulk_create(
fallback_objects, batch_size=500
)
return total_rows
def _create_compliance_summaries(
@@ -1445,9 +1448,13 @@ def _aggregate_findings_by_region(
tenant_id: str, scan_id: str, modeled_threatscore_compliance_id: str
) -> tuple[dict, dict]:
"""
Aggregate findings by region using optimized ORM queries.
Aggregate findings by region using streaming, column-scoped ORM reads.
Replaces nested Python loops with efficient queries and aggregation.
Reads only the consumed columns as tuples via ``values_list`` and streams
them with ``.iterator()``, using the denormalized ``resource_regions`` array
instead of ``prefetch_related("resources")``. ``resource_regions`` mirrors the
regions of a finding's related resources, so it yields the same per-region
tally without joining the resource table.
Args:
tenant_id: Tenant UUID
@@ -1459,12 +1466,12 @@ def _aggregate_findings_by_region(
- check_status_by_region: {region: {check_id: status}}
- findings_count_by_compliance: {region: {normalized_id: {requirement_id: {total, pass}}}}
"""
check_status_by_region = {}
findings_count_by_compliance = {}
check_status_by_region: dict = {}
findings_count_by_compliance: dict = {}
normalized_id = re.sub(r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower())
with rls_transaction(tenant_id, using=READ_REPLICA_ALIAS):
# Fetch only PASS/FAIL findings (optimized query reduces data transfer)
# Other statuses are not needed for check_status or ThreatScore calculation
findings = (
Finding.all_objects.filter(
tenant_id=tenant_id,
@@ -1472,42 +1479,28 @@ def _aggregate_findings_by_region(
muted=False,
status__in=["PASS", "FAIL"],
)
.only("id", "check_id", "status", "compliance")
.prefetch_related(
Prefetch(
"resources",
queryset=Resource.objects.only("id", "region"),
to_attr="small_resources",
)
.values_list("check_id", "status", "resource_regions", "compliance")
.iterator(chunk_size=DJANGO_FINDINGS_BATCH_SIZE)
)
for check_id, status, resource_regions, compliance in findings:
threatscore_requirements = (compliance or {}).get(
modeled_threatscore_compliance_id
)
)
# Process findings in a single pass (more efficient than original nested loops)
normalized_id = re.sub(
r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower()
)
for finding in findings:
status = finding.status
for resource in finding.small_resources:
region = resource.region
# Aggregate check status by region
current_status = check_status_by_region.setdefault(region, {})
for region in resource_regions or ():
# Priority: FAIL > any other status
if current_status.get(finding.check_id) != "FAIL":
current_status[finding.check_id] = status
current_status = check_status_by_region.setdefault(region, {})
if current_status.get(check_id) != "FAIL":
current_status[check_id] = status
# Aggregate ThreatScore compliance counts
if modeled_threatscore_compliance_id in (finding.compliance or {}):
if threatscore_requirements:
compliance_key = findings_count_by_compliance.setdefault(
region, {}
).setdefault(normalized_id, {})
for requirement_id in finding.compliance[
modeled_threatscore_compliance_id
]:
for requirement_id in threatscore_requirements:
requirement_stats = compliance_key.setdefault(
requirement_id, {"total": 0, "pass": 0}
)
@@ -1554,8 +1547,8 @@ def create_compliance_requirements(tenant_id: str, scan_id: str):
(compliance_id, requirement_id)
)
compliance_requirement_rows: list[dict[str, Any]] = []
regions = []
requirements_created = 0
requirement_statuses = defaultdict(
lambda: {"fail_count": 0, "pass_count": 0, "total_count": 0}
)
@@ -1595,44 +1588,93 @@ def create_compliance_requirements(tenant_id: str, scan_id: str):
else:
requirement_stats["failed_checks"] += 1
# Prepare compliance requirement rows and compute summaries in single pass
utc_datetime_now = datetime.now(tz=timezone.utc)
# Pre-compute shared strings (optimization: reduces string conversions)
tenant_id_str = str(tenant_id)
scan_id_str = str(scan_instance.id)
for region in regions:
region_stats = region_requirement_stats.get(region, {})
for compliance_id, compliance in compliance_template.items():
modeled_compliance_id = _normalized_compliance_key(
compliance["framework"], compliance["version"]
# Per-framework constants that don't depend on the region.
compliance_plan = []
for compliance_id, compliance in compliance_template.items():
modeled_compliance_id = _normalized_compliance_key(
compliance["framework"], compliance["version"]
)
framework = compliance["framework"]
version = compliance["version"] or ""
requirements = [
(
requirement_id,
requirement.get("description") or "",
len(requirement["checks"]),
)
compliance_stats = region_stats.get(compliance_id, {})
# Create an overview record for each requirement within each compliance framework
for requirement_id, requirement in compliance[
"requirements"
].items():
stats = compliance_stats.get(requirement_id)
passed_checks = stats["passed_checks"] if stats else 0
failed_checks = stats["failed_checks"] if stats else 0
total_checks = len(requirement["checks"])
if total_checks == 0:
requirement_status = "MANUAL"
elif failed_checks > 0:
requirement_status = "FAIL"
else:
requirement_status = "PASS"
].items()
]
compliance_plan.append(
(
compliance_id,
framework,
version,
modeled_compliance_id,
requirements,
)
)
compliance_requirement_rows.append(
{
# Yield rows lazily (consumed batch-by-batch by COPY) so peak memory
# stays bounded; tally requirement_statuses in the same pass.
def _iter_compliance_requirement_rows():
for region in regions:
region_stats = region_requirement_stats.get(region, {})
region_findings = findings_count_by_compliance.get(region, {})
for (
compliance_id,
framework,
version,
modeled_compliance_id,
requirements,
) in compliance_plan:
compliance_stats = region_stats.get(compliance_id, {})
compliance_findings = region_findings.get(
modeled_compliance_id, {}
)
for requirement_id, description, total_checks in requirements:
stats = compliance_stats.get(requirement_id)
if stats:
passed_checks = stats["passed_checks"]
failed_checks = stats["failed_checks"]
else:
passed_checks = 0
failed_checks = 0
if total_checks == 0:
requirement_status = "MANUAL"
elif failed_checks > 0:
requirement_status = "FAIL"
else:
requirement_status = "PASS"
finding_counts = compliance_findings.get(requirement_id)
if finding_counts:
passed_findings = finding_counts.get("pass", 0)
total_findings = finding_counts.get("total", 0)
else:
passed_findings = 0
total_findings = 0
key = (compliance_id, requirement_id)
requirement_statuses[key]["total_count"] += 1
if requirement_status == "FAIL":
requirement_statuses[key]["fail_count"] += 1
elif requirement_status == "PASS":
requirement_statuses[key]["pass_count"] += 1
yield {
"id": uuid.uuid4(),
"tenant_id": tenant_id_str,
"inserted_at": utc_datetime_now,
"compliance_id": compliance_id,
"framework": compliance["framework"],
"version": compliance["version"] or "",
"description": requirement.get("description") or "",
"framework": framework,
"version": version,
"description": description,
"region": region,
"requirement_id": requirement_id,
"requirement_status": requirement_status,
@@ -1640,41 +1682,23 @@ def create_compliance_requirements(tenant_id: str, scan_id: str):
"failed_checks": failed_checks,
"total_checks": total_checks,
"scan_id": scan_id_str,
"passed_findings": findings_count_by_compliance.get(
region, {}
)
.get(modeled_compliance_id, {})
.get(requirement_id, {})
.get("pass", 0),
"total_findings": findings_count_by_compliance.get(
region, {}
)
.get(modeled_compliance_id, {})
.get(requirement_id, {})
.get("total", 0),
"passed_findings": passed_findings,
"total_findings": total_findings,
}
)
# Update summary tracking (single-pass optimization)
key = (compliance_id, requirement_id)
requirement_statuses[key]["total_count"] += 1
if requirement_status == "FAIL":
requirement_statuses[key]["fail_count"] += 1
elif requirement_status == "PASS":
requirement_statuses[key]["pass_count"] += 1
# Idempotent re-run: COPY can't ON CONFLICT, so clear this scan's rows first.
# Idempotent re-run: clear this scan's rows before re-inserting.
with rls_transaction(tenant_id):
ComplianceRequirementOverview.objects.filter(scan_id=scan_id).delete()
# Bulk create requirement records using PostgreSQL COPY
_persist_compliance_requirement_rows(tenant_id, compliance_requirement_rows)
requirements_created = _persist_compliance_requirement_rows(
tenant_id, _iter_compliance_requirement_rows()
)
# Create pre-aggregated summaries for fast compliance overview lookups
_create_compliance_summaries(tenant_id, scan_id, requirement_statuses)
return {
"requirements_created": len(compliance_requirement_rows),
"requirements_created": requirements_created,
"regions_processed": list(regions),
"compliance_frameworks": (
list(compliance_template.keys()) if regions else []
+173 -74
View File
@@ -3674,19 +3674,19 @@ class TestAggregateFindingsByRegion:
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
# Mock findings with resources
mock_finding1 = MagicMock()
mock_finding1.check_id = "check1"
mock_finding1.status = "FAIL"
mock_finding1.compliance = {modeled_threatscore_compliance_id: ["req1", "req2"]}
mock_resource1 = MagicMock()
mock_resource1.region = "us-east-1"
mock_finding1.small_resources = [mock_resource1]
# (check_id, status, resource_regions, compliance) tuples
finding_rows = [
(
"check1",
"FAIL",
["us-east-1"],
{modeled_threatscore_compliance_id: ["req1", "req2"]},
)
]
mock_queryset = MagicMock()
mock_queryset.only.return_value = mock_queryset
mock_queryset.prefetch_related.return_value = [mock_finding1]
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = finding_rows
ctx = MagicMock()
ctx.__enter__.return_value = None
@@ -3700,6 +3700,12 @@ class TestAggregateFindingsByRegion:
)
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
# Verify structure of check_status_by_region
assert isinstance(check_status_by_region, dict)
assert "us-east-1" in check_status_by_region
@@ -3719,27 +3725,15 @@ class TestAggregateFindingsByRegion:
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
# First finding with PASS status
mock_finding1 = MagicMock()
mock_finding1.check_id = "check1"
mock_finding1.status = "PASS"
mock_finding1.compliance = {}
mock_resource1 = MagicMock()
mock_resource1.region = "us-east-1"
mock_finding1.small_resources = [mock_resource1]
# Second finding with FAIL status for same check/region
mock_finding2 = MagicMock()
mock_finding2.check_id = "check1"
mock_finding2.status = "FAIL"
mock_finding2.compliance = {}
mock_resource2 = MagicMock()
mock_resource2.region = "us-east-1"
mock_finding2.small_resources = [mock_resource2]
# Same check/region: PASS first, then FAIL — FAIL must win
finding_rows = [
("check1", "PASS", ["us-east-1"], {}),
("check1", "FAIL", ["us-east-1"], {}),
]
mock_queryset = MagicMock()
mock_queryset.only.return_value = mock_queryset
mock_queryset.prefetch_related.return_value = [mock_finding1, mock_finding2]
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = finding_rows
ctx = MagicMock()
ctx.__enter__.return_value = None
@@ -3751,6 +3745,12 @@ class TestAggregateFindingsByRegion:
tenant_id, scan_id, modeled_threatscore_compliance_id
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
# FAIL should override PASS
assert check_status_by_region["us-east-1"]["check1"] == "FAIL"
@@ -3765,8 +3765,8 @@ class TestAggregateFindingsByRegion:
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
mock_queryset = MagicMock()
mock_queryset.only.return_value = mock_queryset
mock_queryset.prefetch_related.return_value = []
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = []
ctx = MagicMock()
ctx.__enter__.return_value = None
@@ -3778,6 +3778,12 @@ class TestAggregateFindingsByRegion:
tenant_id, scan_id, modeled_threatscore_compliance_id
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
# Verify filter was called with muted=False
mock_findings_filter.assert_called_once_with(
tenant_id=tenant_id,
@@ -3796,27 +3802,25 @@ class TestAggregateFindingsByRegion:
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
# Finding with PASS status
mock_finding1 = MagicMock()
mock_finding1.check_id = "check1"
mock_finding1.status = "PASS"
mock_finding1.compliance = {modeled_threatscore_compliance_id: ["req1"]}
mock_resource1 = MagicMock()
mock_resource1.region = "us-east-1"
mock_finding1.small_resources = [mock_resource1]
# Finding with FAIL status
mock_finding2 = MagicMock()
mock_finding2.check_id = "check2"
mock_finding2.status = "FAIL"
mock_finding2.compliance = {modeled_threatscore_compliance_id: ["req1"]}
mock_resource2 = MagicMock()
mock_resource2.region = "us-east-1"
mock_finding2.small_resources = [mock_resource2]
# PASS and FAIL findings mapped to the same ThreatScore requirement
finding_rows = [
(
"check1",
"PASS",
["us-east-1"],
{modeled_threatscore_compliance_id: ["req1"]},
),
(
"check2",
"FAIL",
["us-east-1"],
{modeled_threatscore_compliance_id: ["req1"]},
),
]
mock_queryset = MagicMock()
mock_queryset.only.return_value = mock_queryset
mock_queryset.prefetch_related.return_value = [mock_finding1, mock_finding2]
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = finding_rows
ctx = MagicMock()
ctx.__enter__.return_value = None
@@ -3828,6 +3832,12 @@ class TestAggregateFindingsByRegion:
tenant_id, scan_id, modeled_threatscore_compliance_id
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
# Verify compliance counts
normalized_id = re.sub(
r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower()
@@ -3850,27 +3860,15 @@ class TestAggregateFindingsByRegion:
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
# Finding in us-east-1
mock_finding1 = MagicMock()
mock_finding1.check_id = "check1"
mock_finding1.status = "FAIL"
mock_finding1.compliance = {}
mock_resource1 = MagicMock()
mock_resource1.region = "us-east-1"
mock_finding1.small_resources = [mock_resource1]
# Finding in us-west-2
mock_finding2 = MagicMock()
mock_finding2.check_id = "check1"
mock_finding2.status = "PASS"
mock_finding2.compliance = {}
mock_resource2 = MagicMock()
mock_resource2.region = "us-west-2"
mock_finding2.small_resources = [mock_resource2]
# One finding per region
finding_rows = [
("check1", "FAIL", ["us-east-1"], {}),
("check1", "PASS", ["us-west-2"], {}),
]
mock_queryset = MagicMock()
mock_queryset.only.return_value = mock_queryset
mock_queryset.prefetch_related.return_value = [mock_finding1, mock_finding2]
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = finding_rows
ctx = MagicMock()
ctx.__enter__.return_value = None
@@ -3882,6 +3880,12 @@ class TestAggregateFindingsByRegion:
tenant_id, scan_id, modeled_threatscore_compliance_id
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
# Verify both regions are present with correct statuses
assert "us-east-1" in check_status_by_region
assert "us-west-2" in check_status_by_region
@@ -3890,17 +3894,26 @@ class TestAggregateFindingsByRegion:
@patch("tasks.jobs.scan.Finding.all_objects.filter")
@patch("tasks.jobs.scan.rls_transaction")
def test_aggregate_findings_by_region_empty_findings(
def test_aggregate_findings_by_region_multi_region_finding(
self, mock_rls_transaction, mock_findings_filter
):
"""Test with no findings - should return empty dicts."""
"""A finding with multiple resource_regions is tallied in every region."""
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
finding_rows = [
(
"check1",
"FAIL",
["us-east-1", "eu-west-1"],
{modeled_threatscore_compliance_id: ["req1"]},
)
]
mock_queryset = MagicMock()
mock_queryset.only.return_value = mock_queryset
mock_queryset.prefetch_related.return_value = []
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = finding_rows
ctx = MagicMock()
ctx.__enter__.return_value = None
@@ -3914,6 +3927,92 @@ class TestAggregateFindingsByRegion:
)
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
normalized_id = re.sub(
r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower()
)
for region in ("us-east-1", "eu-west-1"):
assert check_status_by_region[region]["check1"] == "FAIL"
req_stats = findings_count_by_compliance[region][normalized_id]["req1"]
assert req_stats == {"total": 1, "pass": 0}
@patch("tasks.jobs.scan.Finding.all_objects.filter")
@patch("tasks.jobs.scan.rls_transaction")
def test_aggregate_findings_by_region_skips_empty_regions(
self, mock_rls_transaction, mock_findings_filter
):
"""A finding with no denormalized regions contributes nothing."""
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
finding_rows = [
("check1", "FAIL", [], {modeled_threatscore_compliance_id: ["req1"]}),
("check2", "PASS", None, {}),
]
mock_queryset = MagicMock()
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = finding_rows
ctx = MagicMock()
ctx.__enter__.return_value = None
ctx.__exit__.return_value = False
mock_rls_transaction.return_value = ctx
mock_findings_filter.return_value = mock_queryset
check_status_by_region, findings_count_by_compliance = (
_aggregate_findings_by_region(
tenant_id, scan_id, modeled_threatscore_compliance_id
)
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
assert check_status_by_region == {}
assert findings_count_by_compliance == {}
@patch("tasks.jobs.scan.Finding.all_objects.filter")
@patch("tasks.jobs.scan.rls_transaction")
def test_aggregate_findings_by_region_empty_findings(
self, mock_rls_transaction, mock_findings_filter
):
"""Test with no findings - should return empty dicts."""
tenant_id = str(uuid.uuid4())
scan_id = str(uuid.uuid4())
modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0"
mock_queryset = MagicMock()
mock_queryset.values_list.return_value = mock_queryset
mock_queryset.iterator.return_value = []
ctx = MagicMock()
ctx.__enter__.return_value = None
ctx.__exit__.return_value = False
mock_rls_transaction.return_value = ctx
mock_findings_filter.return_value = mock_queryset
check_status_by_region, findings_count_by_compliance = (
_aggregate_findings_by_region(
tenant_id, scan_id, modeled_threatscore_compliance_id
)
)
# Streaming query contract: column-scoped values_list + iterator
mock_queryset.values_list.assert_called_once_with(
"check_id", "status", "resource_regions", "compliance"
)
mock_queryset.iterator.assert_called_once()
assert check_status_by_region == {}
assert findings_count_by_compliance == {}
Generated
+70 -4
View File
@@ -4415,8 +4415,8 @@ wheels = [
[[package]]
name = "prowler"
version = "5.27.0"
source = { git = "https://github.com/prowler-cloud/prowler.git?rev=master#0abbb7fc590eaf7de6ed354dd5a217bca261d2b0" }
version = "5.30.0"
source = { git = "https://github.com/prowler-cloud/prowler.git?rev=v5.30#f1d741214a60df17158c3fdc97804fd1fde64f3a" }
dependencies = [
{ name = "alibabacloud-actiontrail20200706" },
{ name = "alibabacloud-credentials" },
@@ -4489,9 +4489,14 @@ dependencies = [
{ name = "pygithub" },
{ name = "python-dateutil" },
{ name = "pytz" },
{ name = "scaleway" },
{ name = "schema" },
{ name = "shodan" },
{ name = "slack-sdk" },
{ name = "stackit-core" },
{ name = "stackit-iaas" },
{ name = "stackit-objectstorage" },
{ name = "stackit-resourcemanager" },
{ name = "tabulate" },
{ name = "tzlocal" },
{ name = "uuid6" },
@@ -4499,7 +4504,7 @@ dependencies = [
[[package]]
name = "prowler-api"
version = "1.31.0"
version = "1.31.3"
source = { virtual = "." }
dependencies = [
{ name = "cartography" },
@@ -4595,7 +4600,7 @@ requires-dist = [
{ name = "matplotlib", specifier = "==3.10.8" },
{ name = "neo4j", specifier = "==6.1.0" },
{ name = "openai", specifier = "==1.109.1" },
{ name = "prowler", git = "https://github.com/prowler-cloud/prowler.git?rev=master" },
{ name = "prowler", git = "https://github.com/prowler-cloud/prowler.git?rev=v5.30" },
{ name = "psycopg2-binary", specifier = "==2.9.9" },
{ name = "pytest-celery", extras = ["redis"], specifier = "==1.3.0" },
{ name = "reportlab", specifier = "==4.4.10" },
@@ -5531,6 +5536,67 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/49/4b/359f28a903c13438ef59ebeee215fb25da53066db67b305c125f1c6d2a25/sqlparse-0.5.5-py3-none-any.whl", hash = "sha256:12a08b3bf3eec877c519589833aed092e2444e68240a3577e8e26148acc7b1ba", size = 46138, upload-time = "2025-12-19T07:17:46.573Z" },
]
[[package]]
name = "stackit-core"
version = "0.2.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cryptography" },
{ name = "pydantic" },
{ name = "pyjwt", extra = ["crypto"] },
{ name = "requests" },
{ name = "urllib3" },
]
sdist = { url = "https://files.pythonhosted.org/packages/24/90/20f9ec7387eec4067cfd3d29055d0e2b5e1e0322c601a7f48125fd8ea35f/stackit_core-0.2.0.tar.gz", hash = "sha256:b8af91877cdb060d6969a303d8cf20bc0b33b345afd91f679c44a987381e2d47", size = 8987, upload-time = "2025-06-12T08:24:45.251Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ab/b4/7b53187ce68956870d864ccb9ccfb68066c9df9de1c9568fd2feb03c4504/stackit_core-0.2.0-py3-none-any.whl", hash = "sha256:04632fc6742790d08ddfcb7f2313e04d1254827397a80250f838a2f81b92645b", size = 10240, upload-time = "2025-06-12T08:24:44.214Z" },
]
[[package]]
name = "stackit-iaas"
version = "1.4.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "pydantic" },
{ name = "python-dateutil" },
{ name = "requests" },
{ name = "stackit-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/52/07/24e65278300d5c3cb19cb1660bff924c80812cf8aad3e715f826bae5aa80/stackit_iaas-1.4.0.tar.gz", hash = "sha256:93523b23442350c7ebefd9129485c4c2a539f694a9c36a0f8edfaba9862057ea", size = 116236, upload-time = "2026-05-13T09:43:15.996Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/08/51/2201164d7bfacf47539888c735f10f6320c188252384957aa1b23121a210/stackit_iaas-1.4.0-py3-none-any.whl", hash = "sha256:3f4a32321b57ac238f73e5d660c6428186b92cc0425c1f0783ba801e377149d9", size = 316588, upload-time = "2026-05-13T09:43:14.943Z" },
]
[[package]]
name = "stackit-objectstorage"
version = "1.4.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "pydantic" },
{ name = "python-dateutil" },
{ name = "requests" },
{ name = "stackit-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/90/80/b790756af40a5c6d979dd688b2557394ac54b594eb4c08edc33157ba890f/stackit_objectstorage-1.4.0.tar.gz", hash = "sha256:4a3812b4de102b199f061706a802909f9e53ae9b0858769d5bd720f814c8bdbe", size = 31814, upload-time = "2026-05-13T09:43:05.027Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/68/f1/ffa8d5e2ec9f818c72a6f045691364eb4e927ee86641993a70882d00205a/stackit_objectstorage-1.4.0-py3-none-any.whl", hash = "sha256:1a3285c6840d95cff591d84fd21803575cb0d010c398e6575ed92987b9c39866", size = 65061, upload-time = "2026-05-13T09:43:04.13Z" },
]
[[package]]
name = "stackit-resourcemanager"
version = "0.8.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "pydantic" },
{ name = "python-dateutil" },
{ name = "requests" },
{ name = "stackit-core" },
]
sdist = { url = "https://files.pythonhosted.org/packages/23/2d/f458f18e48ed2b1c83df52cff7dbdfd5dd904fb2980ffd9385876e47bbd9/stackit_resourcemanager-0.8.0.tar.gz", hash = "sha256:f44542beab4130857f5a7f465cf02defeef657bdf63c1beeb3102f0ba3c003fe", size = 33943, upload-time = "2026-05-13T09:43:08.667Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c7/9c/38a74d0f7a89b4320f6d2366fb660638bda8860daa08748b12c713d84381/stackit_resourcemanager-0.8.0-py3-none-any.whl", hash = "sha256:dd04bb8353d041a137c4dcba190beabded7acfaff1bc98b218fce20a99389ebc", size = 81288, upload-time = "2026-05-13T09:43:07.81Z" },
]
[[package]]
name = "statsd"
version = "4.0.1"
@@ -138,6 +138,10 @@ To keep permissions focused:
4. Continue through the wizard and finish. No principals need to be granted access in step 3 unless you want other identities to impersonate this account.
<Note>
To use this service account with `--organization-id`, additionally grant `roles/cloudasset.viewer` at the organization node and enable the Cloud Asset API in the service account's host project. See [Scanning a Specific GCP Organization](./organization). Without these, organization-wide scans silently fall back to listing only the projects accessible to the service account.
</Note>
### Step 3: Generate a JSON Key
1. Open the newly created service account, move to the **Keys** tab, and choose **Add key > Create new key**.
+13 -2
View File
@@ -11,8 +11,19 @@ prowler gcp --organization-id organization-id
```
<Warning>
Ensure the credentials used have one of the following roles at the organization level:
Cloud Asset Viewer (`roles/cloudasset.viewer`), or Cloud Asset Owner (`roles/cloudasset.owner`).
Ensure the credentials used have one of the following roles bound **at the organization node** (not at a project): Cloud Asset Viewer (`roles/cloudasset.viewer`) or Cloud Asset Owner (`roles/cloudasset.owner`). The role must be bound directly on the organization so the Cloud Asset API can enumerate projects across the whole hierarchy.
```bash
gcloud organizations add-iam-policy-binding <organization-id> \
--member="serviceAccount:<service-account-email>" \
--role="roles/cloudasset.viewer"
```
The Cloud Asset API (`cloudasset.googleapis.com`) must also be enabled in the project that owns the credentials (the service account's host project, or the quota project for user credentials):
```bash
gcloud services enable cloudasset.googleapis.com --project <credentials-project-id>
```
</Warning>
<Note>
+18
View File
@@ -2,6 +2,24 @@
All notable changes to the **Prowler SDK** are documented in this file.
## [5.30.3] (Prowler v5.30.3)
### 🐞 Fixed
- CLI compliance summary tables no longer undercount findings mapped to multiple sections nor double-count a single finding mapped to several requirements within the same group/split, and the Provider column no longer leaks a value from another framework [(#11567)](https://github.com/prowler-cloud/prowler/pull/11567)
---
## [5.30.2] (Prowler v5.30.2)
### 🐞 Fixed
- GCP `logging_log_metric_filter_and_alert_*` checks now credit org-level aggregated sinks filtered to the Admin Activity audit stream [(#11575)](https://github.com/prowler-cloud/prowler/pull/11575)
- A broken built-in provider no longer aborts the CLI when a different provider was invoked [(#11618)](https://github.com/prowler-cloud/prowler/pull/11618)
- GCP organization scans with `--organization-id` no longer silently fall back to the credentials' host project when the Cloud Asset API call fails [(#11280)](https://github.com/prowler-cloud/prowler/pull/11280)
---
## [5.30.0] (Prowler v5.30.0)
### 🚀 Added
+1 -1
View File
@@ -49,7 +49,7 @@ class _MutableTimestamp:
timestamp = _MutableTimestamp(datetime.today())
timestamp_utc = _MutableTimestamp(datetime.now(timezone.utc))
prowler_version = "5.30.0"
prowler_version = "5.30.3"
html_logo_url = "https://github.com/prowler-cloud/prowler/"
square_logo_img = "https://raw.githubusercontent.com/prowler-cloud/prowler/dc7d2d5aeb92fdf12e8604f42ef6472cd3e8e889/docs/img/prowler-logo-black.png"
aws_logo = "https://user-images.githubusercontent.com/38561120/235953920-3e3fba08-0795-41dc-b480-9bea57db9f2e.png"
+9 -7
View File
@@ -15,6 +15,8 @@ from prowler.lib.check.models import Severity
from prowler.lib.cli.redact import warn_sensitive_argument_values
from prowler.lib.outputs.common import Status
from prowler.providers.common.arguments import (
PROVIDER_ALIASES,
enforce_invoked_provider_loaded,
init_providers_parser,
validate_asff_usage,
validate_provider_arguments,
@@ -166,13 +168,13 @@ Detailed documentation at https://docs.prowler.com
if sys.argv[1].startswith("-"):
sys.argv = self.__set_default_provider__(sys.argv)
# Provider aliases mapping
# Microsoft 365
elif sys.argv[1] == "microsoft365":
sys.argv[1] = "m365"
# Oracle Cloud Infrastructure
elif sys.argv[1] == "oci":
sys.argv[1] = "oraclecloud"
# Provider aliases mapping (single source: arguments.PROVIDER_ALIASES)
elif sys.argv[1] in PROVIDER_ALIASES:
sys.argv[1] = PROVIDER_ALIASES[sys.argv[1]]
# Selective fail-loud here (post argv-normalisation, pre parse_args)
# so the invoked-provider check stays correct under parse(args=...).
enforce_invoked_provider_loaded(self)
# Warn about sensitive flags passed with explicit values
# Snapshot argv before parse_args() which may exit on errors
@@ -22,11 +22,14 @@ def get_asd_essential_eight_table(
pass_count = []
fail_count = []
muted_count = []
section_seen = {}
provider = ""
for index, finding in enumerate(findings):
check = bulk_checks_metadata[finding.check_metadata.CheckID]
check_compliances = check.Compliance
for compliance in check_compliances:
if compliance.Framework == "ASD-Essential-Eight":
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
section = attribute.Section
@@ -36,21 +39,33 @@ def get_asd_essential_eight_table(
"PASS": 0,
"Muted": 0,
}
section_seen[section] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
sections[section]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
sections[section]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-section counts: count each finding once per section
# it belongs to (a finding can map to several sections).
if index not in section_seen[section]:
section_seen[section].add(index)
if finding.muted:
sections[section]["Muted"] += 1
elif finding.status == "FAIL":
sections[section]["FAIL"] += 1
elif finding.status == "PASS":
sections[section]["PASS"] += 1
sections = dict(sorted(sections.items()))
for section in sections:
asd_essential_eight_compliance_table["Provider"].append(compliance.Provider)
asd_essential_eight_compliance_table["Provider"].append(provider)
asd_essential_eight_compliance_table["Section"].append(section)
if sections[section]["FAIL"] > 0:
asd_essential_eight_compliance_table["Status"].append(
+20 -6
View File
@@ -22,33 +22,47 @@ def get_c5_table(
fail_count = []
muted_count = []
sections = {}
section_seen = {}
provider = ""
for index, finding in enumerate(findings):
check = bulk_checks_metadata[finding.check_metadata.CheckID]
check_compliances = check.Compliance
for compliance in check_compliances:
if compliance.Framework == "C5":
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
section = attribute.Section
if section not in sections:
sections[section] = {"FAIL": 0, "PASS": 0, "Muted": 0}
section_seen[section] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
sections[section]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
sections[section]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-section counts: count each finding once per section
# it belongs to (a finding can map to several sections).
if index not in section_seen[section]:
section_seen[section].add(index)
if finding.muted:
sections[section]["Muted"] += 1
elif finding.status == "FAIL":
sections[section]["FAIL"] += 1
elif finding.status == "PASS":
sections[section]["PASS"] += 1
sections = dict(sorted(sections.items()))
for section in sections:
section_table["Provider"].append(compliance.Provider)
section_table["Provider"].append(provider)
section_table["Section"].append(section)
if sections[section]["FAIL"] > 0:
section_table["Status"].append(
+20 -6
View File
@@ -22,33 +22,47 @@ def get_ccc_table(
fail_count = []
muted_count = []
sections = {}
section_seen = {}
provider = ""
for index, finding in enumerate(findings):
check = bulk_checks_metadata[finding.check_metadata.CheckID]
check_compliances = check.Compliance
for compliance in check_compliances:
if compliance.Framework == "CCC":
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
section = attribute.Section
if section not in sections:
sections[section] = {"FAIL": 0, "PASS": 0, "Muted": 0}
section_seen[section] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
sections[section]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
sections[section]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-section counts: count each finding once per section
# it belongs to (a finding can map to several sections).
if index not in section_seen[section]:
section_seen[section].add(index)
if finding.muted:
sections[section]["Muted"] += 1
elif finding.status == "FAIL":
sections[section]["FAIL"] += 1
elif finding.status == "PASS":
sections[section]["PASS"] += 1
sections = dict(sorted(sections.items()))
for section in sections:
section_table["Provider"].append(compliance.Provider)
section_table["Provider"].append(provider)
section_table["Section"].append(section)
if sections[section]["FAIL"] > 0:
section_table["Status"].append(
+25 -3
View File
@@ -13,6 +13,9 @@ def get_cis_table(
compliance_overview: bool,
):
sections = {}
section_muted_seen = {}
section_split_seen = {}
provider = ""
cis_compliance_table = {
"Provider": [],
"Section": [],
@@ -29,6 +32,7 @@ def get_cis_table(
for compliance in check_compliances:
version_in_name = compliance_framework.split("_")[1]
if compliance.Framework == "CIS" and version_in_name in compliance.Version:
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
section = attribute.Section
@@ -40,9 +44,19 @@ def get_cis_table(
"Level 2": {"FAIL": 0, "PASS": 0},
"Muted": 0,
}
section_muted_seen[section] = set()
section_split_seen[section] = {
"Level 1": set(),
"Level 2": set(),
}
if finding.muted:
# Overview total: count each finding once per framework
if index not in muted_count:
muted_count.append(index)
# Per-section Muted: count each finding once per section
# it belongs to (a finding can map to several sections).
if index not in section_muted_seen[section]:
section_muted_seen[section].add(index)
sections[section]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
@@ -50,13 +64,21 @@ def get_cis_table(
elif finding.status == "PASS" and index not in pass_count:
pass_count.append(index)
if "Level 1" in attribute.Profile:
if not finding.muted:
if (
not finding.muted
and index not in section_split_seen[section]["Level 1"]
):
section_split_seen[section]["Level 1"].add(index)
if finding.status == "FAIL":
sections[section]["Level 1"]["FAIL"] += 1
else:
sections[section]["Level 1"]["PASS"] += 1
elif "Level 2" in attribute.Profile:
if not finding.muted:
if (
not finding.muted
and index not in section_split_seen[section]["Level 2"]
):
section_split_seen[section]["Level 2"].add(index)
if finding.status == "FAIL":
sections[section]["Level 2"]["FAIL"] += 1
else:
@@ -65,7 +87,7 @@ def get_cis_table(
# Add results to table
sections = dict(sorted(sections.items()))
for section in sections:
cis_compliance_table["Provider"].append(compliance.Provider)
cis_compliance_table["Provider"].append(provider)
cis_compliance_table["Section"].append(section)
if sections[section]["Level 1"]["FAIL"] > 0:
cis_compliance_table["Level 1"].append(
+15 -6
View File
@@ -13,6 +13,8 @@ def get_ens_table(
compliance_overview: bool,
):
marcos = {}
marco_muted_seen = {}
provider = ""
ens_compliance_table = {
"Proveedor": [],
"Marco/Categoria": [],
@@ -31,6 +33,7 @@ def get_ens_table(
check_compliances = check.Compliance
for compliance in check_compliances:
if compliance.Framework == "ENS":
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
marco_categoria = f"{attribute.Marco}/{attribute.Categoria}"
@@ -44,17 +47,23 @@ def get_ens_table(
"Bajo": 0,
"Muted": 0,
}
marco_muted_seen[marco_categoria] = set()
if finding.muted:
# Overview total: count each finding once per framework
if index not in muted_count:
muted_count.append(index)
# Per-marco Muted: count each finding once per marco
# it belongs to (a finding can map to several marcos).
if index not in marco_muted_seen[marco_categoria]:
marco_muted_seen[marco_categoria].add(index)
marcos[marco_categoria]["Muted"] += 1
else:
if finding.status == "FAIL":
if (
attribute.Tipo != "recomendacion"
and index not in fail_count
):
fail_count.append(index)
if attribute.Tipo != "recomendacion":
if index not in fail_count:
fail_count.append(index)
# Mark every marco the finding belongs to as
# NO CUMPLE, not just the first one seen.
marcos[marco_categoria][
"Estado"
] = f"{Fore.RED}NO CUMPLE{Style.RESET_ALL}"
@@ -71,7 +80,7 @@ def get_ens_table(
# Add results to table
for marco in sorted(marcos):
ens_compliance_table["Proveedor"].append(compliance.Provider)
ens_compliance_table["Proveedor"].append(provider)
ens_compliance_table["Marco/Categoria"].append(marco)
ens_compliance_table["Estado"].append(marcos[marco]["Estado"])
ens_compliance_table["Opcional"].append(
@@ -13,7 +13,9 @@ def get_kisa_ismsp_table(
compliance_overview: bool,
):
sections = {}
section_seen = {}
sections_status = {}
provider = ""
kisa_ismsp_compliance_table = {
"Provider": [],
"Section": [],
@@ -31,6 +33,7 @@ def get_kisa_ismsp_table(
compliance.Framework.startswith("KISA")
and compliance.Version in compliance_framework
):
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
section = attribute.Section
@@ -43,16 +46,28 @@ def get_kisa_ismsp_table(
},
"Muted": 0,
}
section_seen[section] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
sections[section]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
sections[section]["Status"]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-section counts: count each finding once per section
# it belongs to (a finding can map to several sections).
if index not in section_seen[section]:
section_seen[section].add(index)
if finding.muted:
sections[section]["Muted"] += 1
elif finding.status == "FAIL":
sections[section]["Status"]["FAIL"] += 1
elif finding.status == "PASS":
sections[section]["Status"]["PASS"] += 1
# Add results to table
@@ -70,7 +85,7 @@ def get_kisa_ismsp_table(
else:
sections_status[section] = f"{Fore.GREEN}PASS{Style.RESET_ALL}"
for section in sections:
kisa_ismsp_compliance_table["Provider"].append(compliance.Provider)
kisa_ismsp_compliance_table["Provider"].append(provider)
kisa_ismsp_compliance_table["Section"].append(section)
kisa_ismsp_compliance_table["Status"].append(sections_status[section])
kisa_ismsp_compliance_table["Muted"].append(
@@ -13,6 +13,8 @@ def get_mitre_attack_table(
compliance_overview: bool,
):
tactics = {}
tactic_seen = {}
provider = ""
mitre_compliance_table = {
"Provider": [],
"Tactic": [],
@@ -30,27 +32,38 @@ def get_mitre_attack_table(
"MITRE-ATTACK" in compliance.Framework
and compliance.Version in compliance_framework
):
provider = compliance.Provider
for requirement in compliance.Requirements:
for tactic in requirement.Tactics:
if tactic not in tactics:
tactics[tactic] = {"FAIL": 0, "PASS": 0, "Muted": 0}
tactic_seen[tactic] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-tactic counts: count each finding once per tactic
# it belongs to (a finding can map to several tactics).
if index not in tactic_seen[tactic]:
tactic_seen[tactic].add(index)
if finding.muted:
tactics[tactic]["Muted"] += 1
else:
if finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
tactics[tactic]["FAIL"] += 1
elif finding.status == "FAIL":
tactics[tactic]["FAIL"] += 1
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
tactics[tactic]["PASS"] += 1
tactics[tactic]["PASS"] += 1
# Add results to table
tactics = dict(sorted(tactics.items()))
for tactic in tactics:
mitre_compliance_table["Provider"].append(compliance.Provider)
mitre_compliance_table["Provider"].append(provider)
mitre_compliance_table["Tactic"].append(tactic)
if tactics[tactic]["FAIL"] > 0:
mitre_compliance_table["Status"].append(
@@ -22,33 +22,47 @@ def get_okta_idaas_stig_table(
fail_count = []
muted_count = []
sections = {}
section_seen = {}
provider = ""
for index, finding in enumerate(findings):
check = bulk_checks_metadata[finding.check_metadata.CheckID]
check_compliances = check.Compliance
for compliance in check_compliances:
if compliance.Framework == "Okta-IDaaS-STIG":
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
section = attribute.Section
if section not in sections:
sections[section] = {"FAIL": 0, "PASS": 0, "Muted": 0}
section_seen[section] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
sections[section]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
sections[section]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-section counts: count each finding once per section
# it belongs to (a finding can map to several sections).
if index not in section_seen[section]:
section_seen[section].add(index)
if finding.muted:
sections[section]["Muted"] += 1
elif finding.status == "FAIL":
sections[section]["FAIL"] += 1
elif finding.status == "PASS":
sections[section]["PASS"] += 1
sections = dict(sorted(sections.items()))
for section in sections:
section_table["Provider"].append(compliance.Provider)
section_table["Provider"].append(provider)
section_table["Section"].append(section)
if sections[section]["FAIL"] > 0:
section_table["Status"].append(
@@ -24,6 +24,8 @@ def get_prowler_threatscore_table(
fail_count = []
muted_count = []
pillars = {}
pillar_seen = {}
provider = ""
generic_score = 0
max_generic_score = 0
counted_findings_generic = []
@@ -35,6 +37,7 @@ def get_prowler_threatscore_table(
check_compliances = check.Compliance
for compliance in check_compliances:
if compliance.Framework == "ProwlerThreatScore":
provider = compliance.Provider
for requirement in compliance.Requirements:
for attribute in requirement.Attributes:
pillar = attribute.Section
@@ -65,17 +68,28 @@ def get_prowler_threatscore_table(
if pillar not in pillars:
pillars[pillar] = {"FAIL": 0, "PASS": 0, "Muted": 0}
pillar_seen[pillar] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
pillars[pillar]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
pillars[pillar]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-pillar counts: count each finding once per pillar
# it belongs to (a finding can map to several pillars).
if index not in pillar_seen[pillar]:
pillar_seen[pillar].add(index)
if finding.muted:
pillars[pillar]["Muted"] += 1
elif finding.status == "FAIL":
pillars[pillar]["FAIL"] += 1
elif finding.status == "PASS":
pillars[pillar]["PASS"] += 1
# Generic score
@@ -90,18 +104,21 @@ def get_prowler_threatscore_table(
counted_findings_generic.append(index)
no_findings_pillars = []
bulk_compliance = Compliance.get_bulk(provider=compliance.Provider.lower()).get(
compliance_framework
bulk_compliance = (
Compliance.get_bulk(provider=provider.lower()).get(compliance_framework)
if provider
else None
)
for requirement in bulk_compliance.Requirements:
for attribute in requirement.Attributes:
pillar = attribute.Section
if pillar not in pillars.keys() and pillar not in no_findings_pillars:
no_findings_pillars.append(pillar)
if bulk_compliance:
for requirement in bulk_compliance.Requirements:
for attribute in requirement.Attributes:
pillar = attribute.Section
if pillar not in pillars.keys() and pillar not in no_findings_pillars:
no_findings_pillars.append(pillar)
pillars = dict(sorted(pillars.items()))
for pillar in pillars:
pillar_table["Provider"].append(compliance.Provider)
pillar_table["Provider"].append(provider)
pillar_table["Pillar"].append(pillar)
if max_score_per_pillar[pillar] == 0:
pillar_score = 100.0
@@ -127,7 +144,7 @@ def get_prowler_threatscore_table(
)
for pillar in no_findings_pillars:
pillar_table["Provider"].append(compliance.Provider)
pillar_table["Provider"].append(provider)
pillar_table["Pillar"].append(pillar)
pillar_table["Score"].append(f"{Style.BRIGHT}{Fore.GREEN}100%{Style.RESET_ALL}")
pillar_table["Status"].append(f"{Fore.GREEN}PASS{Style.RESET_ALL}")
@@ -163,6 +163,7 @@ def _render_grouped(
"""Grouped mode: one row per group with pass/fail counts."""
check_map = _build_requirement_check_map(framework, provider)
groups = {}
group_seen = {}
pass_count = []
fail_count = []
muted_count = []
@@ -176,17 +177,28 @@ def _render_grouped(
for group_key in _get_group_key(req, group_by):
if group_key not in groups:
groups[group_key] = {"FAIL": 0, "PASS": 0, "Muted": 0}
group_seen[group_key] = set()
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
groups[group_key]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
groups[group_key]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-group counts: count each finding once per group it belongs
# to (a finding can map to several groups via several requirements).
if index not in group_seen[group_key]:
group_seen[group_key].add(index)
if finding.muted:
groups[group_key]["Muted"] += 1
elif finding.status == "FAIL":
groups[group_key]["FAIL"] += 1
elif finding.status == "PASS":
groups[group_key]["PASS"] += 1
if not _print_overview(
@@ -258,6 +270,8 @@ def _render_split(
split_field = split_by.field
split_values = split_by.values
groups = {}
group_muted_seen = {}
group_split_seen = {}
pass_count = []
fail_count = []
muted_count = []
@@ -274,12 +288,19 @@ def _render_split(
sv: {"FAIL": 0, "PASS": 0} for sv in split_values
}
groups[group_key]["Muted"] = 0
group_muted_seen[group_key] = set()
group_split_seen[group_key] = {sv: set() for sv in split_values}
split_val = req.attributes.get(split_field, "")
if finding.muted:
# Overview total: count each finding once per framework
if index not in muted_count:
muted_count.append(index)
# Per-group Muted: count each finding once per group it
# belongs to (a finding can map to several groups).
if index not in group_muted_seen[group_key]:
group_muted_seen[group_key].add(index)
groups[group_key]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
@@ -289,7 +310,8 @@ def _render_split(
for sv in split_values:
if sv in str(split_val):
if not finding.muted:
if index not in group_split_seen[group_key][sv]:
group_split_seen[group_key][sv].add(index)
if finding.status == "FAIL":
groups[group_key][sv]["FAIL"] += 1
else:
@@ -364,6 +386,7 @@ def _render_scored(
risk_field = scoring.risk_field
weight_field = scoring.weight_field
groups = {}
group_seen = {}
pass_count = []
fail_count = []
muted_count = []
@@ -388,6 +411,7 @@ def _render_scored(
if group_key not in groups:
groups[group_key] = {"FAIL": 0, "PASS": 0, "Muted": 0}
group_seen[group_key] = set()
score_per_group[group_key] = 0
max_score_per_group[group_key] = 0
counted_per_group[group_key] = []
@@ -398,16 +422,26 @@ def _render_scored(
max_score_per_group[group_key] += risk * weight
counted_per_group[group_key].append(index)
# Overview totals: count each finding once per framework
if finding.muted:
if index not in muted_count:
muted_count.append(index)
groups[group_key]["Muted"] += 1
else:
if finding.status == "FAIL" and index not in fail_count:
elif finding.status == "FAIL":
if index not in fail_count:
fail_count.append(index)
groups[group_key]["FAIL"] += 1
elif finding.status == "PASS" and index not in pass_count:
elif finding.status == "PASS":
if index not in pass_count:
pass_count.append(index)
# Per-group counts: count each finding once per group it belongs
# to (a finding can map to several groups via several requirements).
if index not in group_seen[group_key]:
group_seen[group_key].add(index)
if finding.muted:
groups[group_key]["Muted"] += 1
elif finding.status == "FAIL":
groups[group_key]["FAIL"] += 1
elif finding.status == "PASS":
groups[group_key]["PASS"] += 1
if index not in counted_generic and not finding.muted:
+79 -19
View File
@@ -10,16 +10,43 @@ provider_arguments_lib_path = "lib.arguments.arguments"
validate_provider_arguments_function = "validate_arguments"
init_provider_arguments_function = "init_parser"
# Kept in sync with parser.py's argv normalisation; both consumers import this.
PROVIDER_ALIASES = {
"microsoft365": "m365",
"oci": "oraclecloud",
}
def _invoked_provider_from_argv(available_providers: Sequence[str]) -> Optional[str]:
"""Return the provider name the user invoked, or None.
Mirrors `ProwlerArgumentParser.parse()` resolution: only inspects
`sys.argv[1]`. Scanning the whole argv would misclassify
`prowler --output-directory stackit` as `stackit`.
"""
available = set(available_providers)
if len(sys.argv) < 2:
return "aws" if "aws" in available else None
first = sys.argv[1]
if first in ("-h", "--help", "-v", "--version"):
return None
if first.startswith("-"):
return "aws" if "aws" in available else None
normalized = PROVIDER_ALIASES.get(first, first)
return normalized if normalized in available else None
def init_providers_parser(self):
"""init_providers_parser calls the provider init_parser function to load all the arguments and flags. Receives a ProwlerArgumentParser object"""
# We need to call the arguments parser for each provider
"""Build the subparser of each available provider.
Built-in load failures are captured silently on
`self._builtin_load_failures`; the warn/exit decision is deferred to
`enforce_invoked_provider_loaded()` because `parse(args=...)` can
override `sys.argv` after this function ran.
"""
self._builtin_load_failures = {}
providers = Provider.get_available_providers()
for provider in providers:
# Discriminate built-in vs external upfront via find_spec, so an
# ImportError from a transitive dependency missing inside a built-in
# arguments module surfaces clearly instead of being silently
# re-routed to the entry-point path (which only has external providers).
if Provider.is_builtin(provider):
try:
getattr(
@@ -28,21 +55,9 @@ def init_providers_parser(self):
),
init_provider_arguments_function,
)(self)
except ImportError as e:
logger.critical(
f"Failed to load arguments for built-in provider '{provider}'. "
f"Missing dependency: {e}. "
f"Ensure all required dependencies are installed."
)
logger.debug("Full traceback:", exc_info=True)
sys.exit(1)
except Exception as error:
logger.critical(
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
sys.exit(1)
self._builtin_load_failures[provider] = error
else:
# External provider — init_parser classmethod via entry point
cls = Provider._load_ep_provider(provider)
if cls and hasattr(cls, "init_parser"):
try:
@@ -53,6 +68,51 @@ def init_providers_parser(self):
)
def enforce_invoked_provider_loaded(self):
"""Apply selective fail-loud over the failures captured at init time.
Called by `ProwlerArgumentParser.parse()` AFTER argv normalisation so
the invoked provider matches what argparse will dispatch to — including
the case where `parse(args=...)` overrode the ambient `sys.argv`.
Invoked + failed → critical + `sys.exit(1)`. Others → warning.
"""
failures = getattr(self, "_builtin_load_failures", {})
if not failures:
return
invoked = _invoked_provider_from_argv(Provider.get_available_providers())
for provider, error in failures.items():
if provider == invoked:
continue
if isinstance(error, ImportError):
logger.warning(
f"Skipping built-in provider '{provider}' due to missing "
f"dependency: {error}. It will be unavailable in this "
f"invocation, but the CLI continues because you invoked a "
f"different provider."
)
else:
logger.warning(
f"Skipping built-in provider '{provider}': "
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
if invoked is None or invoked not in failures:
return
error = failures[invoked]
if isinstance(error, ImportError):
logger.critical(
f"Failed to load arguments for built-in provider '{invoked}'. "
f"Missing dependency: {error}. "
f"Ensure all required dependencies are installed."
)
logger.debug("Full traceback:", exc_info=True)
else:
logger.critical(
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
sys.exit(1)
def validate_provider_arguments(arguments: Namespace) -> tuple[bool, str]:
"""validate_provider_arguments returns {True, "} if the provider arguments passed are valid and can be used together"""
try:
+14 -1
View File
@@ -34,11 +34,17 @@ class GCPBaseException(ProwlerException):
"message": "Error loading Service Account Private Key credentials from dictionary",
"remediation": "Check the dictionary and ensure it contains a Service Account Private Key.",
},
(3011, "GCPGetOrganizationProjectsError"): {
"message": "Error retrieving projects under the organization via the Cloud Asset API",
"remediation": "Ensure the Cloud Asset API is enabled in the credentials' project and that the principal has 'roles/cloudasset.viewer' bound at the organization level. See https://cloud.google.com/asset-inventory/docs/access-control.",
},
}
def __init__(self, code, file=None, original_exception=None, message=None):
provider = "GCP"
error_info = self.GCP_ERROR_CODES.get((code, self.__class__.__name__))
# Copy the catalog entry so a custom message does not mutate the
# class-level GCP_ERROR_CODES shared across exception instances.
error_info = dict(self.GCP_ERROR_CODES.get((code, self.__class__.__name__)))
if message:
error_info["message"] = message
super().__init__(
@@ -104,3 +110,10 @@ class GCPLoadServiceAccountKeyFromDictError(GCPCredentialsError):
super().__init__(
3010, file=file, original_exception=original_exception, message=message
)
class GCPGetOrganizationProjectsError(GCPBaseException):
def __init__(self, file=None, original_exception=None, message=None):
super().__init__(
3011, file=file, original_exception=original_exception, message=message
)
+30 -12
View File
@@ -21,6 +21,8 @@ from prowler.providers.common.models import Audit_Metadata, Connection
from prowler.providers.common.provider import Provider
from prowler.providers.gcp.config import DEFAULT_RETRY_ATTEMPTS
from prowler.providers.gcp.exceptions.exceptions import (
GCPBaseException,
GCPGetOrganizationProjectsError,
GCPInvalidProviderIdError,
GCPLoadADCFromDictError,
GCPLoadServiceAccountKeyFromDictError,
@@ -621,10 +623,7 @@ class GcpProvider(Provider):
credentials_file: str
Returns:
dict[str, GCPProject]
Usage:
>>> GcpProvider.get_projects(credentials=credentials, organization_id=organization_id)
dict of project_id and GCPProject object
"""
projects = {}
try:
@@ -632,7 +631,10 @@ class GcpProvider(Provider):
try:
# Initialize Cloud Asset Inventory API for recursive project retrieval
asset_service = discovery.build(
"cloudasset", "v1", credentials=credentials
"cloudasset",
"v1",
credentials=credentials,
num_retries=DEFAULT_RETRY_ATTEMPTS,
)
# Set the scope to the specified organization and filter for projects
scope = f"organizations/{organization_id}"
@@ -643,7 +645,7 @@ class GcpProvider(Provider):
)
while request is not None:
response = request.execute()
response = request.execute(num_retries=DEFAULT_RETRY_ATTEMPTS)
for asset in response.get("assets", []):
# Extract labels and other project details
@@ -688,13 +690,25 @@ class GcpProvider(Provider):
)
except HttpError as http_error:
if "Cloud Asset API has not been used" in str(http_error):
logger.error(
f"Projects cannot be retrieved from the Organization since Cloud Asset API has not been used before or it is disabled [{http_error.__traceback__.tb_lineno}]. Enable it by visiting https://console.developers.google.com/apis/api/cloudasset.googleapis.com/ then retry."
message = (
"Projects cannot be retrieved from the Organization since the Cloud Asset API "
"has not been used before or it is disabled. Enable it by visiting "
"https://console.developers.google.com/apis/api/cloudasset.googleapis.com/ then retry."
)
else:
logger.error(
f"{http_error.__class__.__name__}[{http_error.__traceback__.tb_lineno}]: {http_error}"
message = (
f"Cloud Asset API call failed while listing projects under organization "
f"'{organization_id}': {http_error}. Ensure the credentials' principal has "
"'roles/cloudasset.viewer' bound at the organization level."
)
logger.critical(
f"{http_error.__class__.__name__}[{http_error.__traceback__.tb_lineno}]: {message}"
)
raise GCPGetOrganizationProjectsError(
file=__file__,
original_exception=http_error,
message=message,
)
else:
try:
# Initialize Cloud Resource Manager API for simple project listing
@@ -781,8 +795,10 @@ class GcpProvider(Provider):
labels={},
lifecycle_state="ACTIVE",
)
# If no projects were able to be accessed via API, add them manually from the credentials file
elif credentials_file:
# If no projects were able to be accessed via API, add them manually from the credentials file.
# Skip this fallback when an organization scan was explicitly requested: silently
# downgrading scope to the service account's home project hides permission errors.
elif credentials_file and not organization_id:
with open(credentials_file, "r", encoding="utf-8") as file:
project_id = json.load(file)["project_id"]
# Handle empty or null project names
@@ -798,6 +814,8 @@ class GcpProvider(Provider):
labels={},
lifecycle_state="ACTIVE",
)
except GCPBaseException as gcp_error:
raise gcp_error
except Exception as error:
logger.critical(
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
@@ -1,9 +1,12 @@
import re
from pydantic.v1 import BaseModel
from prowler.lib.logger import logger
from prowler.providers.gcp.config import DEFAULT_RETRY_ATTEMPTS
from prowler.providers.gcp.gcp_provider import GcpProvider
from prowler.providers.gcp.lib.service.service import GCPService
from prowler.providers.gcp.services.monitoring.monitoring_service import Monitoring
class Logging(GCPService):
@@ -121,9 +124,86 @@ class Metric(BaseModel):
bucket_name: str = ""
# A positive selector of the Admin Activity stream: a ``logName`` predicate
# (``:`` has-substring or ``=`` equals) or a ``log_id()`` call. Written verbose
# so each fragment stays legible; ``(?![a-z_])`` keeps a longer stream name
# (``.../activity_v2``) from impersonating Admin Activity.
_ACTIVITY_SELECTOR = re.compile(
r"""
(?: logName \s* [:=] \s* | log_id \s* \( \s* ) # logName: / logName= / log_id(
["']? [^"'\s)]* # optional quote, then path prefix
cloudaudit\.googleapis\.com/activity (?![a-z_]) # the Admin Activity stream itself
""",
re.IGNORECASE | re.VERBOSE,
)
# The same selector for *any* Cloud Audit stream (activity, data_access,
# system_event, policy, access_transparency, …). Used to strip the OR-combined
# audit clauses so we can prove nothing restrictive is left over.
_CLOUDAUDIT_SELECTOR = re.compile(
r"""
(?: logName \s* [:=] \s* | log_id \s* \( \s* ) # logName: / logName= / log_id(
["']? [^"'\s)]* # optional quote, then path prefix
cloudaudit\.googleapis\.com/[a-z_]+ # any cloudaudit stream
["']? \s* \)? # optional closing quote / paren
""",
re.IGNORECASE | re.VERBOSE,
)
# Operators that exclude or narrow coverage. Any of these means we cannot prove
# the sink delivers the *whole* Admin Activity stream, so it is not credited.
_NEGATION_OR_RESTRICTION = re.compile(
r"""
\bNOT\b # NOT exclusion
| \bAND\b # AND conjunction (restriction)
| != | !: # "!=" / "!:" inequality
| (?:^|[\s(]) -\s* [A-Za-z_] # leading "-" exclusion operator
""",
re.IGNORECASE | re.VERBOSE,
)
def _sink_delivers_activity_logs(sink_filter: str) -> bool:
"""True only when a sink's filter *provably* exports the full Admin Activity
audit stream (or everything).
Crediting flips a child project to PASS on a CIS security control, so the
match is deliberately conservative: a false FAIL is safe, a false PASS is
not. A non-``"all"`` filter is credited only when
1. it positively selects the Admin Activity stream
(``logName:.../activity``, ``logName="...activity"`` or
``log_id("...activity")``);
2. it carries no operator that excludes or narrows the stream — ``NOT`` /
``-`` / ``!=`` (negation) or ``AND`` (restriction); and
3. nothing but ``OR``-combined Cloud Audit selectors remains once those are
stripped — an ``OR`` only widens coverage, but any leftover predicate
(``severity>=ERROR``, ``resource.type=...``) could narrow it.
Sink filters encode the stream URL-encoded (``...%2Factivity``) or as a path
— normalize before matching.
"""
if not sink_filter or sink_filter.strip().lower() == "all":
return True
normalized = sink_filter.replace("%2F", "/").replace("%2f", "/")
# 1. The Admin Activity stream must be positively selected.
if not _ACTIVITY_SELECTOR.search(normalized):
return False
# 2. No operator may exclude or narrow that coverage.
if _NEGATION_OR_RESTRICTION.search(normalized):
return False
# 3. Only OR-combined audit selectors may remain — strip them and the OR
# glue; anything left is a predicate we cannot prove is full-coverage.
remainder = _CLOUDAUDIT_SELECTOR.sub(" ", normalized)
remainder = re.sub(r"\bOR\b|[()\s]", " ", remainder, flags=re.IGNORECASE)
return remainder.strip() == ""
def get_projects_covered_by_aggregated_metric(
logging_client, monitoring_client, metric_filter
):
logging_client: Logging,
monitoring_client: Monitoring,
metric_filter: str,
) -> dict[str, str]:
"""Return {project_id: metric_name} for scanned projects whose logs are routed,
via an organization-level sink with includeChildren=True, to a bucket that holds
a bucket-scoped log metric matching ``metric_filter`` that has an alert policy.
@@ -133,6 +213,10 @@ def get_projects_covered_by_aggregated_metric(
every child project's logs into one bucket, where a single bucket-scoped metric
+ alert covers them all. Without crediting that, those child projects are falsely
failed. Mirrors the org-sink handling already in ``logging_sink_created`` (#11355).
A sink is credited when it exports everything (``filter == "all"``) or when its
filter carries the Admin Activity audit stream — the only stream the CIS metric
filters can match (see ``_sink_delivers_activity_logs``).
"""
# Buckets that hold a matching, alerted, bucket-scoped metric -> metric name.
bucket_to_metric = {}
@@ -155,7 +239,7 @@ def get_projects_covered_by_aggregated_metric(
for sink in logging_client.sinks:
if not getattr(sink, "include_children", False):
continue
if getattr(sink, "filter", "all") != "all":
if not _sink_delivers_activity_logs(getattr(sink, "filter", "all")):
continue
for bucket, metric_name in bucket_to_metric.items():
# sink.destination e.g. "logging.googleapis.com/projects/.../buckets/X";
+1 -1
View File
@@ -124,7 +124,7 @@ maintainers = [{name = "Prowler Engineering", email = "engineering@prowler.com"}
name = "prowler"
readme = "README.md"
requires-python = ">=3.10,<3.13"
version = "5.30.0"
version = "5.30.3"
[project.scripts]
prowler = "prowler.__main__:prowler"
@@ -0,0 +1,132 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.asd_essential_eight.asd_essential_eight import (
get_asd_essential_eight_table,
)
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(provider, sections, framework="ASD-Essential-Eight"):
"""Build a per-check compliance covering the given sections."""
return SimpleNamespace(
Framework=framework,
Provider=provider,
Requirements=[
SimpleNamespace(Attributes=[SimpleNamespace(Section=section)])
for section in sections
],
)
class TestASDEssentialEightTable:
"""Test cases verifying multi-section counting and provider-column attribution for the ASD Essential Eight compliance table."""
def test_multi_section_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several sections must show FAIL(1) in
every section, not just the first one seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_asd_essential_eight_table(
findings,
bulk_metadata,
"asd_essential_eight_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both IAM and Logging must report FAIL(1); before the fix Logging was
# undercounted because the per-section count was gated by the global
# dedup list.
assert captured.out.count("FAIL(1)") == 2
def test_multi_section_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several sections must increase the
per-section Muted count in every section, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A real FAIL is needed so the results table is rendered at all.
_make_finding("check_b", "FAIL"),
]
get_asd_essential_eight_table(
findings,
bulk_metadata,
"asd_essential_eight_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# The muted check belongs to both IAM and Logging, so the Muted column
# must read 1 in both rows.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched ASD-Essential-Eight
compliance, never from a different framework that happens to be the last
entry in the check's compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_asd_essential_eight_table(
findings,
bulk_metadata,
"asd_essential_eight_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "leaked_provider" not in captured.out
@@ -0,0 +1,130 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.c5.c5 import get_c5_table
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(provider, sections, framework="C5"):
"""Build a per-check compliance covering the given sections."""
return SimpleNamespace(
Framework=framework,
Provider=provider,
Requirements=[
SimpleNamespace(Attributes=[SimpleNamespace(Section=section)])
for section in sections
],
)
class TestC5Table:
"""Verify multi-section counting and provider-column attribution for the compliance table."""
def test_multi_section_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several sections must show FAIL(1) in
every section, not just the first one seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_c5_table(
findings,
bulk_metadata,
"c5_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both IAM and Logging must report FAIL(1); before the fix Logging was
# undercounted because the per-section count was gated by the global
# dedup list.
assert captured.out.count("FAIL(1)") == 2
def test_multi_section_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several sections must increase the
per-section Muted count in every section, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A real FAIL is needed so the results table is rendered at all.
_make_finding("check_b", "FAIL"),
]
get_c5_table(
findings,
bulk_metadata,
"c5_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# The muted check belongs to both IAM and Logging, so the Muted column
# must read 1 in both rows.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched C5 compliance, never
from a different framework that happens to be the last entry in the
check's compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_c5_table(
findings,
bulk_metadata,
"c5_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "leaked_provider" not in captured.out
@@ -0,0 +1,130 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.ccc.ccc import get_ccc_table
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(provider, sections, framework="CCC"):
"""Build a per-check compliance covering the given sections."""
return SimpleNamespace(
Framework=framework,
Provider=provider,
Requirements=[
SimpleNamespace(Attributes=[SimpleNamespace(Section=section)])
for section in sections
],
)
class TestCCCTable:
"""Test cases verifying multi-section counting and provider-column attribution for the compliance table."""
def test_multi_section_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several sections must show FAIL(1) in
every section, not just the first one seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_ccc_table(
findings,
bulk_metadata,
"ccc_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both IAM and Logging must report FAIL(1); before the fix Logging was
# undercounted because the per-section count was gated by the global
# dedup list.
assert captured.out.count("FAIL(1)") == 2
def test_multi_section_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several sections must increase the
per-section Muted count in every section, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A real FAIL is needed so the results table is rendered at all.
_make_finding("check_b", "FAIL"),
]
get_ccc_table(
findings,
bulk_metadata,
"ccc_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# The muted check belongs to both IAM and Logging, so the Muted column
# must read 1 in both rows.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched CCC compliance, never
from a different framework that happens to be the last entry in the
check's compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_ccc_table(
findings,
bulk_metadata,
"ccc_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "leaked_provider" not in captured.out
@@ -0,0 +1,162 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.cis.cis import get_cis_table
def _strip_ansi(text):
return re.sub(r"\x1b\[[0-9;]*m", "", text)
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _attr(section, profile="Level 1"):
return SimpleNamespace(Section=section, Profile=profile)
def _make_compliance(provider, attributes, version="1.4", framework="CIS"):
"""Build a per-check CIS compliance with the given (section, profile) attrs."""
return SimpleNamespace(
Framework=framework,
Version=version,
Provider=provider,
Requirements=[SimpleNamespace(Attributes=attributes)],
)
def _make_compliance_multi_req(provider, attributes, version="1.4", framework="CIS"):
"""Build a per-check CIS compliance where each attr is its own requirement,
simulating a check that appears in several requirements."""
return SimpleNamespace(
Framework=framework,
Version=version,
Provider=provider,
Requirements=[SimpleNamespace(Attributes=[attr]) for attr in attributes],
)
class TestCISTable:
"""Verify multi-section counting and provider-column attribution for the CIS compliance table."""
def test_muted_multi_section_not_undercounted(self, capsys, tmp_path):
"""A single MUTED finding mapped to several sections must increment the
per-section Muted column for every section, not only the first seen.
CIS counts FAIL/PASS through Level 1/Level 2 buckets, so only the Muted
per-section count was affected by the undercount bug.
"""
bulk_metadata = {
# check_a is muted and belongs to two sections at once.
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("1 IAM"), _attr("2 Logging")])
]
),
# A real (non-muted) finding so the table is rendered.
"check_b": SimpleNamespace(
Compliance=[_make_compliance("aws", [_attr("1 IAM")])]
),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
_make_finding("check_b", "PASS"),
]
get_cis_table(
findings,
bulk_metadata,
"cis_1.4_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# Both section rows must carry a Muted count of 1 in their last cell.
# Before the fix only the first section seen got incremented.
muted_one_rows = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_one_rows) == 2
def test_same_section_level_not_double_counted(self, capsys, tmp_path):
"""A single finding whose check maps to several requirements that share
the same section and profile must count once for that section/level,
not once per requirement (FAIL(1), never FAIL(2))."""
bulk_metadata = {
# check_a is a single FAIL mapped to two requirements, both in the
# same section "1 IAM" and the same profile "Level 1".
"check_a": SimpleNamespace(
Compliance=[
_make_compliance_multi_req("aws", [_attr("1 IAM"), _attr("1 IAM")])
]
),
# A second finding in another section so the table renders.
"check_b": SimpleNamespace(
Compliance=[_make_compliance("aws", [_attr("2 Logging")])]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_cis_table(
findings,
bulk_metadata,
"cis_1.4_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# The "1 IAM" row must show FAIL(1) for Level 1, never FAIL(2).
assert "FAIL(1)" in plain
assert "FAIL(2)" not in plain
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched CIS compliance, not
from a different framework that trails it in the compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("1 IAM")]),
_make_compliance(
"gcp", [_attr("Other")], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("1 IAM")]),
_make_compliance(
"gcp", [_attr("Other")], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_cis_table(
findings,
bulk_metadata,
"cis_1.4_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The trailing unrelated framework's provider must not leak in.
assert "gcp" not in captured.out
@@ -0,0 +1,235 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.ens.ens import get_ens_table
def _strip_ansi(text):
return re.sub(r"\x1b\[[0-9;]*m", "", text)
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _attr(marco, categoria, tipo="requisito", nivel="alto"):
return SimpleNamespace(Marco=marco, Categoria=categoria, Tipo=tipo, Nivel=nivel)
def _make_compliance(provider, attributes, framework="ENS"):
"""Build a per-check ENS compliance with the given marco/categoria attrs."""
return SimpleNamespace(
Framework=framework,
Provider=provider,
Requirements=[SimpleNamespace(Attributes=attributes)],
)
class TestENSTable:
"""Test cases for ENS compliance table rendering.
Verify multi-marco counting and provider-column attribution for the
compliance table.
"""
def test_no_cumple_marked_in_every_marco(self, capsys, tmp_path):
"""A single failing finding mapped to several marcos must mark every
one of them as NO CUMPLE, not only the first marco seen."""
bulk_metadata = {
# check_a fails and belongs to two distinct marcos/categorias.
"check_a": SimpleNamespace(
Compliance=[
_make_compliance(
"aws",
[
_attr("operacional", "control de acceso"),
_attr("organizativo", "politica de seguridad"),
],
)
]
),
# A passing finding so the overview total reaches 2.
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("operacional", "control de acceso")])
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_ens_table(
findings,
bulk_metadata,
"ens_rd2022_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# Both marco rows the failing finding maps to must read NO CUMPLE.
# Before the fix only the first marco was marked, the second stayed
# CUMPLE. Anchor the assertion to the actual marco rows (not the
# overview header line which also mentions NO CUMPLE).
op_row = [
line
for line in plain.splitlines()
if "operacional/control de acceso" in line
]
org_row = [
line
for line in plain.splitlines()
if "organizativo/politica de seguridad" in line
]
assert len(op_row) == 1 and "NO CUMPLE" in op_row[0]
assert len(org_row) == 1 and "NO CUMPLE" in org_row[0]
def test_recomendacion_does_not_set_no_cumple(self, capsys, tmp_path):
"""A FAIL on a 'recomendacion' attribute must not flip a marco to
NO CUMPLE (this path is intentionally excluded from the fix)."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance(
"aws",
[
_attr(
"operacional", "control de acceso", tipo="recomendacion"
)
],
)
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("organizativo", "politica")])
]
),
# A regular (non-recomendacion) check so the results table renders
# at least one marco row and the assertion below is not vacuous.
"check_c": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("operacional", "continuidad")])
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
_make_finding("check_c", "PASS"),
]
get_ens_table(
findings,
bulk_metadata,
"ens_rd2022_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# The recomendacion FAIL must not appear as a NO CUMPLE marco row in the
# results table (the overview header line is allowed to mention it).
marco_rows = [
line
for line in plain.splitlines()
if "operacional" in line or "organizativo" in line
]
# Guard against a vacuous pass: the table must actually render rows.
assert marco_rows
assert all("NO CUMPLE" not in line for line in marco_rows)
def test_muted_multi_marco_not_undercounted(self, capsys, tmp_path):
"""A single MUTED finding mapped to several marcos must increment the
per-marco Muted column for every marco, not only the first seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance(
"aws",
[
_attr("operacional", "control de acceso"),
_attr("organizativo", "politica de seguridad"),
],
)
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", [_attr("operacional", "control de acceso")])
]
),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
_make_finding("check_b", "FAIL"),
]
get_ens_table(
findings,
bulk_metadata,
"ens_rd2022_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# Both marco rows the muted finding maps to must report a Muted count of
# 1 in their last cell.
muted_one_rows = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_one_rows) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Proveedor column must come from the matched ENS compliance, not
from a different framework that trails it in the compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance(
"aws", [_attr("operacional", "control de acceso")]
),
_make_compliance(
"gcp", [_attr("x", "y")], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance(
"aws", [_attr("operacional", "control de acceso")]
),
_make_compliance(
"gcp", [_attr("x", "y")], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_ens_table(
findings,
bulk_metadata,
"ens_rd2022_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
assert "gcp" not in captured.out
@@ -0,0 +1,137 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.kisa_ismsp.kisa_ismsp import get_kisa_ismsp_table
# The generator matches a compliance when its Framework starts with "KISA" and
# its Version is contained in the compliance_framework argument.
COMPLIANCE_FRAMEWORK = "kisa-isms-p-2023_aws"
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(
provider, sections, framework="KISA-ISMS-P", version="kisa-isms-p-2023"
):
"""Build a per-check compliance covering the given sections."""
return SimpleNamespace(
Framework=framework,
Version=version,
Provider=provider,
Requirements=[
SimpleNamespace(Attributes=[SimpleNamespace(Section=section)])
for section in sections
],
)
class TestKISAISMSPTable:
"""Verify multi-section counting and provider-column attribution for the KISA ISMS-P compliance table."""
def test_multi_section_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several sections must show FAIL(1) in
every section, not just the first one seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_kisa_ismsp_table(
findings,
bulk_metadata,
COMPLIANCE_FRAMEWORK,
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both IAM and Logging must report FAIL(1); before the fix Logging was
# undercounted because the per-section count was gated by the global
# dedup list.
assert captured.out.count("FAIL(1)") == 2
def test_multi_section_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several sections must increase the
per-section Muted count in every section, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A real FAIL is needed so the results table is rendered at all.
_make_finding("check_b", "FAIL"),
]
get_kisa_ismsp_table(
findings,
bulk_metadata,
COMPLIANCE_FRAMEWORK,
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# The muted check belongs to both IAM and Logging, so the Muted column
# must read 1 in both rows.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched KISA compliance, never
from a different framework that happens to be the last entry in the
check's compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_kisa_ismsp_table(
findings,
bulk_metadata,
COMPLIANCE_FRAMEWORK,
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "leaked_provider" not in captured.out
@@ -0,0 +1,140 @@
import re
from types import SimpleNamespace
from prowler.lib.outputs.compliance.mitre_attack.mitre_attack import (
get_mitre_attack_table,
)
# The generator matches a compliance when "MITRE-ATTACK" is in its Framework and
# its Version is contained in the compliance_framework argument.
COMPLIANCE_FRAMEWORK = "mitre_attack_aws"
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(
provider, tactics, framework="MITRE-ATTACK", version="mitre_attack"
):
"""Build a per-check compliance covering the given tactics."""
return SimpleNamespace(
Framework=framework,
Version=version,
Provider=provider,
Requirements=[SimpleNamespace(Tactics=tactics)],
)
class TestMitreAttackTable:
"""Test multi-section counting and provider-column attribution for the compliance table."""
def test_multi_tactic_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several tactics must show FAIL(1) in
every tactic, not just the first one seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["Persistence", "Execution"])]
),
"check_b": SimpleNamespace(
Compliance=[_make_compliance("aws", ["Persistence"])]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_mitre_attack_table(
findings,
bulk_metadata,
COMPLIANCE_FRAMEWORK,
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both Persistence and Execution must report FAIL(1); before the fix
# Execution was undercounted because the per-tactic count was gated by
# the global dedup list.
assert captured.out.count("FAIL(1)") == 2
def test_multi_tactic_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several tactics must increase the
per-tactic Muted count in every tactic, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["Persistence", "Execution"])]
),
"check_b": SimpleNamespace(
Compliance=[_make_compliance("aws", ["Persistence"])]
),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A second finding is needed so the table is rendered at all.
_make_finding("check_b", "FAIL"),
]
get_mitre_attack_table(
findings,
bulk_metadata,
COMPLIANCE_FRAMEWORK,
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# The muted check belongs to both Persistence and Execution, so the
# Muted column must read 1 in both rows.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched MITRE-ATTACK
compliance, never from a different framework that happens to be the last
entry in the check's compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["Persistence"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["Persistence"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_mitre_attack_table(
findings,
bulk_metadata,
COMPLIANCE_FRAMEWORK,
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "leaked_provider" not in captured.out
@@ -0,0 +1,136 @@
from types import SimpleNamespace
from prowler.lib.outputs.compliance.okta_idaas_stig.okta_idaas_stig import (
get_okta_idaas_stig_table,
)
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(provider, sections, framework="Okta-IDaaS-STIG"):
"""Build a per-check compliance covering the given sections."""
return SimpleNamespace(
Framework=framework,
Provider=provider,
Requirements=[
SimpleNamespace(Attributes=[SimpleNamespace(Section=section)])
for section in sections
],
)
class TestOktaIDaaSSTIGTable:
"""Test cases for Okta IDaaS STIG compliance table rendering."""
def test_multi_section_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several sections must show FAIL(1) in
every section, not just the first one seen."""
bulk_metadata = {
# check_a belongs to two sections at once.
"check_a": SimpleNamespace(
Compliance=[_make_compliance("okta", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("okta", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_okta_idaas_stig_table(
findings,
bulk_metadata,
"okta_idaas_stig_1r2",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both IAM and Logging must report FAIL(1); before the fix Logging
# was undercounted and rendered as plain PASS.
assert captured.out.count("FAIL(1)") == 2
def test_multi_section_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several sections must increase the
per-section Muted count in every section, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("okta", ["IAM", "Logging"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("okta", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A real FAIL is needed so the results table is rendered at all.
_make_finding("check_b", "FAIL"),
]
get_okta_idaas_stig_table(
findings,
bulk_metadata,
"okta_idaas_stig_1r2",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# The muted check belongs to both IAM and Logging, so the Muted column
# must read 1 in both rows. Before the fix only the first section seen
# was incremented, leaving the second at 0.
# Strip ANSI color codes before counting the bare values per row.
import re
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# Each section row ends with its Muted value in its own cell; both rows
# must carry a Muted count of 1.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched Okta-IDaaS-STIG
compliance, never from a different framework that happens to be the
last entry in the check's compliance list."""
# check_a maps to Okta-IDaaS-STIG (provider "okta") but its compliance
# list ends with a *different* framework whose provider is "aws". With
# the bug the leaked loop variable made the table render "aws".
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("okta", ["IAM"]),
_make_compliance("aws", ["Other"], framework="OtherFramework"),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("okta", ["IAM"]),
_make_compliance("aws", ["Other"], framework="OtherFramework"),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
get_okta_idaas_stig_table(
findings,
bulk_metadata,
"okta_idaas_stig_1r2",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "okta" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "aws" not in captured.out
@@ -0,0 +1,147 @@
import re
from types import SimpleNamespace
from unittest import mock
from prowler.lib.outputs.compliance.prowler_threatscore.prowler_threatscore import (
get_prowler_threatscore_table,
)
# Patch target for the Compliance.get_bulk lookup used to render pillars without
# findings; the tests don't exercise that path so it returns nothing.
COMPLIANCE_PATH = (
"prowler.lib.outputs.compliance.prowler_threatscore.prowler_threatscore.Compliance"
)
def _make_finding(check_id, status="PASS", muted=False):
return SimpleNamespace(
check_metadata=SimpleNamespace(CheckID=check_id),
status=status,
muted=muted,
)
def _make_compliance(provider, pillars, framework="ProwlerThreatScore"):
"""Build a per-check compliance covering the given pillars (Section)."""
return SimpleNamespace(
Framework=framework,
Provider=provider,
Requirements=[
SimpleNamespace(
Attributes=[SimpleNamespace(Section=pillar, LevelOfRisk=5, Weight=100)]
)
for pillar in pillars
],
)
class TestProwlerThreatScoreTable:
"""Verify multi-section counting and provider-column attribution for the compliance table."""
def test_multi_pillar_fail_not_undercounted(self, capsys, tmp_path):
"""A single FAIL check mapped to several pillars must show FAIL(1) in
every pillar, not just the first one seen."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Encryption"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
with mock.patch(COMPLIANCE_PATH) as compliance_mock:
compliance_mock.get_bulk.return_value = {}
get_prowler_threatscore_table(
findings,
bulk_metadata,
"prowler_threatscore_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
# Both IAM and Encryption must report FAIL(1); before the fix Encryption
# was undercounted because the per-pillar count was gated by the global
# dedup list.
assert captured.out.count("FAIL(1)") == 2
def test_multi_pillar_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED check mapped to several pillars must increase the
per-pillar Muted count in every pillar, not only the first one."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[_make_compliance("aws", ["IAM", "Encryption"])]
),
"check_b": SimpleNamespace(Compliance=[_make_compliance("aws", ["IAM"])]),
}
findings = [
_make_finding("check_a", "FAIL", muted=True),
# A real FAIL is needed so the results table is rendered at all.
_make_finding("check_b", "FAIL"),
]
with mock.patch(COMPLIANCE_PATH) as compliance_mock:
compliance_mock.get_bulk.return_value = {}
get_prowler_threatscore_table(
findings,
bulk_metadata,
"prowler_threatscore_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
plain = re.sub(r"\x1b\[[0-9;]*m", "", captured.out)
# The muted check belongs to both IAM and Encryption, so the Muted
# column must read 1 in both rows.
muted_cells = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_cells) == 2
def test_provider_column_not_leaked_from_other_framework(self, capsys, tmp_path):
"""The Provider column must come from the matched ProwlerThreatScore
compliance, never from a different framework that happens to be the last
entry in the check's compliance list."""
bulk_metadata = {
"check_a": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
"check_b": SimpleNamespace(
Compliance=[
_make_compliance("aws", ["IAM"]),
_make_compliance(
"leaked_provider", ["Other"], framework="OtherFramework"
),
]
),
}
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
with mock.patch(COMPLIANCE_PATH) as compliance_mock:
compliance_mock.get_bulk.return_value = {}
get_prowler_threatscore_table(
findings,
bulk_metadata,
"prowler_threatscore_aws",
"output",
str(tmp_path),
False,
)
captured = capsys.readouterr()
assert "aws" in captured.out
# The provider of the unrelated trailing framework must NOT leak into
# the rendered table.
assert "leaked_provider" not in captured.out
@@ -1,3 +1,4 @@
import re
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -26,6 +27,10 @@ def _make_finding(check_id, status="PASS", muted=False):
return finding
def _strip_ansi(text):
return re.sub(r"\x1b\[[0-9;]*m", "", text)
def _make_framework(requirements, table_config, provider="AWS"):
return ComplianceFramework(
framework="TestFW",
@@ -39,6 +44,8 @@ def _make_framework(requirements, table_config, provider="AWS"):
class TestBuildRequirementCheckMap:
"""Test cases for building the requirement-to-check map of a framework."""
def test_basic(self):
reqs = [
UniversalComplianceRequirement(
@@ -103,6 +110,8 @@ class TestBuildRequirementCheckMap:
class TestGetGroupKey:
"""Test cases for resolving the group key of a requirement."""
def test_normal_field(self):
req = UniversalComplianceRequirement(
id="1.1",
@@ -124,7 +133,9 @@ class TestGetGroupKey:
class TestGroupedMode:
def test_grouped_rendering(self, capsys):
"""Test cases for grouped-mode universal compliance table rendering."""
def test_grouped_rendering(self, capsys, tmp_path):
reqs = [
UniversalComplianceRequirement(
id="1.1",
@@ -156,7 +167,7 @@ class TestGroupedMode:
bulk_metadata,
"test_fw",
"output",
"/tmp",
str(tmp_path),
False,
framework=fw,
)
@@ -167,9 +178,118 @@ class TestGroupedMode:
assert "PASS" in captured.out
assert "FAIL" in captured.out
def test_grouped_multi_section_no_undercount(self, capsys, tmp_path):
"""A single check mapped to several sections must be counted in
every section it belongs to, not only the first one seen."""
reqs = [
UniversalComplianceRequirement(
id="1.1",
description="test",
attributes={"Section": "IAM"},
checks={"aws": ["check_a", "check_b"]},
),
UniversalComplianceRequirement(
id="2.1",
description="test2",
attributes={"Section": "Logging"},
checks={"aws": ["check_a"]},
),
]
tc = TableConfig(group_by="Section")
fw = _make_framework(reqs, tc)
# check_a (FAIL) belongs to both IAM and Logging sections; check_b
# (PASS, IAM only) is added so the overview total reaches 2 and the
# results table is rendered.
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
bulk_metadata = {
"check_a": MagicMock(Compliance=[]),
"check_b": MagicMock(Compliance=[]),
}
get_universal_table(
findings,
bulk_metadata,
"test_fw",
"output",
str(tmp_path),
False,
framework=fw,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# Both the IAM and Logging rows must report FAIL(1). Before the fix the
# second section seen (Logging) was undercounted to FAIL(0) and rendered
# as PASS. Anchor each occurrence to its own table row so an unrelated
# "FAIL(1)" elsewhere cannot mask an undercount.
iam_row = [
line for line in plain.splitlines() if "IAM" in line and "FAIL(1)" in line
]
logging_row = [
line
for line in plain.splitlines()
if "Logging" in line and "FAIL(1)" in line
]
assert len(iam_row) == 1
assert len(logging_row) == 1
def test_grouped_multi_section_muted_not_undercounted(self, capsys, tmp_path):
"""A single MUTED finding mapped to several groups must be counted in
the per-group Muted column of every group it belongs to."""
reqs = [
UniversalComplianceRequirement(
id="1.1",
description="test",
attributes={"Section": "IAM"},
checks={"aws": ["check_a", "check_b"]},
),
UniversalComplianceRequirement(
id="2.1",
description="test2",
attributes={"Section": "Logging"},
checks={"aws": ["check_a"]},
),
]
tc = TableConfig(group_by="Section")
fw = _make_framework(reqs, tc)
# check_a is MUTED and belongs to both IAM and Logging; check_b is a
# plain FAIL so the overview total reaches 2 and the table is rendered.
findings = [
_make_finding("check_a", "FAIL", muted=True),
_make_finding("check_b", "FAIL"),
]
bulk_metadata = {
"check_a": MagicMock(Compliance=[]),
"check_b": MagicMock(Compliance=[]),
}
get_universal_table(
findings,
bulk_metadata,
"test_fw",
"output",
str(tmp_path),
False,
framework=fw,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# The muted finding belongs to both sections, so both the IAM row and
# the Logging row must carry a Muted count of 1 in their last cell.
muted_one_rows = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_one_rows) == 2
class TestSplitMode:
def test_split_rendering(self, capsys):
"""Test cases for split-mode universal compliance table rendering."""
def test_split_rendering(self, capsys, tmp_path):
reqs = [
UniversalComplianceRequirement(
id="1.1",
@@ -204,7 +324,7 @@ class TestSplitMode:
bulk_metadata,
"test_fw",
"output",
"/tmp",
str(tmp_path),
False,
framework=fw,
)
@@ -214,9 +334,119 @@ class TestSplitMode:
assert "Level 1" in captured.out
assert "Level 2" in captured.out
def test_split_muted_multi_section_not_undercounted(self, capsys, tmp_path):
"""In split mode a single MUTED finding mapped to several groups must
be counted in the Muted column of every group it belongs to."""
reqs = [
UniversalComplianceRequirement(
id="1.1",
description="test",
attributes={"Section": "Storage", "Profile": "Level 1"},
checks={"aws": ["check_a", "check_b"]},
),
UniversalComplianceRequirement(
id="2.1",
description="test2",
attributes={"Section": "Logging", "Profile": "Level 1"},
checks={"aws": ["check_a"]},
),
]
tc = TableConfig(
group_by="Section",
split_by=SplitByConfig(field="Profile", values=["Level 1", "Level 2"]),
)
fw = _make_framework(reqs, tc)
# check_a is MUTED and belongs to both Storage and Logging; check_b is a
# plain FAIL so the table is rendered.
findings = [
_make_finding("check_a", "FAIL", muted=True),
_make_finding("check_b", "FAIL"),
]
bulk_metadata = {
"check_a": MagicMock(Compliance=[]),
"check_b": MagicMock(Compliance=[]),
}
get_universal_table(
findings,
bulk_metadata,
"test_fw",
"output",
str(tmp_path),
False,
framework=fw,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# Both section rows must carry a Muted count of 1 (last cell). Before the
# fix only the first group seen incremented Muted, leaving the other 0.
muted_one_rows = re.findall(r"│\s*1\s*│\s*$", plain, flags=re.MULTILINE)
assert len(muted_one_rows) == 2
def test_split_same_group_value_not_double_counted(self, capsys, tmp_path):
"""A single finding whose check maps to several requirements that share
the same group and split value must count once for that group/split,
not once per requirement (FAIL(1), never FAIL(2))."""
reqs = [
# check_a appears in two requirements, both Storage / Level 1.
UniversalComplianceRequirement(
id="1.1",
description="test",
attributes={"Section": "Storage", "Profile": "Level 1"},
checks={"aws": ["check_a"]},
),
UniversalComplianceRequirement(
id="1.2",
description="test2",
attributes={"Section": "Storage", "Profile": "Level 1"},
checks={"aws": ["check_a"]},
),
# A second group so the table renders with more than one finding.
UniversalComplianceRequirement(
id="2.1",
description="test3",
attributes={"Section": "Logging", "Profile": "Level 1"},
checks={"aws": ["check_b"]},
),
]
tc = TableConfig(
group_by="Section",
split_by=SplitByConfig(field="Profile", values=["Level 1", "Level 2"]),
)
fw = _make_framework(reqs, tc)
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
bulk_metadata = {
"check_a": MagicMock(Compliance=[]),
"check_b": MagicMock(Compliance=[]),
}
get_universal_table(
findings,
bulk_metadata,
"test_fw",
"output",
str(tmp_path),
False,
framework=fw,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
# The Storage row must show FAIL(1) for Level 1, never FAIL(2).
assert "FAIL(1)" in plain
assert "FAIL(2)" not in plain
class TestScoredMode:
def test_scored_rendering(self, capsys):
"""Test cases for scored-mode universal compliance table rendering."""
def test_scored_rendering(self, capsys, tmp_path):
reqs = [
UniversalComplianceRequirement(
id="1.1",
@@ -251,7 +481,7 @@ class TestScoredMode:
bulk_metadata,
"test_fw",
"output",
"/tmp",
str(tmp_path),
False,
framework=fw,
)
@@ -261,9 +491,68 @@ class TestScoredMode:
assert "Score" in captured.out
assert "Threat Score" in captured.out
def test_scored_multi_section_fail_not_undercounted(self, capsys, tmp_path):
"""In scored mode a single FAIL finding mapped to several groups must
show FAIL(1) in every group it belongs to, not only the first one."""
reqs = [
UniversalComplianceRequirement(
id="1.1",
description="test",
attributes={"Section": "IAM", "LevelOfRisk": 5, "Weight": 100},
checks={"aws": ["check_a", "check_b"]},
),
UniversalComplianceRequirement(
id="2.1",
description="test2",
attributes={"Section": "Logging", "LevelOfRisk": 3, "Weight": 50},
checks={"aws": ["check_a"]},
),
]
tc = TableConfig(
group_by="Section",
scoring=ScoringConfig(risk_field="LevelOfRisk", weight_field="Weight"),
)
fw = _make_framework(reqs, tc)
# check_a (FAIL) belongs to both IAM and Logging; check_b (PASS, IAM
# only) raises the overview total to 2 so the table is rendered.
findings = [
_make_finding("check_a", "FAIL"),
_make_finding("check_b", "PASS"),
]
bulk_metadata = {
"check_a": MagicMock(Compliance=[]),
"check_b": MagicMock(Compliance=[]),
}
get_universal_table(
findings,
bulk_metadata,
"test_fw",
"output",
str(tmp_path),
False,
framework=fw,
)
captured = capsys.readouterr()
plain = _strip_ansi(captured.out)
iam_row = [
line for line in plain.splitlines() if "IAM" in line and "FAIL(1)" in line
]
logging_row = [
line
for line in plain.splitlines()
if "Logging" in line and "FAIL(1)" in line
]
assert len(iam_row) == 1
assert len(logging_row) == 1
class TestCustomLabels:
def test_ens_spanish_labels(self, capsys):
"""Test cases for custom-label universal compliance table rendering."""
def test_ens_spanish_labels(self, capsys, tmp_path):
reqs = [
UniversalComplianceRequirement(
id="1.1",
@@ -300,7 +589,7 @@ class TestCustomLabels:
bulk_metadata,
"test_fw",
"output",
"/tmp",
str(tmp_path),
False,
framework=fw,
)
@@ -311,7 +600,9 @@ class TestCustomLabels:
class TestMultiProviderDictChecks:
def test_only_aws_checks_matched(self, capsys):
"""Test cases for multi-provider dict checks in the universal table."""
def test_only_aws_checks_matched(self, capsys, tmp_path):
"""With dict checks and provider='aws', only AWS checks match findings."""
reqs = [
UniversalComplianceRequirement(
@@ -352,7 +643,7 @@ class TestMultiProviderDictChecks:
bulk_metadata,
"multi_cloud",
"output",
"/tmp",
str(tmp_path),
False,
framework=fw,
provider="aws",
@@ -366,7 +657,9 @@ class TestMultiProviderDictChecks:
class TestNoTableConfig:
def test_returns_early_without_table_config(self, capsys):
"""Test cases for the universal table when no table config is present."""
def test_returns_early_without_table_config(self, capsys, tmp_path):
fw = ComplianceFramework(
framework="TestFW",
name="Test",
@@ -374,11 +667,11 @@ class TestNoTableConfig:
description="Test",
requirements=[],
)
get_universal_table([], {}, "test", "out", "/tmp", False, framework=fw)
get_universal_table([], {}, "test", "out", str(tmp_path), False, framework=fw)
captured = capsys.readouterr()
assert captured.out == ""
def test_returns_early_without_framework(self, capsys):
get_universal_table([], {}, "test", "out", "/tmp", False, framework=None)
def test_returns_early_without_framework(self, capsys, tmp_path):
get_universal_table([], {}, "test", "out", str(tmp_path), False, framework=None)
captured = capsys.readouterr()
assert captured.out == ""
+297 -11
View File
@@ -417,17 +417,19 @@ class TestIsBuiltinProvider:
class TestInitProvidersParserBuiltinDependencyFailure:
"""Tests the critical behavior fix: when a built-in provider's arguments
module exists but its imports fail (e.g. boto3 not installed), we must
fail loudly with a clear message — not silently fall through to entry
points as if the provider were external."""
"""Selective fail-loud: init captures failures silently, enforce emits
warning for non-invoked and exits for the invoked broken provider."""
@patch("sys.argv", ["prowler", "aws"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_builtin_with_missing_transitive_dep_fails_loudly(
self, mock_import, mock_is_builtin
):
from prowler.providers.common.arguments import init_providers_parser
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
mock_import.side_effect = ImportError("No module named 'boto3'")
@@ -435,14 +437,14 @@ class TestInitProvidersParserBuiltinDependencyFailure:
parser = MagicMock()
parser._providers = ["aws"]
with (
patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws"],
),
pytest.raises(SystemExit),
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws"],
):
init_providers_parser(parser)
assert "aws" in parser._builtin_load_failures
with pytest.raises(SystemExit):
enforce_invoked_provider_loaded(parser)
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.Provider._load_ep_provider")
@@ -466,6 +468,290 @@ class TestInitProvidersParserBuiltinDependencyFailure:
ext_cls.init_parser.assert_called_once_with(parser)
@patch("sys.argv", ["prowler", "aws"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_unrelated_builtin_failure_does_not_abort_when_other_provider_invoked(
self, mock_import, mock_is_builtin
):
"""Broken stackit + invoked aws → warning, no abort."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
aws_module = MagicMock()
def import_side_effect(module_path):
if "stackit" in module_path:
raise ImportError("No module named 'stackit.objectstorage'")
return aws_module
mock_import.side_effect = import_side_effect
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws", "stackit"],
):
init_providers_parser(parser)
assert "stackit" in parser._builtin_load_failures
enforce_invoked_provider_loaded(parser)
aws_module.init_parser.assert_called_once_with(parser)
@patch("sys.argv", ["prowler", "-h"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_no_provider_invoked_failure_does_not_abort(
self, mock_import, mock_is_builtin
):
"""`prowler -h` + broken built-in → warning, help still renders."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
mock_import.side_effect = ImportError("No module named 'stackit.objectstorage'")
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["stackit"],
):
init_providers_parser(parser)
enforce_invoked_provider_loaded(parser)
@patch("sys.argv", ["prowler", "microsoft365"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_invoked_microsoft365_alias_still_triggers_fail_loud(
self, mock_import, mock_is_builtin
):
"""Alias `microsoft365 → m365` must be normalised before matching."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
mock_import.side_effect = ImportError("No module named 'msgraph'")
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["m365"],
):
init_providers_parser(parser)
with pytest.raises(SystemExit):
enforce_invoked_provider_loaded(parser)
@patch("sys.argv", ["prowler", "oci"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_invoked_oci_alias_still_triggers_fail_loud(
self, mock_import, mock_is_builtin
):
"""Alias `oci → oraclecloud` must be normalised before matching."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
mock_import.side_effect = ImportError("No module named 'oci'")
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["oraclecloud"],
):
init_providers_parser(parser)
with pytest.raises(SystemExit):
enforce_invoked_provider_loaded(parser)
@patch("sys.argv", ["prowler", "--output-directory", "stackit"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_flag_value_matching_provider_name_not_treated_as_invoked(
self, mock_import, mock_is_builtin
):
"""Flag-first invocation → invoked is 'aws' (default), not the flag's value."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
aws_module = MagicMock()
def import_side_effect(module_path):
if "stackit" in module_path:
raise ImportError("No module named 'stackit.objectstorage'")
return aws_module
mock_import.side_effect = import_side_effect
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws", "stackit"],
):
init_providers_parser(parser)
enforce_invoked_provider_loaded(parser)
aws_module.init_parser.assert_called_once_with(parser)
@patch("sys.argv", ["prowler", "aws"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_invoked_builtin_non_import_error_fails_loudly(
self, mock_import, mock_is_builtin
):
"""Non-ImportError in invoked provider → still fail-loud."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
mock_import.side_effect = RuntimeError("Unexpected error in aws init_parser")
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws"],
):
init_providers_parser(parser)
with pytest.raises(SystemExit):
enforce_invoked_provider_loaded(parser)
@patch("sys.argv", ["prowler", "aws"])
@patch("prowler.providers.common.arguments.Provider.is_builtin")
@patch("prowler.providers.common.arguments.import_module")
def test_unrelated_builtin_non_import_error_does_not_abort(
self, mock_import, mock_is_builtin
):
"""Non-ImportError in unrelated provider → warning, no abort."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
mock_is_builtin.return_value = True
aws_module = MagicMock()
def import_side_effect(module_path):
if "stackit" in module_path:
raise RuntimeError("Unexpected error in stackit init_parser")
return aws_module
mock_import.side_effect = import_side_effect
parser = MagicMock()
with patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws", "stackit"],
):
init_providers_parser(parser)
enforce_invoked_provider_loaded(parser)
aws_module.init_parser.assert_called_once_with(parser)
class TestParseArgsOverrideAlignment:
"""Regression: `parse(args=...)` overrides sys.argv AFTER __init__ ran;
the selective fail-loud must read argv at enforce time, not init time."""
def test_enforce_reads_current_sys_argv_not_init_time_sys_argv(self):
"""Init with argv=['prowler','-h'] (no provider) captures stackit
failure silently. Enforce with argv=['prowler','stackit'] must
fail-loud — proving alignment under parse(args=...)."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
def import_side_effect(path):
if "stackit" in path:
raise ImportError("No module named 'stackit.objectstorage'")
return MagicMock()
parser = MagicMock()
with (
patch(
"prowler.providers.common.arguments.Provider.is_builtin",
return_value=True,
),
patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws", "stackit"],
),
patch(
"prowler.providers.common.arguments.import_module",
side_effect=import_side_effect,
),
):
# Phase 1: __init__ with ambient argv = ['prowler', '-h']
with patch("sys.argv", ["prowler", "-h"]):
init_providers_parser(parser)
# Failure captured silently — no SystemExit during init
assert "stackit" in parser._builtin_load_failures
# Phase 2: parse(args=...) overrode sys.argv → stackit invoked
with patch("sys.argv", ["prowler", "stackit"]):
with pytest.raises(SystemExit):
enforce_invoked_provider_loaded(parser)
def test_enforce_reads_current_sys_argv_for_no_invocation(self):
"""Inverse: init's argv invokes stackit, but parse(args=['prowler',
'-h']) overrides. Enforce must NOT fail-loud."""
from prowler.providers.common.arguments import (
enforce_invoked_provider_loaded,
init_providers_parser,
)
def import_side_effect(path):
if "stackit" in path:
raise ImportError("No module named 'stackit.objectstorage'")
return MagicMock()
parser = MagicMock()
with (
patch(
"prowler.providers.common.arguments.Provider.is_builtin",
return_value=True,
),
patch(
"prowler.providers.common.arguments.Provider.get_available_providers",
return_value=["aws", "stackit"],
),
patch(
"prowler.providers.common.arguments.import_module",
side_effect=import_side_effect,
),
):
# Phase 1: __init__ with ambient argv pretending stackit invoked
with patch("sys.argv", ["prowler", "stackit"]):
init_providers_parser(parser)
assert "stackit" in parser._builtin_load_failures
# Phase 2: parse(args=['prowler', '-h']) overrode sys.argv →
# no provider invoked anymore → enforce must NOT exit
with patch("sys.argv", ["prowler", "-h"]):
enforce_invoked_provider_loaded(parser)
class TestInitGlobalProviderBuiltinDependencyFailure:
"""Same contract as TestInitProvidersParserBuiltinDependencyFailure but
+64
View File
@@ -13,6 +13,7 @@ from prowler.config.config import (
)
from prowler.providers.common.models import Connection
from prowler.providers.gcp.exceptions.exceptions import (
GCPGetOrganizationProjectsError,
GCPInvalidProviderIdError,
GCPNoAccesibleProjectsError,
GCPTestConnectionError,
@@ -1077,3 +1078,66 @@ class TestGCPProvider:
assert gcp_provider.skip_api_check is True
mocked_is_api_active.assert_not_called()
def test_get_projects_organization_id_permission_denied_raises(self):
"""When --organization-id is set and the Cloud Asset API returns a 403,
get_projects must raise GCPGetOrganizationProjectsError instead of
silently falling back to the service account's home project.
Regression test for https://github.com/prowler-cloud/prowler/issues/11250.
"""
from googleapiclient.errors import HttpError
forbidden_response = MagicMock(status=403, reason="Forbidden")
http_error = HttpError(
resp=forbidden_response,
content=b'{"error": {"code": 403, "message": "Permission denied on resource organization"}}',
uri="https://cloudasset.googleapis.com/v1/organizations/123:listAssets",
)
asset_service = MagicMock()
asset_service.assets.return_value.list.return_value.execute.side_effect = (
http_error
)
with patch(
"prowler.providers.gcp.gcp_provider.discovery.build",
return_value=asset_service,
):
with pytest.raises(GCPGetOrganizationProjectsError):
GcpProvider.get_projects(
credentials=MagicMock(),
organization_id="test-organization-id",
credentials_file="test_credentials_file",
)
def test_get_projects_organization_id_cloud_asset_api_disabled_raises(self):
"""When --organization-id is set and the Cloud Asset API is disabled,
get_projects must raise GCPGetOrganizationProjectsError with the
enable-API remediation rather than swallowing the error."""
from googleapiclient.errors import HttpError
disabled_response = MagicMock(status=403, reason="Forbidden")
http_error = HttpError(
resp=disabled_response,
content=b'{"error": {"message": "Cloud Asset API has not been used in project 123 before or it is disabled."}}',
uri="https://cloudasset.googleapis.com/v1/organizations/123:listAssets",
)
asset_service = MagicMock()
asset_service.assets.return_value.list.return_value.execute.side_effect = (
http_error
)
with patch(
"prowler.providers.gcp.gcp_provider.discovery.build",
return_value=asset_service,
):
with pytest.raises(GCPGetOrganizationProjectsError) as exc_info:
GcpProvider.get_projects(
credentials=MagicMock(),
organization_id="test-organization-id",
credentials_file="test_credentials_file",
)
assert "Cloud Asset API" in str(exc_info.value)
@@ -1,5 +1,7 @@
from unittest.mock import MagicMock, patch
import pytest
from prowler.providers.gcp.services.logging.logging_service import Logging
from tests.providers.gcp.gcp_fixtures import (
GCP_PROJECT_ID,
@@ -291,6 +293,93 @@ class TestGetProjectsCoveredByAggregatedMetric:
)
assert self._run(logging_client, monitoring_client) == {}
def test_not_covered_when_sink_filter_omits_activity_stream(self):
"""A sink that routes cloudaudit streams but NOT Admin Activity (here,
data_access only) does not deliver the entries the CIS metric filters
match, so it must not be credited — right service, wrong stream."""
logging_client, monitoring_client = self._clients(
sink_filter="logName: /logs/cloudaudit.googleapis.com%2Fdata_access"
)
assert self._run(logging_client, monitoring_client) == {}
def test_covered_when_sink_filter_carries_activity_stream_encoded(self):
"""A sink filtered to the cloudaudit streams (URL-encoded logName form,
as returned by the Logging API) delivers every Admin Activity entry the
CIS metric filters can match, so it must be credited."""
logging_client, monitoring_client = self._clients(
sink_filter=(
"logName: /logs/cloudaudit.googleapis.com%2Factivity OR "
"logName: /logs/cloudaudit.googleapis.com%2Fdata_access"
)
)
assert self._run(logging_client, monitoring_client) == {
GCP_PROJECT_ID: "central-metric"
}
def test_covered_when_sink_filter_carries_activity_stream_plain(self):
logging_client, monitoring_client = self._clients(
sink_filter='logName="projects/p/logs/cloudaudit.googleapis.com/activity"'
)
assert self._run(logging_client, monitoring_client) == {
GCP_PROJECT_ID: "central-metric"
}
@pytest.mark.parametrize(
"sink_filter",
[
# --- Negation: the stream is named but excluded. ---
'NOT logName:"projects/p/logs/cloudaudit.googleapis.com%2Factivity"',
'-logName:"projects/p/logs/cloudaudit.googleapis.com%2Factivity"',
'NOT log_id("cloudaudit.googleapis.com/activity")',
# "!=" inequality (and its spaced form) excludes the stream.
'logName!="projects/p/logs/cloudaudit.googleapis.com%2Factivity"',
'logName != "projects/p/logs/cloudaudit.googleapis.com/activity"',
# Activity negated inside a compound filter.
'resource.type="gce_instance" AND '
'NOT logName:"projects/p/logs/cloudaudit.googleapis.com%2Factivity"',
# --- Restriction: the stream is named but AND-narrowed, so only a
# subset of Admin Activity entries reaches the bucket. ---
'logName:"projects/p/logs/cloudaudit.googleapis.com%2Factivity" '
'AND resource.type="gce_instance"',
'log_id("cloudaudit.googleapis.com/activity") '
'AND resource.type="gce_instance"',
'logName="projects/p/logs/cloudaudit.googleapis.com/activity" '
"AND severity>=ERROR",
'logName:"projects/p/logs/cloudaudit.googleapis.com%2Factivity" '
'AND protoPayload.methodName="SetIamPolicy"',
# --- OR-ed with a non-audit predicate: fail closed, since we credit
# only unions of provable Cloud Audit stream selectors. ---
'logName:"projects/p/logs/cloudaudit.googleapis.com%2Factivity" '
"OR severity>=ERROR",
],
)
def test_not_covered_when_sink_filter_negated_or_restrictive(self, sink_filter):
"""A filter that names the Admin Activity stream but negates, narrows, or
mixes in an unprovable predicate is not credited — we credit only filters
we can prove deliver every Admin Activity entry the CIS metrics match."""
logging_client, monitoring_client = self._clients(sink_filter=sink_filter)
assert self._run(logging_client, monitoring_client) == {}
def test_covered_when_activity_logname_has_hyphenated_path(self):
"""A hyphen in the project path must not be mistaken for the ``-`` (NOT)
negation operator — the activity stream is still delivered."""
logging_client, monitoring_client = self._clients(
sink_filter='logName="projects/my-project/logs/cloudaudit.googleapis.com/activity"'
)
assert self._run(logging_client, monitoring_client) == {
GCP_PROJECT_ID: "central-metric"
}
def test_covered_when_sink_filter_uses_log_id_selector(self):
"""The ``log_id()`` form is an equivalent positive full-coverage selector
of the Admin Activity stream and is credited like the ``logName`` form."""
logging_client, monitoring_client = self._clients(
sink_filter='log_id("cloudaudit.googleapis.com/activity")'
)
assert self._run(logging_client, monitoring_client) == {
GCP_PROJECT_ID: "central-metric"
}
def test_not_covered_when_sink_destination_bucket_differs(self):
logging_client, monitoring_client = self._clients(
sink_destination="logging.googleapis.com/projects/x/locations/eu/buckets/other"
+10 -1
View File
@@ -2,6 +2,15 @@
All notable changes to the **Prowler UI** are documented in this file.
## [1.30.1] (Prowler v5.30.1)
### 🐞 Fixed
- Threat Map no longer shows an empty map for accounts that only have Okta or Google Workspace scans [(#11542)](https://github.com/prowler-cloud/prowler/pull/11542)
- Compliance attributes requests now pass the selected scan, so multi-provider universal frameworks (e.g. CSA CCM) load the check IDs of the scan's provider and Azure/GCP requirement details show their findings instead of appearing empty [(#11546)](https://github.com/prowler-cloud/prowler/pull/11546)
---
## [1.30.0] (Prowler v5.30.0)
### 🚀 Added
@@ -12,7 +21,7 @@ All notable changes to the **Prowler UI** are documented in this file.
### 🔄 Changed
- Renamed "Customer Support" to "Support Desk" in the side menu, showing it only in Prowler Cloud/Enterprise, while "Community Support" now shows only in Prowler OSS [(#11508)](https://github.com/prowler-cloud/prowler/pull/11508)
- Compliance detail page now shows a "still loading" retry state while the API warms its compliance catalog, instead of rendering an empty page [(#4554)](https://github.com/prowler-cloud/prowler-cloud/pull/4554)
- Compliance detail page now shows a "still loading" retry state while the API warms its compliance catalog, instead of rendering an empty page [(#11530)](https://github.com/prowler-cloud/prowler/pull/11530)
### 🐞 Fixed
+10 -1
View File
@@ -73,12 +73,21 @@ export const getComplianceOverviewMetadataInfo = async ({
}
};
export const getComplianceAttributes = async (complianceId: string) => {
export const getComplianceAttributes = async (
complianceId: string,
scanId?: string,
) => {
const headers = await getAuthHeaders({ contentType: false });
try {
const url = new URL(`${apiBaseUrl}/compliance-overviews/attributes`);
url.searchParams.append("filter[compliance_id]", complianceId);
// Pass the scan so multi-provider universal frameworks (e.g. CSA CCM)
// resolve the check IDs for the scan's provider instead of defaulting to
// the first provider that declares the framework.
if (scanId) {
url.searchParams.append("filter[scan_id]", scanId);
}
const response = await fetch(url.toString(), {
headers,
@@ -0,0 +1,62 @@
import { describe, expect, it } from "vitest";
import { adaptRegionsOverviewToThreatMap } from "./threat-map.adapter";
import type { RegionsOverviewResponse } from "./types";
function buildRegionsResponse(
rows: Array<{ providerType: string; region: string }>,
): RegionsOverviewResponse {
return {
data: rows.map(({ providerType, region }, index) => ({
type: "regions-overview",
id: `region-${index}`,
attributes: {
provider_type: providerType,
region,
total: 10,
fail: 4,
muted: 0,
pass: 6,
},
})),
meta: { version: "v1" },
};
}
describe("adaptRegionsOverviewToThreatMap", () => {
it("maps okta regions to a global location", () => {
const response = buildRegionsResponse([
{ providerType: "okta", region: "global" },
]);
const result = adaptRegionsOverviewToThreatMap(response);
expect(result.locations).toHaveLength(1);
expect(result.locations[0]).toMatchObject({
providerType: "okta",
region: "global",
name: "Okta - Global",
totalFindings: 10,
failFindings: 4,
});
expect(result.regions).toEqual(["global"]);
});
it("maps googleworkspace regions to a global location", () => {
const response = buildRegionsResponse([
{ providerType: "googleworkspace", region: "global" },
]);
const result = adaptRegionsOverviewToThreatMap(response);
expect(result.locations).toHaveLength(1);
expect(result.locations[0]).toMatchObject({
providerType: "googleworkspace",
region: "global",
name: "Google Workspace - Global",
totalFindings: 10,
failFindings: 4,
});
expect(result.regions).toEqual(["global"]);
});
});
@@ -261,6 +261,19 @@ const ALIBABACLOUD_COORDINATES: Record<string, { lat: number; lng: number }> = {
global: { lat: 30.3, lng: 120.2 }, // Global fallback (Hangzhou HQ)
};
// Okta is a SaaS identity platform without user-facing regions
const OKTA_COORDINATES: Record<string, { lat: number; lng: number }> = {
global: { lat: 37.8, lng: -122.4 }, // Global fallback (San Francisco HQ)
};
// Google Workspace is a SaaS suite without user-facing regions
const GOOGLEWORKSPACE_COORDINATES: Record<
string,
{ lat: number; lng: number }
> = {
global: { lat: 37.4, lng: -122.1 }, // Global fallback (Mountain View HQ)
};
const PROVIDER_COORDINATES: Record<
string,
Record<string, { lat: number; lng: number }>
@@ -277,6 +290,8 @@ const PROVIDER_COORDINATES: Record<
oraclecloud: ORACLECLOUD_COORDINATES,
mongodbatlas: MONGODBATLAS_COORDINATES,
alibabacloud: ALIBABACLOUD_COORDINATES,
okta: OKTA_COORDINATES,
googleworkspace: GOOGLEWORKSPACE_COORDINATES,
};
// Returns [lng, lat] format for D3/GeoJSON compatibility
@@ -87,7 +87,7 @@ export default async function ComplianceDetail({
"filter[scan_id]": selectedScanId ?? undefined,
},
}),
getComplianceAttributes(complianceId),
getComplianceAttributes(complianceId, selectedScanId ?? undefined),
selectedScanId
? getScan(selectedScanId, { include: "provider" })
: Promise.resolve(null),
+93
View File
@@ -0,0 +1,93 @@
import { render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import { ThreatMap } from "./threat-map";
import type { ThreatMapData } from "./threat-map.types";
vi.mock("next/navigation", () => ({
useRouter: () => ({ push: vi.fn() }),
useSearchParams: () => new URLSearchParams(),
}));
vi.mock("./horizontal-bar-chart", () => ({
HorizontalBarChart: () => <div data-testid="bar-chart" />,
}));
function buildLocation(providerType: string, region: string) {
return {
id: `${providerType}-${region}`,
name: `${providerType} - ${region}`,
region,
regionCode: region,
providerType,
coordinates: [-122.4, 37.8] as [number, number],
totalFindings: 10,
failFindings: 4,
riskLevel: "high" as const,
severityData: [
{ name: "Fail", value: 4, percentage: 40 },
{ name: "Pass", value: 6, percentage: 60 },
],
};
}
describe("ThreatMap region selector", () => {
it("auto-selects the region when it is the only one available", () => {
const data: ThreatMapData = {
locations: [
buildLocation("okta", "global"),
buildLocation("googleworkspace", "global"),
],
regions: ["global"],
};
render(<ThreatMap data={data} />);
const select = screen.getByRole("combobox", {
name: "Filter threat map by region",
});
expect(select).toHaveValue("global");
expect(screen.getByText("Global Regions")).toBeInTheDocument();
expect(
screen.queryByText("Select a location on the map to view details"),
).not.toBeInTheDocument();
});
it("keeps All Regions as default when there are multiple regions", () => {
const data: ThreatMapData = {
locations: [
buildLocation("aws", "us-east-1"),
buildLocation("okta", "global"),
],
regions: ["global", "us-east-1"],
};
render(<ThreatMap data={data} />);
const select = screen.getByRole("combobox", {
name: "Filter threat map by region",
});
expect(select).toHaveValue("All Regions");
expect(
screen.getByRole("option", { name: "All Regions" }),
).toBeInTheDocument();
});
it("shows the global option capitalized while keeping its filter value", () => {
const data: ThreatMapData = {
locations: [
buildLocation("aws", "us-east-1"),
buildLocation("okta", "global"),
],
regions: ["global", "us-east-1"],
};
render(<ThreatMap data={data} />);
const globalOption = screen.getByRole("option", { name: "Global" });
expect(globalOption).toHaveValue("global");
expect(
screen.getByRole("option", { name: "us-east-1" }),
).toBeInTheDocument();
});
});
+10 -4
View File
@@ -124,7 +124,11 @@ export function ThreatMap({
x: number;
y: number;
} | null>(null);
const [selectedRegion, setSelectedRegion] = useState("All Regions");
// With a single region "All Regions" adds nothing, so it starts selected
const hasSingleRegion = data.regions.length === 1;
const [selectedRegion, setSelectedRegion] = useState(
hasSingleRegion ? data.regions[0] : "All Regions",
);
const [worldData, setWorldData] = useState<FeatureCollection | null>(null);
const [isLoadingMap, setIsLoadingMap] = useState(true);
const [dimensions, setDimensions] = useState<{
@@ -424,10 +428,12 @@ export function ThreatMap({
onChange={(e) => setSelectedRegion(e.target.value)}
className="border-border-neutral-primary bg-bg-neutral-secondary text-text-neutral-primary appearance-none rounded-lg border px-4 py-2 pr-10 text-sm focus:outline-none focus-visible:ring-2 focus-visible:ring-offset-2"
>
<option value="All Regions">All Regions</option>
{!hasSingleRegion && (
<option value="All Regions">All Regions</option>
)}
{sortedRegions.map((region) => (
<option key={region} value={region}>
{region}
{region.toLowerCase() === "global" ? "Global" : region}
</option>
))}
</select>
@@ -467,7 +473,7 @@ export function ThreatMap({
<div className="border-border-neutral-primary bg-bg-neutral-secondary absolute bottom-4 left-4 flex items-center gap-2 rounded-full border px-3 py-1.5">
<div
aria-hidden="true"
className="bg-data-critical h-3 w-3 rounded"
className="bg-bg-data-critical h-3 w-3 rounded"
/>
<span className="text-text-neutral-primary text-sm font-medium">
{locationCount} Locations
Generated
+1 -1
View File
@@ -3245,7 +3245,7 @@ wheels = [
[[package]]
name = "prowler"
version = "5.30.0"
version = "5.30.3"
source = { editable = "." }
dependencies = [
{ name = "alibabacloud-actiontrail20200706" },