From db125b9c1fe5efebe0416888d31e329e77f04e47 Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 24 Sep 2026 17:52:14 +0100 Subject: [PATCH 01/13] feat(graph): add node_service to upsert graph nodes and link findings Signed-off-by: Tanvir Farhad --- scanner/graph/node_service.py | 85 +++++++++++++++++++++++ tests/test_graph_node_service.py | 111 +++++++++++++++++++++++++++++++ 2 files changed, 196 insertions(+) create mode 100644 scanner/graph/node_service.py create mode 100644 tests/test_graph_node_service.py diff --git a/scanner/graph/node_service.py b/scanner/graph/node_service.py new file mode 100644 index 00000000..c1154c1a --- /dev/null +++ b/scanner/graph/node_service.py @@ -0,0 +1,85 @@ +"""Upsert graph nodes from an InventorySnapshot and link findings to nodes.""" +from __future__ import annotations + +import json +import logging +import uuid +from typing import TYPE_CHECKING + +import psycopg2 +import psycopg2.extras + +if TYPE_CHECKING: + from scanner.arg_inventory import InventorySnapshot + +logger = logging.getLogger(__name__) + +_UPSERT_NODE_SQL = """ +INSERT INTO graph_nodes ( + node_id, tenant_id, subscription_id, resource_id, resource_type, + name, location, resource_group, snapshot_id, properties, created_at, updated_at +) +VALUES ( + %(node_id)s, %(tenant_id)s, %(subscription_id)s, %(resource_id)s, %(resource_type)s, + %(name)s, %(location)s, %(resource_group)s, %(snapshot_id)s, %(properties)s, + now(), now() +) +ON CONFLICT (tenant_id, resource_id) +DO UPDATE SET + resource_type = EXCLUDED.resource_type, + name = EXCLUDED.name, + location = EXCLUDED.location, + resource_group = EXCLUDED.resource_group, + snapshot_id = EXCLUDED.snapshot_id, + properties = EXCLUDED.properties, + updated_at = now() +""" + +_LINK_FINDINGS_SQL = """ +INSERT INTO finding_graph_nodes (finding_id, node_id) +SELECT f.id, n.node_id +FROM findings f +JOIN graph_nodes n ON lower(f.resource_id) = lower(n.resource_id) + AND n.tenant_id = ( + SELECT tenant_id FROM graph_nodes + WHERE snapshot_id = ( + SELECT snapshot_id FROM graph_nodes ORDER BY updated_at DESC LIMIT 1 + ) + LIMIT 1 + ) +WHERE f.scan_id = %(scan_id)s +ON CONFLICT DO NOTHING +""" + + +def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: + """Upsert graph_nodes from snapshot resources. Returns count of rows written.""" + if not snapshot.resources: + return 0 + + written = 0 + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + for resource in snapshot.resources: + cur.execute(_UPSERT_NODE_SQL, { + "node_id": str(uuid.uuid4()), + "tenant_id": resource.tenant_id, + "subscription_id": resource.subscription_id, + "resource_id": resource.resource_id, + "resource_type": resource.resource_type, + "name": resource.name, + "location": resource.location, + "resource_group": resource.resource_group, + "snapshot_id": resource.snapshot_id, + "properties": json.dumps(resource.properties), + }) + written += 1 + return written + + +def link_findings_to_nodes(scan_id: str, dsn: str) -> int: + """Link findings from this scan to their graph nodes by resource_id.""" + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute(_LINK_FINDINGS_SQL, {"scan_id": scan_id}) + return cur.rowcount if cur.rowcount >= 0 else 0 diff --git a/tests/test_graph_node_service.py b/tests/test_graph_node_service.py new file mode 100644 index 00000000..12f70b8b --- /dev/null +++ b/tests/test_graph_node_service.py @@ -0,0 +1,111 @@ +"""Tests for graph node upsert and finding-to-node linking.""" +from unittest.mock import MagicMock, patch + +from scanner.arg_inventory import InventorySnapshot, InventoryStatus, InventoryResource +from scanner.graph.node_service import populate_nodes, link_findings_to_nodes + + +def _make_resource(resource_id: str, resource_type: str = "microsoft.network/virtualnetworks", + subscription_id: str = "00000000-0000-0000-0000-000000000002") -> InventoryResource: + return InventoryResource( + snapshot_id="snap-1", + tenant_id="00000000-0000-0000-0000-000000000001", + subscription_id=subscription_id, + resource_id=resource_id, + resource_type=resource_type, + name=resource_id.split("/")[-1], + location="eastus", + resource_group="rg", + tags={}, + properties={"key": "val"}, + ) + + +def _make_snapshot(resources=()): + return InventorySnapshot( + snapshot_id="snap-1", + tenant_id="00000000-0000-0000-0000-000000000001", + requested_subscriptions=("00000000-0000-0000-0000-000000000002",), + status=InventoryStatus.COMPLETE, + collected_at="2026-09-24T00:00:00+00:00", + duration_ms=10, + pages=1, + resources=tuple(resources), + errors=(), + ) + + +def test_populate_nodes_upserts_each_resource(): + resources = [ + _make_resource("/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet1"), + _make_resource("/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet2"), + ] + snapshot = _make_snapshot(resources) + + mock_conn = MagicMock() + mock_cur = MagicMock() + mock_conn.__enter__ = MagicMock(return_value=mock_conn) + mock_conn.__exit__ = MagicMock(return_value=False) + mock_cur.__enter__ = MagicMock(return_value=mock_cur) + mock_cur.__exit__ = MagicMock(return_value=False) + mock_conn.cursor.return_value = mock_cur + + with patch("scanner.graph.node_service.psycopg2.connect", return_value=mock_conn): + count = populate_nodes(snapshot, "postgresql://test/db") + + assert count == 2 + assert mock_cur.execute.call_count == 2 + + +def test_populate_nodes_empty_snapshot_returns_zero(): + snapshot = _make_snapshot(resources=[]) + + with patch("scanner.graph.node_service.psycopg2.connect") as mock_connect: + count = populate_nodes(snapshot, "postgresql://test/db") + + assert count == 0 + mock_connect.assert_not_called() + + +def test_populate_nodes_does_not_write_cross_tenant_resources(): + resource = _make_resource("/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet1") + snapshot = _make_snapshot([resource]) + + captured_args = [] + mock_conn = MagicMock() + mock_cur = MagicMock() + mock_conn.__enter__ = MagicMock(return_value=mock_conn) + mock_conn.__exit__ = MagicMock(return_value=False) + mock_cur.__enter__ = MagicMock(return_value=mock_cur) + mock_cur.__exit__ = MagicMock(return_value=False) + mock_conn.cursor.return_value = mock_cur + + def capture_execute(sql, params=None): + if params: + captured_args.append(params) + + mock_cur.execute.side_effect = capture_execute + + with patch("scanner.graph.node_service.psycopg2.connect", return_value=mock_conn): + populate_nodes(snapshot, "postgresql://test/db") + + for params in captured_args: + assert "00000000-0000-0000-0000-000000000001" in str(params), "tenant_id must be in every write" + + +def test_link_findings_to_nodes_executes_insert(): + mock_conn = MagicMock() + mock_cur = MagicMock() + mock_conn.__enter__ = MagicMock(return_value=mock_conn) + mock_conn.__exit__ = MagicMock(return_value=False) + mock_cur.__enter__ = MagicMock(return_value=mock_cur) + mock_cur.__exit__ = MagicMock(return_value=False) + mock_conn.cursor.return_value = mock_cur + mock_cur.rowcount = 3 + + with patch("scanner.graph.node_service.psycopg2.connect", return_value=mock_conn): + link_findings_to_nodes("scan-uuid-1", "postgresql://test/db") + + assert mock_cur.execute.called + sql_called = mock_cur.execute.call_args[0][0] + assert "finding_graph_nodes" in sql_called From 4d195fd9cdf9dee803cc734ac39446bf209c794b Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 24 Sep 2026 17:54:17 +0100 Subject: [PATCH 02/13] feat(graph): add edge detectors for NSG, subnet, public IP, identity, and storage Signed-off-by: Tanvir Farhad --- scanner/graph/edge_detector.py | 173 ++++++++++++++++++++++++++++++ tests/test_graph_edge_detector.py | 120 +++++++++++++++++++++ 2 files changed, 293 insertions(+) create mode 100644 scanner/graph/edge_detector.py create mode 100644 tests/test_graph_edge_detector.py diff --git a/scanner/graph/edge_detector.py b/scanner/graph/edge_detector.py new file mode 100644 index 00000000..61cf0ad5 --- /dev/null +++ b/scanner/graph/edge_detector.py @@ -0,0 +1,173 @@ +"""Typed edge detectors that infer relationships between Azure resources in an InventorySnapshot.""" +from __future__ import annotations + +import logging +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from scanner.arg_inventory import InventorySnapshot + +logger = logging.getLogger(__name__) + +_NSG_TYPE = "microsoft.network/networksecuritygroups" +_PUBLIC_IP_TYPE = "microsoft.network/publicipaddresses" +_STORAGE_TYPE = "microsoft.storage/storageaccounts" + + +@dataclass +class GraphEdge: + """A directed relationship between two Azure resources.""" + + source_resource_id: str + target_resource_id: str + relationship_type: str + evidence_source: str + confidence: float + + +class EdgeDetector(ABC): + """Base class for relationship detectors.""" + + @abstractmethod + def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: + """Return edges inferred from this snapshot.""" + + +class NsgToSubnetDetector(EdgeDetector): + """NSG -> Subnet: PROTECTS (ARG-confirmed, confidence 1.0).""" + + def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: + edges = [] + for resource in snapshot.resources: + if resource.resource_type.lower() != _NSG_TYPE: + continue + subnets = resource.properties.get("subnets") or [] + for subnet in subnets: + subnet_id = subnet.get("id") if isinstance(subnet, dict) else None + if not subnet_id: + continue + edges.append(GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=subnet_id, + relationship_type="PROTECTS", + evidence_source="arg:properties.subnets", + confidence=1.0, + )) + return edges + + +class SubnetToResourceDetector(EdgeDetector): + """Resource -> Subnet: MEMBER_OF (inferred from properties, confidence 0.8).""" + + def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: + edges = [] + resource_ids = {r.resource_id.lower() for r in snapshot.resources} + for resource in snapshot.resources: + subnet_id = (resource.properties.get("subnet") or {}).get("id") + if not subnet_id: + continue + if subnet_id.lower() not in resource_ids: + continue + edges.append(GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=subnet_id, + relationship_type="MEMBER_OF", + evidence_source="arg:properties.subnet.id", + confidence=0.8, + )) + return edges + + +class PublicIpToResourceDetector(EdgeDetector): + """PublicIP -> Resource: EXPOSES (ARG-confirmed, confidence 1.0).""" + + def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: + edges = [] + for resource in snapshot.resources: + if resource.resource_type.lower() != _PUBLIC_IP_TYPE: + continue + ip_config = resource.properties.get("ipConfiguration") or {} + target_id = ip_config.get("id") + if not target_id: + continue + # Strip the NIC sub-path to get the parent VM resource ID (4 provider path segments) + parts = target_id.split("/providers/") + if len(parts) >= 2: + provider_path = parts[-1].split("/") + if len(provider_path) >= 4: + target_id = "/providers/".join(parts[:-1]) + "/providers/" + "/".join(provider_path[:4]) + edges.append(GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=target_id, + relationship_type="EXPOSES", + evidence_source="arg:properties.ipConfiguration.id", + confidence=1.0, + )) + return edges + + +class IdentityToResourceDetector(EdgeDetector): + """Identity -> Resource: HAS_IDENTITY (inferred from properties, confidence 0.8).""" + + def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: + edges = [] + for resource in snapshot.resources: + identity = resource.properties.get("identity") or {} + user_assigned = identity.get("userAssignedIdentities") or {} + for identity_id in user_assigned: + if not identity_id: + continue + edges.append(GraphEdge( + source_resource_id=identity_id, + target_resource_id=resource.resource_id, + relationship_type="HAS_IDENTITY", + evidence_source="arg:properties.identity.userAssignedIdentities", + confidence=0.8, + )) + return edges + + +class StoragePrivateEndpointDetector(EdgeDetector): + """Storage -> PrivateEndpoint: REACHABLE_VIA (ARG-confirmed, confidence 1.0).""" + + def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: + edges = [] + for resource in snapshot.resources: + if resource.resource_type.lower() != _STORAGE_TYPE: + continue + connections = resource.properties.get("privateEndpointConnections") or [] + for conn in connections: + pe_id = ((conn.get("properties") or {}).get("privateEndpoint") or {}).get("id") + if not pe_id: + continue + edges.append(GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=pe_id, + relationship_type="REACHABLE_VIA", + evidence_source="arg:properties.privateEndpointConnections", + confidence=1.0, + )) + return edges + + +def detect_all_edges(snapshot: InventorySnapshot) -> list[GraphEdge]: + """Run all detectors and return the combined edge list. + + Per-detector exceptions are caught and logged; the function never raises. + """ + detectors: list[EdgeDetector] = [ + NsgToSubnetDetector(), + SubnetToResourceDetector(), + PublicIpToResourceDetector(), + IdentityToResourceDetector(), + StoragePrivateEndpointDetector(), + ] + edges: list[GraphEdge] = [] + for detector in detectors: + try: + edges.extend(detector.detect(snapshot)) + except Exception as exc: + logger.warning("edge_detector: %s failed: %s", type(detector).__name__, exc) + return edges diff --git a/tests/test_graph_edge_detector.py b/tests/test_graph_edge_detector.py new file mode 100644 index 00000000..247e09d9 --- /dev/null +++ b/tests/test_graph_edge_detector.py @@ -0,0 +1,120 @@ +"""Tests for graph edge detectors.""" +from scanner.arg_inventory import InventorySnapshot, InventoryStatus, InventoryResource +from scanner.graph.edge_detector import ( + GraphEdge, # noqa: F401 — imported to verify public API surface + NsgToSubnetDetector, + SubnetToResourceDetector, # noqa: F401 — imported to verify public API surface + PublicIpToResourceDetector, + IdentityToResourceDetector, + StoragePrivateEndpointDetector, + detect_all_edges, +) + + +def _resource(resource_id, resource_type, properties=None, subscription_id="00000000-0000-0000-0000-000000000002"): + return InventoryResource( + snapshot_id="snap-1", + tenant_id="00000000-0000-0000-0000-000000000001", + subscription_id=subscription_id, + resource_id=resource_id, + resource_type=resource_type, + name=resource_id.split("/")[-1], + location="eastus", + resource_group="rg", + tags={}, + properties=properties or {}, + ) + + +def _snapshot(*resources): + return InventorySnapshot( + snapshot_id="snap-1", + tenant_id="00000000-0000-0000-0000-000000000001", + requested_subscriptions=("00000000-0000-0000-0000-000000000002",), + status=InventoryStatus.COMPLETE, + collected_at="2026-09-24T00:00:00+00:00", + duration_ms=10, + pages=1, + resources=tuple(resources), + errors=(), + ) + + +def test_nsg_to_subnet_detects_protects_edge(): + subnet_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet1/subnets/default" + nsg_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/networkSecurityGroups/nsg1" + nsg = _resource(nsg_id, "microsoft.network/networksecuritygroups", { + "subnets": [{"id": subnet_id}] + }) + snapshot = _snapshot(nsg) + edges = NsgToSubnetDetector().detect(snapshot) + assert len(edges) == 1 + assert edges[0].relationship_type == "PROTECTS" + assert edges[0].source_resource_id == nsg_id + assert edges[0].target_resource_id == subnet_id + assert edges[0].confidence == 1.0 + + +def test_nsg_to_subnet_no_subnets_returns_empty(): + nsg = _resource("/nsg1", "microsoft.network/networksecuritygroups", {"subnets": []}) + snapshot = _snapshot(nsg) + assert NsgToSubnetDetector().detect(snapshot) == [] + + +def test_public_ip_to_resource_detects_exposes_edge(): + vm_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Compute/virtualMachines/vm1" + pip_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/publicIPAddresses/pip1" + pip = _resource(pip_id, "microsoft.network/publicipaddresses", { + "ipConfiguration": {"id": vm_id + "/networkInterfaces/nic1/ipConfigurations/ipconfig1"} + }) + snapshot = _snapshot(pip) + edges = PublicIpToResourceDetector().detect(snapshot) + assert len(edges) == 1 + assert edges[0].relationship_type == "EXPOSES" + assert edges[0].source_resource_id == pip_id + assert edges[0].confidence == 1.0 + + +def test_identity_to_resource_detects_has_identity_edge(): + vm_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Compute/virtualMachines/vm1" + identity_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.ManagedIdentity/userAssignedIdentities/id1" + vm = _resource(vm_id, "microsoft.compute/virtualmachines", { + "identity": { + "userAssignedIdentities": {identity_id: {}} + } + }) + snapshot = _snapshot(vm) + edges = IdentityToResourceDetector().detect(snapshot) + assert len(edges) == 1 + assert edges[0].relationship_type == "HAS_IDENTITY" + assert edges[0].target_resource_id == vm_id + assert edges[0].confidence == 0.8 + + +def test_storage_private_endpoint_detects_reachable_via_edge(): + storage_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/sa1" + pe_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/privateEndpoints/pe1" + storage = _resource(storage_id, "microsoft.storage/storageaccounts", { + "privateEndpointConnections": [{"properties": {"privateEndpoint": {"id": pe_id}}}] + }) + snapshot = _snapshot(storage) + edges = StoragePrivateEndpointDetector().detect(snapshot) + assert len(edges) == 1 + assert edges[0].relationship_type == "REACHABLE_VIA" + assert edges[0].confidence == 1.0 + + +def test_detect_all_edges_runs_all_detectors(): + nsg_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/networkSecurityGroups/nsg1" + subnet_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet1/subnets/default" + nsg = _resource(nsg_id, "microsoft.network/networksecuritygroups", { + "subnets": [{"id": subnet_id}] + }) + snapshot = _snapshot(nsg) + edges = detect_all_edges(snapshot) + assert any(e.relationship_type == "PROTECTS" for e in edges) + + +def test_empty_snapshot_returns_no_edges(): + snapshot = _snapshot() + assert detect_all_edges(snapshot) == [] From ef8a7acd28100afc66c296b9d5120cff2304df2a Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 24 Sep 2026 17:55:25 +0100 Subject: [PATCH 03/13] fix(graph): pass tenant_id explicitly to link_findings_to_nodes Replaces ORDER BY updated_at DESC LIMIT 1 subquery with a direct bound parameter to prevent cross-tenant finding linkage under concurrent scans. Signed-off-by: Tanvir Farhad --- scanner/graph/node_service.py | 12 +++--------- tests/test_graph_node_service.py | 2 +- 2 files changed, 4 insertions(+), 10 deletions(-) diff --git a/scanner/graph/node_service.py b/scanner/graph/node_service.py index c1154c1a..3f466218 100644 --- a/scanner/graph/node_service.py +++ b/scanner/graph/node_service.py @@ -40,13 +40,7 @@ SELECT f.id, n.node_id FROM findings f JOIN graph_nodes n ON lower(f.resource_id) = lower(n.resource_id) - AND n.tenant_id = ( - SELECT tenant_id FROM graph_nodes - WHERE snapshot_id = ( - SELECT snapshot_id FROM graph_nodes ORDER BY updated_at DESC LIMIT 1 - ) - LIMIT 1 - ) + AND n.tenant_id = %(tenant_id)s WHERE f.scan_id = %(scan_id)s ON CONFLICT DO NOTHING """ @@ -77,9 +71,9 @@ def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: return written -def link_findings_to_nodes(scan_id: str, dsn: str) -> int: +def link_findings_to_nodes(scan_id: str, tenant_id: str, dsn: str) -> int: """Link findings from this scan to their graph nodes by resource_id.""" with psycopg2.connect(dsn) as conn: with conn.cursor() as cur: - cur.execute(_LINK_FINDINGS_SQL, {"scan_id": scan_id}) + cur.execute(_LINK_FINDINGS_SQL, {"scan_id": scan_id, "tenant_id": tenant_id}) return cur.rowcount if cur.rowcount >= 0 else 0 diff --git a/tests/test_graph_node_service.py b/tests/test_graph_node_service.py index 12f70b8b..c643344e 100644 --- a/tests/test_graph_node_service.py +++ b/tests/test_graph_node_service.py @@ -104,7 +104,7 @@ def test_link_findings_to_nodes_executes_insert(): mock_cur.rowcount = 3 with patch("scanner.graph.node_service.psycopg2.connect", return_value=mock_conn): - link_findings_to_nodes("scan-uuid-1", "postgresql://test/db") + link_findings_to_nodes("scan-uuid-1", "00000000-0000-0000-0000-000000000001", "postgresql://test/db") assert mock_cur.execute.called sql_called = mock_cur.execute.call_args[0][0] From fe44a1a8b49f11340cd2e8e68c6e5a68d58a2d34 Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 24 Sep 2026 18:01:39 +0100 Subject: [PATCH 04/13] feat(graph): wire post-scan node and edge population into ScanEngine Signed-off-by: Tanvir Farhad --- scanner/engine.py | 10 ++++ scanner/graph/graph_populator.py | 85 ++++++++++++++++++++++++++++ tests/test_graph_engine_post_scan.py | 63 +++++++++++++++++++++ 3 files changed, 158 insertions(+) create mode 100644 scanner/graph/graph_populator.py create mode 100644 tests/test_graph_engine_post_scan.py diff --git a/scanner/engine.py b/scanner/engine.py index a43c1847..25ff6ae6 100644 --- a/scanner/engine.py +++ b/scanner/engine.py @@ -3,6 +3,7 @@ import importlib.util import inspect import logging +import os import uuid from datetime import datetime, timezone from pathlib import Path @@ -12,6 +13,7 @@ from openshield.severity import CONTRACT_VERSION, SeverityContractError, normalize_severity, score_findings from scanner.azure_client import AzureClient from scanner.evaluation import EvaluationStatus, RuleEvaluation, subscription_scope_id +from scanner.graph.graph_populator import populate_graph from scanner.graph.snapshot_bridge import collect_snapshot logger = logging.getLogger(__name__) @@ -210,6 +212,14 @@ def run_scan(self, scan_id: Optional[str] = None) -> Dict[str, Any]: logger.info("Scan %s complete — %d total finding(s). Normalising results...", scan_id, len(findings)) + # Populate the attack graph as a non-fatal post-scan step. + dsn = os.environ.get("DATABASE_URL") + if snapshot and dsn: + try: + populate_graph(scan_id, snapshot, dsn) + except Exception as exc: + logger.warning("run_scan: graph population raised unexpectedly: %s", exc) + return make_serializable(result) def _evaluate_rule(self, rule: Any, rule_id: str) -> List[RuleEvaluation]: diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py new file mode 100644 index 00000000..540bb821 --- /dev/null +++ b/scanner/graph/graph_populator.py @@ -0,0 +1,85 @@ +"""Orchestrate post-scan graph population: nodes, edges, finding links.""" +from __future__ import annotations + +import logging +import uuid +from typing import TYPE_CHECKING + +import psycopg2 +import psycopg2.extras + +from scanner.graph.node_service import link_findings_to_nodes, populate_nodes +from scanner.graph.edge_detector import detect_all_edges + +if TYPE_CHECKING: + from scanner.arg_inventory import InventorySnapshot + +logger = logging.getLogger(__name__) + +_UPSERT_EDGE_SQL = """ +INSERT INTO graph_edges ( + edge_id, source_node_id, target_node_id, relationship_type, + evidence_source, evidence_snapshot_id, confidence, collected_at, properties +) +SELECT + %(edge_id)s, + src.node_id, + tgt.node_id, + %(relationship_type)s, + %(evidence_source)s, + %(evidence_snapshot_id)s, + %(confidence)s, + now(), + '{}'::jsonb +FROM graph_nodes src, graph_nodes tgt +WHERE lower(src.resource_id) = lower(%(source_resource_id)s) + AND lower(tgt.resource_id) = lower(%(target_resource_id)s) +ON CONFLICT (source_node_id, target_node_id, relationship_type) DO UPDATE SET + confidence = EXCLUDED.confidence, + evidence_source = EXCLUDED.evidence_source, + evidence_snapshot_id = EXCLUDED.evidence_snapshot_id, + collected_at = now() +""" + + +def _write_edges(edges: list, snapshot_id: str, dsn: str) -> int: + if not edges: + return 0 + written = 0 + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + for edge in edges: + cur.execute(_UPSERT_EDGE_SQL, { + "edge_id": str(uuid.uuid4()), + "source_resource_id": edge.source_resource_id, + "target_resource_id": edge.target_resource_id, + "relationship_type": edge.relationship_type, + "evidence_source": edge.evidence_source, + "evidence_snapshot_id": snapshot_id, + "confidence": edge.confidence, + }) + written += 1 + return written + + +def populate_graph(scan_id: str, snapshot: InventorySnapshot, dsn: str) -> None: + """Populate nodes, edges, and finding links for one scan. Failure is non-fatal.""" + try: + node_count = populate_nodes(snapshot, dsn) + logger.info("graph: upserted %d nodes for scan %s", node_count, scan_id) + except Exception as exc: + logger.warning("graph: node population failed for scan %s: %s", scan_id, exc) + return + + try: + edges = detect_all_edges(snapshot) + edge_count = _write_edges(edges, snapshot.snapshot_id, dsn) + logger.info("graph: wrote %d edges for scan %s", edge_count, scan_id) + except Exception as exc: + logger.warning("graph: edge population failed for scan %s: %s", scan_id, exc) + + try: + link_count = link_findings_to_nodes(scan_id, snapshot.tenant_id, dsn) + logger.info("graph: linked %d findings to nodes for scan %s", link_count, scan_id) + except Exception as exc: + logger.warning("graph: finding link failed for scan %s: %s", scan_id, exc) diff --git a/tests/test_graph_engine_post_scan.py b/tests/test_graph_engine_post_scan.py new file mode 100644 index 00000000..71bc51b6 --- /dev/null +++ b/tests/test_graph_engine_post_scan.py @@ -0,0 +1,63 @@ +"""Tests for post-scan graph population wiring in ScanEngine.""" +from unittest.mock import patch +import pytest + +from scanner.engine import ScanEngine +from scanner.arg_inventory import InventorySnapshot, InventoryStatus + + +def _make_snapshot(): + return InventorySnapshot( + snapshot_id="snap-1", + tenant_id="00000000-0000-0000-0000-000000000001", + requested_subscriptions=("00000000-0000-0000-0000-000000000002",), + status=InventoryStatus.COMPLETE, + collected_at="2026-09-24T00:00:00+00:00", + duration_ms=10, + pages=1, + resources=(), + errors=(), + ) + + +@pytest.fixture +def engine(): + with patch("scanner.engine.AzureClient"), patch("scanner.engine.ScanEngine.load_rules"): + eng = ScanEngine("00000000-0000-0000-0000-000000000002") + eng.rules = [] + return eng + + +def test_populate_graph_called_when_snapshot_and_dsn_available(engine): + snapshot = _make_snapshot() + with patch("scanner.engine.collect_snapshot", return_value=snapshot), \ + patch("scanner.engine.populate_graph") as mock_populate, \ + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}): + engine.run_scan() + mock_populate.assert_called_once() + + +def test_populate_graph_not_called_when_snapshot_none(engine): + with patch("scanner.engine.collect_snapshot", return_value=None), \ + patch("scanner.engine.populate_graph") as mock_populate, \ + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}): + engine.run_scan() + mock_populate.assert_not_called() + + +def test_populate_graph_not_called_when_no_database_url(engine): + snapshot = _make_snapshot() + with patch("scanner.engine.collect_snapshot", return_value=snapshot), \ + patch("scanner.engine.populate_graph") as mock_populate, \ + patch.dict("os.environ", {}, clear=True): + engine.run_scan() + mock_populate.assert_not_called() + + +def test_scan_succeeds_even_if_populate_graph_raises(engine): + snapshot = _make_snapshot() + with patch("scanner.engine.collect_snapshot", return_value=snapshot), \ + patch("scanner.engine.populate_graph", side_effect=Exception("DB down")), \ + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}): + result = engine.run_scan() + assert result["status"] == "completed" From 46887672351c6497a4b2ab6362d428b5cbc9aa77 Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Fri, 25 Sep 2026 00:45:04 +0100 Subject: [PATCH 05/13] fix(lint): apply ruff format to node/edge population files Signed-off-by: Tanvir Farhad --- scanner/graph/edge_detector.py | 81 ++++++++++++++++------------ scanner/graph/graph_populator.py | 22 ++++---- scanner/graph/node_service.py | 28 +++++----- tests/test_graph_edge_detector.py | 33 ++++++------ tests/test_graph_engine_post_scan.py | 33 +++++++----- tests/test_graph_node_service.py | 8 ++- 6 files changed, 118 insertions(+), 87 deletions(-) diff --git a/scanner/graph/edge_detector.py b/scanner/graph/edge_detector.py index 61cf0ad5..ae07eae8 100644 --- a/scanner/graph/edge_detector.py +++ b/scanner/graph/edge_detector.py @@ -1,4 +1,5 @@ """Typed edge detectors that infer relationships between Azure resources in an InventorySnapshot.""" + from __future__ import annotations import logging @@ -48,13 +49,15 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: subnet_id = subnet.get("id") if isinstance(subnet, dict) else None if not subnet_id: continue - edges.append(GraphEdge( - source_resource_id=resource.resource_id, - target_resource_id=subnet_id, - relationship_type="PROTECTS", - evidence_source="arg:properties.subnets", - confidence=1.0, - )) + edges.append( + GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=subnet_id, + relationship_type="PROTECTS", + evidence_source="arg:properties.subnets", + confidence=1.0, + ) + ) return edges @@ -70,13 +73,15 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: continue if subnet_id.lower() not in resource_ids: continue - edges.append(GraphEdge( - source_resource_id=resource.resource_id, - target_resource_id=subnet_id, - relationship_type="MEMBER_OF", - evidence_source="arg:properties.subnet.id", - confidence=0.8, - )) + edges.append( + GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=subnet_id, + relationship_type="MEMBER_OF", + evidence_source="arg:properties.subnet.id", + confidence=0.8, + ) + ) return edges @@ -98,13 +103,15 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: provider_path = parts[-1].split("/") if len(provider_path) >= 4: target_id = "/providers/".join(parts[:-1]) + "/providers/" + "/".join(provider_path[:4]) - edges.append(GraphEdge( - source_resource_id=resource.resource_id, - target_resource_id=target_id, - relationship_type="EXPOSES", - evidence_source="arg:properties.ipConfiguration.id", - confidence=1.0, - )) + edges.append( + GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=target_id, + relationship_type="EXPOSES", + evidence_source="arg:properties.ipConfiguration.id", + confidence=1.0, + ) + ) return edges @@ -119,13 +126,15 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: for identity_id in user_assigned: if not identity_id: continue - edges.append(GraphEdge( - source_resource_id=identity_id, - target_resource_id=resource.resource_id, - relationship_type="HAS_IDENTITY", - evidence_source="arg:properties.identity.userAssignedIdentities", - confidence=0.8, - )) + edges.append( + GraphEdge( + source_resource_id=identity_id, + target_resource_id=resource.resource_id, + relationship_type="HAS_IDENTITY", + evidence_source="arg:properties.identity.userAssignedIdentities", + confidence=0.8, + ) + ) return edges @@ -142,13 +151,15 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: pe_id = ((conn.get("properties") or {}).get("privateEndpoint") or {}).get("id") if not pe_id: continue - edges.append(GraphEdge( - source_resource_id=resource.resource_id, - target_resource_id=pe_id, - relationship_type="REACHABLE_VIA", - evidence_source="arg:properties.privateEndpointConnections", - confidence=1.0, - )) + edges.append( + GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=pe_id, + relationship_type="REACHABLE_VIA", + evidence_source="arg:properties.privateEndpointConnections", + confidence=1.0, + ) + ) return edges diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index 540bb821..c5c47565 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -1,4 +1,5 @@ """Orchestrate post-scan graph population: nodes, edges, finding links.""" + from __future__ import annotations import logging @@ -49,15 +50,18 @@ def _write_edges(edges: list, snapshot_id: str, dsn: str) -> int: with psycopg2.connect(dsn) as conn: with conn.cursor() as cur: for edge in edges: - cur.execute(_UPSERT_EDGE_SQL, { - "edge_id": str(uuid.uuid4()), - "source_resource_id": edge.source_resource_id, - "target_resource_id": edge.target_resource_id, - "relationship_type": edge.relationship_type, - "evidence_source": edge.evidence_source, - "evidence_snapshot_id": snapshot_id, - "confidence": edge.confidence, - }) + cur.execute( + _UPSERT_EDGE_SQL, + { + "edge_id": str(uuid.uuid4()), + "source_resource_id": edge.source_resource_id, + "target_resource_id": edge.target_resource_id, + "relationship_type": edge.relationship_type, + "evidence_source": edge.evidence_source, + "evidence_snapshot_id": snapshot_id, + "confidence": edge.confidence, + }, + ) written += 1 return written diff --git a/scanner/graph/node_service.py b/scanner/graph/node_service.py index 3f466218..2b571e7f 100644 --- a/scanner/graph/node_service.py +++ b/scanner/graph/node_service.py @@ -1,4 +1,5 @@ """Upsert graph nodes from an InventorySnapshot and link findings to nodes.""" + from __future__ import annotations import json @@ -55,18 +56,21 @@ def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: with psycopg2.connect(dsn) as conn: with conn.cursor() as cur: for resource in snapshot.resources: - cur.execute(_UPSERT_NODE_SQL, { - "node_id": str(uuid.uuid4()), - "tenant_id": resource.tenant_id, - "subscription_id": resource.subscription_id, - "resource_id": resource.resource_id, - "resource_type": resource.resource_type, - "name": resource.name, - "location": resource.location, - "resource_group": resource.resource_group, - "snapshot_id": resource.snapshot_id, - "properties": json.dumps(resource.properties), - }) + cur.execute( + _UPSERT_NODE_SQL, + { + "node_id": str(uuid.uuid4()), + "tenant_id": resource.tenant_id, + "subscription_id": resource.subscription_id, + "resource_id": resource.resource_id, + "resource_type": resource.resource_type, + "name": resource.name, + "location": resource.location, + "resource_group": resource.resource_group, + "snapshot_id": resource.snapshot_id, + "properties": json.dumps(resource.properties), + }, + ) written += 1 return written diff --git a/tests/test_graph_edge_detector.py b/tests/test_graph_edge_detector.py index 247e09d9..4649cb7d 100644 --- a/tests/test_graph_edge_detector.py +++ b/tests/test_graph_edge_detector.py @@ -1,4 +1,5 @@ """Tests for graph edge detectors.""" + from scanner.arg_inventory import InventorySnapshot, InventoryStatus, InventoryResource from scanner.graph.edge_detector import ( GraphEdge, # noqa: F401 — imported to verify public API surface @@ -43,9 +44,7 @@ def _snapshot(*resources): def test_nsg_to_subnet_detects_protects_edge(): subnet_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet1/subnets/default" nsg_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/networkSecurityGroups/nsg1" - nsg = _resource(nsg_id, "microsoft.network/networksecuritygroups", { - "subnets": [{"id": subnet_id}] - }) + nsg = _resource(nsg_id, "microsoft.network/networksecuritygroups", {"subnets": [{"id": subnet_id}]}) snapshot = _snapshot(nsg) edges = NsgToSubnetDetector().detect(snapshot) assert len(edges) == 1 @@ -64,9 +63,11 @@ def test_nsg_to_subnet_no_subnets_returns_empty(): def test_public_ip_to_resource_detects_exposes_edge(): vm_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Compute/virtualMachines/vm1" pip_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/publicIPAddresses/pip1" - pip = _resource(pip_id, "microsoft.network/publicipaddresses", { - "ipConfiguration": {"id": vm_id + "/networkInterfaces/nic1/ipConfigurations/ipconfig1"} - }) + pip = _resource( + pip_id, + "microsoft.network/publicipaddresses", + {"ipConfiguration": {"id": vm_id + "/networkInterfaces/nic1/ipConfigurations/ipconfig1"}}, + ) snapshot = _snapshot(pip) edges = PublicIpToResourceDetector().detect(snapshot) assert len(edges) == 1 @@ -78,11 +79,9 @@ def test_public_ip_to_resource_detects_exposes_edge(): def test_identity_to_resource_detects_has_identity_edge(): vm_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Compute/virtualMachines/vm1" identity_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.ManagedIdentity/userAssignedIdentities/id1" - vm = _resource(vm_id, "microsoft.compute/virtualmachines", { - "identity": { - "userAssignedIdentities": {identity_id: {}} - } - }) + vm = _resource( + vm_id, "microsoft.compute/virtualmachines", {"identity": {"userAssignedIdentities": {identity_id: {}}}} + ) snapshot = _snapshot(vm) edges = IdentityToResourceDetector().detect(snapshot) assert len(edges) == 1 @@ -94,9 +93,11 @@ def test_identity_to_resource_detects_has_identity_edge(): def test_storage_private_endpoint_detects_reachable_via_edge(): storage_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Storage/storageAccounts/sa1" pe_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/privateEndpoints/pe1" - storage = _resource(storage_id, "microsoft.storage/storageaccounts", { - "privateEndpointConnections": [{"properties": {"privateEndpoint": {"id": pe_id}}}] - }) + storage = _resource( + storage_id, + "microsoft.storage/storageaccounts", + {"privateEndpointConnections": [{"properties": {"privateEndpoint": {"id": pe_id}}}]}, + ) snapshot = _snapshot(storage) edges = StoragePrivateEndpointDetector().detect(snapshot) assert len(edges) == 1 @@ -107,9 +108,7 @@ def test_storage_private_endpoint_detects_reachable_via_edge(): def test_detect_all_edges_runs_all_detectors(): nsg_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/networkSecurityGroups/nsg1" subnet_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Network/virtualNetworks/vnet1/subnets/default" - nsg = _resource(nsg_id, "microsoft.network/networksecuritygroups", { - "subnets": [{"id": subnet_id}] - }) + nsg = _resource(nsg_id, "microsoft.network/networksecuritygroups", {"subnets": [{"id": subnet_id}]}) snapshot = _snapshot(nsg) edges = detect_all_edges(snapshot) assert any(e.relationship_type == "PROTECTS" for e in edges) diff --git a/tests/test_graph_engine_post_scan.py b/tests/test_graph_engine_post_scan.py index 71bc51b6..8c4fe17c 100644 --- a/tests/test_graph_engine_post_scan.py +++ b/tests/test_graph_engine_post_scan.py @@ -1,4 +1,5 @@ """Tests for post-scan graph population wiring in ScanEngine.""" + from unittest.mock import patch import pytest @@ -30,34 +31,42 @@ def engine(): def test_populate_graph_called_when_snapshot_and_dsn_available(engine): snapshot = _make_snapshot() - with patch("scanner.engine.collect_snapshot", return_value=snapshot), \ - patch("scanner.engine.populate_graph") as mock_populate, \ - patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}): + with ( + patch("scanner.engine.collect_snapshot", return_value=snapshot), + patch("scanner.engine.populate_graph") as mock_populate, + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), + ): engine.run_scan() mock_populate.assert_called_once() def test_populate_graph_not_called_when_snapshot_none(engine): - with patch("scanner.engine.collect_snapshot", return_value=None), \ - patch("scanner.engine.populate_graph") as mock_populate, \ - patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}): + with ( + patch("scanner.engine.collect_snapshot", return_value=None), + patch("scanner.engine.populate_graph") as mock_populate, + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), + ): engine.run_scan() mock_populate.assert_not_called() def test_populate_graph_not_called_when_no_database_url(engine): snapshot = _make_snapshot() - with patch("scanner.engine.collect_snapshot", return_value=snapshot), \ - patch("scanner.engine.populate_graph") as mock_populate, \ - patch.dict("os.environ", {}, clear=True): + with ( + patch("scanner.engine.collect_snapshot", return_value=snapshot), + patch("scanner.engine.populate_graph") as mock_populate, + patch.dict("os.environ", {}, clear=True), + ): engine.run_scan() mock_populate.assert_not_called() def test_scan_succeeds_even_if_populate_graph_raises(engine): snapshot = _make_snapshot() - with patch("scanner.engine.collect_snapshot", return_value=snapshot), \ - patch("scanner.engine.populate_graph", side_effect=Exception("DB down")), \ - patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}): + with ( + patch("scanner.engine.collect_snapshot", return_value=snapshot), + patch("scanner.engine.populate_graph", side_effect=Exception("DB down")), + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), + ): result = engine.run_scan() assert result["status"] == "completed" diff --git a/tests/test_graph_node_service.py b/tests/test_graph_node_service.py index c643344e..f4204001 100644 --- a/tests/test_graph_node_service.py +++ b/tests/test_graph_node_service.py @@ -1,12 +1,16 @@ """Tests for graph node upsert and finding-to-node linking.""" + from unittest.mock import MagicMock, patch from scanner.arg_inventory import InventorySnapshot, InventoryStatus, InventoryResource from scanner.graph.node_service import populate_nodes, link_findings_to_nodes -def _make_resource(resource_id: str, resource_type: str = "microsoft.network/virtualnetworks", - subscription_id: str = "00000000-0000-0000-0000-000000000002") -> InventoryResource: +def _make_resource( + resource_id: str, + resource_type: str = "microsoft.network/virtualnetworks", + subscription_id: str = "00000000-0000-0000-0000-000000000002", +) -> InventoryResource: return InventoryResource( snapshot_id="snap-1", tenant_id="00000000-0000-0000-0000-000000000001", From 8302ef58128dc8b3cac75c3799c069829a8a9af3 Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Sun, 27 Sep 2026 11:35:15 +0100 Subject: [PATCH 06/13] fix(graph): post-save graph population, fix detectors for ARG schema, tenant filter on edge upsert Signed-off-by: Tanvir Farhad --- scanner/arg_inventory.py | 1 + scanner/engine.py | 12 ++------ scanner/graph/edge_detector.py | 37 ++++++++++++++-------- scanner/graph/graph_populator.py | 7 +++-- scanner/worker.py | 15 +++++++++ tests/test_graph_engine_post_scan.py | 46 +++++++++++++++------------- 6 files changed, 71 insertions(+), 47 deletions(-) diff --git a/scanner/arg_inventory.py b/scanner/arg_inventory.py index feaefb54..92574470 100644 --- a/scanner/arg_inventory.py +++ b/scanner/arg_inventory.py @@ -17,6 +17,7 @@ DEFAULT_QUERY = """ Resources +| extend properties = bag_merge(properties, pack('identity', identity)) | project id, name, type, location, subscriptionId, resourceGroup, tenantId, tags, properties | order by id asc """.strip() diff --git a/scanner/engine.py b/scanner/engine.py index 25ff6ae6..92906b8e 100644 --- a/scanner/engine.py +++ b/scanner/engine.py @@ -3,7 +3,6 @@ import importlib.util import inspect import logging -import os import uuid from datetime import datetime, timezone from pathlib import Path @@ -13,7 +12,6 @@ from openshield.severity import CONTRACT_VERSION, SeverityContractError, normalize_severity, score_findings from scanner.azure_client import AzureClient from scanner.evaluation import EvaluationStatus, RuleEvaluation, subscription_scope_id -from scanner.graph.graph_populator import populate_graph from scanner.graph.snapshot_bridge import collect_snapshot logger = logging.getLogger(__name__) @@ -63,6 +61,7 @@ def __init__(self, subscription_id: str) -> None: self.subscription_id = subscription_id self.client = AzureClient(subscription_id) self.rules: List[Any] = [] + self.snapshot: Optional[Any] = None self.load_rules() # ------------------------------------------------------------------ # @@ -119,6 +118,7 @@ def run_scan(self, scan_id: Optional[str] = None) -> Dict[str, Any]: # Collect an ARG inventory snapshot for graph population and rule enrichment. # Failure is non-fatal: rules fall back to direct SDK calls. snapshot = collect_snapshot(self.client, self.subscription_id) + self.snapshot = snapshot logger.info( "Scan %s starting against subscription %s — %d rules loaded", @@ -212,14 +212,6 @@ def run_scan(self, scan_id: Optional[str] = None) -> Dict[str, Any]: logger.info("Scan %s complete — %d total finding(s). Normalising results...", scan_id, len(findings)) - # Populate the attack graph as a non-fatal post-scan step. - dsn = os.environ.get("DATABASE_URL") - if snapshot and dsn: - try: - populate_graph(scan_id, snapshot, dsn) - except Exception as exc: - logger.warning("run_scan: graph population raised unexpectedly: %s", exc) - return make_serializable(result) def _evaluate_rule(self, rule: Any, rule_id: str) -> List[RuleEvaluation]: diff --git a/scanner/graph/edge_detector.py b/scanner/graph/edge_detector.py index ae07eae8..ce5878d1 100644 --- a/scanner/graph/edge_detector.py +++ b/scanner/graph/edge_detector.py @@ -68,20 +68,31 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: edges = [] resource_ids = {r.resource_id.lower() for r in snapshot.resources} for resource in snapshot.resources: - subnet_id = (resource.properties.get("subnet") or {}).get("id") - if not subnet_id: - continue - if subnet_id.lower() not in resource_ids: - continue - edges.append( - GraphEdge( - source_resource_id=resource.resource_id, - target_resource_id=subnet_id, - relationship_type="MEMBER_OF", - evidence_source="arg:properties.subnet.id", - confidence=0.8, + subnet_ids: list[str] = [] + + # Direct subnet property (VMs attached directly) + direct = (resource.properties.get("subnet") or {}).get("id") + if direct: + subnet_ids.append(direct) + + # NIC ipConfigurations pattern (most common path for NICs) + for ip_cfg in resource.properties.get("ipConfigurations", []): + sid = (ip_cfg.get("properties", {}).get("subnet") or {}).get("id") + if sid: + subnet_ids.append(sid) + + for subnet_id in subnet_ids: + if not subnet_id or subnet_id.lower() not in resource_ids: + continue + edges.append( + GraphEdge( + source_resource_id=resource.resource_id, + target_resource_id=subnet_id, + relationship_type="MEMBER_OF", + evidence_source="arg:properties.subnet.id", + confidence=0.8, + ) ) - ) return edges diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index c5c47565..2da6479e 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -34,7 +34,9 @@ '{}'::jsonb FROM graph_nodes src, graph_nodes tgt WHERE lower(src.resource_id) = lower(%(source_resource_id)s) + AND src.tenant_id = %(tenant_id)s AND lower(tgt.resource_id) = lower(%(target_resource_id)s) + AND tgt.tenant_id = %(tenant_id)s ON CONFLICT (source_node_id, target_node_id, relationship_type) DO UPDATE SET confidence = EXCLUDED.confidence, evidence_source = EXCLUDED.evidence_source, @@ -43,7 +45,7 @@ """ -def _write_edges(edges: list, snapshot_id: str, dsn: str) -> int: +def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str) -> int: if not edges: return 0 written = 0 @@ -59,6 +61,7 @@ def _write_edges(edges: list, snapshot_id: str, dsn: str) -> int: "relationship_type": edge.relationship_type, "evidence_source": edge.evidence_source, "evidence_snapshot_id": snapshot_id, + "tenant_id": tenant_id, "confidence": edge.confidence, }, ) @@ -77,7 +80,7 @@ def populate_graph(scan_id: str, snapshot: InventorySnapshot, dsn: str) -> None: try: edges = detect_all_edges(snapshot) - edge_count = _write_edges(edges, snapshot.snapshot_id, dsn) + edge_count = _write_edges(edges, snapshot.snapshot_id, snapshot.tenant_id, dsn) logger.info("graph: wrote %d edges for scan %s", edge_count, scan_id) except Exception as exc: logger.warning("graph: edge population failed for scan %s: %s", scan_id, exc) diff --git a/scanner/worker.py b/scanner/worker.py index 2ae431ae..5c0d388a 100644 --- a/scanner/worker.py +++ b/scanner/worker.py @@ -27,6 +27,7 @@ ) from scanner.engine import ScanEngine from scanner.enrichment_worker import process_enrichment_job +from scanner.graph.graph_populator import populate_graph configure_logging() logger = logging.getLogger("scanner.worker") @@ -216,6 +217,20 @@ def run_worker(): if heartbeat.lost.is_set(): raise LostLease(f"Scan {scan_id} lost its lease before completion") db.save_scan(result, worker_id, fencing_token) + + # Graph population runs after findings are persisted so + # link_findings_to_nodes can join against the saved rows. + dsn = os.environ.get("DATABASE_URL") + if engine.snapshot and dsn: + try: + populate_graph(scan_id, engine.snapshot, dsn) + except Exception as exc: + logger.warning( + "worker: graph population raised unexpectedly for scan %s: %s", + scan_id, + exc, + ) + SCANS_TOTAL.labels(status="completed").inc() logger.info( "Successfully completed scan %s", diff --git a/tests/test_graph_engine_post_scan.py b/tests/test_graph_engine_post_scan.py index 8c4fe17c..dd4dc643 100644 --- a/tests/test_graph_engine_post_scan.py +++ b/tests/test_graph_engine_post_scan.py @@ -1,6 +1,6 @@ -"""Tests for post-scan graph population wiring in ScanEngine.""" +"""Tests for post-scan graph population wiring in ScanEngine and worker.""" -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest from scanner.engine import ScanEngine @@ -29,44 +29,46 @@ def engine(): return eng -def test_populate_graph_called_when_snapshot_and_dsn_available(engine): +def test_snapshot_stored_on_engine_after_run_scan(engine): snapshot = _make_snapshot() - with ( - patch("scanner.engine.collect_snapshot", return_value=snapshot), - patch("scanner.engine.populate_graph") as mock_populate, - patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), - ): + with patch("scanner.engine.collect_snapshot", return_value=snapshot): engine.run_scan() - mock_populate.assert_called_once() + assert engine.snapshot is snapshot -def test_populate_graph_not_called_when_snapshot_none(engine): - with ( - patch("scanner.engine.collect_snapshot", return_value=None), - patch("scanner.engine.populate_graph") as mock_populate, - patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), - ): +def test_snapshot_is_none_on_engine_when_collection_fails(engine): + with patch("scanner.engine.collect_snapshot", return_value=None): engine.run_scan() - mock_populate.assert_not_called() + assert engine.snapshot is None -def test_populate_graph_not_called_when_no_database_url(engine): +def test_scan_result_does_not_call_populate_graph(engine): + """populate_graph must NOT be called inside run_scan — it belongs in worker.""" snapshot = _make_snapshot() with ( patch("scanner.engine.collect_snapshot", return_value=snapshot), - patch("scanner.engine.populate_graph") as mock_populate, - patch.dict("os.environ", {}, clear=True), + patch("scanner.graph.graph_populator.populate_graph") as mock_populate, + patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), ): engine.run_scan() mock_populate.assert_not_called() -def test_scan_succeeds_even_if_populate_graph_raises(engine): +def test_scan_succeeds_and_stores_snapshot_even_without_database_url(engine): snapshot = _make_snapshot() with ( patch("scanner.engine.collect_snapshot", return_value=snapshot), - patch("scanner.engine.populate_graph", side_effect=Exception("DB down")), - patch.dict("os.environ", {"DATABASE_URL": "postgresql://test/db"}), + patch.dict("os.environ", {}, clear=True), ): result = engine.run_scan() assert result["status"] == "completed" + assert engine.snapshot is snapshot + + +def test_worker_imports_populate_graph(): + """Sanity-check that worker imports populate_graph for post-save invocation.""" + import scanner.worker as worker_module # noqa: PLC0415 + + assert hasattr(worker_module, "populate_graph"), ( + "worker.py must import populate_graph so it can call it after db.save_scan" + ) From d48da0a5d6e2f42f66ba6635e04a3a61cd49d69d Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Sun, 27 Sep 2026 23:05:13 +0100 Subject: [PATCH 07/13] fix(lint): remove unused MagicMock import from post-scan graph test Signed-off-by: Tanvir Farhad --- tests/test_graph_engine_post_scan.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_graph_engine_post_scan.py b/tests/test_graph_engine_post_scan.py index dd4dc643..520614f0 100644 --- a/tests/test_graph_engine_post_scan.py +++ b/tests/test_graph_engine_post_scan.py @@ -1,6 +1,6 @@ """Tests for post-scan graph population wiring in ScanEngine and worker.""" -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest from scanner.engine import ScanEngine From 627abaf494d202869aab502814566e618f194ac3 Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 1 Oct 2026 00:37:39 +0100 Subject: [PATCH 08/13] fix(graph): synthesise subnet nodes from VNet properties.subnets before edge detection ARG Resources has no top-level rows for subnets; they are nested inside the parent VNet's properties.subnets. Without this, SubnetToResourceDetector and NsgToSubnetDetector produce edges whose target has no graph_node row, and _UPSERT_EDGE_SQL silently drops them (INSERT ... SELECT JOIN graph_nodes). _synthesise_subnet_resources() walks VNet resources, extracts each entry in properties.subnets, and returns synthetic InventoryResource objects inheriting the VNet's tenant_id, subscription_id, location and resource_group. populate_graph builds an augmented snapshot (dataclasses.replace on the frozen dataclass) before calling populate_nodes and detect_all_edges, so subnets get graph_nodes and the NSG->subnet->NIC path is traversable by BFS in #354. 7 new tests in test_graph_populator_subnet_synthesis.py use realistic ARG response fixtures (VNet with subnet, NSG, NIC, VM with user-assigned identity) and assert both the synthesis behaviour and the original bug (edge dropped without synthesis). Signed-off-by: Tanvir Farhad --- scanner/graph/graph_populator.py | 53 ++++- .../test_graph_populator_subnet_synthesis.py | 209 ++++++++++++++++++ 2 files changed, 260 insertions(+), 2 deletions(-) create mode 100644 tests/test_graph_populator_subnet_synthesis.py diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index 2da6479e..c3a073dc 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -4,11 +4,13 @@ import logging import uuid +from dataclasses import replace as dc_replace from typing import TYPE_CHECKING import psycopg2 import psycopg2.extras +from scanner.arg_inventory import InventoryResource from scanner.graph.node_service import link_findings_to_nodes, populate_nodes from scanner.graph.edge_detector import detect_all_edges @@ -69,17 +71,64 @@ def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str) -> int return written +_VNET_TYPE = "microsoft.network/virtualnetworks" +_SUBNET_TYPE = "microsoft.network/virtualnetworks/subnets" + + +def _synthesise_subnet_resources(snapshot: InventorySnapshot) -> list[InventoryResource]: + """Return synthetic InventoryResource entries for subnets nested inside VNets. + + ARG Resources has no top-level rows for subnets; they appear only as + properties.subnets on the parent VNet. Without this step, SubnetToResourceDetector + and NsgToSubnetDetector produce edges whose target has no matching graph_node, + and _UPSERT_EDGE_SQL silently drops them (INSERT ... SELECT JOIN graph_nodes). + """ + subnets: list[InventoryResource] = [] + for resource in snapshot.resources: + if resource.resource_type.lower() != _VNET_TYPE: + continue + for subnet in resource.properties.get("subnets") or []: + subnet_id = subnet.get("id") if isinstance(subnet, dict) else None + if not subnet_id: + continue + subnet_name = subnet.get("name", subnet_id.split("/")[-1]) + subnets.append( + InventoryResource( + snapshot_id=resource.snapshot_id, + tenant_id=resource.tenant_id, + subscription_id=resource.subscription_id, + resource_id=subnet_id, + resource_type=_SUBNET_TYPE, + name=subnet_name, + location=resource.location, + resource_group=resource.resource_group, + tags={}, + properties=subnet.get("properties") or {}, + ) + ) + return subnets + + def populate_graph(scan_id: str, snapshot: InventorySnapshot, dsn: str) -> None: """Populate nodes, edges, and finding links for one scan. Failure is non-fatal.""" + subnet_resources = _synthesise_subnet_resources(snapshot) + if subnet_resources: + logger.debug("graph: synthesised %d subnet nodes from VNet properties", len(subnet_resources)) + + augmented_snapshot = dc_replace( + snapshot, + resources=snapshot.resources + tuple(subnet_resources), + ) + try: - node_count = populate_nodes(snapshot, dsn) + node_count = populate_nodes(augmented_snapshot, dsn) logger.info("graph: upserted %d nodes for scan %s", node_count, scan_id) except Exception as exc: logger.warning("graph: node population failed for scan %s: %s", scan_id, exc) return try: - edges = detect_all_edges(snapshot) + edges = detect_all_edges(augmented_snapshot) edge_count = _write_edges(edges, snapshot.snapshot_id, snapshot.tenant_id, dsn) logger.info("graph: wrote %d edges for scan %s", edge_count, scan_id) except Exception as exc: diff --git a/tests/test_graph_populator_subnet_synthesis.py b/tests/test_graph_populator_subnet_synthesis.py new file mode 100644 index 00000000..06688896 --- /dev/null +++ b/tests/test_graph_populator_subnet_synthesis.py @@ -0,0 +1,209 @@ +"""Tests for subnet node synthesis in graph_populator. + +ARG Resources has no top-level rows for subnets — they are nested under the +parent VNet's properties.subnets. _synthesise_subnet_resources extracts them +so populate_graph can write graph_nodes for subnets and edge upserts succeed. +""" + +from scanner.arg_inventory import InventoryResource, InventorySnapshot, InventoryStatus +from scanner.graph.graph_populator import _synthesise_subnet_resources + +# Realistic ARG response shapes ------------------------------------------------- +# These mirror what Azure Resource Graph returns for a VNet with one subnet, +# a NIC attached to that subnet, and a VM with a user-assigned identity. + +_TENANT = "00000000-0000-0000-0000-000000000001" +_SUB = "00000000-0000-0000-0000-000000000002" +_SNAP = "snap-test-1" + +_VNET_ID = ( + "/subscriptions/00000000-0000-0000-0000-000000000002" + "/resourceGroups/rg-prod" + "/providers/Microsoft.Network/virtualNetworks/vnet-prod" +) +_SUBNET_ID = _VNET_ID + "/subnets/default" +_NSG_ID = ( + "/subscriptions/00000000-0000-0000-0000-000000000002" + "/resourceGroups/rg-prod" + "/providers/Microsoft.Network/networkSecurityGroups/nsg-prod" +) +_NIC_ID = ( + "/subscriptions/00000000-0000-0000-0000-000000000002" + "/resourceGroups/rg-prod" + "/providers/Microsoft.Network/networkInterfaces/nic-prod" +) +_VM_ID = ( + "/subscriptions/00000000-0000-0000-0000-000000000002" + "/resourceGroups/rg-prod" + "/providers/Microsoft.Compute/virtualMachines/vm-prod" +) +_IDENTITY_ID = ( + "/subscriptions/00000000-0000-0000-0000-000000000002" + "/resourceGroups/rg-prod" + "/providers/Microsoft.ManagedIdentity/userAssignedIdentities/id-prod" +) + + +def _resource(resource_id, resource_type, properties=None): + return InventoryResource( + snapshot_id=_SNAP, + tenant_id=_TENANT, + subscription_id=_SUB, + resource_id=resource_id, + resource_type=resource_type, + name=resource_id.split("/")[-1], + location="uksouth", + resource_group="rg-prod", + tags={}, + properties=properties or {}, + ) + + +# ARG returns subnets nested inside the VNet properties.subnets array +_VNET = _resource( + _VNET_ID, + "microsoft.network/virtualnetworks", + { + "subnets": [ + { + "id": _SUBNET_ID, + "name": "default", + "properties": { + "addressPrefix": "10.0.0.0/24", + "networkSecurityGroup": {"id": _NSG_ID}, + }, + } + ], + "addressSpace": {"addressPrefixes": ["10.0.0.0/16"]}, + }, +) + +_NSG = _resource( + _NSG_ID, + "microsoft.network/networksecuritygroups", + {"subnets": [{"id": _SUBNET_ID}]}, +) + +# ARG returns NIC with ipConfigurations.properties.subnet.id +_NIC = _resource( + _NIC_ID, + "microsoft.network/networkinterfaces", + { + "ipConfigurations": [ + { + "properties": { + "subnet": {"id": _SUBNET_ID}, + "privateIPAddress": "10.0.0.4", + } + } + ] + }, +) + +_VM = _resource( + _VM_ID, + "microsoft.compute/virtualmachines", + { + "identity": { + "type": "UserAssigned", + "userAssignedIdentities": {_IDENTITY_ID: {}}, + } + }, +) + + +def _snapshot(*resources): + return InventorySnapshot( + snapshot_id=_SNAP, + tenant_id=_TENANT, + requested_subscriptions=(_SUB,), + status=InventoryStatus.COMPLETE, + collected_at="2026-09-30T00:00:00+00:00", + duration_ms=42, + pages=1, + resources=tuple(resources), + errors=(), + ) + + +# Tests ------------------------------------------------------------------------- + + +def test_synthesise_extracts_subnet_from_vnet(): + """Subnet nested in VNet properties.subnets must produce a synthetic node.""" + snapshot = _snapshot(_VNET, _NSG, _NIC, _VM) + subnets = _synthesise_subnet_resources(snapshot) + assert len(subnets) == 1 + subnet = subnets[0] + assert subnet.resource_id == _SUBNET_ID + assert subnet.resource_type == "microsoft.network/virtualnetworks/subnets" + assert subnet.tenant_id == _TENANT + assert subnet.subscription_id == _SUB + assert subnet.name == "default" + assert subnet.location == "uksouth" + + +def test_synthesise_subnet_inherits_vnet_metadata(): + """Synthesised subnet must carry VNet's tenant_id, subscription_id and location.""" + snapshot = _snapshot(_VNET) + subnets = _synthesise_subnet_resources(snapshot) + assert subnets[0].snapshot_id == _SNAP + assert subnets[0].resource_group == "rg-prod" + + +def test_synthesise_subnet_properties_from_arg(): + """Subnet properties block from ARG (addressPrefix, NSG link) must be preserved.""" + snapshot = _snapshot(_VNET) + subnets = _synthesise_subnet_resources(snapshot) + assert "addressPrefix" in subnets[0].properties + assert subnets[0].properties["addressPrefix"] == "10.0.0.0/24" + + +def test_synthesise_no_subnets_returns_empty(): + """VNet with empty subnets list must not produce any synthetic nodes.""" + vnet_no_subnets = _resource(_VNET_ID, "microsoft.network/virtualnetworks", {"subnets": []}) + snapshot = _snapshot(vnet_no_subnets) + assert _synthesise_subnet_resources(snapshot) == [] + + +def test_synthesise_non_vnet_resources_ignored(): + """NSGs, NICs, and VMs must not produce subnet synthetics.""" + snapshot = _snapshot(_NSG, _NIC, _VM) + assert _synthesise_subnet_resources(snapshot) == [] + + +def test_synthesise_missing_subnet_id_skipped(): + """Subnet entry with no id field must be silently skipped.""" + vnet = _resource( + _VNET_ID, + "microsoft.network/virtualnetworks", + {"subnets": [{"name": "broken", "properties": {}}]}, + ) + snapshot = _snapshot(vnet) + assert _synthesise_subnet_resources(snapshot) == [] + + +def test_nsg_to_subnet_edge_requires_subnet_in_resource_ids(): + """Confirm the original problem: SubnetToResourceDetector drops NIC->subnet + when subnet is not in resource_ids. After synthesis it must be present.""" + from scanner.graph.edge_detector import SubnetToResourceDetector + + # Without synthesis: NIC->subnet edge is dropped + snapshot_no_subnet = _snapshot(_NSG, _NIC, _VM) + edges_without = SubnetToResourceDetector().detect(snapshot_no_subnet) + assert not any(e.target_resource_id == _SUBNET_ID for e in edges_without), ( + "subnet should not appear in edges when it has no node" + ) + + # With synthesis: subnet is in resource_ids, edge is emitted + subnets = _synthesise_subnet_resources(_snapshot(_VNET, _NSG, _NIC, _VM)) + from dataclasses import replace as dc_replace + + augmented = dc_replace( + snapshot_no_subnet, + resources=snapshot_no_subnet.resources + tuple(subnets), + ) + edges_with = SubnetToResourceDetector().detect(augmented) + assert any(e.target_resource_id == _SUBNET_ID for e in edges_with), ( + "NIC->subnet MEMBER_OF edge must be present after subnet synthesis" + ) From 4b478ae115d11e852e73d2a35dc311f5e0f334cb Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Sun, 4 Oct 2026 02:42:29 +0100 Subject: [PATCH 09/13] test(graph): verify DEFAULT_QUERY bag_merge and identity field parsing by IdentityToResourceDetector Signed-off-by: Tanvir Farhad --- tests/test_arg_inventory.py | 60 +++++++++++++++++++++++++++++++++++++ 1 file changed, 60 insertions(+) diff --git a/tests/test_arg_inventory.py b/tests/test_arg_inventory.py index 17fe54ec..55f95f56 100644 --- a/tests/test_arg_inventory.py +++ b/tests/test_arg_inventory.py @@ -227,3 +227,63 @@ def test_context_manager_closes_sdk_client(): pass client.close.assert_called_once_with() + + +def test_default_query_contains_bag_merge_and_identity(): + """DEFAULT_QUERY must merge the top-level identity column into properties. + + ARG does not include the identity field in the properties column by default. + The bag_merge call ensures it is accessible as properties.identity so that + IdentityToResourceDetector can read it without special-casing the query caller. + """ + from scanner.arg_inventory import DEFAULT_QUERY + + assert "bag_merge" in DEFAULT_QUERY + assert "identity" in DEFAULT_QUERY + + +def test_identity_field_merged_into_properties_is_parsed_by_identity_detector(): + """A row whose identity is merged into properties by bag_merge is detected correctly. + + Simulates the ARG response shape produced by DEFAULT_QUERY: the top-level + identity column is merged into the properties dict before ArgInventoryClient + returns it, so IdentityToResourceDetector should find userAssignedIdentities there. + """ + from scanner.arg_inventory import InventoryResource, InventorySnapshot, InventoryStatus + from scanner.graph.edge_detector import IdentityToResourceDetector + + vm_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.Compute/virtualMachines/vm1" + identity_id = "/subscriptions/sub/resourceGroups/rg/providers/Microsoft.ManagedIdentity/userAssignedIdentities/mid1" + + # Simulate what ARG returns after bag_merge(properties, pack('identity', identity)): + # the resource row has properties.identity set with userAssignedIdentities. + resource = InventoryResource( + snapshot_id="snap-kql", + tenant_id=TENANT_ID, + subscription_id=SUBSCRIPTION_A, + resource_id=vm_id, + resource_type="microsoft.compute/virtualmachines", + name="vm1", + location="eastus", + resource_group="rg", + tags={}, + properties={"identity": {"userAssignedIdentities": {identity_id: {}}}}, + ) + snapshot = InventorySnapshot( + snapshot_id="snap-kql", + tenant_id=TENANT_ID, + requested_subscriptions=(SUBSCRIPTION_A,), + status=InventoryStatus.COMPLETE, + collected_at="2026-09-24T00:00:00+00:00", + duration_ms=1, + pages=1, + resources=(resource,), + errors=(), + ) + + edges = IdentityToResourceDetector().detect(snapshot) + + assert len(edges) == 1 + assert edges[0].relationship_type == "HAS_IDENTITY" + assert edges[0].source_resource_id == identity_id + assert edges[0].target_resource_id == vm_id From c4f8b81f09cada4902bbf4a36738718c619b87eb Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Sun, 4 Oct 2026 14:26:46 +0100 Subject: [PATCH 10/13] fix: close psycopg2 connections explicitly, fix edge rowcount, correct PublicIp path trimming to 3 segments Signed-off-by: Tanvir Farhad --- scanner/graph/edge_detector.py | 6 +++--- scanner/graph/graph_populator.py | 13 ++++++++++--- scanner/graph/node_service.py | 25 ++++++++++++++++++++----- 3 files changed, 33 insertions(+), 11 deletions(-) diff --git a/scanner/graph/edge_detector.py b/scanner/graph/edge_detector.py index ce5878d1..2c494a20 100644 --- a/scanner/graph/edge_detector.py +++ b/scanner/graph/edge_detector.py @@ -108,12 +108,12 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: target_id = ip_config.get("id") if not target_id: continue - # Strip the NIC sub-path to get the parent VM resource ID (4 provider path segments) + # Trim to the NIC resource ID (namespace/type/name = 3 segments). parts = target_id.split("/providers/") if len(parts) >= 2: provider_path = parts[-1].split("/") - if len(provider_path) >= 4: - target_id = "/providers/".join(parts[:-1]) + "/providers/" + "/".join(provider_path[:4]) + if len(provider_path) >= 3: + target_id = "/providers/".join(parts[:-1]) + "/providers/" + "/".join(provider_path[:3]) edges.append( GraphEdge( source_resource_id=resource.resource_id, diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index c3a073dc..1c62ef87 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING import psycopg2 -import psycopg2.extras + from scanner.arg_inventory import InventoryResource from scanner.graph.node_service import link_findings_to_nodes, populate_nodes @@ -51,7 +51,8 @@ def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str) -> int if not edges: return 0 written = 0 - with psycopg2.connect(dsn) as conn: + conn = psycopg2.connect(dsn) + try: with conn.cursor() as cur: for edge in edges: cur.execute( @@ -67,7 +68,13 @@ def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str) -> int "confidence": edge.confidence, }, ) - written += 1 + written += max(cur.rowcount, 0) + conn.commit() + except Exception: + conn.rollback() + raise + finally: + conn.close() return written diff --git a/scanner/graph/node_service.py b/scanner/graph/node_service.py index 2b571e7f..50355b8c 100644 --- a/scanner/graph/node_service.py +++ b/scanner/graph/node_service.py @@ -8,7 +8,7 @@ from typing import TYPE_CHECKING import psycopg2 -import psycopg2.extras + if TYPE_CHECKING: from scanner.arg_inventory import InventorySnapshot @@ -53,7 +53,8 @@ def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: return 0 written = 0 - with psycopg2.connect(dsn) as conn: + conn = psycopg2.connect(dsn) + try: with conn.cursor() as cur: for resource in snapshot.resources: cur.execute( @@ -71,13 +72,27 @@ def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: "properties": json.dumps(resource.properties), }, ) - written += 1 + written += max(cur.rowcount, 0) + conn.commit() + except Exception: + conn.rollback() + raise + finally: + conn.close() return written def link_findings_to_nodes(scan_id: str, tenant_id: str, dsn: str) -> int: """Link findings from this scan to their graph nodes by resource_id.""" - with psycopg2.connect(dsn) as conn: + conn = psycopg2.connect(dsn) + try: with conn.cursor() as cur: cur.execute(_LINK_FINDINGS_SQL, {"scan_id": scan_id, "tenant_id": tenant_id}) - return cur.rowcount if cur.rowcount >= 0 else 0 + result = cur.rowcount if cur.rowcount >= 0 else 0 + conn.commit() + return result + except Exception: + conn.rollback() + raise + finally: + conn.close() From 55c52fadadfa2b4761cf1ee78b809e48622c910b Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 8 Oct 2026 14:10:16 +0100 Subject: [PATCH 11/13] fix(graph): publish atomic current scopes and expire complete snapshots Signed-off-by: Tanvir Farhad --- .../e2c4a6b8d013_graph_current_scope.py | 48 ++++++ scanner/graph/graph_populator.py | 93 ++++++++---- scanner/graph/node_service.py | 54 ++++--- tests/test_graph_freshness_postgres.py | 139 ++++++++++++++++++ tests/test_graph_freshness_schema.py | 17 +++ tests/test_graph_node_service.py | 13 ++ 6 files changed, 313 insertions(+), 51 deletions(-) create mode 100644 alembic/versions/e2c4a6b8d013_graph_current_scope.py create mode 100644 tests/test_graph_freshness_postgres.py create mode 100644 tests/test_graph_freshness_schema.py diff --git a/alembic/versions/e2c4a6b8d013_graph_current_scope.py b/alembic/versions/e2c4a6b8d013_graph_current_scope.py new file mode 100644 index 00000000..cf4ca8d2 --- /dev/null +++ b/alembic/versions/e2c4a6b8d013_graph_current_scope.py @@ -0,0 +1,48 @@ +"""Track the explicit current snapshot for each tenant/subscription graph. + +Revision ID: e2c4a6b8d013 +Revises: e1f2a3b4c5d6 +""" + +from typing import Sequence, Union +from alembic import op + +revision: str = "e2c4a6b8d013" +down_revision: str = "e1f2a3b4c5d6" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.execute(""" + CREATE TABLE graph_snapshot_scopes ( + tenant_id text NOT NULL, + subscription_id text NOT NULL, + snapshot_id text NOT NULL, + status text NOT NULL CHECK (status IN ('COMPLETE', 'PARTIAL')), + collected_at timestamptz NOT NULL, + PRIMARY KEY (tenant_id, subscription_id) + ) + """) + op.execute(""" + CREATE VIEW current_graph_nodes AS + SELECT n.* FROM graph_nodes n + JOIN graph_snapshot_scopes s + ON s.tenant_id = n.tenant_id AND s.subscription_id = n.subscription_id + AND s.snapshot_id = n.snapshot_id + """) + op.execute(""" + CREATE VIEW current_graph_edges AS + SELECT e.* FROM graph_edges e + JOIN current_graph_nodes src ON src.node_id = e.source_node_id + JOIN current_graph_nodes tgt ON tgt.node_id = e.target_node_id + WHERE src.tenant_id = tgt.tenant_id + AND e.evidence_snapshot_id = src.snapshot_id + AND e.evidence_snapshot_id = tgt.snapshot_id + """) + + +def downgrade() -> None: + op.execute("DROP VIEW current_graph_edges") + op.execute("DROP VIEW current_graph_nodes") + op.execute("DROP TABLE graph_snapshot_scopes") diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index 1c62ef87..eb8e4095 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -7,11 +7,9 @@ from dataclasses import replace as dc_replace from typing import TYPE_CHECKING -import psycopg2 - -from scanner.arg_inventory import InventoryResource -from scanner.graph.node_service import link_findings_to_nodes, populate_nodes +from scanner.arg_inventory import InventoryResource, InventoryStatus +from scanner.graph.node_service import graph_connection, link_findings_to_nodes, populate_nodes from scanner.graph.edge_detector import detect_all_edges if TYPE_CHECKING: @@ -39,6 +37,8 @@ AND src.tenant_id = %(tenant_id)s AND lower(tgt.resource_id) = lower(%(target_resource_id)s) AND tgt.tenant_id = %(tenant_id)s + AND src.snapshot_id = %(evidence_snapshot_id)s + AND tgt.snapshot_id = %(evidence_snapshot_id)s ON CONFLICT (source_node_id, target_node_id, relationship_type) DO UPDATE SET confidence = EXCLUDED.confidence, evidence_source = EXCLUDED.evidence_source, @@ -47,12 +47,11 @@ """ -def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str) -> int: +def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str, *, connection=None) -> int: if not edges: return 0 written = 0 - conn = psycopg2.connect(dsn) - try: + with graph_connection(dsn, connection) as conn: with conn.cursor() as cur: for edge in edges: cur.execute( @@ -69,12 +68,6 @@ def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str) -> int }, ) written += max(cur.rowcount, 0) - conn.commit() - except Exception: - conn.rollback() - raise - finally: - conn.close() return written @@ -127,22 +120,62 @@ def populate_graph(scan_id: str, snapshot: InventorySnapshot, dsn: str) -> None: resources=snapshot.resources + tuple(subnet_resources), ) - try: - node_count = populate_nodes(augmented_snapshot, dsn) - logger.info("graph: upserted %d nodes for scan %s", node_count, scan_id) - except Exception as exc: - logger.warning("graph: node population failed for scan %s: %s", scan_id, exc) + if snapshot.status == InventoryStatus.FAILED: return - - try: - edges = detect_all_edges(augmented_snapshot) - edge_count = _write_edges(edges, snapshot.snapshot_id, snapshot.tenant_id, dsn) - logger.info("graph: wrote %d edges for scan %s", edge_count, scan_id) - except Exception as exc: - logger.warning("graph: edge population failed for scan %s: %s", scan_id, exc) - try: - link_count = link_findings_to_nodes(scan_id, snapshot.tenant_id, dsn) - logger.info("graph: linked %d findings to nodes for scan %s", link_count, scan_id) - except Exception as exc: - logger.warning("graph: finding link failed for scan %s: %s", scan_id, exc) + # Publish the new scope only after every write succeeds. Partial snapshots + # retain historical rows, but the views expose only explicit current evidence. + with graph_connection(dsn) as conn: + node_count = populate_nodes(augmented_snapshot, dsn, connection=conn) + edges = detect_all_edges(augmented_snapshot) + edge_count = _write_edges(edges, snapshot.snapshot_id, snapshot.tenant_id, dsn, connection=conn) + with conn.cursor() as cur: + params = { + "tenant": snapshot.tenant_id, + "subscriptions": list(snapshot.requested_subscriptions), + "snapshot": snapshot.snapshot_id, + } + if snapshot.status == InventoryStatus.COMPLETE: + cur.execute( + """ + DELETE FROM graph_edges e USING graph_nodes src, graph_nodes tgt + WHERE src.node_id = e.source_node_id AND tgt.node_id = e.target_node_id + AND src.tenant_id = %(tenant)s AND tgt.tenant_id = %(tenant)s + AND (src.subscription_id = ANY(%(subscriptions)s) + OR tgt.subscription_id = ANY(%(subscriptions)s)) + AND e.evidence_snapshot_id <> %(snapshot)s + """, + params, + ) + cur.execute( + """ + DELETE FROM graph_nodes + WHERE tenant_id = %(tenant)s AND subscription_id = ANY(%(subscriptions)s) + AND snapshot_id <> %(snapshot)s + """, + params, + ) + for subscription in snapshot.requested_subscriptions: + cur.execute( + """ + INSERT INTO graph_snapshot_scopes + (tenant_id, subscription_id, snapshot_id, status, collected_at) + VALUES (%s, %s, %s, %s, %s) + ON CONFLICT (tenant_id, subscription_id) DO UPDATE SET + snapshot_id = EXCLUDED.snapshot_id, status = EXCLUDED.status, + collected_at = EXCLUDED.collected_at + """, + ( + snapshot.tenant_id, + subscription, + snapshot.snapshot_id, + snapshot.status.value, + snapshot.collected_at, + ), + ) + link_count = link_findings_to_nodes(scan_id, snapshot.tenant_id, dsn, connection=conn) + logger.info( + "graph: wrote %d nodes, %d edges, %d links for scan %s", node_count, edge_count, link_count, scan_id + ) + except Exception: + logger.warning("graph: population failed for scan %s", scan_id, exc_info=True) diff --git a/scanner/graph/node_service.py b/scanner/graph/node_service.py index 50355b8c..0ae48676 100644 --- a/scanner/graph/node_service.py +++ b/scanner/graph/node_service.py @@ -5,6 +5,7 @@ import json import logging import uuid +from contextlib import contextmanager from typing import TYPE_CHECKING import psycopg2 @@ -40,30 +41,54 @@ INSERT INTO finding_graph_nodes (finding_id, node_id) SELECT f.id, n.node_id FROM findings f -JOIN graph_nodes n ON lower(f.resource_id) = lower(n.resource_id) +JOIN scans s ON s.scan_id = f.scan_id +JOIN current_graph_nodes n ON lower(f.resource_id) = lower(n.resource_id) AND n.tenant_id = %(tenant_id)s + AND n.subscription_id = s.subscription_id WHERE f.scan_id = %(scan_id)s ON CONFLICT DO NOTHING """ -def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: +@contextmanager +def graph_connection(dsn: str, existing=None): + """Own a transaction only when no caller transaction was provided.""" + if existing is not None: + yield existing + return + conn = psycopg2.connect(dsn) + try: + yield conn + conn.commit() + except Exception: + conn.rollback() + raise + finally: + conn.close() + + +def populate_nodes(snapshot: InventorySnapshot, dsn: str, *, connection=None) -> int: """Upsert graph_nodes from snapshot resources. Returns count of rows written.""" if not snapshot.resources: return 0 written = 0 - conn = psycopg2.connect(dsn) - try: + with graph_connection(dsn, connection) as conn: with conn.cursor() as cur: for resource in snapshot.resources: + if ( + resource.tenant_id != snapshot.tenant_id + or resource.subscription_id not in snapshot.requested_subscriptions + or resource.snapshot_id != snapshot.snapshot_id + ): + raise ValueError("resource is outside its snapshot scope") cur.execute( _UPSERT_NODE_SQL, { "node_id": str(uuid.uuid4()), "tenant_id": resource.tenant_id, "subscription_id": resource.subscription_id, - "resource_id": resource.resource_id, + "resource_id": resource.resource_id.lower(), "resource_type": resource.resource_type, "name": resource.name, "location": resource.location, @@ -73,26 +98,13 @@ def populate_nodes(snapshot: InventorySnapshot, dsn: str) -> int: }, ) written += max(cur.rowcount, 0) - conn.commit() - except Exception: - conn.rollback() - raise - finally: - conn.close() return written -def link_findings_to_nodes(scan_id: str, tenant_id: str, dsn: str) -> int: +def link_findings_to_nodes(scan_id: str, tenant_id: str, dsn: str, *, connection=None) -> int: """Link findings from this scan to their graph nodes by resource_id.""" - conn = psycopg2.connect(dsn) - try: + with graph_connection(dsn, connection) as conn: with conn.cursor() as cur: cur.execute(_LINK_FINDINGS_SQL, {"scan_id": scan_id, "tenant_id": tenant_id}) result = cur.rowcount if cur.rowcount >= 0 else 0 - conn.commit() - return result - except Exception: - conn.rollback() - raise - finally: - conn.close() + return result diff --git a/tests/test_graph_freshness_postgres.py b/tests/test_graph_freshness_postgres.py new file mode 100644 index 00000000..5c802249 --- /dev/null +++ b/tests/test_graph_freshness_postgres.py @@ -0,0 +1,139 @@ +"""Current graph evidence across successive complete and partial scans.""" + +import os +import uuid +from dataclasses import replace + +import psycopg2 +import pytest + +from scanner.arg_inventory import InventoryResource, InventorySnapshot, InventoryStatus +from scanner.graph.graph_populator import populate_graph + +pytestmark = pytest.mark.skipif(not os.environ.get("DATABASE_URL"), reason="requires PostgreSQL") + + +@pytest.fixture +def graph_scope(): + tenant, subscription = str(uuid.uuid4()), str(uuid.uuid4()) + dsn = os.environ["DATABASE_URL"] + yield tenant, subscription, dsn + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("DELETE FROM graph_nodes WHERE tenant_id = %s", (tenant,)) + cur.execute("DELETE FROM graph_snapshot_scopes WHERE tenant_id = %s", (tenant,)) + + +def snapshot(tenant, subscription, *, status=InventoryStatus.COMPLETE, linked=True, empty=False): + sid = str(uuid.uuid4()) + prefix = f"/subscriptions/{subscription}/resourceGroups/rg/providers/Microsoft.Network/" + subnet = prefix + "virtualNetworks/v/subnets/s" + nic = prefix + "networkInterfaces/n" + + def resource(rid, kind, properties): + return InventoryResource( + sid, tenant, subscription, rid, kind, rid.split("/")[-1], "uksouth", "rg", {}, properties + ) + + resources = ( + () + if empty + else ( + resource(subnet, "microsoft.network/virtualnetworks/subnets", {}), + resource( + nic, + "microsoft.network/networkinterfaces", + {"ipConfigurations": [{"properties": {"subnet": {"id": subnet}}}]} if linked else {}, + ), + ) + ) + return InventorySnapshot(sid, tenant, (subscription,), status, "2026-10-07T00:00:00Z", 1, 1, resources, ()) + + +def counts(dsn, tenant, subscription, current=False): + table = "current_graph_nodes" if current else "graph_nodes" + edges = "current_graph_edges" if current else "graph_edges" + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute( + f"SELECT count(*) FROM {table} WHERE tenant_id=%s AND subscription_id=%s", (tenant, subscription) + ) + nodes = cur.fetchone()[0] + cur.execute( + f"SELECT count(*) FROM {edges} e JOIN graph_nodes n ON n.node_id=e.source_node_id " + "WHERE n.tenant_id=%s AND n.subscription_id=%s", + (tenant, subscription), + ) + return nodes, cur.fetchone()[0] + + +def test_complete_snapshot_removes_detached_relationship(graph_scope): + tenant, sub, dsn = graph_scope + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub), dsn) + assert counts(dsn, tenant, sub) == (2, 1) + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub, linked=False), dsn) + assert counts(dsn, tenant, sub) == (2, 0) + assert counts(dsn, tenant, sub, current=True) == (2, 0) + + +def test_empty_complete_snapshot_expires_only_its_scope(graph_scope): + tenant, sub, dsn = graph_scope + other = str(uuid.uuid4()) + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub), dsn) + populate_graph(str(uuid.uuid4()), snapshot(tenant, other), dsn) + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub, empty=True), dsn) + assert counts(dsn, tenant, sub) == (0, 0) + assert counts(dsn, tenant, other, current=True) == (2, 1) + + +def test_partial_snapshot_preserves_history_but_hides_unobserved_evidence(graph_scope): + tenant, sub, dsn = graph_scope + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub), dsn) + partial = snapshot(tenant, sub, status=InventoryStatus.PARTIAL, linked=False) + partial = replace(partial, resources=partial.resources[:1]) + populate_graph(str(uuid.uuid4()), partial, dsn) + assert counts(dsn, tenant, sub) == (2, 1) + assert counts(dsn, tenant, sub, current=True) == (1, 0) + + +def test_failed_snapshot_keeps_last_published_scope(graph_scope): + tenant, sub, dsn = graph_scope + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub), dsn) + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub, status=InventoryStatus.FAILED, empty=True), dsn) + assert counts(dsn, tenant, sub, current=True) == (2, 1) + + +def test_foreign_resource_cannot_overwrite_other_tenant(graph_scope): + tenant, sub, dsn = graph_scope + good = snapshot(tenant, sub) + populate_graph(str(uuid.uuid4()), good, dsn) + foreign = replace(good, tenant_id=str(uuid.uuid4()), snapshot_id=str(uuid.uuid4())) + populate_graph(str(uuid.uuid4()), foreign, dsn) + assert counts(dsn, tenant, sub, current=True) == (2, 1) + + +def test_publication_failure_keeps_previous_graph_atomically(graph_scope, monkeypatch): + tenant, sub, dsn = graph_scope + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub), dsn) + + def fail_edges(_snapshot): + raise RuntimeError("detector failure") + + monkeypatch.setattr("scanner.graph.graph_populator.detect_all_edges", fail_edges) + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub, linked=False), dsn) + assert counts(dsn, tenant, sub, current=True) == (2, 1) + + +def test_complete_snapshot_does_not_prune_another_tenant(graph_scope): + tenant, sub, dsn = graph_scope + other = str(uuid.uuid4()) + try: + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub), dsn) + populate_graph(str(uuid.uuid4()), snapshot(other, sub), dsn) + populate_graph(str(uuid.uuid4()), snapshot(tenant, sub, empty=True), dsn) + assert counts(dsn, other, sub, current=True) == (2, 1) + finally: + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("DELETE FROM graph_nodes WHERE tenant_id=%s", (other,)) + cur.execute("DELETE FROM graph_snapshot_scopes WHERE tenant_id=%s", (other,)) diff --git a/tests/test_graph_freshness_schema.py b/tests/test_graph_freshness_schema.py new file mode 100644 index 00000000..b9b71b74 --- /dev/null +++ b/tests/test_graph_freshness_schema.py @@ -0,0 +1,17 @@ +"""Validate current-evidence schema can represent empty and partial scopes.""" + +from pathlib import Path + +from alembic.config import Config +from alembic.script import ScriptDirectory + + +def test_current_graph_schema_follows_foundation(): + script = ScriptDirectory.from_config(Config("alembic.ini")) + revision = script.get_revision("e2c4a6b8d013") + assert revision.revision == "e2c4a6b8d013" + assert revision.down_revision == "e1f2a3b4c5d6" + text = Path(revision.path).read_text() + assert "graph_snapshot_scopes" in text + assert "current_graph_nodes" in text + assert "current_graph_edges" in text diff --git a/tests/test_graph_node_service.py b/tests/test_graph_node_service.py index f4204001..862ba102 100644 --- a/tests/test_graph_node_service.py +++ b/tests/test_graph_node_service.py @@ -48,6 +48,7 @@ def test_populate_nodes_upserts_each_resource(): mock_conn = MagicMock() mock_cur = MagicMock() + mock_cur.rowcount = 1 mock_conn.__enter__ = MagicMock(return_value=mock_conn) mock_conn.__exit__ = MagicMock(return_value=False) mock_cur.__enter__ = MagicMock(return_value=mock_cur) @@ -78,6 +79,7 @@ def test_populate_nodes_does_not_write_cross_tenant_resources(): captured_args = [] mock_conn = MagicMock() mock_cur = MagicMock() + mock_cur.rowcount = 1 mock_conn.__enter__ = MagicMock(return_value=mock_conn) mock_conn.__exit__ = MagicMock(return_value=False) mock_cur.__enter__ = MagicMock(return_value=mock_cur) @@ -100,6 +102,7 @@ def capture_execute(sql, params=None): def test_link_findings_to_nodes_executes_insert(): mock_conn = MagicMock() mock_cur = MagicMock() + mock_cur.rowcount = 1 mock_conn.__enter__ = MagicMock(return_value=mock_conn) mock_conn.__exit__ = MagicMock(return_value=False) mock_cur.__enter__ = MagicMock(return_value=mock_cur) @@ -113,3 +116,13 @@ def test_link_findings_to_nodes_executes_insert(): assert mock_cur.execute.called sql_called = mock_cur.execute.call_args[0][0] assert "finding_graph_nodes" in sql_called + + +def test_populate_nodes_reports_zero_actual_writes(): + mock_conn = MagicMock() + mock_cur = mock_conn.cursor.return_value.__enter__.return_value + mock_cur.rowcount = 0 + with patch("scanner.graph.node_service.psycopg2.connect", return_value=mock_conn): + assert ( + populate_nodes(_make_snapshot([_make_resource("/subscriptions/sub/resource")]), "postgresql://test/db") == 0 + ) From 0dc393a1df4f63d038a250c57de1006fb249de8d Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 8 Oct 2026 14:48:14 +0100 Subject: [PATCH 12/13] fix(graph): preserve published evidence when an edge detector fails Signed-off-by: Tanvir Farhad --- scanner/graph/edge_detector.py | 4 ++- tests/test_graph_detector_failure_postgres.py | 32 +++++++++++++++++++ 2 files changed, 35 insertions(+), 1 deletion(-) create mode 100644 tests/test_graph_detector_failure_postgres.py diff --git a/scanner/graph/edge_detector.py b/scanner/graph/edge_detector.py index 2c494a20..79dfed3a 100644 --- a/scanner/graph/edge_detector.py +++ b/scanner/graph/edge_detector.py @@ -177,7 +177,8 @@ def detect(self, snapshot: InventorySnapshot) -> list[GraphEdge]: def detect_all_edges(snapshot: InventorySnapshot) -> list[GraphEdge]: """Run all detectors and return the combined edge list. - Per-detector exceptions are caught and logged; the function never raises. + A detector failure aborts graph publication. Publishing a partial edge list + as complete would remove valid relationships from the preceding snapshot. """ detectors: list[EdgeDetector] = [ NsgToSubnetDetector(), @@ -192,4 +193,5 @@ def detect_all_edges(snapshot: InventorySnapshot) -> list[GraphEdge]: edges.extend(detector.detect(snapshot)) except Exception as exc: logger.warning("edge_detector: %s failed: %s", type(detector).__name__, exc) + raise return edges diff --git a/tests/test_graph_detector_failure_postgres.py b/tests/test_graph_detector_failure_postgres.py new file mode 100644 index 00000000..f1dc79d9 --- /dev/null +++ b/tests/test_graph_detector_failure_postgres.py @@ -0,0 +1,32 @@ +"""Detector failure must preserve the last published graph evidence.""" + +import uuid +from dataclasses import replace + +from tests.test_graph_freshness_postgres import counts, snapshot +from tests.test_graph_freshness_postgres import graph_scope as _graph_scope +from tests.test_graph_freshness_postgres import pytestmark as pytestmark + +from scanner.graph.graph_populator import populate_graph + +graph_scope = _graph_scope + + +def test_malformed_detector_input_preserves_published_graph(graph_scope): + tenant, subscription, dsn = graph_scope + good = replace(snapshot(tenant, subscription), collected_at="2026-10-07T01:00:00Z") + populate_graph(str(uuid.uuid4()), good, dsn) + assert counts(dsn, tenant, subscription, current=True) == (2, 1) + + malformed = replace(snapshot(tenant, subscription), collected_at="2026-10-07T02:00:00Z") + malformed = replace( + malformed, + resources=( + malformed.resources[0], + replace(malformed.resources[1], properties={"ipConfigurations": [None]}), + ), + ) + populate_graph(str(uuid.uuid4()), malformed, dsn) + + assert counts(dsn, tenant, subscription, current=True) == (2, 1) + assert counts(dsn, tenant, subscription) == (2, 1) From a01afb5eb7c808537349401c59c1a7465190344a Mon Sep 17 00:00:00 2001 From: Tanvir Farhad Date: Thu, 8 Oct 2026 15:58:14 +0100 Subject: [PATCH 13/13] fix(graph): serialize scope publication and reject delayed snapshots Signed-off-by: Tanvir Farhad --- scanner/graph/graph_populator.py | 12 ++++++- scanner/graph/node_service.py | 10 ++++++ tests/test_graph_freshness_postgres.py | 47 ++++++++++++++++++++++++++ 3 files changed, 68 insertions(+), 1 deletion(-) diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index eb8e4095..d26e63ef 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -9,7 +9,7 @@ from scanner.arg_inventory import InventoryResource, InventoryStatus -from scanner.graph.node_service import graph_connection, link_findings_to_nodes, populate_nodes +from scanner.graph.node_service import graph_connection, link_findings_to_nodes, lock_graph_scopes, populate_nodes from scanner.graph.edge_detector import detect_all_edges if TYPE_CHECKING: @@ -126,6 +126,16 @@ def populate_graph(scan_id: str, snapshot: InventorySnapshot, dsn: str) -> None: # Publish the new scope only after every write succeeds. Partial snapshots # retain historical rows, but the views expose only explicit current evidence. with graph_connection(dsn) as conn: + lock_graph_scopes(conn, snapshot.tenant_id, snapshot.requested_subscriptions) + with conn.cursor() as cur: + cur.execute( + "SELECT 1 FROM graph_snapshot_scopes WHERE tenant_id=%s " + "AND subscription_id=ANY(%s) AND collected_at > %s::timestamptz LIMIT 1", + (snapshot.tenant_id, list(snapshot.requested_subscriptions), snapshot.collected_at), + ) + if cur.fetchone() is not None: + logger.info("graph: ignored older snapshot for scan %s", scan_id) + return node_count = populate_nodes(augmented_snapshot, dsn, connection=conn) edges = detect_all_edges(augmented_snapshot) edge_count = _write_edges(edges, snapshot.snapshot_id, snapshot.tenant_id, dsn, connection=conn) diff --git a/scanner/graph/node_service.py b/scanner/graph/node_service.py index 0ae48676..bc60ae49 100644 --- a/scanner/graph/node_service.py +++ b/scanner/graph/node_service.py @@ -108,3 +108,13 @@ def link_findings_to_nodes(scan_id: str, tenant_id: str, dsn: str, *, connection cur.execute(_LINK_FINDINGS_SQL, {"scan_id": scan_id, "tenant_id": tenant_id}) result = cur.rowcount if cur.rowcount >= 0 else 0 return result + + +def lock_graph_scopes(conn, tenant_id: str, subscriptions) -> None: + """Serialize all publications and traversals within the same scope.""" + with conn.cursor() as cur: + for subscription in sorted(set(subscriptions)): + cur.execute( + "SELECT pg_advisory_xact_lock(hashtextextended(%s, 0))", + (f"openshield-graph:{tenant_id}:{subscription}",), + ) diff --git a/tests/test_graph_freshness_postgres.py b/tests/test_graph_freshness_postgres.py index 5c802249..b4a897ee 100644 --- a/tests/test_graph_freshness_postgres.py +++ b/tests/test_graph_freshness_postgres.py @@ -137,3 +137,50 @@ def test_complete_snapshot_does_not_prune_another_tenant(graph_scope): with conn.cursor() as cur: cur.execute("DELETE FROM graph_nodes WHERE tenant_id=%s", (other,)) cur.execute("DELETE FROM graph_snapshot_scopes WHERE tenant_id=%s", (other,)) + + +def test_delayed_older_snapshot_cannot_replace_or_prune_current_scope(graph_scope): + tenant, sub, dsn = graph_scope + newer = replace(snapshot(tenant, sub), collected_at="2026-10-07T00:02:00Z") + older = replace(snapshot(tenant, sub, empty=True), collected_at="2026-10-07T00:01:00Z") + populate_graph(str(uuid.uuid4()), newer, dsn) + populate_graph(str(uuid.uuid4()), older, dsn) + assert counts(dsn, tenant, sub, current=True) == (2, 1) + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute( + "SELECT snapshot_id FROM graph_snapshot_scopes WHERE tenant_id=%s AND subscription_id=%s", (tenant, sub) + ) + assert cur.fetchone()[0] == newer.snapshot_id + + +def test_publication_waits_for_scope_transaction_lock(graph_scope): + from concurrent.futures import ThreadPoolExecutor + from threading import Event + + tenant, sub, dsn = graph_scope + entered = Event() + pending = snapshot(tenant, sub) + blocker = psycopg2.connect(dsn) + try: + with blocker.cursor() as cur: + cur.execute("SELECT pg_advisory_xact_lock(hashtextextended(%s, 0))", (f"openshield-graph:{tenant}:{sub}",)) + + def publish(): + entered.set() + populate_graph(str(uuid.uuid4()), pending, dsn) + + with ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(publish) + assert entered.wait(2) + try: + future.result(timeout=0.3) + pytest.fail("publication completed while its scope lock was held") + except TimeoutError: + pass + finally: + blocker.rollback() + future.result(timeout=5) + assert counts(dsn, tenant, sub, current=True) == (2, 1) + finally: + blocker.close()