from __future__ import annotations import importlib.util import json import os import subprocess import stat import sys import tempfile import unittest from contextlib import redirect_stderr, redirect_stdout from io import StringIO from datetime import UTC, datetime from pathlib import Path from unittest import mock ROOT = Path(__file__).resolve().parents[1] SCRIPTS = ROOT / "scripts" TESTS = ROOT / "tests" for path in (SCRIPTS, TESTS): if str(path) not in sys.path: sys.path.insert(0, str(path)) from test_custody_contract import broker_receipt, projection_contract, projection_receipt SPEC = importlib.util.spec_from_file_location( "custody_projection", SCRIPTS / "custody-projection.py" ) assert SPEC and SPEC.loader module = importlib.util.module_from_spec(SPEC) SPEC.loader.exec_module(module) NOW = datetime(2026, 8, 22, 22, 1, tzinfo=UTC) class CustodyProjectionTests(unittest.TestCase): def test_policy_and_manifest_are_derived_from_one_contract(self) -> None: contract = projection_contract() policy = module.build_policy(contract) manifest = json.loads(module.build_manifest(contract, "custody:" + "9" * 32)) rendered = json.dumps(manifest, sort_keys=True) self.assertEqual(4, policy.count('capabilities = ["read"]')) self.assertIn(contract["engagement_id"], policy) self.assertIn(contract["engagement_id"], rendered) self.assertNotIn("WH-ENG-20260822-AUDIT-E2-01", policy + rendered) self.assertEqual( ["token-a", "token-b"], [item["secretKey"] for item in manifest["items"][1]["spec"]["data"]], ) def test_identity_overlay_is_exact_and_removable(self) -> None: contract = projection_contract() existing = [{"name": "production", "tokens": ["not-inspected"]}] updated = module.add_temporary_identities( existing, contract, {"token-a": "a", "token-b": "b"} ) self.assertEqual(3, len(updated)) temporary = updated[1:] self.assertTrue(all(item["expires_at"] == contract["window"]["expires_at"] for item in temporary)) restored, removed = module.remove_temporary_identities(updated, contract) self.assertEqual(existing, restored) self.assertEqual(sorted(module.identity_names(contract)), removed) def test_broker_gate_runs_before_operator_or_token_generation(self) -> None: contract = projection_contract() operator = object() with mock.patch.object( module, "current_broker_receipt", side_effect=module.ContractError("not connected"), ), mock.patch.object(module.secrets, "token_urlsafe") as token_urlsafe: with self.assertRaises(module.ContractError): module.project(operator, contract, now=NOW) token_urlsafe.assert_not_called() def test_cli_broker_gate_runs_before_platform_authority_is_opened(self) -> None: contract = projection_contract() with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "contract.json" path.write_text(json.dumps(contract), encoding="utf-8") argv = [ "custody-projection.py", "project", "--contract", str(path), "--confirm", f"{contract['engagement_id']}:attended", ] with mock.patch.object(sys, "argv", argv), mock.patch.object( module, "current_broker_receipt", side_effect=module.ContractError("broker missing"), ), mock.patch.object(module, "Operator") as operator, redirect_stderr(StringIO()): self.assertEqual(1, module.main()) operator.assert_not_called() def test_validate_and_render_cli_require_no_authority(self) -> None: contract = projection_contract() with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "contract.json" path.write_text(json.dumps(contract), encoding="utf-8") for command in ("validate", "render"): with mock.patch.object( sys, "argv", ["custody-projection.py", command, "--contract", str(path)] ), mock.patch.object(module, "Operator") as operator, redirect_stdout(StringIO()): self.assertEqual(0, module.main()) operator.assert_not_called() def test_transaction_rolls_back_after_every_mutation_boundary(self) -> None: for failed_index in range(4): events: list[str] = [] steps = [] for index in range(4): def step(index=index) -> None: events.append(f"step-{index}") if index == failed_index: raise module.ProcedureError(f"fail-{index}") steps.append(step) with self.assertRaisesRegex(module.ProcedureError, f"fail-{failed_index}"): module.transactional(steps, lambda: events.append("rollback")) self.assertEqual("rollback", events[-1]) self.assertEqual(1, events.count("rollback")) self.assertNotIn(f"step-{failed_index + 1}", events) def test_cleanup_failure_is_reported_without_original_output(self) -> None: def fail() -> None: raise module.ProcedureError("projection-stage") def cleanup_fail() -> None: raise module.ProcedureError("provider-secret-output") with self.assertRaisesRegex( module.ProcedureError, "cleanup could not be proven" ) as caught: module.transactional([fail], cleanup_fail) self.assertNotIn("provider-secret-output", str(caught.exception)) def test_receipt_write_is_mode_0600_and_canonical_receipt_validates(self) -> None: contract = projection_contract() receipt = projection_receipt(contract) with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "projection.json" module.write_receipt(path, receipt) self.assertEqual(0o600, stat.S_IMODE(path.stat().st_mode)) self.assertEqual(receipt, json.loads(path.read_text(encoding="utf-8"))) self.assertIs(receipt, module.validate_projection_receipt(receipt, contract)) def test_status_cannot_infer_absence_after_connectivity_failure(self) -> None: class Disconnected: def registry(self): raise module.ProcedureError("registry connectivity failed") with self.assertRaisesRegex(module.ProcedureError, "connectivity"): module.projection_status( Disconnected(), projection_contract(), projection_receipt(projection_contract()) ) def test_cleanup_refuses_recreated_resource_uid(self) -> None: contract = projection_contract() receipt = projection_receipt(contract) class RecreatedResource: def kubectl(self, args, *, label, **kwargs): if "clustersecretstore" in args: value = "replacement-store-uid" else: value = "" return subprocess.CompletedProcess(args, 0, stdout=value, stderr="") with self.assertRaisesRegex(module.ProcedureError, "refusing deletion"): module.verify_receipt_resource_scope(RecreatedResource(), contract, receipt) def test_expiry_state_is_fail_closed(self) -> None: contract = projection_contract() self.assertEqual("before-window", module.time_state(contract, datetime(2026, 8, 22, 21, 59, tzinfo=UTC))) self.assertEqual("projection-window-open", module.time_state(contract, NOW)) self.assertEqual("projection-cutoff-passed", module.time_state(contract, datetime(2026, 8, 22, 22, 4, tzinfo=UTC))) with self.assertRaisesRegex(module.ProcedureError, "before receipt expiry"): module.assert_expired_cleanup(contract, NOW) module.assert_expired_cleanup(contract, datetime(2026, 8, 22, 22, 16, tzinfo=UTC)) if __name__ == "__main__": unittest.main()