"""Locate a specific resource block in a directory's .tf files.""" from __future__ import annotations from dataclasses import dataclass from pathlib import Path import re from scripts.hcl_diff import find_resource_blocks _KEY_ATTRIBUTES = { "kms_key_id", "kms_key_arn", "bucket_key_enabled", "publicly_accessible", "acl", "versioning", "tags", "deletion_protection", } _VAR_REF_RE = re.compile(r"\bvar\.([A-Za-z0-9_]+)") _LOCAL_REF_RE = re.compile(r"\blocal\.([A-Za-z0-9_]+)") _DATA_REF_RE = re.compile(r"\bdata\.aws_iam_policy_document\.([A-Za-z0-9_]+)") _SG_REF_TEMPLATE = 'security_group_id = aws_security_group.{name}.id' @dataclass(frozen=True) class BlockLocation: file: Path start_line: int end_line: int text: str header: str evidence_line: str key_attributes: dict[str, str | bool | int | list[str]] review_context: dict[str, object] def _parse_value(raw: str) -> str | bool | int | list[str]: value = raw.strip().rstrip(",") if value.lower() in {"true", "false"}: return value.lower() == "true" if re.fullmatch(r"-?\d+", value): return int(value) if value.startswith('"') and value.endswith('"'): return value[1:-1] if value.startswith("[") and value.endswith("]"): inner = value[1:-1].strip() if not inner: return [] parts = [part.strip().strip('"') for part in inner.split(",")] return [part for part in parts if part] return value def _extract_key_attributes(lines: list[str]) -> dict[str, str | bool | int | list[str]]: out: dict[str, str | bool | int | list[str]] = {} for line in lines: stripped = line.strip() if "=" not in stripped or stripped.startswith("#"): continue key, _, raw_value = stripped.partition("=") key = key.strip() if key not in _KEY_ATTRIBUTES: continue out[key] = _parse_value(raw_value) return out def _find_named_block(lines: list[str], header_re: re.Pattern[str], name: str) -> list[str]: i = 0 while i < len(lines): line = lines[i] match = header_re.match(line) if not match or match.group("name") != name: i += 1 continue depth = line.count("{") - line.count("}") j = i while depth > 0 and j + 1 < len(lines): j += 1 depth += lines[j].count("{") - lines[j].count("}") return lines[i:j + 1] return [] def _collect_variable_defaults(root: Path, names: set[str]) -> dict[str, str | bool | int | list[str]]: out: dict[str, str | bool | int | list[str]] = {} header_re = re.compile(r'^\s*variable\s+"(?P[^"]+)"\s*\{') for tf in sorted(root.glob("*.tf")): lines = tf.read_text().splitlines() for name in names: if name in out: continue block = _find_named_block(lines, header_re, name) for line in block[1:]: stripped = line.strip() if stripped.startswith("default"): _, _, raw = stripped.partition("=") out[name] = _parse_value(raw) break return out def _collect_locals(root: Path, names: set[str]) -> dict[str, str | bool | int | list[str]]: out: dict[str, str | bool | int | list[str]] = {} header_re = re.compile(r"^\s*locals\s*\{") for tf in sorted(root.glob("*.tf")): lines = tf.read_text().splitlines() i = 0 while i < len(lines): line = lines[i] if not header_re.match(line): i += 1 continue depth = line.count("{") - line.count("}") j = i while depth > 0 and j + 1 < len(lines): j += 1 depth += lines[j].count("{") - lines[j].count("}") for block_line in lines[i + 1:j]: stripped = block_line.strip() if "=" not in stripped: continue key, _, raw = stripped.partition("=") key = key.strip() if key in names and key not in out: out[key] = _parse_value(raw) i = j + 1 return out def _collect_related_policy_docs(root: Path, names: set[str]) -> list[str]: out: list[str] = [] header_re = re.compile( r'^\s*data\s+"aws_iam_policy_document"\s+"(?P[^"]+)"\s*\{' ) for tf in sorted(root.glob("*.tf")): lines = tf.read_text().splitlines() for name in names: block = _find_named_block(lines, header_re, name) if block: out.append(block[0].strip()) return out def _collect_related_sg_rules(root: Path, name: str) -> list[str]: out: list[str] = [] target = _SG_REF_TEMPLATE.format(name=name) header_re = re.compile( r'^\s*resource\s+"aws_security_group_rule"\s+"(?P[^"]+)"\s*\{' ) for tf in sorted(root.glob("*.tf")): lines = tf.read_text().splitlines() i = 0 while i < len(lines): line = lines[i] if not header_re.match(line): i += 1 continue depth = line.count("{") - line.count("}") j = i while depth > 0 and j + 1 < len(lines): j += 1 depth += lines[j].count("{") - lines[j].count("}") block = lines[i:j + 1] if any(target in block_line for block_line in block[1:]): out.append(block[0].strip()) i = j + 1 return out def _build_review_context(root: Path, rtype: str, rname: str, block_lines: list[str]) -> dict[str, object]: joined = "\n".join(block_lines) variable_names = set(_VAR_REF_RE.findall(joined)) local_names = set(_LOCAL_REF_RE.findall(joined)) policy_doc_names = set(_DATA_REF_RE.findall(joined)) related_blocks: list[str] = [] related_blocks.extend(_collect_related_policy_docs(root, policy_doc_names)) if rtype == "aws_security_group": related_blocks.extend(_collect_related_sg_rules(root, rname)) return { "variables": _collect_variable_defaults(root, variable_names), "locals": _collect_locals(root, local_names), "related_blocks": related_blocks, } def find_block(source_dir: Path | str, local_address: str) -> BlockLocation | None: """Search .tf files in `source_dir` for a resource matching `local_address`. `local_address` is `.` (e.g. `aws_iam_role.svc`). Returns the first match found. Returns None if no match. """ if "." not in local_address: return None rtype, _, rname = local_address.partition(".") root = Path(source_dir) for tf in sorted(root.glob("*.tf")): for blk in find_resource_blocks(tf): if blk.type == rtype and blk.name == rname: lines = tf.read_text().splitlines() block_lines = lines[blk.start - 1: blk.end] header = block_lines[0] if block_lines else "" evidence_line = header for line in block_lines[1:]: if line.strip(): evidence_line = line.strip() break return BlockLocation( file=tf, start_line=blk.start, end_line=blk.end, text="\n".join(block_lines), header=header, evidence_line=evidence_line, key_attributes=_extract_key_attributes(block_lines[1:]), review_context=_build_review_context(root, rtype, rname, block_lines), ) return None