diff --git a/alembic/versions/f2a3b4c5d6e7_attack_paths.py b/alembic/versions/f2a3b4c5d6e7_attack_paths.py new file mode 100644 index 0000000..e312c06 --- /dev/null +++ b/alembic/versions/f2a3b4c5d6e7_attack_paths.py @@ -0,0 +1,65 @@ +"""Add attack_paths table for pre-computed BFS traversal results. + +Revision ID: f2a3b4c5d6e7 +Revises: e1f2a3b4c5d6 +Create Date: 2026-09-24 00:00:00.000000 +""" + +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +revision: str = "f2a3b4c5d6e7" +down_revision: Union[str, Sequence[str], None] = "e1f2a3b4c5d6" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.create_table( + "attack_paths", + sa.Column("path_id", postgresql.UUID(), nullable=False), + sa.Column("tenant_id", sa.Text(), nullable=False), + sa.Column("scan_id", sa.Text(), nullable=False), + sa.Column("source_node_id", postgresql.UUID(), nullable=False), + sa.Column("target_node_id", postgresql.UUID(), nullable=False), + sa.Column("path_node_ids", postgresql.ARRAY(postgresql.UUID()), nullable=False), + sa.Column("path_length", sa.Integer(), nullable=False), + sa.Column("min_confidence", sa.Float(), nullable=False, server_default=sa.text("1.0")), + sa.Column( + "relationship_types", postgresql.ARRAY(sa.Text()), nullable=False, server_default=sa.text("ARRAY[]::text[]") + ), + sa.Column("computed_at", sa.DateTime(timezone=True), nullable=False, server_default=sa.text("now()")), + sa.ForeignKeyConstraint( + ["source_node_id"], + ["graph_nodes.node_id"], + name="attack_paths_source_fkey", + ondelete="CASCADE", + ), + sa.ForeignKeyConstraint( + ["target_node_id"], + ["graph_nodes.node_id"], + name="attack_paths_target_fkey", + ondelete="CASCADE", + ), + sa.PrimaryKeyConstraint("path_id", name="attack_paths_pkey"), + ) + op.create_index("idx_attack_paths_tenant_scan", "attack_paths", ["tenant_id", "scan_id"]) + op.create_index("idx_attack_paths_source", "attack_paths", ["source_node_id"]) + op.create_index("idx_attack_paths_target", "attack_paths", ["target_node_id"]) + op.create_index( + "uq_attack_paths_source_target_scan", + "attack_paths", + ["source_node_id", "target_node_id", "scan_id"], + unique=True, + ) + + +def downgrade() -> None: + op.drop_index("uq_attack_paths_source_target_scan", table_name="attack_paths") + op.drop_index("idx_attack_paths_target", table_name="attack_paths") + op.drop_index("idx_attack_paths_source", table_name="attack_paths") + op.drop_index("idx_attack_paths_tenant_scan", table_name="attack_paths") + op.drop_table("attack_paths") diff --git a/api/app.py b/api/app.py index 8069951..a135a24 100644 --- a/api/app.py +++ b/api/app.py @@ -256,6 +256,7 @@ def verify_jwt() -> None: # ------------------------------------------------------------------ # from api.routes.ai import ai_bp from api.routes.assurance import assurance_bp + from api.routes.attack_graph import attack_graph_bp from api.routes.cbom import cbom_bp from api.routes.compliance import compliance_bp from api.routes.drift import drift_bp @@ -268,6 +269,7 @@ def verify_jwt() -> None: app.register_blueprint(ai_bp) app.register_blueprint(assurance_bp) + app.register_blueprint(attack_graph_bp) app.register_blueprint(cbom_bp) app.register_blueprint(compliance_bp) app.register_blueprint(drift_bp) diff --git a/api/routes/attack_graph.py b/api/routes/attack_graph.py new file mode 100644 index 0000000..748234c --- /dev/null +++ b/api/routes/attack_graph.py @@ -0,0 +1,232 @@ +"""Attack graph API: resource nodes, edges, and pre-computed attack paths.""" + +import logging +import os + +import psycopg2.extras +from flask import Blueprint, g, jsonify, request + +from api.models.finding import DatabaseManager +from api.validation import ValidationError, positive_integer, uuid_string + +attack_graph_bp = Blueprint("attack_graph", __name__) +logger = logging.getLogger(__name__) + +_DEFAULT_LIMIT = 100 +_MAX_LIMIT = 500 + + +def _get_db() -> DatabaseManager: + if "db" not in g: + g.db = DatabaseManager(os.environ["DATABASE_URL"]) + g.db.connect() + return g.db + + +def _tenant_id() -> str | None: + """Resolve tenant_id from the verified principal. + + The 'tenant' field comes from the verified token's tenant claim. + For tokens without a tenant claim, admins may supply + X-Tenant-Id as a request header (never a query param, which leaks into + logs and caches). Non-admin tokens cannot override the header. + """ + user = getattr(g, "user", {}) or {} + # OIDC path: tid claim decoded by the verifier into user["tenant"] + tid = user.get("tenant") + if tid: + return tid + # Shared-secret path: admin-only header override for multi-tenant deployments + if user.get("role") == "admin": + return request.headers.get("X-Tenant-Id") or None + return None + + +@attack_graph_bp.teardown_request +def _close_db(exc): + db = g.pop("db", None) + if db is not None: + db.close() + + +@attack_graph_bp.get("/api/v1/attack-graph") +def get_attack_graph(): + """Return graph nodes and edges for the caller's tenant (latest snapshot). + + Query params: subscription_id (optional), limit (default 100, max 500) + """ + try: + limit = positive_integer(int(request.args.get("limit", _DEFAULT_LIMIT)), "limit") + if limit > _MAX_LIMIT: + limit = _MAX_LIMIT + subscription_id = request.args.get("subscription_id") + except (ValidationError, ValueError): + return jsonify({"error": "Invalid request parameters"}), 400 + + tenant_id = _tenant_id() + if not tenant_id: + # A caller without a verified tenant or admin fallback cannot access + # tenant-scoped graph evidence. + return jsonify({"error": "tenant_id not available; OIDC authentication required"}), 403 + + try: + conn = _get_db().conn + with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur: + if subscription_id: + cur.execute( + """ + SELECT n.node_id::text, n.resource_id, n.resource_type, n.name, + n.location, n.resource_group, n.subscription_id, n.snapshot_id, + n.updated_at + FROM current_graph_nodes n + WHERE n.tenant_id = %(tenant_id)s + AND n.subscription_id = %(subscription_id)s + ORDER BY n.updated_at DESC + LIMIT %(limit)s + """, + {"tenant_id": tenant_id, "subscription_id": subscription_id, "limit": limit}, + ) + else: + cur.execute( + """ + SELECT n.node_id::text, n.resource_id, n.resource_type, n.name, + n.location, n.resource_group, n.subscription_id, n.snapshot_id, + n.updated_at + FROM current_graph_nodes n + WHERE n.tenant_id = %(tenant_id)s + ORDER BY n.updated_at DESC + LIMIT %(limit)s + """, + {"tenant_id": tenant_id, "limit": limit}, + ) + nodes = cur.fetchall() + + node_ids = [row["node_id"] for row in nodes] + edges: list = [] + if node_ids: + cur.execute( + """ + SELECT e.edge_id::text, e.source_node_id::text, e.target_node_id::text, + e.relationship_type, e.confidence, e.evidence_source, e.collected_at + FROM current_graph_edges e + WHERE e.source_node_id = ANY(%(node_ids)s::uuid[]) + AND e.target_node_id = ANY(%(node_ids)s::uuid[]) + """, + {"node_ids": node_ids}, + ) + edges = cur.fetchall() + except Exception: + logger.exception("get_attack_graph failed for tenant %s", tenant_id) + return jsonify({"error": "internal server error"}), 500 + + return jsonify({"nodes": [dict(r) for r in nodes], "edges": [dict(r) for r in edges]}) + + +@attack_graph_bp.get("/api/v1/attack-paths") +def list_attack_paths(): + """Return pre-computed attack paths for a scan. + + Query params: scan_id (required), limit (default 100, max 500) + """ + scan_id = request.args.get("scan_id") + if not scan_id: + return jsonify({"error": "scan_id is required"}), 400 + try: + scan_id = uuid_string(scan_id, "scan_id") + limit = positive_integer(int(request.args.get("limit", _DEFAULT_LIMIT)), "limit") + if limit > _MAX_LIMIT: + limit = _MAX_LIMIT + except (ValidationError, ValueError): + return jsonify({"error": "Invalid request parameters"}), 400 + + tenant_id = _tenant_id() + if not tenant_id: + # A caller without a verified tenant or admin fallback cannot access + # tenant-scoped graph evidence. + return jsonify({"error": "tenant_id not available; OIDC authentication required"}), 403 + + try: + conn = _get_db().conn + with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur: + cur.execute( + """ + SELECT ap.path_id::text, ap.source_node_id::text, ap.target_node_id::text, + ap.path_node_ids, ap.path_length, ap.min_confidence, + ap.relationship_types, ap.computed_at, + src.resource_type AS source_type, src.name AS source_name, + tgt.resource_type AS target_type, tgt.name AS target_name + FROM attack_paths ap + JOIN graph_nodes src ON src.node_id = ap.source_node_id + JOIN graph_nodes tgt ON tgt.node_id = ap.target_node_id + WHERE ap.scan_id = %(scan_id)s + AND ap.tenant_id = %(tenant_id)s + ORDER BY ap.path_length ASC, ap.min_confidence DESC + LIMIT %(limit)s + """, + {"scan_id": scan_id, "tenant_id": tenant_id, "limit": limit}, + ) + rows = cur.fetchall() + except Exception: + logger.exception("list_attack_paths failed for scan %s", scan_id) + return jsonify({"error": "internal server error"}), 500 + + return jsonify({"scan_id": scan_id, "paths": [dict(r) for r in rows]}) + + +@attack_graph_bp.get("/api/v1/attack-paths/") +def get_attack_path(path_id: str): + """Return a single attack path with full node detail for each hop.""" + try: + path_id = uuid_string(path_id, "path_id") + except (ValidationError, ValueError): + return jsonify({"error": "Invalid request parameters"}), 400 + + tenant_id = _tenant_id() + if not tenant_id: + # A caller without a verified tenant or admin fallback cannot access + # tenant-scoped graph evidence. + return jsonify({"error": "tenant_id not available; OIDC authentication required"}), 403 + + try: + conn = _get_db().conn + with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur: + cur.execute( + """ + SELECT ap.path_id::text, ap.scan_id, ap.source_node_id::text, + ap.target_node_id::text, ap.path_node_ids, ap.path_length, + ap.min_confidence, ap.relationship_types, ap.computed_at + FROM attack_paths ap + WHERE ap.path_id = %(path_id)s::uuid + AND ap.tenant_id = %(tenant_id)s + """, + {"path_id": path_id, "tenant_id": tenant_id}, + ) + row = cur.fetchone() + + if row is None: + return jsonify({"error": "not found"}), 404 + + path = dict(row) + + # Fetch full node detail for each hop + node_ids = [str(nid) for nid in path["path_node_ids"]] + if node_ids: + with conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur: + cur.execute( + """ + SELECT node_id::text, resource_id, resource_type, name, + location, resource_group, subscription_id + FROM graph_nodes + WHERE node_id = ANY(%(ids)s::uuid[]) + AND tenant_id = %(tenant_id)s + """, + {"ids": node_ids, "tenant_id": tenant_id}, + ) + nodes_by_id = {r["node_id"]: dict(r) for r in cur.fetchall()} + path["hops"] = [nodes_by_id.get(str(nid), {"node_id": str(nid)}) for nid in path["path_node_ids"]] + except Exception: + logger.exception("get_attack_path failed for path_id %s", path_id) + return jsonify({"error": "internal server error"}), 500 + + path["path_node_ids"] = [str(nid) for nid in path["path_node_ids"]] + return jsonify(path) diff --git a/docs/api-reference.md b/docs/api-reference.md index 7192e14..a1771df 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -697,3 +697,17 @@ The following endpoints are called by the frontend but have no backend implement | Endpoint | Used by | Status | |---|---|---| | `GET /api/monitoring` | Monitoring page — score trend chart, category distribution | Deferred. Score and findings data come from `GET /api/score` and `GET /api/findings` instead. | + +--- + +## Attack graph endpoints + +`GET /api/attack-graph` returns explicitly observed nodes and relationships from the latest published inventory snapshot for each subscription in the caller's tenant. Complete snapshots expire absent resources and relationships within their tenant/subscription scope. Partial snapshots retain historical evidence in storage, while current graph queries and traversal exclude unobserved resources and relationships. Failed collection keeps the previous published snapshot. `GET /api/attack-paths` and `GET /api/attack-paths/` expose paths computed for a requested scan. Successful subsequent scans replace prior paths for the same tenant/subscription, including scans with no findings. + +### Authentication requirement + +Callers use the tenant claim preserved by the token verifier in either authentication mode. A supplied `X-Tenant-Id` header cannot override that claim. An authenticated admin without a tenant claim may supply `X-Tenant-Id` explicitly; viewer and operator tokens cannot select a tenant. Requests without a usable tenant return `403 {"error": "tenant_id not available; OIDC authentication required"}`. Tenant query parameters never determine graph scope. + +### Attack-path retention + +Attack paths are replaced by a newer successful scan within the same tenant/subscription. Clean scans remove prior paths even when there are no findings. Failed scans and delayed older scans retain the newer published evidence. Scope publication and traversal serialize on the same database lock. diff --git a/scanner/graph/graph_populator.py b/scanner/graph/graph_populator.py index 456f98e..d621f1c 100644 --- a/scanner/graph/graph_populator.py +++ b/scanner/graph/graph_populator.py @@ -11,6 +11,7 @@ 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 +from scanner.graph.path_traversal import compute_attack_paths if TYPE_CHECKING: from scanner.arg_inventory import InventorySnapshot @@ -189,3 +190,10 @@ def populate_graph(scan_id: str, snapshot: InventorySnapshot, dsn: str) -> None: ) except Exception: logger.warning("graph: population failed for scan %s", scan_id, exc_info=True) + return + + try: + path_count = compute_attack_paths(scan_id, snapshot.tenant_id, dsn) + logger.info("graph: computed %d attack paths for scan %s", path_count, scan_id) + except Exception as exc: + logger.warning("graph: path traversal failed for scan %s: %s", scan_id, exc) diff --git a/scanner/graph/path_traversal.py b/scanner/graph/path_traversal.py new file mode 100644 index 0000000..19e5ae7 --- /dev/null +++ b/scanner/graph/path_traversal.py @@ -0,0 +1,218 @@ +"""BFS path traversal over the attack graph stored in PostgreSQL. + +Computes shortest paths from every finding-linked node to every reachable node +and persists them in attack_paths for API consumption. +""" + +from __future__ import annotations + +import logging +import uuid +from collections import deque +from typing import Any + +import psycopg2 +import psycopg2.extras + +from scanner.graph.node_service import lock_graph_scopes + +logger = logging.getLogger(__name__) + +# Paths longer than this are not persisted because they are rarely actionable and +# keeping them would inflate the table for large graphs. +_MAX_PATH_LENGTH = 8 + +# Relationships whose edges are also traversed in reverse during BFS. +# MEMBER_OF and PROTECTS are excluded: reversing MEMBER_OF would let any two +# VMs in the same subnet reach each other via the subnet node. +_REVERSE_RELS: frozenset[str] = frozenset({"EXPOSES", "HAS_IDENTITY"}) + + +def _load_adjacency(conn: Any, tenant_id: str) -> dict[str, list[tuple[str, str, float]]]: + """Return {source_node_id: [(target_node_id, relationship_type, confidence)]}.""" + with conn.cursor() as cur: + cur.execute( + """ + SELECT e.source_node_id::text, e.target_node_id::text, + e.relationship_type, e.confidence + FROM current_graph_edges e + JOIN current_graph_nodes src ON src.node_id = e.source_node_id + WHERE src.tenant_id = %(tenant_id)s + """, + {"tenant_id": tenant_id}, + ) + # Only traverse in reverse for relationships where the reverse direction + # is semantically meaningful. MEMBER_OF and PROTECTS are not reversed: + # reversing MEMBER_OF would let any two VMs in the same subnet reach + # each other via the subnet node, producing spurious lateral-movement paths. + adj: dict[str, list[tuple[str, str, float]]] = {} + for src, tgt, rel, conf in cur.fetchall(): + adj.setdefault(src, []).append((tgt, rel, conf)) + if rel in _REVERSE_RELS: + adj.setdefault(tgt, []).append((src, rel + "_REV", conf)) + return adj + + +def _load_finding_nodes(conn: Any, scan_id: str, tenant_id: str) -> list[str]: + """Return node_ids linked to findings from this scan (same tenant).""" + with conn.cursor() as cur: + cur.execute( + """ + SELECT DISTINCT fgn.node_id::text + FROM finding_graph_nodes fgn + JOIN findings f ON f.id = fgn.finding_id + JOIN current_graph_nodes n ON n.node_id = fgn.node_id + WHERE f.scan_id = %(scan_id)s + AND n.tenant_id = %(tenant_id)s + """, + {"scan_id": scan_id, "tenant_id": tenant_id}, + ) + return [row[0] for row in cur.fetchall()] + + +def _bfs_from( + start: str, + adj: dict[str, list[tuple[str, str, float]]], +) -> list[dict[str, Any]]: + """BFS from start. Returns one path record per reachable node (shortest path).""" + paths: list[dict[str, Any]] = [] + # queue: (current_node_id, path_so_far, min_confidence, relationship_types) + queue: deque[tuple[str, list[str], float, list[str]]] = deque() + queue.append((start, [start], 1.0, [])) + visited: set[str] = {start} + + while queue: + node, path, min_conf, rels = queue.popleft() + for neighbour, rel, conf in adj.get(node, []): + if neighbour in visited: + continue + visited.add(neighbour) + new_path = path + [neighbour] + new_rels = rels + [rel] + new_conf = min(min_conf, conf) + paths.append( + { + "target_node_id": neighbour, + "path_node_ids": new_path, + "path_length": len(new_path) - 1, + "min_confidence": new_conf, + "relationship_types": new_rels, + } + ) + if len(new_path) - 1 < _MAX_PATH_LENGTH: + queue.append((neighbour, new_path, new_conf, new_rels)) + + return paths + + +def _delete_stale_paths(conn: Any, scan_id: str, tenant_id: str) -> bool: + """Delete attack paths from previous scans for the same subscription, keeping only the current scan.""" + with conn.cursor() as cur: + cur.execute("SELECT subscription_id FROM scans WHERE scan_id=%s::uuid", (scan_id,)) + scope = cur.fetchone() + if scope is None: + return False + lock_graph_scopes(conn, tenant_id, [scope[0]]) + with conn.cursor() as cur: + cur.execute( + "SELECT current.status='completed' AND NOT EXISTS (" + "SELECT 1 FROM scans newer WHERE newer.subscription_id=current.subscription_id " + "AND newer.status='completed' AND newer.started_at > current.started_at) " + "FROM scans current WHERE current.scan_id=%s::uuid", + (scan_id,), + ) + row = cur.fetchone() + if row is None or not row[0]: + return False + cur.execute( + """ + DELETE FROM attack_paths ap + USING scans previous, scans current + WHERE ap.tenant_id = %(tenant_id)s + AND ap.scan_id <> %(scan_id)s + AND previous.scan_id::text = ap.scan_id + AND current.scan_id = %(scan_id)s::uuid + AND current.status = 'completed' + AND previous.subscription_id = current.subscription_id + AND previous.started_at <= current.started_at + """, + {"tenant_id": tenant_id, "scan_id": scan_id}, + ) + + return True + + +def _write_paths( + conn: Any, + scan_id: str, + tenant_id: str, + source_node_id: str, + paths: list[dict[str, Any]], +) -> int: + if not paths: + return 0 + with conn.cursor() as cur: + rows = [ + ( + str(uuid.uuid4()), + tenant_id, + scan_id, + source_node_id, + p["target_node_id"], + p["path_node_ids"], + p["path_length"], + p["min_confidence"], + p["relationship_types"], + ) + for p in paths + ] + psycopg2.extras.execute_values( + cur, + """ + INSERT INTO attack_paths + (path_id, tenant_id, scan_id, source_node_id, target_node_id, + path_node_ids, path_length, min_confidence, relationship_types) + VALUES %s + ON CONFLICT (source_node_id, target_node_id, scan_id) DO NOTHING + """, + rows, + template="(%s, %s, %s, %s::uuid, %s::uuid, %s::uuid[], %s, %s, %s)", + ) + inserted = cur.rowcount + # cur.rowcount reflects actual inserts after ON CONFLICT DO NOTHING; + # fall back to attempted count if the driver reports -1. + return inserted if inserted >= 0 else len(rows) + + +def compute_attack_paths(scan_id: str, tenant_id: str, dsn: str) -> int: + """Run BFS from every finding-linked node and persist paths. Returns path count.""" + try: + conn = psycopg2.connect(dsn) + conn.autocommit = False + try: + if not _delete_stale_paths(conn, scan_id, tenant_id): + conn.commit() + return 0 + adj = _load_adjacency(conn, tenant_id) + source_nodes = _load_finding_nodes(conn, scan_id, tenant_id) + if not source_nodes: + logger.info("graph path traversal: no finding-linked nodes for scan %s", scan_id) + conn.commit() + return 0 + + total = 0 + for src in source_nodes: + paths = _bfs_from(src, adj) + total += _write_paths(conn, scan_id, tenant_id, src, paths) + + conn.commit() + logger.info("graph path traversal: %d paths written for scan %s", total, scan_id) + return total + except Exception: + conn.rollback() + raise + finally: + conn.close() + except Exception: + logger.error("graph path traversal failed for scan %s", scan_id, exc_info=True) + return 0 diff --git a/tests/test_attack_graph_api.py b/tests/test_attack_graph_api.py new file mode 100644 index 0000000..453b08a --- /dev/null +++ b/tests/test_attack_graph_api.py @@ -0,0 +1,139 @@ +"""Unit tests for the attack graph API routes.""" + +import time +from unittest.mock import MagicMock, patch + +import pytest + +_TENANT = "00000000-0000-0000-0000-000000000001" + + +@pytest.fixture() +def tenant_auth_headers(app): + """Admin JWT headers with X-Tenant-Id for shared-secret mode tests.""" + import jwt + + secret = app.config["JWT_SECRET"] + payload = { + "sub": "test-user", + "role": "admin", + "iat": int(time.time()), + "exp": int(time.time()) + 3600, + } + token = jwt.encode(payload, secret, algorithm="HS256") + return { + "Authorization": f"Bearer {token}", + "Content-Type": "application/json", + "X-Tenant-Id": _TENANT, + } + + +def _mock_conn(rows_sequence=None): + """Return a mock psycopg2 connection returning preset rows per cursor call.""" + rows_sequence = list(rows_sequence or []) + conn = MagicMock() + conn.autocommit = True + call_idx = [0] + + def make_cursor(*args, **kwargs): + cur = MagicMock() + cur.__enter__ = lambda s: s + cur.__exit__ = MagicMock(return_value=False) + idx = call_idx[0] + call_idx[0] += 1 + rows = rows_sequence[idx] if idx < len(rows_sequence) else [] + cur.fetchall.return_value = rows + cur.fetchone.return_value = rows[0] if rows else None + return cur + + conn.cursor.side_effect = make_cursor + return conn + + +def test_get_attack_graph_returns_empty_nodes_and_edges(client, tenant_auth_headers): + conn = _mock_conn([[]]) + with patch("api.routes.attack_graph.psycopg2.connect", return_value=conn): + with patch.dict("os.environ", {"DATABASE_URL": "postgresql://fake/db"}): + resp = client.get("/api/v1/attack-graph", headers=tenant_auth_headers) + assert resp.status_code == 200 + data = resp.get_json() + assert "nodes" in data + assert "edges" in data + + +def test_get_attack_graph_viewer_cannot_supply_tenant_header(client, app): + """Viewer-role tokens must not use X-Tenant-Id header (admin-only).""" + import jwt + + secret = app.config["JWT_SECRET"] + payload = { + "sub": "viewer-user", + "role": "viewer", + "iat": int(time.time()), + "exp": int(time.time()) + 3600, + } + token = jwt.encode(payload, secret, algorithm="HS256") + headers = { + "Authorization": f"Bearer {token}", + "X-Tenant-Id": _TENANT, + } + resp = client.get("/api/v1/attack-graph", headers=headers) + assert resp.status_code == 403 + + +def test_list_attack_paths_missing_scan_id_returns_400(client, tenant_auth_headers): + resp = client.get("/api/v1/attack-paths", headers=tenant_auth_headers) + assert resp.status_code == 400 + assert b"scan_id" in resp.data + + +def test_list_attack_paths_invalid_uuid_returns_400(client, tenant_auth_headers): + resp = client.get("/api/v1/attack-paths?scan_id=not-a-uuid", headers=tenant_auth_headers) + assert resp.status_code == 400 + + +def test_get_attack_path_invalid_uuid_returns_400(client, tenant_auth_headers): + resp = client.get("/api/v1/attack-paths/not-a-uuid", headers=tenant_auth_headers) + assert resp.status_code == 400 + + +def test_get_attack_path_not_found_returns_404(client, tenant_auth_headers): + valid_uuid = "00000000-0000-0000-0000-000000000002" + conn = _mock_conn([[]]) + with patch("api.routes.attack_graph.psycopg2.connect", return_value=conn): + with patch.dict("os.environ", {"DATABASE_URL": "postgresql://fake/db"}): + resp = client.get(f"/api/v1/attack-paths/{valid_uuid}", headers=tenant_auth_headers) + assert resp.status_code == 404 + + +@pytest.mark.parametrize("limit", ["abc", "", "1.5", "0", "-1"]) +@pytest.mark.parametrize( + "endpoint", ["/api/v1/attack-graph", "/api/v1/attack-paths?scan_id=00000000-0000-0000-0000-000000000002"] +) +def test_malformed_limit_returns_400(client, tenant_auth_headers, endpoint, limit): + separator = "&" if "?" in endpoint else "?" + response = client.get(f"{endpoint}{separator}limit={limit}", headers=tenant_auth_headers) + assert response.status_code == 400 + + +def test_verified_viewer_tenant_cannot_be_overridden_by_header_or_query(client, app): + import jwt + + token = jwt.encode( + {"sub": "viewer", "role": "viewer", "tid": _TENANT, "exp": int(time.time()) + 60}, + app.config["JWT_SECRET"], + algorithm="HS256", + ) + foreign = "00000000-0000-0000-0000-000000000099" + conn = MagicMock() + cursor = conn.cursor.return_value.__enter__.return_value + cursor.fetchall.return_value = [] + db = MagicMock() + db.conn = conn + with patch("api.routes.attack_graph._get_db", return_value=db): + response = client.get( + f"/api/v1/attack-graph?tenant_id={foreign}", + headers={"Authorization": f"Bearer {token}", "X-Tenant-Id": foreign}, + ) + assert response.status_code == 200 + assert cursor.execute.call_args[0][1]["tenant_id"] == _TENANT diff --git a/tests/test_graph_path_traversal.py b/tests/test_graph_path_traversal.py new file mode 100644 index 0000000..969cf89 --- /dev/null +++ b/tests/test_graph_path_traversal.py @@ -0,0 +1,107 @@ +"""Unit tests for BFS path traversal (scanner/graph/path_traversal.py).""" + +from unittest.mock import MagicMock, patch + +from scanner.graph.path_traversal import _bfs_from, _load_adjacency, compute_attack_paths + + +def test_bfs_from_returns_empty_for_isolated_node(): + adj = {} + result = _bfs_from("node-A", adj) + assert result == [] + + +def test_bfs_from_single_hop(): + adj = {"node-A": [("node-B", "PROTECTS", 1.0)]} + result = _bfs_from("node-A", adj) + assert len(result) == 1 + path = result[0] + assert path["target_node_id"] == "node-B" + assert path["path_length"] == 1 + assert path["min_confidence"] == 1.0 + assert path["relationship_types"] == ["PROTECTS"] + assert path["path_node_ids"] == ["node-A", "node-B"] + + +def test_bfs_from_two_hops(): + adj = { + "A": [("B", "PROTECTS", 1.0)], + "B": [("C", "EXPOSES", 0.8)], + } + result = _bfs_from("A", adj) + assert len(result) == 2 + two_hop = next(r for r in result if r["target_node_id"] == "C") + assert two_hop["path_length"] == 2 + assert two_hop["min_confidence"] == 0.8 + assert two_hop["relationship_types"] == ["PROTECTS", "EXPOSES"] + + +def test_bfs_from_does_not_revisit_nodes(): + adj = { + "A": [("B", "PROTECTS", 1.0), ("C", "EXPOSES", 1.0)], + "B": [("C", "MEMBER_OF", 1.0)], + } + result = _bfs_from("A", adj) + target_nodes = [r["target_node_id"] for r in result] + assert target_nodes.count("C") == 1 + + +def test_member_of_not_reversed_prevents_subnet_bridging(): + # Two VMs sharing a subnet. MEMBER_OF is forward-only (VM -> subnet), + # so BFS from vm-a must not reach vm-b through the subnet. + adj = { + "vm-a": [("subnet-1", "MEMBER_OF", 0.9)], + "vm-b": [("subnet-1", "MEMBER_OF", 0.9)], + # subnet-1 has no outgoing edges: MEMBER_OF is not reversed + } + result = _bfs_from("vm-a", adj) + target_ids = {r["target_node_id"] for r in result} + assert "vm-b" not in target_ids + assert "subnet-1" in target_ids # forward hop is still present + + +def test_load_adjacency_reverses_exposes_not_member_of(): + conn = MagicMock() + cur = MagicMock() + cur.__enter__ = lambda s: s + cur.__exit__ = MagicMock(return_value=False) + conn.cursor.return_value = cur + # One EXPOSES edge and one MEMBER_OF edge + cur.fetchall.return_value = [ + ("pub-ip", "vm-1", "EXPOSES", 1.0), + ("vm-1", "subnet-1", "MEMBER_OF", 0.9), + ] + + adj = _load_adjacency(conn, "tenant-x") + + # EXPOSES should have a reverse entry (vm-1 -> pub-ip via EXPOSES_REV) + rev_exposes = [n for n, rel, _ in adj.get("vm-1", []) if rel == "EXPOSES_REV"] + assert rev_exposes, "EXPOSES_REV edge missing from vm-1" + + # MEMBER_OF should NOT have a reverse entry (subnet-1 should not point to vm-1) + rev_member = [n for n, rel, _ in adj.get("subnet-1", []) if rel == "MEMBER_OF_REV"] + assert not rev_member, "MEMBER_OF_REV edge should not exist (spurious bridging)" + + +def test_compute_attack_paths_returns_zero_when_no_finding_nodes(): + with patch("scanner.graph.path_traversal.psycopg2.connect") as mock_connect: + conn = MagicMock() + mock_connect.return_value = conn + cur = MagicMock() + cur.__enter__ = lambda s: s + cur.__exit__ = MagicMock(return_value=False) + conn.cursor.return_value = cur + # adjacency returns empty; finding_nodes returns empty + cur.fetchall.return_value = [] + + result = compute_attack_paths("scan-1", "tenant-1", "postgresql://x") + + assert result == 0 + + +def test_compute_attack_paths_handles_db_error_gracefully(): + with patch("scanner.graph.path_traversal.psycopg2.connect") as mock_connect: + mock_connect.side_effect = Exception("DB unavailable") + result = compute_attack_paths("scan-1", "tenant-1", "postgresql://x") + + assert result == 0 diff --git a/tests/test_graph_paths_postgres.py b/tests/test_graph_paths_postgres.py new file mode 100644 index 0000000..7b3f288 --- /dev/null +++ b/tests/test_graph_paths_postgres.py @@ -0,0 +1,225 @@ +"""Database regressions for current traversal and successful clean-scan retention.""" + +import os +import uuid +from dataclasses import replace + +import psycopg2 +import pytest + +from scanner.graph.graph_populator import populate_graph +from scanner.graph.path_traversal import _delete_stale_paths, _load_adjacency, compute_attack_paths +from tests.test_graph_freshness_postgres import snapshot + +pytestmark = pytest.mark.skipif(not os.environ.get("DATABASE_URL"), reason="requires PostgreSQL") + + +@pytest.fixture +def path_scope(): + tenant, sub = str(uuid.uuid4()), str(uuid.uuid4()) + dsn = os.environ["DATABASE_URL"] + scans = [] + + def scan(subscription=sub): + sid = str(uuid.uuid4()) + scans.append(sid) + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute( + "INSERT INTO scans (scan_id, subscription_id, started_at, status) VALUES (%s,%s,now(),'completed')", + (sid, subscription), + ) + return sid + + yield tenant, sub, dsn, scan + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("DELETE FROM attack_paths WHERE tenant_id=%s", (tenant,)) + cur.execute("DELETE FROM graph_nodes WHERE tenant_id=%s", (tenant,)) + cur.execute("DELETE FROM graph_snapshot_scopes WHERE tenant_id=%s", (tenant,)) + cur.execute("DELETE FROM findings WHERE scan_id=ANY(%s::uuid[])", (scans,)) + cur.execute("DELETE FROM scans WHERE scan_id=ANY(%s::uuid[])", (scans,)) + + +def risky_scan(tenant, sub, dsn, sid): + snap = snapshot(tenant, sub) + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute( + "INSERT INTO findings (scan_id, rule_id, rule_name, severity, resource_id, detected_at, finding_key) " + "VALUES (%s,'AZ-TEST-001','test','HIGH',%s,now(),%s)", + (sid, snap.resources[1].resource_id, uuid.uuid4().hex), + ) + populate_graph(sid, snap, dsn) + return snap + + +def path_count(dsn, tenant, sid): + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("SELECT count(*) FROM attack_paths WHERE tenant_id=%s AND scan_id=%s", (tenant, sid)) + return cur.fetchone()[0] + + +def test_successful_clean_scan_clears_previous_paths_from_authoritative_scope(path_scope): + tenant, sub, dsn, scan = path_scope + risky = scan() + risky_scan(tenant, sub, dsn, risky) + # Seed one prior path so the cleanup regression is independent of traversal. + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("SELECT node_id::text FROM graph_nodes WHERE tenant_id=%s ORDER BY resource_id", (tenant,)) + nodes = [row[0] for row in cur.fetchall()] + cur.execute( + "INSERT INTO attack_paths " + "(path_id,tenant_id,scan_id,source_node_id,target_node_id," + "path_node_ids,path_length,relationship_types) " + "VALUES (%s,%s,%s,%s,%s,%s::uuid[],1,ARRAY['MEMBER_OF']) ON CONFLICT DO NOTHING", + (str(uuid.uuid4()), tenant, risky, nodes[0], nodes[1], nodes), + ) + clean = scan() + assert path_count(dsn, tenant, risky) == 1 + with psycopg2.connect(dsn) as conn: + _delete_stale_paths(conn, clean, tenant) + assert path_count(dsn, tenant, risky) == 0 + + +def test_partial_snapshot_does_not_traverse_retained_old_edge(path_scope): + from scanner.arg_inventory import InventoryStatus + + tenant, sub, dsn, scan = path_scope + first = scan() + risky_scan(tenant, sub, dsn, first) + with psycopg2.connect(dsn) as conn: + assert _load_adjacency(conn, tenant) + partial = snapshot(tenant, sub, status=InventoryStatus.PARTIAL, linked=False) + populate_graph(scan(), replace(partial, resources=partial.resources[:1]), dsn) + with psycopg2.connect(dsn) as conn: + assert _load_adjacency(conn, tenant) == {} + + +def test_risky_then_clean_scan_computes_zero_and_removes_old_paths(path_scope): + tenant, sub, dsn, scan = path_scope + first = scan() + risky_scan(tenant, sub, dsn, first) + assert path_count(dsn, tenant, first) == 1 + clean = scan() + populate_graph(clean, snapshot(tenant, sub, linked=False), dsn) + assert compute_attack_paths(clean, tenant, dsn) == 0 + assert path_count(dsn, tenant, first) == 0 + + +def test_clean_scan_retention_preserves_other_subscription_paths(path_scope): + tenant, sub, dsn, scan = path_scope + first = scan() + risky_scan(tenant, sub, dsn, first) + other_sub = str(uuid.uuid4()) + other = scan(other_sub) + risky_scan(tenant, other_sub, dsn, other) + clean = scan() + assert compute_attack_paths(clean, tenant, dsn) == 0 + assert path_count(dsn, tenant, first) == 0 + assert path_count(dsn, tenant, other) == 1 + + +def test_unsuccessful_scan_cannot_clear_previous_paths(path_scope): + tenant, sub, dsn, scan = path_scope + first = scan() + risky_scan(tenant, sub, dsn, first) + failed = scan() + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("UPDATE scans SET status='failed' WHERE scan_id=%s", (failed,)) + _delete_stale_paths(conn, failed, tenant) + assert path_count(dsn, tenant, first) == 1 + + +def test_cleanup_keeps_other_tenants_paths(path_scope): + tenant, sub, dsn, scan = path_scope + other_tenant = str(uuid.uuid4()) + first, other = scan(), scan() + try: + risky_scan(tenant, sub, dsn, first) + risky_scan(other_tenant, sub, dsn, other) + clean = scan() + compute_attack_paths(clean, tenant, dsn) + assert path_count(dsn, tenant, first) == 0 + assert path_count(dsn, other_tenant, other) == 1 + finally: + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("DELETE FROM graph_nodes WHERE tenant_id=%s", (other_tenant,)) + cur.execute("DELETE FROM graph_snapshot_scopes WHERE tenant_id=%s", (other_tenant,)) + + +def test_authenticated_graph_api_selects_only_current_verified_tenant_evidence(path_scope, client, app): + import time + import jwt + from scanner.arg_inventory import InventoryStatus + + tenant, sub, dsn, scan = path_scope + first = scan() + risky_scan(tenant, sub, dsn, first) + partial = snapshot(tenant, sub, status=InventoryStatus.PARTIAL, linked=False) + populate_graph(scan(), replace(partial, resources=partial.resources[:1]), dsn) + token = jwt.encode( + {"sub": "graph-viewer", "role": "viewer", "tid": tenant, "exp": int(time.time()) + 60}, + app.config["JWT_SECRET"], + algorithm="HS256", + ) + foreign = str(uuid.uuid4()) + response = client.get( + f"/api/v1/attack-graph?subscription_id={sub}&tenant_id={foreign}", + headers={"Authorization": f"Bearer {token}", "X-Tenant-Id": foreign}, + ) + assert response.status_code == 200 + result = response.get_json() + assert len(result["nodes"]) == 1 + assert result["nodes"][0]["resource_id"] == partial.resources[0].resource_id.lower() + assert result["edges"] == [] + + +def test_delayed_older_traversal_cannot_delete_newer_paths(path_scope): + tenant, sub, dsn, scan = path_scope + older, newer = scan(), scan() + with psycopg2.connect(dsn) as conn: + with conn.cursor() as cur: + cur.execute("UPDATE scans SET started_at=now()-interval '2 minutes' WHERE scan_id=%s", (older,)) + cur.execute("UPDATE scans SET started_at=now()-interval '1 minute' WHERE scan_id=%s", (newer,)) + risky_scan(tenant, sub, dsn, newer) + assert path_count(dsn, tenant, newer) == 1 + assert compute_attack_paths(older, tenant, dsn) == 0 + assert path_count(dsn, tenant, newer) == 1 + + +def test_traversal_waits_for_scope_publication_lock(path_scope): + from concurrent.futures import ThreadPoolExecutor + from threading import Event + + tenant, sub, dsn, scan = path_scope + sid = scan() + risky_scan(tenant, sub, dsn, sid) + entered = Event() + 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 traverse(): + entered.set() + return compute_attack_paths(sid, tenant, dsn) + + with ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(traverse) + assert entered.wait(2) + try: + future.result(timeout=0.3) + pytest.fail("traversal completed while its scope lock was held") + except TimeoutError: + pass + finally: + blocker.rollback() + assert future.result(timeout=5) == 0 + assert path_count(dsn, tenant, sid) == 1 + finally: + blocker.close()