fix(api): scope Attack Paths predefined queries with provider label (#12167)

This commit is contained in:
Daniel Barranquero
2026-07-30 11:55:24 +02:00
committed by GitHub
parent f3b8ac1dbb
commit 5c4b0ba1fe
4 changed files with 199 additions and 1 deletions
@@ -0,0 +1 @@
Attack Paths predefined queries on migrated graphs are now scoped with the provider label, letting the graph database seed from its label index instead of a global label scan and preventing query timeouts on Neptune
@@ -115,7 +115,26 @@ def execute_query(
# TODO: drop after Neptune cutover
# Route reads by the scan row's recorded sink, not by current settings.
backend = sink_module.get_backend_for_scan(scan)
graph = backend.execute_read_query(database_name, definition.cypher, parameters)
cypher = definition.cypher
# Every synced node carries a `_Provider_{uuid}` isolation label (the
# sync labels the whole provider subgraph). Injecting it into the
# predefined query's node patterns gives the planner a selective label
# index to seed from instead of a global label scan (`:AWSRole` across
# every tenant), which on Neptune is the difference between a sub-second
# plan and a query that times out. The custom-query path relies on this
# same injection.
#
# Restrict it to migrated scans: that catalog runs on the Neptune sink
# where the plan blowup happens, while the pre-cutover legacy catalog
# runs on the old sink and is dropped after the cutover, so leave it
# byte-for-byte unchanged. This only affects the query plan, not
# isolation - `_serialize_graph` already label-filters both catalogs.
# TODO: drop the is_migrated guard after Neptune cutover
if scan.is_migrated:
cypher = inject_provider_label(cypher, provider_id)
graph = backend.execute_read_query(database_name, cypher, parameters)
return _serialize_graph(graph, provider_id)
except graph_database.WriteQueryNotAllowedException:
@@ -154,6 +154,88 @@ def test_execute_query_serializes_graph(
assert result["relationships"][0]["label"] == "OWNS"
def test_execute_query_injects_provider_label_when_migrated(
attack_paths_query_definition_factory,
sink_backend_stub,
):
# On migrated graphs the predefined cypher must be scoped with the
# provider label so the planner seeds from the label index instead of a
# global label scan (the Neptune cartesian/timeout fix).
definition = attack_paths_query_definition_factory(
id="aws-iam",
name="IAM",
short_description="Short desc",
description="",
cypher="MATCH (aws:AWSAccount)--(target_role:AWSRole) RETURN target_role",
parameters=[],
)
provider_id = "test-provider-123"
plabel = get_provider_label(provider_id)
parameters = {"provider_uid": "123"}
graph_result = MagicMock()
graph_result.nodes = []
graph_result.relationships = []
sink_backend_stub.execute_read_query.return_value = graph_result
# Injection is gated on `is_migrated`, not the sink (it is a pure string
# transform), so `neo4j` exercises the same code path as Neptune here.
views_helpers.execute_query(
"db-tenant-test",
definition,
parameters,
provider_id=provider_id,
scan=MagicMock(is_migrated=True, sink_backend="neo4j"),
)
executed_cypher = sink_backend_stub.execute_read_query.call_args[0][1]
assert executed_cypher != definition.cypher
# Both node patterns are scoped - not just one. Asserting the exact rewrite
# (rather than `f":{plabel}" in executed_cypher`, which a partial injection
# would still satisfy) proves every node got the label and that injection
# inserted labels and nothing else.
assert executed_cypher == (
f"MATCH (aws:AWSAccount:{plabel})--(target_role:AWSRole:{plabel}) "
"RETURN target_role"
)
# Parameters are passed through untouched.
assert sink_backend_stub.execute_read_query.call_args[0][2] == parameters
def test_execute_query_does_not_inject_label_when_deprecated(
attack_paths_query_definition_factory,
sink_backend_stub,
):
# The pre-cutover legacy catalog runs on the old sink and is removed after
# the Neptune cutover, so it must run verbatim (no injection).
definition = attack_paths_query_definition_factory(
id="aws-iam",
name="IAM",
short_description="Short desc",
description="",
cypher="MATCH (aws:AWSAccount)--(target_role:AWSRole) RETURN target_role",
parameters=[],
)
parameters = {"provider_uid": "123"}
graph_result = MagicMock()
graph_result.nodes = []
graph_result.relationships = []
sink_backend_stub.execute_read_query.return_value = graph_result
views_helpers.execute_query(
"db-tenant-test",
definition,
parameters,
provider_id="test-provider-123",
scan=MagicMock(is_migrated=False, sink_backend="neo4j"),
)
sink_backend_stub.execute_read_query.assert_called_once_with(
"db-tenant-test", definition.cypher, parameters
)
def test_execute_query_wraps_graph_errors(
attack_paths_query_definition_factory,
sink_backend_stub,
@@ -1,5 +1,6 @@
"""Unit tests for the Cypher sanitizer (validation + provider-label injection)."""
import re
from unittest.mock import patch
import pytest
@@ -22,6 +23,38 @@ def _inject(cypher: str) -> str:
return inject_provider_label(cypher, PROVIDER_ID)
# String literals and line comments can contain parentheses that look like node
# patterns; strip them first. Implemented here independently of the sanitizer so
# the node count is an oracle for the injector rather than a copy of its regexes.
_STRING_OR_COMMENT_RE = re.compile(r"'(?:[^'\\]|\\.)*'|\"(?:[^\"\\]|\\.)*\"|//[^\n]*")
# A node pattern is `(`, not preceded by a word char (which would make it a
# function call), wrapping an optional variable, zero or more `:Label`s and an
# optional `{property map}` - and nothing else, which excludes parenthesized
# expressions such as `(a OR b)` in a WHERE clause.
_NODE_PATTERN_RE = re.compile(
r"(?<![\w`])\("
r"\s*(?:[a-zA-Z_]\w*)?"
r"(?:\s*:\s*(?:`[^`]*`|[a-zA-Z_]\w*))*"
r"(?:\s*\{[^{}]*\})?"
r"\s*\)"
)
def _count_node_patterns(cypher: str) -> int:
"""Count node patterns in a query, independently of the injector.
Injection appends exactly one provider label per node pattern, so the
number of injected labels must equal this count - proving *every* node is
scoped, not just one."""
stripped = _STRING_OR_COMMENT_RE.sub("", cypher)
return sum(
1
for match in _NODE_PATTERN_RE.finditer(stripped)
if match.group(0)[1:-1].strip()
)
def test_generic_inject_label_reuses_provider_injection_pipeline():
result = inject_label("MATCH (n:AWSRole)--(m) RETURN n, m", "_Tenant_test")
@@ -427,3 +460,66 @@ class TestValidation:
)
def test_allows_clean_queries(self, cypher):
validate_custom_query(cypher)
# ---------------------------------------------------------------------------
# Predefined-catalog injection (Option 1: label-scoped predefined queries)
# ---------------------------------------------------------------------------
def _all_predefined_queries():
"""Every predefined query in the migrated catalog, as (id, cypher)."""
from api.attack_paths.queries.registry import _QUERY_DEFINITIONS
return [
(definition.id, definition.cypher)
for definitions in _QUERY_DEFINITIONS.values()
for definition in definitions
]
_PREDEFINED_QUERIES = _all_predefined_queries()
class TestPredefinedCatalogInjection:
"""`execute_query` injects the provider label into predefined queries on
migrated graphs. The injection must be *lossless* for every catalog query:
it may only insert `:_Provider_{uuid}` tokens and must not otherwise alter
the cypher (which would corrupt a hand-authored query). This runs over the
whole catalog so a regex regression is caught for all queries at once.
Injection is a pure string transform, so it is sink-independent (the same
result is sent to Neo4j and Neptune)."""
def test_catalog_is_not_empty(self):
# Guard against the parametrized tests silently covering nothing.
assert len(_PREDEFINED_QUERIES) > 0
@pytest.mark.parametrize(
"cypher",
[cypher for _, cypher in _PREDEFINED_QUERIES],
ids=[query_id for query_id, _ in _PREDEFINED_QUERIES],
)
def test_injection_is_lossless(self, cypher):
injected = _inject(cypher)
# Every node pattern is scoped - not just one. A partial-injection
# regression that missed some nodes would still satisfy a bare
# `f":{LABEL}" in injected` check, so assert the label count matches the
# number of node patterns.
assert injected.count(f":{LABEL}") == _count_node_patterns(cypher)
# Stripping the injected tokens restores the query verbatim, proving
# injection changed nothing but the labels.
assert injected.replace(f":{LABEL}", "") == cypher
@pytest.mark.parametrize(
"cypher",
[cypher for _, cypher in _PREDEFINED_QUERIES],
ids=[query_id for query_id, _ in _PREDEFINED_QUERIES],
)
def test_injection_preserves_parameter_placeholders(self, cypher):
# Label injection must never touch `$param` bindings.
original_params = sorted(set(re.findall(r"\$\w+", cypher)))
injected_params = sorted(set(re.findall(r"\$\w+", _inject(cypher))))
assert injected_params == original_params