diff --git a/api/src/backend/tasks/jobs/attack_paths/__init__.py b/api/src/backend/tasks/jobs/attack_paths/__init__.py index 0f586cdcf6..8fb57bc907 100644 --- a/api/src/backend/tasks/jobs/attack_paths/__init__.py +++ b/api/src/backend/tasks/jobs/attack_paths/__init__.py @@ -1,5 +1,7 @@ +from tasks.jobs.attack_paths.db_utils import can_provider_run_attack_paths_scan from tasks.jobs.attack_paths.scan import run as attack_paths_scan __all__ = [ "attack_paths_scan", + "can_provider_run_attack_paths_scan", ] diff --git a/api/src/backend/tasks/jobs/attack_paths/db_utils.py b/api/src/backend/tasks/jobs/attack_paths/db_utils.py index 2565dfa872..63451ef74d 100644 --- a/api/src/backend/tasks/jobs/attack_paths/db_utils.py +++ b/api/src/backend/tasks/jobs/attack_paths/db_utils.py @@ -1,6 +1,5 @@ from datetime import datetime, timezone from typing import Any -from uuid import UUID from cartography.config import Config as CartographyConfig @@ -13,15 +12,19 @@ from api.models import ( from tasks.jobs.attack_paths.providers import is_provider_available +def can_provider_run_attack_paths_scan(tenant_id: str, provider_id: int) -> bool: + with rls_transaction(tenant_id): + prowler_api_provider = ProwlerAPIProvider.objects.get(id=provider_id) + + return is_provider_available(prowler_api_provider.provider) + + def create_attack_paths_scan( tenant_id: str, scan_id: str, provider_id: int, ) -> ProwlerAPIAttackPathsScan | None: - with rls_transaction(tenant_id): - prowler_api_provider = ProwlerAPIProvider.objects.get(id=provider_id) - - if not is_provider_available(prowler_api_provider.provider): + if not can_provider_run_attack_paths_scan(tenant_id, provider_id): return None with rls_transaction(tenant_id): diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index 4e97c581f7..73fab9ffed 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -20,7 +20,7 @@ from config.django.base import DJANGO_FINDINGS_BATCH_SIZE, DJANGO_TMP_OUTPUT_DIR from prowler.lib.check.compliance_models import Compliance from prowler.lib.outputs.compliance.generic.generic import GenericCompliance from prowler.lib.outputs.finding import Finding as FindingOutput -from tasks.jobs.attack_paths import attack_paths_scan +from tasks.jobs.attack_paths import attack_paths_scan, can_provider_run_attack_paths_scan from tasks.jobs.backfill import ( backfill_compliance_summaries, backfill_daily_severity_summaries, @@ -153,9 +153,11 @@ def _perform_scan_complete_tasks(tenant_id: str, scan_id: str, provider_id: str) ), ), ).apply_async() - perform_attack_paths_scan_task.apply_async( - kwargs={"tenant_id": tenant_id, "scan_id": scan_id} - ) + + if can_provider_run_attack_paths_scan(tenant_id, provider_id): + perform_attack_paths_scan_task.apply_async( + kwargs={"tenant_id": tenant_id, "scan_id": scan_id} + ) @shared_task(base=RLSTask, name="provider-connection-check")