feat(attack-paths): fixing deletion of old graph before loading dump

This commit is contained in:
Josema Camacho
2026-01-14 14:52:40 +01:00
parent fdc3b18542
commit 068efb3318
11 changed files with 68 additions and 6287 deletions
+31 -12
View File
@@ -103,15 +103,18 @@ def drop_database(database: str) -> None:
session.run(query)
def drop_subgraph(database: str, root_node_label: str, root_node_id: str) -> int:
def drop_subgraph(database: str, root_node_label: str, root_node_id: str, provider_id: str) -> int:
query = """
MATCH (rn:__ROOT_NODE_LABEL__ {id: $root_node_id})
MATCH (rn:__ROOT_NODE_LABEL__ {id: $root_node_id, prowler_provider_id: $prowler_provider_id})
CALL apoc.path.subgraphNodes(rn, {})
YIELD node
DETACH DELETE node
RETURN COUNT(node) AS deleted_nodes_count
""".replace("__ROOT_NODE_LABEL__", root_node_label)
parameters = {"root_node_id": root_node_id}
parameters = {
"root_node_id": root_node_id,
"prowler_provider_id": provider_id,
}
with get_session(database) as session:
result = session.run(query, parameters)
@@ -175,7 +178,7 @@ def create_database_dump(database: str, root_node_label: str, root_node_id: str)
return dump_filename_path
def load_database_dump(dump_filename_path: str, database: str) -> str:
def load_database_dump(dump_filename_path: str, database: str, provider_id: str) -> None:
BATCH_SIZE = 1000
query = """
@@ -188,9 +191,9 @@ def load_database_dump(dump_filename_path: str, database: str) -> str:
UNWIND rows AS row
WITH row
WHERE row.type = 'node'
MERGE (n {piid: row.id})
MERGE (n {piid: row.id, prowler_provider_id: $prowler_provider_id})
SET n += COALESCE(row.properties, {})
FOREACH (l IN COALESCE(row.labels, []) | SET n:$(l))
SET n:$(COALESCE(row.labels, []))
}
// Create relationships from the batch
@@ -198,11 +201,21 @@ def load_database_dump(dump_filename_path: str, database: str) -> str:
WITH rows
UNWIND rows AS row
WITH row
WHERE row.type = 'relationship'
MATCH (s {piid: row.start}), (t {piid: row.end})
CREATE (s)-[r:$(row.label)]->(t)
SET r += COALESCE(row.properties, {})
};
WHERE row.type = 'relationship' AND row.label IS NOT NULL
MATCH (s {piid: row.start, prowler_provider_id: $prowler_provider_id}),
(t {piid: row.end, prowler_provider_id: $prowler_provider_id})
CALL apoc.merge.relationship(
s,
row.label,
{},
COALESCE(row.properties, {}),
t
) YIELD rel
RETURN TRUE AS relationship_created
}
RETURN relationship_created AS export_finished;
// It needs to return something because of the use of `apoc.merge.relationship` inside a CALL
"""
def chunks(iterable, size):
@@ -216,7 +229,13 @@ def load_database_dump(dump_filename_path: str, database: str) -> str:
with get_session(database) as neo4j_session:
with open(dump_filename_path, "r", encoding="utf-8") as f:
for batch in chunks(f, BATCH_SIZE):
neo4j_session.run(query, {"lines": batch}).consume()
neo4j_session.run(
query,
{
"prowler_provider_id": provider_id,
"lines": batch,
}
).consume()
cartography_create_indexes.run(neo4j_session, None)
attack_paths_prowler.create_indexes(neo4j_session)
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -141,18 +141,15 @@ def run(tenant_id: str, scan_id: str, task_id: str) -> dict[str, Any]:
graph_database.drop_database(cartography_config.neo4j_database)
db_utils.update_attack_paths_scan_progress(attack_paths_scan, 98)
logger.info(f"Dump created at {dump_filename_path} for Attack Paths scan {attack_paths_scan.id}")
# Load dump into tenant's main database
graph_database.create_database(tenant_database)
graph_database.drop_subgraph(tenant_database, root_node_label, root_node_id)
graph_database.create_database(tenant_database) # Create it if not exists
graph_database.drop_subgraph(tenant_database, root_node_label, root_node_id, str(prowler_api_provider.id))
db_utils.update_attack_paths_scan_progress(attack_paths_scan, 99)
logger.info(f"Tenant database {tenant_database} ready, loading dump now")
graph_database.load_database_dump(dump_filename_path, tenant_database)
graph_database.load_database_dump(dump_filename_path, tenant_database, str(prowler_api_provider.id))
logger.info(f"Dump loaded into tenant database {tenant_database} for Attack Paths scan {attack_paths_scan.id}")
db_utils.finish_attack_paths_scan(
+1 -1
View File
@@ -58,7 +58,7 @@ def delete_provider(tenant_id: str, pk: str):
tenant_graph_database = graph_database.get_database_name(tenant_id)
root_node_label = providers.get_root_node_label(instance.provider)
root_node_id = str(instance.uid)
graph_database.drop_subgraph(tenant_graph_database, root_node_label, root_node_id)
graph_database.drop_subgraph(tenant_graph_database, root_node_label, root_node_id, str(instance.id))
# Finally, delete the provider instance itself
try:
+5
View File
@@ -164,6 +164,11 @@ def perform_scan_task(
Returns:
dict: The result of the scan execution, typically including the status and results of the performed checks.
"""
perform_attack_paths_scan_task.apply_async(
kwargs={"tenant_id": tenant_id, "scan_id": scan_id}
)
return # TODO: Delete this block
result = perform_prowler_scan(
tenant_id=tenant_id,
scan_id=scan_id,
@@ -157,8 +157,8 @@ class TestAttackPathsRun:
]
mock_create_dump.assert_called_once_with("temp-db", "AWSAccount", str(provider.uid))
mock_drop_db.assert_called_once_with("temp-db")
mock_drop_subgraph.assert_called_once_with("tenant-db", "AWSAccount", str(provider.uid))
mock_load_dump.assert_called_once_with("/tmp/dump", "tenant-db")
mock_drop_subgraph.assert_called_once_with("tenant-db", "AWSAccount", str(provider.uid), str(provider.id))
mock_load_dump.assert_called_once_with("/tmp/dump", "tenant-db", str(provider.id))
def test_run_failure_marks_scan_failed(
self, tenants_fixture, providers_fixture, scans_fixture
+3 -1
View File
@@ -28,6 +28,7 @@ class TestDeleteProvider:
"tenant-db",
providers.get_root_node_label(instance.provider),
str(instance.uid),
str(instance.id),
)
def test_delete_provider_does_not_exist(self, tenants_fixture):
@@ -70,7 +71,8 @@ class TestDeleteTenant:
call(
f"db-{tenant.id}",
providers.get_root_node_label(provider.provider),
provider.uid,
str(provider.uid),
str(provider.id),
)
for provider in provider_list
]