Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions scanner/arg_inventory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
2 changes: 2 additions & 0 deletions scanner/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

# ------------------------------------------------------------------ #
Expand Down Expand Up @@ -117,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",
Expand Down
195 changes: 195 additions & 0 deletions scanner/graph/edge_detector.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
"""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.

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
148 changes: 148 additions & 0 deletions scanner/graph/graph_populator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,148 @@
"""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

import psycopg2


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

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
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) -> int:
if not edges:
return 0
written = 0
conn = psycopg2.connect(dsn)
try:
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)
conn.commit()
except Exception:
conn.rollback()
raise
finally:
conn.close()
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(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(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)
Loading
Loading