fix(api): reap orphaned attack paths temp Neo4j databases (#12832)

This commit is contained in:
César Arroba
2026-09-24 13:23:36 +02:00
committed by GitHub
parent bf179212a5
commit 576433d85d
11 changed files with 499 additions and 0 deletions
@@ -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
@@ -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)
@@ -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",
]
@@ -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
@@ -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),
]
@@ -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")
+12
View File
@@ -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).
@@ -0,0 +1,107 @@
"""Periodic reaper for orphaned temp Neo4j scan databases.
`scan.py` creates a throw-away `db-tmp-scan-<attack_paths_scan_id>` 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
@@ -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",
}
+8
View File
@@ -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)."""
@@ -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()