diff --git a/api/changelog.d/attack-paths-graph-mutation-retries.fixed.md b/api/changelog.d/attack-paths-graph-mutation-retries.fixed.md new file mode 100644 index 0000000000..f48e99462c --- /dev/null +++ b/api/changelog.d/attack-paths-graph-mutation-retries.fixed.md @@ -0,0 +1 @@ +Attack Paths graph mutations now retry transient Neptune concurrency and deadline failures, while Neo4j mutations use managed transaction retries diff --git a/api/src/backend/api/attack_paths/retryable_session.py b/api/src/backend/api/attack_paths/retryable_session.py index 16f0d9e31a..2cd32f54ee 100644 --- a/api/src/backend/api/attack_paths/retryable_session.py +++ b/api/src/backend/api/attack_paths/retryable_session.py @@ -1,4 +1,6 @@ import logging +import random +import time from collections.abc import Callable from typing import Any @@ -9,17 +11,19 @@ logger = logging.getLogger(__name__) class RetryableSession: - """ - Wrapper around `neo4j.Session` that retries `neo4j.exceptions.ServiceUnavailable` errors. - """ + """Wrapper around ``neo4j.Session`` with a refreshable retry policy.""" def __init__( self, session_factory: Callable[[], neo4j.Session], max_retries: int, + retry_if: Callable[[Exception], bool] | None = None, + initial_retry_delay_seconds: float = 0, ) -> None: self._session_factory = session_factory self._max_retries = max(0, max_retries) + self._retry_if = retry_if + self._initial_retry_delay_seconds = max(0.0, initial_retry_delay_seconds) self._session = self._session_factory() def close(self) -> None: @@ -56,24 +60,47 @@ class RetryableSession: method = getattr(self._session, method_name) return method(*args, **kwargs) - except ( - BrokenPipeError, - ConnectionResetError, - neo4j.exceptions.ServiceUnavailable, - ) as exc: # pragma: no cover - depends on infra + except Exception as exc: + if not self._should_retry(exc): + raise + last_exc = exc attempt += 1 if attempt > self._max_retries: raise + delay = self._retry_delay(attempt) logger.warning( - f"Neo4j session {method_name} failed with {type(exc).__name__} ({attempt}/{self._max_retries} attempts). Retrying..." + "Graph session %s failed with %s; retry %s/%s in %.3fs", + method_name, + type(exc).__name__, + attempt, + self._max_retries, + delay, ) self._refresh_session() + if delay: + time.sleep(delay) raise last_exc if last_exc else RuntimeError("Unexpected retry loop exit") + def _should_retry(self, exc: Exception) -> bool: + if isinstance( + exc, + ( + BrokenPipeError, + ConnectionResetError, + neo4j.exceptions.ServiceUnavailable, + ), + ): + return True + return self._retry_if(exc) if self._retry_if else False + + def _retry_delay(self, attempt: int) -> float: + max_delay = self._initial_retry_delay_seconds * (2**attempt) + return random.uniform(max_delay / 2, max_delay) if max_delay else 0 + def _refresh_session(self) -> None: if self._session is not None: try: diff --git a/api/src/backend/api/attack_paths/sink/drop.py b/api/src/backend/api/attack_paths/sink/drop.py index 9b4044a8f0..be13a394f5 100644 --- a/api/src/backend/api/attack_paths/sink/drop.py +++ b/api/src/backend/api/attack_paths/sink/drop.py @@ -42,6 +42,10 @@ def delete_batches( batch_size: int, drop_t0: float, ) -> tuple[int, int]: + def delete_batch(tx: Any) -> int: + record = tx.run(query, {"batch_size": batch_size}).single() + return (record[count_key] if record else 0) or 0 + deleted_total = initial_total batches = 0 while True: @@ -56,8 +60,7 @@ def delete_batches( deleted_total, time.perf_counter() - drop_t0, ) - record = session.run(query, {"batch_size": batch_size}).single() - deleted = (record[count_key] if record else 0) or 0 + deleted = session.execute_write(delete_batch) if deleted == 0: return deleted_total, batches diff --git a/api/src/backend/api/attack_paths/sink/neo4j.py b/api/src/backend/api/attack_paths/sink/neo4j.py index c248237f01..cb7d4889b0 100644 --- a/api/src/backend/api/attack_paths/sink/neo4j.py +++ b/api/src/backend/api/attack_paths/sink/neo4j.py @@ -355,7 +355,7 @@ class Neo4jSink(SinkDatabase): f"ON (n.`{PROVIDER_ELEMENT_ID_PROPERTY}`)" ) with self.get_session(database) as session: - session.run(query).consume() + session.execute_write(lambda tx: tx.run(query).consume()) def write_nodes( self, @@ -377,7 +377,7 @@ class Neo4jSink(SinkDatabase): SET n += row.props """ with self.get_session(database) as session: - session.run(query, {"rows": rows}).consume() + session.execute_write(lambda tx: tx.run(query, {"rows": rows}).consume()) def write_relationships( self, @@ -403,7 +403,7 @@ class Neo4jSink(SinkDatabase): SET r += row.props """ with self.get_session(database) as session: - session.run(query, {"rows": rows}).consume() + session.execute_write(lambda tx: tx.run(query, {"rows": rows}).consume()) # For compatibility with test harnesses that patch the concrete driver def get_driver(self) -> neo4j.Driver: diff --git a/api/src/backend/api/attack_paths/sink/neptune.py b/api/src/backend/api/attack_paths/sink/neptune.py index b0d12069a3..dbb687e89c 100644 --- a/api/src/backend/api/attack_paths/sink/neptune.py +++ b/api/src/backend/api/attack_paths/sink/neptune.py @@ -59,17 +59,28 @@ CONNECTION_TIMEOUT = env.int("NEPTUNE_CONNECTION_TIMEOUT", default=10) # Roll connections hourly so SigV4 rotations and cert refreshes don't strand long-lived pool entries MAX_CONNECTION_LIFETIME = env.int("NEPTUNE_MAX_CONNECTION_LIFETIME", default=3600) MAX_CONNECTION_POOL_SIZE = env.int("NEPTUNE_MAX_CONNECTION_POOL_SIZE", default=50) +NEPTUNE_WRITE_RETRY_DELAY_SECONDS = 2 READ_EXCEPTION_CODES = [ "Neo.ClientError.Statement.AccessMode", "Neo.ClientError.Procedure.ProcedureNotFound", ] CLIENT_STATEMENT_EXCEPTION_PREFIX = "Neo.ClientError.Statement." +RETRYABLE_WRITE_ERROR_PREFIXES = ( + "Operation failed due to conflicting concurrent operations", + "Operation terminated (deadline exceeded)", +) # Refresh 60s before the 5-minute SigV4 window closes SIGV4_TOKEN_LIFETIME_MINUTES = 4 +def _is_retryable_write_error(exc: Exception) -> bool: + if not isinstance(exc, neo4j.exceptions.Neo4jError): + return False + return bool(exc.message and exc.message.startswith(RETRYABLE_WRITE_ERROR_PREFIXES)) + + class NeptuneSink(SinkDatabase): """Neptune-backed sink. Single database; isolation is label-based.""" @@ -205,11 +216,16 @@ class NeptuneSink(SinkDatabase): session_wrapper: RetryableSession | None = None try: + is_write_session = default_access_mode != neo4j.READ_ACCESS session_wrapper = RetryableSession( session_factory=lambda: driver.session( default_access_mode=default_access_mode ), max_retries=SERVICE_UNAVAILABLE_MAX_RETRIES, + retry_if=_is_retryable_write_error if is_write_session else None, + initial_retry_delay_seconds=( + NEPTUNE_WRITE_RETRY_DELAY_SECONDS if is_write_session else 0 + ), ) yield session_wrapper @@ -405,7 +421,7 @@ class NeptuneSink(SinkDatabase): SET n.`{PROVIDER_ELEMENT_ID_PROPERTY}` = row.provider_element_id """ with self.get_session() as session: - session.run(query, {"rows": rows}).consume() + session.execute_write(lambda tx: tx.run(query, {"rows": rows}).consume()) def write_relationships( self, @@ -429,7 +445,7 @@ class NeptuneSink(SinkDatabase): SET r += row.props """ with self.get_session() as session: - session.run(query, {"rows": rows}).consume() + session.execute_write(lambda tx: tx.run(query, {"rows": rows}).consume()) # Test helpers diff --git a/api/src/backend/api/tests/test_retryable_session.py b/api/src/backend/api/tests/test_retryable_session.py new file mode 100644 index 0000000000..8bb457b8b6 --- /dev/null +++ b/api/src/backend/api/tests/test_retryable_session.py @@ -0,0 +1,85 @@ +from unittest.mock import MagicMock, patch + +import pytest +from api.attack_paths.retryable_session import RetryableSession +from neo4j.exceptions import ServiceUnavailable + + +class TestRetryableSession: + @patch("api.attack_paths.retryable_session.time.sleep") + @patch("api.attack_paths.retryable_session.random.uniform", return_value=3.0) + def test_custom_retry_uses_backoff_and_a_fresh_session( + self, mock_uniform, mock_sleep + ): + retryable_error = RuntimeError("retryable") + first_session = MagicMock() + first_session.execute_write.side_effect = retryable_error + second_session = MagicMock() + second_session.execute_write.return_value = "success" + session_factory = MagicMock(side_effect=[first_session, second_session]) + work = MagicMock() + + session = RetryableSession( + session_factory=session_factory, + max_retries=3, + retry_if=lambda exc: exc is retryable_error, + initial_retry_delay_seconds=2, + ) + + assert session.execute_write(work) == "success" + assert session_factory.call_count == 2 + first_session.close.assert_called_once_with() + mock_uniform.assert_called_once_with(2.0, 4.0) + mock_sleep.assert_called_once_with(3.0) + + def test_connection_errors_remain_retryable(self): + first_session = MagicMock() + first_session.run.side_effect = ServiceUnavailable("unavailable") + second_session = MagicMock() + second_session.run.return_value = "success" + session_factory = MagicMock(side_effect=[first_session, second_session]) + + session = RetryableSession(session_factory=session_factory, max_retries=1) + + assert session.run("RETURN 1") == "success" + first_session.close.assert_called_once_with() + + def test_non_retryable_error_is_raised_without_refreshing_session(self): + error = RuntimeError("do not retry") + driver_session = MagicMock() + driver_session.execute_write.side_effect = error + session_factory = MagicMock(return_value=driver_session) + session = RetryableSession( + session_factory=session_factory, + max_retries=3, + retry_if=lambda _: False, + initial_retry_delay_seconds=2, + ) + + with pytest.raises(RuntimeError) as exc_info: + session.execute_write(MagicMock()) + + assert exc_info.value is error + session_factory.assert_called_once_with() + driver_session.close.assert_not_called() + + def test_retry_exhaustion_raises_the_last_error(self): + error = RuntimeError("still retryable") + driver_sessions = [MagicMock() for _ in range(3)] + for driver_session in driver_sessions: + driver_session.execute_write.side_effect = error + session_factory = MagicMock(side_effect=driver_sessions) + session = RetryableSession( + session_factory=session_factory, + max_retries=2, + retry_if=lambda _: True, + ) + + with pytest.raises(RuntimeError) as exc_info: + session.execute_write(MagicMock()) + + assert exc_info.value is error + assert session_factory.call_count == 3 + driver_sessions[0].close.assert_called_once_with() + driver_sessions[1].close.assert_called_once_with() + driver_sessions[2].close.assert_not_called() diff --git a/api/src/backend/api/tests/test_sink.py b/api/src/backend/api/tests/test_sink.py index 4bb302d492..3636771ddc 100644 --- a/api/src/backend/api/tests/test_sink.py +++ b/api/src/backend/api/tests/test_sink.py @@ -6,18 +6,20 @@ builds dual writer/reader Bolt drivers. """ import json -from importlib import import_module from unittest.mock import MagicMock, patch +import neo4j import pytest - -# Prime patch-target resolution. `api.attack_paths.sink/__init__.py` doesn't -# eagerly import these submodules (they're loaded on demand inside the -# factory), so `mock.patch("api.attack_paths.sink..…")` would fail with -# AttributeError on first call. Importing here registers them as attributes -# of the package before any decorator runs. -import_module("api.attack_paths.sink.neo4j") -import_module("api.attack_paths.sink.neptune") +from api.attack_paths import sink as sink_module +from api.attack_paths.database import GraphDatabaseQueryException +from api.attack_paths.sink import factory +from api.attack_paths.sink.neo4j import DATABASE_NOT_FOUND_CODE, Neo4jSink +from api.attack_paths.sink.neptune import ( + NEPTUNE_WRITE_RETRY_DELAY_SECONDS, + NeptuneSink, + _is_retryable_write_error, + _NeptuneAuthToken, +) @pytest.fixture(autouse=True) @@ -26,8 +28,6 @@ def reset_sink_state(): The cache lives in `api.attack_paths.sink.factory`, not on the package. """ - from api.attack_paths.sink import factory - original_backend = factory._backend original_secondary = dict(factory._secondary_backends) factory._backend = None @@ -40,29 +40,20 @@ def reset_sink_state(): class TestSinkFactory: def test_default_resolves_to_neo4j(self, settings): - from api.attack_paths.sink import factory - settings.ATTACK_PATHS_SINK_DATABASE = "neo4j" assert factory._resolve_setting() == "neo4j" def test_neptune_resolves_correctly(self, settings): - from api.attack_paths.sink import factory - settings.ATTACK_PATHS_SINK_DATABASE = "neptune" assert factory._resolve_setting() == "neptune" def test_invalid_value_raises(self, settings): - from api.attack_paths.sink import factory - settings.ATTACK_PATHS_SINK_DATABASE = "foo" with pytest.raises(RuntimeError, match="ATTACK_PATHS_SINK_DATABASE"): factory._resolve_setting() @patch("api.attack_paths.sink.neo4j.neo4j.GraphDatabase.driver") def test_init_builds_neo4j_backend_by_default(self, mock_driver, settings): - from api.attack_paths import sink as sink_module - from api.attack_paths.sink.neo4j import Neo4jSink - settings.ATTACK_PATHS_SINK_DATABASE = "neo4j" settings.DATABASES = { **settings.DATABASES, @@ -85,9 +76,6 @@ class TestSinkFactory: def test_init_builds_neptune_backend( self, mock_driver, mock_auth_provider, settings ): - from api.attack_paths import sink as sink_module - from api.attack_paths.sink.neptune import NeptuneSink - settings.ATTACK_PATHS_SINK_DATABASE = "neptune" settings.DATABASES = { **settings.DATABASES, @@ -116,8 +104,6 @@ class TestSinkFactory: def test_neptune_reader_falls_back_to_writer( self, mock_driver, mock_auth_provider, settings ): - from api.attack_paths import sink as sink_module - settings.ATTACK_PATHS_SINK_DATABASE = "neptune" settings.DATABASES = { **settings.DATABASES, @@ -144,8 +130,6 @@ class TestGetBackendForScan: def test_legacy_scan_in_neo4j_process_uses_active_backend( self, mock_driver, settings ): - from api.attack_paths import sink as sink_module - settings.ATTACK_PATHS_SINK_DATABASE = "neo4j" settings.DATABASES = { **settings.DATABASES, @@ -164,8 +148,6 @@ class TestGetBackendForScan: assert backend is sink_module.get_backend() def test_neptune_scan_on_neo4j_process_uses_neptune_secondary(self, settings): - from api.attack_paths.sink import factory - settings.ATTACK_PATHS_SINK_DATABASE = "neo4j" active_neo4j = MagicMock(name="neo4j-active") factory._backend = active_neo4j @@ -190,6 +172,29 @@ def _count_result(key: str, count: int) -> MagicMock: return MagicMock(single=MagicMock(return_value={key: count})) +def _run_managed_write(session: MagicMock) -> MagicMock: + transaction = MagicMock() + session.execute_write.call_args.args[0](transaction) + return transaction + + +def _managed_write_session( + results: list[MagicMock], +) -> tuple[MagicMock, list[MagicMock]]: + session = MagicMock() + transactions: list[MagicMock] = [] + result_iter = iter(results) + + def execute_write(work): + transaction = MagicMock() + transaction.run.return_value = next(result_iter) + transactions.append(transaction) + return work(transaction) + + session.execute_write.side_effect = execute_write + return session, transactions + + def _directed_drop_results( outgoing_rels: int, incoming_rels: int, @@ -207,31 +212,26 @@ def _directed_drop_results( class TestNeo4jSinkSyncWrites: def test_ensure_sync_indexes_runs_create_index_idempotent(self): - from api.attack_paths.sink.neo4j import Neo4jSink - sink = Neo4jSink() session = MagicMock() - session.run.return_value = MagicMock() with patch.object(sink, "get_session", return_value=_session_ctx(session)): sink.ensure_sync_indexes("db-tenant-x") - query = session.run.call_args.args[0] + transaction = _run_managed_write(session) + query = transaction.run.call_args.args[0] assert "CREATE INDEX" in query assert "IF NOT EXISTS" in query assert "`_ProviderResource`" in query assert "`_provider_element_id`" in query + transaction.run.return_value.consume.assert_called_once_with() def test_write_nodes_skips_empty_batch(self): - from api.attack_paths.sink.neo4j import Neo4jSink - sink = Neo4jSink() with patch.object(sink, "get_session") as get_session: sink.write_nodes("db-tenant-x", "`AWSUser`", []) get_session.assert_not_called() def test_write_nodes_merges_on_provider_resource_label(self): - from api.attack_paths.sink.neo4j import Neo4jSink - sink = Neo4jSink() session = MagicMock() with patch.object(sink, "get_session", return_value=_session_ctx(session)): @@ -241,15 +241,15 @@ class TestNeo4jSinkSyncWrites: [{"provider_element_id": "p:e", "props": {"k": "v"}}], ) - query, params = session.run.call_args.args + transaction = _run_managed_write(session) + query, params = transaction.run.call_args.args assert "MERGE (n:`_ProviderResource`" in query assert "`_provider_element_id`: row.provider_element_id" in query assert "SET n:`AWSUser`:`_ProviderResource`" in query assert params == {"rows": [{"provider_element_id": "p:e", "props": {"k": "v"}}]} + transaction.run.return_value.consume.assert_called_once_with() def test_write_relationships_scopes_endpoints_by_provider_label(self): - from api.attack_paths.sink.neo4j import Neo4jSink - sink = Neo4jSink() session = MagicMock() provider_id = "00000000-0000-0000-0000-000000000abc" @@ -268,24 +268,22 @@ class TestNeo4jSinkSyncWrites: ], ) - query = session.run.call_args.args[0] + transaction = _run_managed_write(session) + query = transaction.run.call_args.args[0] assert ":`_Provider_00000000000000000000000000000abc`" in query assert ":RESOURCE" in query.replace("`", "") assert "MERGE (s)-[r:`RESOURCE`" in query + transaction.run.return_value.consume.assert_called_once_with() class TestNeptuneSinkSyncWrites: def test_ensure_sync_indexes_is_noop(self): - from api.attack_paths.sink.neptune import NeptuneSink - sink = NeptuneSink() with patch.object(sink, "get_session") as get_session: sink.ensure_sync_indexes("ignored") get_session.assert_not_called() def test_write_nodes_merges_on_neptune_id_with_provider_resource_label(self): - from api.attack_paths.sink.neptune import NeptuneSink - sink = NeptuneSink() session = MagicMock() with patch.object(sink, "get_session", return_value=_session_ctx(session)): @@ -295,16 +293,16 @@ class TestNeptuneSinkSyncWrites: [{"provider_element_id": "p:e", "props": {"k": "v"}}], ) - query = session.run.call_args.args[0] + transaction = _run_managed_write(session) + query = transaction.run.call_args.args[0] # Neptune assigns a default `vertex` label to any unlabeled node, # so the MERGE must pin a real label at creation time. assert "MERGE (n:`_ProviderResource` {`~id`: row.provider_element_id})" in query assert "SET n:`AWSUser`" in query assert "SET n.`_provider_element_id` = row.provider_element_id" in query + transaction.run.return_value.consume.assert_called_once_with() def test_write_relationships_matches_endpoints_by_id(self): - from api.attack_paths.sink.neptune import NeptuneSink - sink = NeptuneSink() session = MagicMock() with patch.object(sink, "get_session", return_value=_session_ctx(session)): @@ -322,30 +320,86 @@ class TestNeptuneSinkSyncWrites: ], ) - query = session.run.call_args.args[0] + transaction = _run_managed_write(session) + query = transaction.run.call_args.args[0] assert "MATCH (s) WHERE id(s) = row.start_element_id" in query assert "MATCH (e) WHERE id(e) = row.end_element_id" in query assert "MERGE (s)-[r:`RESOURCE`" in query + transaction.run.return_value.consume.assert_called_once_with() + + +class TestNeptuneRetryPolicy: + @pytest.mark.parametrize( + "message", + [ + "Operation failed due to conflicting concurrent operations " + + "(please retry), 0 transactions are currently rolling back.", + "Operation terminated (deadline exceeded)", + ], + ) + def test_observed_transient_write_errors_are_retryable(self, message): + error = MagicMock(spec=neo4j.exceptions.Neo4jError) + error.message = message + + assert _is_retryable_write_error(error) is True + + def test_unrelated_database_error_is_not_retryable(self): + error = MagicMock(spec=neo4j.exceptions.Neo4jError) + error.message = "Operation terminated (out of memory)" + + assert _is_retryable_write_error(error) is False + + def test_non_neo4j_error_is_not_retryable(self): + error = RuntimeError( + "Operation failed due to conflicting concurrent operations" + ) + + assert _is_retryable_write_error(error) is False + + @patch("api.attack_paths.sink.neptune.RetryableSession") + def test_writer_session_enables_neptune_retry_policy(self, retryable_session): + sink = NeptuneSink() + driver = MagicMock() + with patch.object(sink, "_get_writer", return_value=driver): + with sink.get_session(): + pass + + kwargs = retryable_session.call_args.kwargs + assert kwargs["retry_if"] is _is_retryable_write_error + assert ( + kwargs["initial_retry_delay_seconds"] == NEPTUNE_WRITE_RETRY_DELAY_SECONDS + ) + + @patch("api.attack_paths.sink.neptune.RetryableSession") + def test_reader_session_does_not_enable_write_retry_policy(self, retryable_session): + sink = NeptuneSink() + driver = MagicMock() + with patch.object(sink, "_get_reader", return_value=driver): + with sink.get_session(default_access_mode=neo4j.READ_ACCESS): + pass + + kwargs = retryable_session.call_args.kwargs + assert kwargs["retry_if"] is None + assert kwargs["initial_retry_delay_seconds"] == 0 class TestNeptuneSinkDropSubgraph: def test_drop_subgraph_deletes_directed_rels_before_nodes_in_bounded_batches(self): - from api.attack_paths.sink.neptune import NeptuneSink - sink = NeptuneSink() - session = MagicMock() - session.run.side_effect = _directed_drop_results( - outgoing_rels=50, - incoming_rels=30, - nodes=10, + session, transactions = _managed_write_session( + _directed_drop_results( + outgoing_rels=50, + incoming_rels=30, + nodes=10, + ) ) with patch.object(sink, "get_session", return_value=_session_ctx(session)): deleted = sink.drop_subgraph("ignored", "provider-1") assert deleted == 10 - assert session.run.call_count == 6 - queries = [call.args[0] for call in session.run.call_args_list] + assert session.execute_write.call_count == 6 + queries = [transaction.run.call_args.args[0] for transaction in transactions] assert ")-[r]->()" in queries[0] assert ")<-[r]-()" in queries[2] @@ -362,14 +416,13 @@ class TestNeo4jSinkDropSubgraph: """Neo4j drop deletes relationships then nodes in batches (no ``DETACH DELETE``).""" def test_drop_subgraph_deletes_directed_rels_before_nodes_in_bounded_batches(self): - from api.attack_paths.sink.neo4j import Neo4jSink - sink = Neo4jSink() - session = MagicMock() - session.run.side_effect = _directed_drop_results( - outgoing_rels=50, - incoming_rels=30, - nodes=10, + session, transactions = _managed_write_session( + _directed_drop_results( + outgoing_rels=50, + incoming_rels=30, + nodes=10, + ) ) provider_id = "00000000-0000-0000-0000-000000000abc" @@ -378,9 +431,9 @@ class TestNeo4jSinkDropSubgraph: # Only phase-2 node counts contribute to the return value. assert deleted == 10 - assert session.run.call_count == 6 + assert session.execute_write.call_count == 6 - queries = [call.args[0] for call in session.run.call_args_list] + queries = [transaction.run.call_args.args[0] for transaction in transactions] # Regression guard: the memory blow-up was caused by DETACH DELETE. assert all("DETACH DELETE" not in query for query in queries) assert all("DISTINCT r" not in query for query in queries) @@ -399,12 +452,9 @@ class TestNeo4jSinkDropSubgraph: assert last_rel < first_node def test_drop_subgraph_returns_zero_when_database_does_not_exist(self): - from api.attack_paths.database import GraphDatabaseQueryException - from api.attack_paths.sink.neo4j import DATABASE_NOT_FOUND_CODE, Neo4jSink - sink = Neo4jSink() session = MagicMock() - session.run.side_effect = GraphDatabaseQueryException( + session.execute_write.side_effect = GraphDatabaseQueryException( message="db missing", code=DATABASE_NOT_FOUND_CODE ) @@ -418,8 +468,6 @@ class TestSinkHasProviderData: """``has_provider_data`` is the read-path probe used by API views.""" def test_neo4j_returns_true_when_provider_node_exists(self): - from api.attack_paths.sink.neo4j import Neo4jSink - sink = Neo4jSink() session = MagicMock() session.run.return_value.single.return_value = MagicMock() @@ -433,9 +481,6 @@ class TestSinkHasProviderData: assert ":`_Provider_00000000000000000000000000000abc`" in query def test_neo4j_returns_false_when_database_does_not_exist(self): - from api.attack_paths.database import GraphDatabaseQueryException - from api.attack_paths.sink.neo4j import DATABASE_NOT_FOUND_CODE, Neo4jSink - sink = Neo4jSink() session = MagicMock() session.run.side_effect = GraphDatabaseQueryException( @@ -448,8 +493,6 @@ class TestSinkHasProviderData: assert present is False def test_neptune_returns_true_when_provider_node_exists(self): - from api.attack_paths.sink.neptune import NeptuneSink - sink = NeptuneSink() session = MagicMock() session.run.return_value.single.return_value = MagicMock() @@ -463,8 +506,6 @@ class TestGetBackendForScanCutover: """``get_backend_for_scan`` keeps old-sink scans queryable after cutover.""" def test_legacy_scan_on_neptune_process_uses_neo4j_secondary(self, settings): - from api.attack_paths.sink import factory - settings.ATTACK_PATHS_SINK_DATABASE = "neptune" active_neptune = MagicMock(name="neptune-active") factory._backend = active_neptune @@ -487,8 +528,6 @@ class TestSinkVerifyConnectivity: @patch("api.attack_paths.sink.neo4j.neo4j.GraphDatabase.driver") def test_neo4j_verifies_its_driver(self, mock_driver, settings): - from api.attack_paths.sink.neo4j import Neo4jSink - settings.DATABASES = { **settings.DATABASES, "neo4j": { @@ -513,8 +552,6 @@ class TestSinkVerifyConnectivity: def test_neptune_verifies_reader_not_writer( self, mock_driver, mock_auth_provider, settings ): - from api.attack_paths.sink.neptune import NeptuneSink - settings.DATABASES = { **settings.DATABASES, "neptune": { @@ -548,8 +585,6 @@ class TestSinkInitToleratesUnreachableSink: @patch("api.attack_paths.sink.neo4j.neo4j.GraphDatabase.driver") def test_neo4j_init_continues_when_verify_fails(self, mock_driver, settings): - from api.attack_paths.sink.neo4j import Neo4jSink - settings.DATABASES = { **settings.DATABASES, "neo4j": { @@ -573,8 +608,6 @@ class TestSinkInitToleratesUnreachableSink: def test_neptune_init_continues_when_verify_fails( self, mock_driver, mock_auth_provider, settings ): - from api.attack_paths.sink.neptune import NeptuneSink - settings.DATABASES = { **settings.DATABASES, "neptune": { @@ -601,8 +634,6 @@ class TestNeptuneAdminNoOps: @pytest.mark.parametrize("method", ["create_database", "drop_database"]) def test_admin_ops_return_none_without_touching_a_session(self, method): - from api.attack_paths.sink.neptune import NeptuneSink - sink = NeptuneSink() with patch.object(sink, "get_session") as get_session: assert getattr(sink, method)("ignored") is None @@ -617,8 +648,6 @@ class TestNeptuneAuthToken: def test_host_header_includes_non_default_port(self, mock_boto, mock_sigv4): # Neptune runs on 8182; the SigV4 canonical Host must keep the port or # the signature is rejected. - from api.attack_paths.sink.neptune import _NeptuneAuthToken - credentials = MagicMock() credentials.get_frozen_credentials.return_value = MagicMock() mock_boto.return_value.get_credentials.return_value = credentials