diff --git a/api/changelog.d/attack-paths-tmp-db-reaper.fixed.md b/api/changelog.d/attack-paths-tmp-db-reaper.fixed.md new file mode 100644 index 0000000000..b165fc6d40 --- /dev/null +++ b/api/changelog.d/attack-paths-tmp-db-reaper.fixed.md @@ -0,0 +1 @@ +Adds a periodic sweep that drops orphaned Attack Paths temp Neo4j scan databases left behind when a worker or Neo4j crashes mid-scan, before they accumulate unbounded diff --git a/api/src/backend/api/attack_paths/database.py b/api/src/backend/api/attack_paths/database.py index 3ef55b7eca..1db76c4307 100644 --- a/api/src/backend/api/attack_paths/database.py +++ b/api/src/backend/api/attack_paths/database.py @@ -207,6 +207,11 @@ def drop_database(database: str) -> None: sink_module.get_backend().drop_database(database) +def list_databases() -> list[str]: + """List database names on the ingest cluster. Temp scan DBs always live here.""" + return ingest.list_databases() + + def drop_subgraph(database: str, provider_id: str) -> int: return sink_module.get_backend().drop_subgraph(database, provider_id) diff --git a/api/src/backend/api/attack_paths/ingest/__init__.py b/api/src/backend/api/attack_paths/ingest/__init__.py index 5833b8b373..d95482ed85 100644 --- a/api/src/backend/api/attack_paths/ingest/__init__.py +++ b/api/src/backend/api/attack_paths/ingest/__init__.py @@ -13,6 +13,7 @@ from api.attack_paths.ingest.driver import ( get_session, get_uri, init_driver, + list_databases, run_cypher, ) @@ -25,5 +26,6 @@ __all__ = [ "get_session", "get_uri", "init_driver", + "list_databases", "run_cypher", ] diff --git a/api/src/backend/api/attack_paths/ingest/driver.py b/api/src/backend/api/attack_paths/ingest/driver.py index 1b05c721e7..5c8b573ad7 100644 --- a/api/src/backend/api/attack_paths/ingest/driver.py +++ b/api/src/backend/api/attack_paths/ingest/driver.py @@ -165,6 +165,14 @@ def drop_database(database: str) -> None: session.run(f"DROP DATABASE `{database}` IF EXISTS DESTROY DATA") +def list_databases() -> list[str]: + """List every database name on the Neo4j temp-database cluster.""" + # A cluster returns one row per hosting server, so dedupe on name + with get_session() as session: + result = session.run("SHOW DATABASES YIELD name RETURN DISTINCT name") + return [record["name"] for record in result] + + def clear_cache(database: str) -> None: """Best-effort cache clear for a Neo4j database.""" from api.attack_paths.database import GraphDatabaseQueryException diff --git a/api/src/backend/api/migrations/0100_attack_paths_tmp_db_reap_periodic_task.py b/api/src/backend/api/migrations/0100_attack_paths_tmp_db_reap_periodic_task.py new file mode 100644 index 0000000000..ff27aff249 --- /dev/null +++ b/api/src/backend/api/migrations/0100_attack_paths_tmp_db_reap_periodic_task.py @@ -0,0 +1,48 @@ +from django.db import migrations + +TASK_NAME = "attack-paths-reap-orphaned-tmp-databases" +INTERVAL_HOURS = 6 + + +def create_periodic_task(apps, schema_editor): + IntervalSchedule = apps.get_model("django_celery_beat", "IntervalSchedule") + PeriodicTask = apps.get_model("django_celery_beat", "PeriodicTask") + + schedule, _ = IntervalSchedule.objects.get_or_create( + every=INTERVAL_HOURS, + period="hours", + ) + + PeriodicTask.objects.update_or_create( + name=TASK_NAME, + defaults={ + "task": TASK_NAME, + "interval": schedule, + "enabled": True, + }, + ) + + +def delete_periodic_task(apps, schema_editor): + IntervalSchedule = apps.get_model("django_celery_beat", "IntervalSchedule") + PeriodicTask = apps.get_model("django_celery_beat", "PeriodicTask") + + PeriodicTask.objects.filter(name=TASK_NAME).delete() + + # Clean up the schedule if no other task references it + IntervalSchedule.objects.filter( + every=INTERVAL_HOURS, + period="hours", + periodictask__isnull=True, + ).delete() + + +class Migration(migrations.Migration): + dependencies = [ + ("api", "0099_delete_tenant_onboarding_profile"), + ("django_celery_beat", "0019_alter_periodictasks_options"), + ] + + operations = [ + migrations.RunPython(create_periodic_task, delete_periodic_task), + ] diff --git a/api/src/backend/api/tests/test_attack_paths_database.py b/api/src/backend/api/tests/test_attack_paths_database.py index c4aca45928..a563667852 100644 --- a/api/src/backend/api/tests/test_attack_paths_database.py +++ b/api/src/backend/api/tests/test_attack_paths_database.py @@ -187,6 +187,27 @@ class TestRoutingByDatabasePrefix: sink_backend_stub.drop_database.assert_called_once_with("db-tenant-abc") mock_ingest.drop_database.assert_not_called() + def test_list_databases_always_routes_to_ingest(self, sink_backend_stub): + with patch("api.attack_paths.database.ingest") as mock_ingest: + mock_ingest.list_databases.return_value = ["db-tmp-scan-uuid-1"] + + assert db_module.list_databases() == ["db-tmp-scan-uuid-1"] + + mock_ingest.list_databases.assert_called_once_with() + + def test_ingest_list_databases_dedupes_cluster_rows(self): + from api.attack_paths.ingest import driver as ingest_driver + + with patch.object(ingest_driver, "get_session") as mock_get_session: + session = mock_get_session.return_value.__enter__.return_value + session.run.return_value = [{"name": "db-tmp-scan-uuid-1"}] + + assert ingest_driver.list_databases() == ["db-tmp-scan-uuid-1"] + + session.run.assert_called_once_with( + "SHOW DATABASES YIELD name RETURN DISTINCT name" + ) + def test_clear_cache_routes_temp_to_ingest(self, sink_backend_stub): with patch("api.attack_paths.database.ingest") as mock_ingest: db_module.clear_cache("db-tmp-scan-uuid-1") diff --git a/api/src/backend/config/django/base.py b/api/src/backend/config/django/base.py index 202476bd4b..664fcd1c4a 100644 --- a/api/src/backend/config/django/base.py +++ b/api/src/backend/config/django/base.py @@ -7,6 +7,7 @@ from config.settings.eventstream import * # noqa from config.settings.partitions import * # noqa from config.settings.sentry import * # noqa from config.settings.social_login import * # noqa +from django.core.exceptions import ImproperlyConfigured SECRET_KEY = env("SECRET_KEY", default="secret") DEBUG = env.bool("DJANGO_DEBUG", default=False) @@ -328,6 +329,17 @@ ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES = env.int( "ATTACK_PATHS_SCAN_STALE_THRESHOLD_MINUTES", 960 ) # 16h +# Minimum age (of the scan row, or of the scan id itself when the row is gone) before +# the periodic reaper will drop an orphaned temp Neo4j database. Keeps a scan that is +# still legitimately in flight from ever losing its staging database mid-run. +ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS = env.int( + "ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS", 6 +) +if ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS <= 0: + raise ImproperlyConfigured( + "ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS must be a positive number of hours" + ) + # Selects where the persistent attack-paths graph is stored. The scan # temporary database is always Neo4j; only the sink is configurable. # Valid values: "neo4j" (default, OSS and local dev), "neptune" (hosted). diff --git a/api/src/backend/tasks/jobs/attack_paths/tmp_db_reaper.py b/api/src/backend/tasks/jobs/attack_paths/tmp_db_reaper.py new file mode 100644 index 0000000000..3d575d41d4 --- /dev/null +++ b/api/src/backend/tasks/jobs/attack_paths/tmp_db_reaper.py @@ -0,0 +1,107 @@ +"""Periodic reaper for orphaned temp Neo4j scan databases. + +`scan.py` creates a throw-away `db-tmp-scan-` database per +scan and drops it once the scan finishes, success or failure. When the worker +or Neo4j itself dies mid-scan, that drop never runs and nothing else ever +revisits the database - it sits there forever. This sweep lists every temp +database on the ingest cluster and drops the ones whose scan is gone or has +been finished for longer than the configured safety margin. +""" + +from datetime import UTC, datetime, timedelta + +from api.attack_paths import database as graph_database +from api.db_router import MainRouter +from api.models import AttackPathsScan, StateChoices +from api.uuid_utils import datetime_from_uuid7 +from celery.utils.log import get_task_logger +from config.django.base import ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS +from uuid6 import UUID as UUID7 + +logger = get_task_logger(__name__) + +TERMINAL_STATES = ( + StateChoices.COMPLETED, + StateChoices.FAILED, + StateChoices.CANCELLED, +) + + +def reap_orphaned_tmp_databases() -> dict: + """Drop temp Neo4j scan databases whose scan is gone or long finished. + + A failure listing databases aborts the whole sweep (nothing to iterate). + A failure reaping one database is logged and skipped so the rest of the + sweep still runs. + """ + now = datetime.now(tz=UTC) + safety_margin = timedelta(hours=ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS) + + try: + databases = graph_database.list_databases() + except Exception: + logger.exception("Failed to list ingest Neo4j databases for temp-db reap") + return {"dropped_count": 0, "databases": []} + + tmp_databases = [ + name for name in databases if name.startswith(graph_database.TEMP_DB_PREFIX) + ] + + dropped: list[str] = [] + for database in tmp_databases: + try: + if _is_orphaned(database, now, safety_margin): + graph_database.drop_database(database) + dropped.append(database) + logger.info(f"Dropped orphaned temp Neo4j database `{database}`") + except Exception: + logger.exception(f"Failed to reap temp Neo4j database `{database}`") + + logger.info(f"Temp Neo4j database reap: {len(dropped)} dropped") + return {"dropped_count": len(dropped), "databases": dropped} + + +def _is_orphaned(database: str, now: datetime, safety_margin: timedelta) -> bool: + """Decide whether a temp database is safe to drop. + + No scan row: the row was hard-deleted (tenant/provider cleanup) or was + never created. Falls back to the scan id's own UUIDv7 timestamp so a + database created moments ago is never touched even without a row to check. + + Scan row present: only reapable once it reached a terminal state and has + been finished for longer than the safety margin, so a scan still + legitimately executing is never touched. + """ + scan_id = database[len(graph_database.TEMP_DB_PREFIX) :] + + try: + scan_uuid = UUID7(scan_id) + except ValueError: + logger.warning( + f"Temp database `{database}` has an unparseable scan id, skipping" + ) + return False + + # Global sweep with no tenant context: admin_db bypasses RLS on purpose, the same + # way cleanup_stale_attack_paths_scans finds stale scans across every tenant. + scan = ( + AttackPathsScan.all_objects.using(MainRouter.admin_db) + .filter(id=scan_uuid) + .first() + ) + + if scan is None: + if scan_uuid.version != 7: + logger.warning( + f"Temp database `{database}` has no scan row and a non-UUIDv7 id, " + "skipping" + ) + return False + return now - datetime_from_uuid7(scan_uuid) >= safety_margin + + if scan.state not in TERMINAL_STATES: + return False + + # `mark_scan_finished` does not touch `updated_at`, so prefer `completed_at` + finished_at = scan.completed_at or scan.updated_at + return now - finished_at >= safety_margin diff --git a/api/src/backend/tasks/jobs/orphan_recovery.py b/api/src/backend/tasks/jobs/orphan_recovery.py index c8cda54cd2..7290fb313f 100644 --- a/api/src/backend/tasks/jobs/orphan_recovery.py +++ b/api/src/backend/tasks/jobs/orphan_recovery.py @@ -84,6 +84,7 @@ _SKIP_RECOVERY = { "scan-perform-scheduled", "attack-paths-scan-perform", "attack-paths-cleanup-stale-scans", + "attack-paths-reap-orphaned-tmp-databases", "reconcile-orphan-tasks", } diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index 332f893aa2..9ec6526963 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -43,6 +43,7 @@ from tasks.jobs.attack_paths import ( ) from tasks.jobs.attack_paths import db_utils as attack_paths_db_utils from tasks.jobs.attack_paths.cleanup import cleanup_stale_attack_paths_scans +from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases from tasks.jobs.backfill import ( aggregate_scan_category_summaries, aggregate_scan_resource_group_summaries, @@ -727,6 +728,13 @@ def cleanup_stale_attack_paths_scans_task(): return cleanup_stale_attack_paths_scans() +@shared_task( + name="attack-paths-reap-orphaned-tmp-databases", queue="attack-paths-scans" +) +def reap_orphaned_attack_paths_tmp_databases_task(): + return reap_orphaned_tmp_databases() + + @shared_task(name="reconcile-orphan-tasks", queue="celery") def reconcile_orphan_tasks_task(): """Periodic watchdog: recover tasks whose worker is gone (deploys, crashes).""" diff --git a/api/src/backend/tasks/tests/test_attack_paths_tmp_db_reaper.py b/api/src/backend/tasks/tests/test_attack_paths_tmp_db_reaper.py new file mode 100644 index 0000000000..ab8d478383 --- /dev/null +++ b/api/src/backend/tasks/tests/test_attack_paths_tmp_db_reaper.py @@ -0,0 +1,286 @@ +from datetime import UTC, datetime, timedelta +from unittest.mock import patch +from uuid import uuid4 + +import pytest +from api.attack_paths.database import TEMP_DB_PREFIX +from api.models import AttackPathsScan, StateChoices +from api.uuid_utils import datetime_to_uuid7 +from config.django.base import ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS + +MARGIN = timedelta(hours=ATTACK_PATHS_TMP_DB_REAP_SAFETY_MARGIN_HOURS) + + +def _tmp_db_name(scan_uuid) -> str: + return f"{TEMP_DB_PREFIX}{scan_uuid}" + + +@pytest.mark.django_db +class TestReapOrphanedTmpDatabases: + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_ignores_databases_without_the_temp_prefix(self, mock_list, mock_drop): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + mock_list.return_value = ["db-tenant-abc123", "system", "neo4j"] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_temp_db_with_no_scan_row_past_safety_margin( + self, mock_list, mock_drop + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + old_scan_id = datetime_to_uuid7( + datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + ) + database = _tmp_db_name(old_scan_id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [database]} + mock_drop.assert_called_once_with(database) + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_temp_db_with_no_scan_row_inside_safety_margin( + self, mock_list, mock_drop + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + recent_scan_id = datetime_to_uuid7(datetime.now(tz=UTC) - timedelta(minutes=5)) + database = _tmp_db_name(recent_scan_id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_temp_db_with_unparseable_scan_id(self, mock_list, mock_drop): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + database = f"{TEMP_DB_PREFIX}not-a-uuid" + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_terminal_scan_past_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old_updated_at = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.COMPLETED, + ) + AttackPathsScan.objects.filter(id=scan.id).update(updated_at=old_updated_at) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [database]} + mock_drop.assert_called_once_with(database) + + @pytest.mark.parametrize( + "state", + [StateChoices.FAILED, StateChoices.CANCELLED], + ) + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_other_terminal_states_past_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider, state + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old_updated_at = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=state, + ) + AttackPathsScan.objects.filter(id=scan.id).update(updated_at=old_updated_at) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result["dropped_count"] == 1 + mock_drop.assert_called_once_with(database) + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_terminal_scan_inside_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.COMPLETED, + ) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_margin_counts_from_completed_at_not_updated_at( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.COMPLETED, + ) + AttackPathsScan.objects.filter(id=scan.id).update( + updated_at=old, completed_at=datetime.now(tz=UTC) + ) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_drops_scan_completed_past_safety_margin( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + old = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.FAILED, + ) + AttackPathsScan.objects.filter(id=scan.id).update( + updated_at=old, completed_at=old + ) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [database]} + mock_drop.assert_called_once_with(database) + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_never_drops_an_executing_scan_regardless_of_age( + self, mock_list, mock_drop, tenants_fixture, aws_provider + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + tenant = tenants_fixture[0] + very_old = datetime.now(tz=UTC) - timedelta(days=30) + scan = AttackPathsScan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + state=StateChoices.EXECUTING, + ) + AttackPathsScan.objects.filter(id=scan.id).update(updated_at=very_old) + + database = _tmp_db_name(scan.id) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_one_failed_drop_does_not_stop_the_rest_of_the_sweep( + self, mock_list, mock_drop + ): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + old_time = datetime.now(tz=UTC) - MARGIN - timedelta(hours=1) + failing_scan_id = datetime_to_uuid7(old_time) + succeeding_scan_id = datetime_to_uuid7(old_time) + failing_db = _tmp_db_name(failing_scan_id) + succeeding_db = _tmp_db_name(succeeding_scan_id) + mock_list.return_value = [failing_db, succeeding_db] + mock_drop.side_effect = [Exception("boom"), None] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 1, "databases": [succeeding_db]} + assert mock_drop.call_count == 2 + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_returns_empty_result_when_listing_databases_fails(self, mock_list): + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + mock_list.side_effect = Exception("neo4j unreachable") + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.drop_database") + @patch("tasks.jobs.attack_paths.tmp_db_reaper.graph_database.list_databases") + def test_preserves_temp_db_with_random_uuid_and_no_row(self, mock_list, mock_drop): + """A non-UUIDv7 id with no matching row has no reliable timestamp, so it + must be left alone rather than guessed at.""" + from tasks.jobs.attack_paths.tmp_db_reaper import reap_orphaned_tmp_databases + + database = _tmp_db_name(uuid4()) + mock_list.return_value = [database] + + result = reap_orphaned_tmp_databases() + + assert result == {"dropped_count": 0, "databases": []} + mock_drop.assert_not_called() + + +class TestReapOrphanedTmpDatabasesTask: + @patch( + "tasks.tasks.reap_orphaned_tmp_databases", + return_value={"dropped_count": 2, "databases": ["db-tmp-scan-a"]}, + ) + def test_task_invokes_the_reaper(self, mock_reap): + from tasks.tasks import reap_orphaned_attack_paths_tmp_databases_task + + result = reap_orphaned_attack_paths_tmp_databases_task.run() + + assert result == {"dropped_count": 2, "databases": ["db-tmp-scan-a"]} + mock_reap.assert_called_once_with()