diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index 0d488576c1..5ddad0522d 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -28,6 +28,7 @@ All notable changes to the **Prowler API** are documented in this file. ### 🐞 Fixed - Attack Paths: Orphaned temporary Neo4j databases are now cleaned up on scan failure and provider deletion [(#10101)](https://github.com/prowler-cloud/prowler/pull/10101) +- Attack Paths: scan no longer raises `DatabaseError` when provider is deleted mid-scan [(#10116)](https://github.com/prowler-cloud/prowler/pull/10116) ### 🔐 Security diff --git a/api/src/backend/api/decorators.py b/api/src/backend/api/decorators.py index d2330a6a06..f9b165ef20 100644 --- a/api/src/backend/api/decorators.py +++ b/api/src/backend/api/decorators.py @@ -2,7 +2,7 @@ import uuid from functools import wraps from django.core.exceptions import ObjectDoesNotExist -from django.db import IntegrityError, connection, transaction +from django.db import DatabaseError, connection, transaction from rest_framework_json_api.serializers import ValidationError from api.db_router import READ_REPLICA_ALIAS @@ -74,12 +74,13 @@ def set_tenant(func=None, *, keep_tenant=False): def handle_provider_deletion(func): """ - Decorator that raises ProviderDeletedException if provider was deleted during execution. + Decorator that raises `ProviderDeletedException` if provider was deleted during execution. - Catches ObjectDoesNotExist and IntegrityError, checks if provider still exists, - and raises ProviderDeletedException if not. Otherwise, re-raises original exception. + Catches `ObjectDoesNotExist` and `DatabaseError` (including `IntegrityError`), checks if + provider still exists, and raises `ProviderDeletedException` if not. Otherwise, + re-raises original exception. - Requires tenant_id and provider_id in kwargs. + Requires `tenant_id` and `provider_id` in kwargs. Example: @shared_task @@ -92,7 +93,7 @@ def handle_provider_deletion(func): def wrapper(*args, **kwargs): try: return func(*args, **kwargs) - except (ObjectDoesNotExist, IntegrityError): + except (ObjectDoesNotExist, DatabaseError): tenant_id = kwargs.get("tenant_id") provider_id = kwargs.get("provider_id") diff --git a/api/src/backend/api/tests/test_decorators.py b/api/src/backend/api/tests/test_decorators.py index 9a113abad8..2d09a40734 100644 --- a/api/src/backend/api/tests/test_decorators.py +++ b/api/src/backend/api/tests/test_decorators.py @@ -3,7 +3,7 @@ from unittest.mock import call, patch import pytest from django.core.exceptions import ObjectDoesNotExist -from django.db import IntegrityError +from django.db import DatabaseError, IntegrityError from api.db_utils import POSTGRES_TENANT_VAR, SET_CONFIG_QUERY from api.decorators import handle_provider_deletion, set_tenant @@ -165,6 +165,46 @@ class TestHandleProviderDeletionDecorator: with pytest.raises(ProviderDeletedException): task_func(tenant_id=str(tenant.id), provider_id=deleted_provider_id) + @patch("api.decorators.rls_transaction") + @patch("api.decorators.Provider.objects.filter") + def test_database_error_provider_deleted( + self, mock_filter, mock_rls, tenants_fixture + ): + """Raises ProviderDeletedException on DatabaseError when provider deleted.""" + tenant = tenants_fixture[0] + deleted_provider_id = str(uuid.uuid4()) + + mock_rls.return_value.__enter__ = lambda s: None + mock_rls.return_value.__exit__ = lambda s, *args: None + mock_filter.return_value.exists.return_value = False + + @handle_provider_deletion + def task_func(**kwargs): + raise DatabaseError("Save with update_fields did not affect any rows") + + with pytest.raises(ProviderDeletedException): + task_func(tenant_id=str(tenant.id), provider_id=deleted_provider_id) + + @patch("api.decorators.rls_transaction") + @patch("api.decorators.Provider.objects.filter") + def test_database_error_provider_exists_reraises( + self, mock_filter, mock_rls, tenants_fixture, providers_fixture + ): + """Re-raises original DatabaseError when provider still exists.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + + mock_rls.return_value.__enter__ = lambda s: None + mock_rls.return_value.__exit__ = lambda s, *args: None + mock_filter.return_value.exists.return_value = True + + @handle_provider_deletion + def task_func(**kwargs): + raise DatabaseError("Save with update_fields did not affect any rows") + + with pytest.raises(DatabaseError): + task_func(tenant_id=str(tenant.id), provider_id=str(provider.id)) + def test_missing_provider_and_scan_raises_assertion(self, tenants_fixture): """Raises AssertionError when neither provider_id nor scan_id in kwargs.""" diff --git a/api/src/backend/tasks/jobs/attack_paths/scan.py b/api/src/backend/tasks/jobs/attack_paths/scan.py index da70b77383..cd39700dd4 100644 --- a/api/src/backend/tasks/jobs/attack_paths/scan.py +++ b/api/src/backend/tasks/jobs/attack_paths/scan.py @@ -204,8 +204,8 @@ def run(tenant_id: str, scan_id: str, task_id: str) -> dict[str, Any]: return ingestion_exceptions except Exception as e: - exception_message = utils.stringify_exception(e, "Cartography failed") - logger.error(exception_message) + exception_message = utils.stringify_exception(e, "Attack Paths scan failed") + logger.exception(exception_message) ingestion_exceptions["global_error"] = exception_message # Handling databases changes @@ -213,11 +213,17 @@ def run(tenant_id: str, scan_id: str, task_id: str) -> dict[str, Any]: graph_database.drop_database(tmp_cartography_config.neo4j_database) except Exception: - logger.exception( + logger.error( f"Failed to drop temporary Neo4j database {tmp_cartography_config.neo4j_database} during cleanup" ) - db_utils.finish_attack_paths_scan( - attack_paths_scan, StateChoices.FAILED, ingestion_exceptions - ) + try: + db_utils.finish_attack_paths_scan( + attack_paths_scan, StateChoices.FAILED, ingestion_exceptions + ) + except Exception: + logger.warning( + f"Could not mark attack paths scan {attack_paths_scan.id} as FAILED (row may have been deleted)" + ) + raise diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index 30cc0b09c4..721fba9d07 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -383,6 +383,7 @@ class AttackPathsScanRLSTask(RLSTask): name="attack-paths-scan-perform", queue="attack-paths-scans", ) +@handle_provider_deletion def perform_attack_paths_scan_task(self, tenant_id: str, scan_id: str): """ Execute an Attack Paths scan for the given provider within the current tenant RLS context.