fix(api): retry transient attack paths graph mutations (#11961)

This commit is contained in:
Josema Camacho
2026-07-14 09:45:52 +02:00
committed by GitHub
parent 1a1e804c00
commit 8debf70d5c
7 changed files with 269 additions and 108 deletions
@@ -0,0 +1 @@
Attack Paths graph mutations now retry transient Neptune concurrency and deadline failures, while Neo4j mutations use managed transaction retries
@@ -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:
@@ -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
@@ -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:
@@ -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
@@ -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()
+121 -92
View File
@@ -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.<sub>.…")` 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