Skip to content
17 changes: 15 additions & 2 deletions backend/engines/drift.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,10 +114,19 @@ def parse(self, state_content: dict) -> dict[str, dict]:
continue
rtype = resource["type"]
rname = resource["name"]
for instance in resource.get("instances", []):
instances = resource.get("instances", [])
for instance in instances:
attrs = instance.get("attributes", {})
rid = attrs.get("id") or f"{rtype}.{rname}"
key = f"{rtype}.{rname}"
base_key = f"{rtype}.{rname}"
index_key = instance.get("index_key")
if index_key is not None:
suffix = f"[{index_key}]" if isinstance(index_key, int) else f"[{json.dumps(index_key)}]"
key = f"{base_key}{suffix}"
elif len(instances) > 1:
raise ValueError(f"Multiple unindexed Terraform instances: {base_key}")
else:
key = base_key
resources[key] = {"type": rtype, "name": rname, "id": rid, "attributes": attrs}
return resources

Expand Down Expand Up @@ -148,6 +157,7 @@ def collect_ec2_instances(self) -> dict[str, dict]:
}
except Exception as e:
log.error("Failed to collect EC2 instances", error=str(e))
raise RuntimeError("AWS EC2 instances inventory incomplete") from e
return resources

def collect_s3_buckets(self) -> dict[str, dict]:
Expand Down Expand Up @@ -180,6 +190,7 @@ def collect_s3_buckets(self) -> dict[str, dict]:
resources[name] = data
except Exception as e:
log.error("Failed to collect S3 buckets", error=str(e))
raise RuntimeError("AWS S3 buckets inventory incomplete") from e
return resources

def collect_security_groups(self) -> dict[str, dict]:
Expand All @@ -197,6 +208,7 @@ def collect_security_groups(self) -> dict[str, dict]:
}
except Exception as e:
log.error("Failed to collect security groups", error=str(e))
raise RuntimeError("AWS security groups inventory incomplete") from e
return resources

def collect_rds_instances(self) -> dict[str, dict]:
Expand All @@ -215,6 +227,7 @@ def collect_rds_instances(self) -> dict[str, dict]:
}
except Exception as e:
log.error("Failed to collect RDS instances", error=str(e))
raise RuntimeError("AWS RDS instances inventory incomplete") from e
return resources

def collect_all(self) -> dict[str, dict[str, dict]]:
Expand Down
43 changes: 43 additions & 0 deletions backend/tests/test_drift_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -298,3 +298,46 @@ def test_fingerprint_differs_for_different_drift(analyzer, parser, clean_tfstate
f1 = analyzer.analyze(tf_resources, live_a, "us-east-1")[0]
f2 = analyzer.analyze(tf_resources, live_b, "us-east-1")[0]
assert f1.fingerprint != f2.fingerprint


# A failed AWS enumeration must not be converted into a false deletion.
@pytest.mark.parametrize(
("collector_method", "service", "api_method"),
[
("collect_ec2_instances", "ec2", "get_paginator"),
("collect_security_groups", "ec2", "get_paginator"),
("collect_rds_instances", "rds", "get_paginator"),
("collect_s3_buckets", "s3", "list_buckets"),
],
)
def test_aws_inventory_enumeration_failure_is_fatal(
monkeypatch, collector_method, service, api_method
):
from unittest.mock import Mock

from backend.engines.drift import AWSStateCollector

client = Mock()
getattr(client, api_method).side_effect = RuntimeError("injected AWS failure")
session = Mock()
session.client.return_value = client
collector = AWSStateCollector(session, "us-east-1")

with pytest.raises(RuntimeError, match="inventory incomplete"):
getattr(collector, collector_method)()


def test_terraform_parser_preserves_multiple_instances(parser):
state = {
"resources": [{
"mode": "managed",
"type": "aws_instance",
"name": "web",
"instances": [
{"index_key": 0, "attributes": {"id": "i-one"}},
{"index_key": 1, "attributes": {"id": "i-two"}},
],
}],
}
result = parser.parse(state)
assert {item["id"] for item in result.values()} == {"i-one", "i-two"}
4 changes: 2 additions & 2 deletions cli/driftguard_cli/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@

import json
import os
from dataclasses import dataclass, asdict
from dataclasses import asdict, dataclass
from pathlib import Path

CONFIG_DIR = Path.home() / ".driftguard"
Expand All @@ -24,7 +24,7 @@ class Config:
api_key: str | None = None

@classmethod
def load(cls) -> "Config":
def load(cls) -> Config:
cfg = cls()
if CONFIG_FILE.exists():
try:
Expand Down
2 changes: 1 addition & 1 deletion cli/driftguard_cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,8 @@

import typer
from rich.console import Console
from rich.table import Table
from rich.panel import Panel
from rich.table import Table

from .client import DriftGuardAPIError, DriftGuardClient
from .config import Config
Expand Down
1 change: 0 additions & 1 deletion cli/tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,6 @@

import httpx
import pytest

from driftguard_cli.client import DriftGuardAPIError, DriftGuardClient


Expand Down
Loading