shoc-frontend-new/scripts/check-terraform-release-plan.py

522 lines
17 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Reject HCP Terraform plans that are not a frontend content-release update.
Accepts exactly:
- an update of the release pointer (content, plus computed etag/version_id)
- an update of the distribution with only origin[*].origin_path changed
- exactly one action invocation for the CloudFront invalidation
after origin_path values must match the expected labels. before origin_path
values must match the pointer's prior current/previous. This script may read a
local plan JSON file or download plan JSON from the documented HashiCorp
endpoint:
GET https://app.terraform.io/api/v2/plans/:id/json-output
The download follows exactly one redirect, and only to archivist.terraform.io.
It does not create, apply, discard, or poll runs.
"""
from __future__ import annotations
import argparse
import json
import os
import re
import ssl
import sys
import urllib.error
import urllib.request
from pathlib import Path
from typing import Any, Callable
from urllib.parse import urlparse
POINTER_ADDRESS = "module.environment_owned.aws_s3_object.release_pointer"
DISTRIBUTION_ADDRESS = "module.environment_owned.aws_cloudfront_distribution.site"
ACTION_ADDRESS = (
"module.environment_owned.action.aws_cloudfront_create_invalidation.release"
)
API_HOST = "app.terraform.io"
ARCHIVE_HOST = "archivist.terraform.io"
PLAN_ID_RE = re.compile(r"^plan-[A-Za-z0-9]+$")
VERSION_LABEL_RE = re.compile(r"^[0-9a-f]{40}-[0-9]+-[0-9]+$")
IGNORED_ACTIONS = {"no-op", "read"}
UNSAFE_ACTIONS = {"create", "delete"}
POINTER_UNKNOWN_ATTRIBUTES = frozenset({"etag", "version_id"})
DISTRIBUTION_UNKNOWN_ATTRIBUTES = frozenset(
{
"etag",
"last_modified_time",
"status",
"in_progress_validation_batches",
}
)
REDIRECT_STATUSES = {301, 302, 303, 307, 308}
UrlOpen = Callable[..., Any]
class _NoRedirectHandler(urllib.request.HTTPRedirectHandler):
"""Return the redirect response instead of following it."""
def http_error_301(self, req, fp, code, msg, headers):
return self._capture(req, fp, code, headers)
http_error_302 = http_error_303 = http_error_307 = http_error_308 = http_error_301
@staticmethod
def _capture(req, fp, code, headers):
response = urllib.response.addinfourl(fp, headers, req.full_url, code=code)
response.msg = "Redirect"
return response
def _urlopen_without_redirects(
*handlers: urllib.request.BaseHandler,
) -> UrlOpen:
context = ssl.create_default_context()
opener = urllib.request.build_opener(
urllib.request.HTTPSHandler(context=context),
_NoRedirectHandler,
*handlers,
)
return opener.open
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
source = parser.add_mutually_exclusive_group(required=True)
source.add_argument(
"plan_json",
type=Path,
nargs="?",
help="Local Terraform plan JSON. Mutually exclusive with --plan-id.",
)
source.add_argument(
"--plan-id",
help="HCP Terraform plan ID. Downloads JSON from app.terraform.io.",
)
parser.add_argument(
"--expected-version-label",
required=True,
help="Immutable current release the plan must apply. Empty string is the legacy root.",
)
parser.add_argument(
"--expected-previous-version-label",
default="",
help="Previous release label the origin group must fail over to.",
)
parser.add_argument(
"--evidence-out",
type=Path,
help="Write machine-readable proof after every assertion passes.",
)
return parser.parse_args()
def download_plan_json(
plan_id: str,
token: str,
*,
urlopen: UrlOpen | None = None,
handlers: tuple[urllib.request.BaseHandler, ...] = (),
) -> dict[str, Any]:
if not PLAN_ID_RE.fullmatch(plan_id):
raise ValueError(f"plan id {plan_id!r} is not a valid HCP plan id")
if not token:
raise ValueError("TF_API_TOKEN is required to download plan JSON")
opener = urlopen or _urlopen_without_redirects(*handlers)
api_url = f"https://{API_HOST}/api/v2/plans/{plan_id}/json-output"
request = urllib.request.Request(
api_url,
method="GET",
headers={
"Authorization": f"Bearer {token}",
"Content-Type": "application/vnd.api+json",
"Accept": "application/json",
},
)
first = _open_pinned(opener, request, allowed_host=API_HOST)
try:
if first.status == 204:
raise ValueError(
"plan JSON is not ready; refusing to poll the plans endpoint"
)
if first.status not in REDIRECT_STATUSES:
raise ValueError(
f"expected a redirect from {API_HOST}, got HTTP {first.status}"
)
location = first.headers.get("Location")
if not location:
raise ValueError(f"{API_HOST} redirect is missing a Location header")
archive = urlparse(location)
if archive.scheme != "https" or archive.hostname != ARCHIVE_HOST:
raise ValueError(
"refusing redirect that is not https://"
f"{ARCHIVE_HOST}/"
)
archive_request = urllib.request.Request(location, method="GET")
second = _open_pinned(opener, archive_request, allowed_host=ARCHIVE_HOST)
try:
if second.status in REDIRECT_STATUSES:
raise ValueError(
f"refusing a second redirect from {ARCHIVE_HOST}"
)
if second.status != 200:
raise ValueError(
f"plan JSON download from {ARCHIVE_HOST} returned "
f"HTTP {second.status}"
)
payload = second.read()
finally:
second.close()
finally:
first.close()
plan = json.loads(payload.decode("utf-8"))
if not isinstance(plan, dict):
raise ValueError("plan JSON must be an object")
return plan
def _open_pinned(urlopen: UrlOpen, request: urllib.request.Request, *, allowed_host: str):
parsed = urlparse(request.full_url)
if parsed.scheme != "https" or parsed.hostname != allowed_host:
raise ValueError(
f"refusing to contact {parsed.scheme}://{parsed.hostname} "
f"(pinned host is {allowed_host})"
)
context = ssl.create_default_context()
try:
return urlopen(request, context=context, timeout=30)
except TypeError:
return urlopen(request, timeout=30)
def _is_nested_unknown(value: Any) -> bool:
if isinstance(value, dict):
return any(item is True or _is_nested_unknown(item) for item in value.values())
if isinstance(value, list):
return any(item is True or _is_nested_unknown(item) for item in value)
return False
def changed_attributes(
change: dict[str, Any],
*,
computed_unknown: frozenset[str],
) -> set[str]:
before = change.get("before") or {}
after = change.get("after") or {}
unknown = change.get("after_unknown") or {}
keys = set(before) | set(after) | set(unknown)
changed: set[str] = set()
for key in keys:
unknown_value = unknown.get(key)
if unknown_value is True:
if key in computed_unknown:
continue
changed.add(key)
continue
if _is_nested_unknown(unknown_value):
changed.add(key)
continue
if before.get(key) != after.get(key):
changed.add(key)
return changed
def _label_ok(label: str) -> bool:
return label == "" or bool(VERSION_LABEL_RE.fullmatch(label))
def origin_path_for_label(label: str) -> str:
return "" if label == "" else f"/releases/{label}"
def _origin_map(origins: Any) -> dict[str, dict[str, Any]]:
if not isinstance(origins, list):
return {}
mapped: dict[str, dict[str, Any]] = {}
for origin in origins:
if not isinstance(origin, dict):
continue
origin_id = origin.get("origin_id")
if not isinstance(origin_id, str) or not origin_id:
continue
mapped[origin_id] = origin
return mapped
def _origin_paths(origins: Any) -> dict[str, str]:
return {
origin_id: origin.get("origin_path") or ""
for origin_id, origin in _origin_map(origins).items()
}
def _decode_pointer(content: Any) -> dict[str, str]:
if not isinstance(content, str) or not content:
return {}
try:
payload = json.loads(content)
except json.JSONDecodeError:
return {}
if not isinstance(payload, dict):
return {}
return {
"current": payload.get("current") or "",
"previous": payload.get("previous") or "",
}
def _validate_pointer(
resource: dict[str, Any],
expected_current: str,
expected_previous: str,
) -> list[str]:
violations: list[str] = []
change = resource.get("change") or {}
changed = changed_attributes(change, computed_unknown=POINTER_UNKNOWN_ATTRIBUTES)
if changed != {"content"}:
violations.append(
f"{POINTER_ADDRESS}: expected only content to change, found "
f"{sorted(changed) if changed else 'no attribute changes'}"
)
after = _decode_pointer((change.get("after") or {}).get("content"))
if after.get("current") != expected_current:
violations.append(
f"{POINTER_ADDRESS}: after current {after.get('current')!r} does not match "
f"{expected_current!r}"
)
if after.get("previous") != expected_previous:
violations.append(
f"{POINTER_ADDRESS}: after previous {after.get('previous')!r} does not match "
f"{expected_previous!r}"
)
unknown = change.get("after_unknown") or {}
if unknown.get("content") is True:
violations.append(f"{POINTER_ADDRESS}: content after value is unknown")
return violations
def _origin_non_path_fields_changed(before: dict[str, Any], after: dict[str, Any]) -> bool:
before_rest = {key: value for key, value in before.items() if key != "origin_path"}
after_rest = {key: value for key, value in after.items() if key != "origin_path"}
return before_rest != after_rest
def _validate_distribution(
resource: dict[str, Any],
pointer_before: dict[str, str],
expected_current: str,
expected_previous: str,
) -> list[str]:
violations: list[str] = []
change = resource.get("change") or {}
changed = changed_attributes(
change, computed_unknown=DISTRIBUTION_UNKNOWN_ATTRIBUTES
)
if changed != {"origin"}:
violations.append(
f"{DISTRIBUTION_ADDRESS}: expected only origin to change, found "
f"{sorted(changed) if changed else 'no attribute changes'}"
)
return violations
before_origins = _origin_map((change.get("before") or {}).get("origin"))
after_origins = _origin_map((change.get("after") or {}).get("origin"))
if set(before_origins) != set(after_origins):
violations.append(
f"{DISTRIBUTION_ADDRESS}: origin IDs changed "
f"from {sorted(before_origins)} to {sorted(after_origins)}"
)
return violations
for origin_id, before_origin in before_origins.items():
if _origin_non_path_fields_changed(before_origin, after_origins[origin_id]):
violations.append(
f"{DISTRIBUTION_ADDRESS}: origin {origin_id!r} changed a field other than origin_path"
)
after_paths = sorted(_origin_paths((change.get("after") or {}).get("origin")).values())
expected_after = sorted(
[
origin_path_for_label(expected_current),
origin_path_for_label(expected_previous),
]
)
if after_paths != expected_after:
violations.append(
f"{DISTRIBUTION_ADDRESS}: after origin_path {after_paths} does not match "
f"{expected_after}"
)
before_paths = sorted(_origin_paths((change.get("before") or {}).get("origin")).values())
expected_before = sorted(
[
origin_path_for_label(pointer_before.get("current", "")),
origin_path_for_label(pointer_before.get("previous", "")),
]
)
if before_paths != expected_before:
violations.append(
f"{DISTRIBUTION_ADDRESS}: before origin_path {before_paths} does not match "
f"pointer prior values {expected_before}"
)
return violations
def _validate_actions(plan: dict[str, Any]) -> list[str]:
invocations = plan.get("action_invocations")
if invocations is None:
return ["plan is missing action_invocations"]
if not isinstance(invocations, list):
return ["action_invocations must be a list"]
addresses = [
item.get("address")
for item in invocations
if isinstance(item, dict)
]
if addresses != [ACTION_ADDRESS]:
return [
"expected exactly one action_invocations entry "
f"{ACTION_ADDRESS}, found {addresses}"
]
return []
def validate_plan(
plan: dict[str, Any],
expected_current: str,
expected_previous: str,
) -> list[str]:
violations: list[str] = []
if not _label_ok(expected_current):
violations.append(
"expected version label must be empty or <full-sha>-<run-id>-<attempt>"
)
return violations
if not _label_ok(expected_previous):
violations.append(
"expected previous version label must be empty or <full-sha>-<run-id>-<attempt>"
)
return violations
updates: dict[str, dict[str, Any]] = {}
for resource in plan.get("resource_changes", []):
if resource.get("mode", "managed") != "managed":
continue
address = resource.get("address", "<unknown>")
change = resource.get("change") or {}
actions = list(change.get("actions") or [])
action_set = set(actions)
if action_set <= IGNORED_ACTIONS:
continue
if change.get("importing"):
violations.append(f"{address}: import actions are not allowed")
unsafe = sorted(action_set & UNSAFE_ACTIONS)
if unsafe:
violations.append(f"{address}: unsafe actions {unsafe}")
if "replace" in action_set or actions in (
["delete", "create"],
["create", "delete"],
):
violations.append(f"{address}: replacement is not allowed")
if "update" in action_set:
updates[address] = resource
if action_set != {"update"}:
violations.append(
f"{address}: update must be the only action, got {actions}"
)
if address not in {POINTER_ADDRESS, DISTRIBUTION_ADDRESS} and (
action_set - IGNORED_ACTIONS
):
violations.append(
f"{address}: managed address is outside the content-release update"
)
if set(updates) != {POINTER_ADDRESS, DISTRIBUTION_ADDRESS}:
violations.append(
"expected exactly the pointer and distribution updates, found "
f"{sorted(updates)}"
)
violations.extend(_validate_actions(plan))
return violations
pointer_change = updates[POINTER_ADDRESS].get("change") or {}
pointer_before = _decode_pointer((pointer_change.get("before") or {}).get("content"))
violations.extend(
_validate_pointer(updates[POINTER_ADDRESS], expected_current, expected_previous)
)
violations.extend(
_validate_distribution(
updates[DISTRIBUTION_ADDRESS],
pointer_before,
expected_current,
expected_previous,
)
)
violations.extend(_validate_actions(plan))
return violations
def main() -> int:
args = parse_args()
if args.plan_id:
try:
plan = download_plan_json(args.plan_id, os.environ.get("TF_API_TOKEN", ""))
except (OSError, ValueError, json.JSONDecodeError, urllib.error.URLError) as exc:
print(f"FAIL: could not download plan JSON: {exc}", file=sys.stderr)
return 1
else:
if args.plan_json is None:
print("FAIL: plan JSON path or --plan-id is required", file=sys.stderr)
return 1
plan = json.loads(args.plan_json.read_text(encoding="utf-8"))
violations = validate_plan(
plan,
args.expected_version_label,
args.expected_previous_version_label,
)
if violations:
print("FAIL: Terraform plan is not a content-release update", file=sys.stderr)
for violation in violations:
print(f" - {violation}", file=sys.stderr)
return 1
if args.evidence_out:
evidence = {
"pointer_address": POINTER_ADDRESS,
"distribution_address": DISTRIBUTION_ADDRESS,
"action_address": ACTION_ADDRESS,
"expected_version_label": args.expected_version_label,
"expected_previous_version_label": args.expected_previous_version_label,
"managed_updates": 2,
"action_invocations": 1,
"creates": 0,
"deletes": 0,
"replacements": 0,
}
args.evidence_out.write_text(
json.dumps(evidence, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
print(
"PASS: content-release plan updates "
f"{POINTER_ADDRESS} and {DISTRIBUTION_ADDRESS} to "
f"{args.expected_version_label} (previous {args.expected_previous_version_label!r})"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())