diff --git a/alembic/versions/e2c4a6b8d013_graph_current_scope.py b/alembic/versions/e2c4a6b8d013_graph_current_scope.py new file mode 100644 index 0000000..cf4ca8d --- /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/arg_inventory.py b/scanner/arg_inventory.py index feaefb5..9257447 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 89ff6f6..5bf08da 100644 --- a/scanner/engine.py +++ b/scanner/engine.py @@ -61,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() # ------------------------------------------------------------------ # @@ -118,6 +119,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", diff --git a/scanner/graph/edge_detector.py b/scanner/graph/edge_detector.py new file mode 100644 index 0000000..79dfed3 --- /dev/null +++ b/scanner/graph/edge_detector.py @@ -0,0 +1,197 @@ +"""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_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 + + +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 + # 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) >= 3: + target_id = "/providers/".join(parts[:-1]) + "/providers/" + "/".join(provider_path[:3]) + 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. + + 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(), + 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) + raise + return edges diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py new file mode 100644 index 0000000..d26e63e --- /dev/null +++ b/scanner/graph/graph_populator.py @@ -0,0 +1,191 @@ +"""Orchestrate post-scan graph population: nodes, edges, finding links.""" + +from __future__ import annotations + +import logging +import uuid +from dataclasses import replace as dc_replace +from typing import TYPE_CHECKING + + +from scanner.arg_inventory import InventoryResource, InventoryStatus +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: + 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 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, + evidence_snapshot_id = EXCLUDED.evidence_snapshot_id, + collected_at = now() +""" + + +def _write_edges(edges: list, snapshot_id: str, tenant_id: str, dsn: str, *, connection=None) -> int: + if not edges: + return 0 + written = 0 + with graph_connection(dsn, connection) 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, + "tenant_id": tenant_id, + "confidence": edge.confidence, + }, + ) + written += max(cur.rowcount, 0) + 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), + ) + + if snapshot.status == InventoryStatus.FAILED: + return + try: + # 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) + 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 new file mode 100644 index 0000000..bc60ae4 --- /dev/null +++ b/scanner/graph/node_service.py @@ -0,0 +1,120 @@ +"""Upsert graph nodes from an InventorySnapshot and link findings to nodes.""" + +from __future__ import annotations + +import json +import logging +import uuid +from contextlib import contextmanager +from typing import TYPE_CHECKING + +import psycopg2 + + +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 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 +""" + + +@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 + 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.lower(), + "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 += max(cur.rowcount, 0) + return written + + +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.""" + 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 + 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/scanner/worker.py b/scanner/worker.py index 21fd5bf..3aab059 100644 --- a/scanner/worker.py +++ b/scanner/worker.py @@ -29,6 +29,7 @@ from api.services.pattern_service import PatternService 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") @@ -219,6 +220,18 @@ def run_worker(): 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, + ) # Apply lifecycle tracking. Lifecycle failures are non-fatal: # the scan is already persisted, so we log and continue rather # than marking the scan as failed. diff --git a/tests/test_arg_inventory.py b/tests/test_arg_inventory.py index 17fe54e..55f95f5 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 diff --git a/tests/test_graph_detector_failure_postgres.py b/tests/test_graph_detector_failure_postgres.py new file mode 100644 index 0000000..f1dc79d --- /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) diff --git a/tests/test_graph_edge_detector.py b/tests/test_graph_edge_detector.py new file mode 100644 index 0000000..4649cb7 --- /dev/null +++ b/tests/test_graph_edge_detector.py @@ -0,0 +1,119 @@ +"""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) == [] diff --git a/tests/test_graph_engine_post_scan.py b/tests/test_graph_engine_post_scan.py new file mode 100644 index 0000000..520614f --- /dev/null +++ b/tests/test_graph_engine_post_scan.py @@ -0,0 +1,74 @@ +"""Tests for post-scan graph population wiring in ScanEngine and worker.""" + +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_snapshot_stored_on_engine_after_run_scan(engine): + snapshot = _make_snapshot() + with patch("scanner.engine.collect_snapshot", return_value=snapshot): + engine.run_scan() + assert engine.snapshot is snapshot + + +def test_snapshot_is_none_on_engine_when_collection_fails(engine): + with patch("scanner.engine.collect_snapshot", return_value=None): + engine.run_scan() + assert engine.snapshot is None + + +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.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_and_stores_snapshot_even_without_database_url(engine): + snapshot = _make_snapshot() + with ( + patch("scanner.engine.collect_snapshot", return_value=snapshot), + 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" + ) diff --git a/tests/test_graph_freshness_postgres.py b/tests/test_graph_freshness_postgres.py new file mode 100644 index 0000000..b4a897e --- /dev/null +++ b/tests/test_graph_freshness_postgres.py @@ -0,0 +1,186 @@ +"""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,)) + + +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() diff --git a/tests/test_graph_freshness_schema.py b/tests/test_graph_freshness_schema.py new file mode 100644 index 0000000..b9b71b7 --- /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 new file mode 100644 index 0000000..862ba10 --- /dev/null +++ b/tests/test_graph_node_service.py @@ -0,0 +1,128 @@ +"""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_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) + 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_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) + 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_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) + 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", "00000000-0000-0000-0000-000000000001", "postgresql://test/db") + + 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 + ) diff --git a/tests/test_graph_populator_subnet_synthesis.py b/tests/test_graph_populator_subnet_synthesis.py new file mode 100644 index 0000000..0668889 --- /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" + )