#!/usr/bin/env python3
from __future__ import annotations

import importlib.util
import json
import os
import shutil
import subprocess
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch

MODULE_PATH = Path(__file__).resolve().parents[1] / "phase1-control-plane.py"
SPEC = importlib.util.spec_from_file_location("phase1_control_plane", MODULE_PATH)
assert SPEC and SPEC.loader
orchestrator = importlib.util.module_from_spec(SPEC)
SPEC.loader.exec_module(orchestrator)


class Phase1OrchestratorTest(unittest.TestCase):
    def test_remote_helper_error_payload_is_preserved_in_receipt(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            instance = object.__new__(orchestrator.Orchestrator)
            instance.ssh = ["ssh", "debian3"]
            instance.remote_helper_path = "/tmp/phase1-server.py"
            instance.receipt_dir = root
            payload = {
                "schema_version": 2,
                "status": "error",
                "error": "trusted K3s launcher could not be located",
                "diagnostics": {"secret_values_recorded": False, "service": {"ActiveState": "active"}},
            }
            completed = subprocess.CompletedProcess(["ssh"], 2, json.dumps(payload), "")
            with patch.object(orchestrator, "run_command", return_value=completed):
                with self.assertRaises(orchestrator.Phase1Error):
                    instance.remote_call("inspect")
            self.assertEqual(json.loads((root / "remote-inspect-error.json").read_text()), payload)

    def test_systemd_escape_preserves_spaces_and_literal_percent(self) -> None:
        rendered = orchestrator.Orchestrator.systemd_escape('/tmp/a path/100%/"quoted"')
        self.assertEqual(rendered, "/tmp/a\\x20path/100%%/\\x22quoted\\x22")
        with self.assertRaises(orchestrator.Phase1Error):
            orchestrator.Orchestrator.systemd_escape("/tmp/line\nbreak")

    def test_endpoint_loopback_remap_preserves_explicit_api_port(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp); kubeconfig = root / "kubeconfig"; kubeconfig.write_text("placeholder")
            instance = object.__new__(orchestrator.Orchestrator); instance.kubeconfig = kubeconfig
            payload = '{"clusters":[{"cluster":{"server":"https://127.0.0.1:7443"}}]}'
            completed = subprocess.CompletedProcess(["kubectl"], 0, payload, "")
            with patch.object(orchestrator, "require_commands"), patch.object(orchestrator, "run_command", return_value=completed):
                endpoint = instance.endpoint_from_kubeconfig("100.101.104.41")
            self.assertEqual(endpoint, "https://100.101.104.41:7443")

    def test_age_identity_symlink_is_rejected_before_key_use(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp); target = root / "target"; target.write_text("AGE-SECRET-KEY-TEST\n"); target.chmod(0o600); link = root / "identity"; link.symlink_to(target)
            instance = object.__new__(orchestrator.Orchestrator); instance.age_identity = link; instance.args = SimpleNamespace(install_age=False)
            with patch.object(orchestrator.shutil, "which", return_value="/usr/bin/fake"):
                with self.assertRaises(orchestrator.Phase1Error): instance.ensure_age()

    def test_capture_timer_rollback_removes_partial_transaction_on_error(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp); instance = object.__new__(orchestrator.Orchestrator); instance.receipt_dir = root / "receipt"; instance.receipt_dir.mkdir(); unsafe = root / "managed-directory"; unsafe.mkdir()
            with self.assertRaises(orchestrator.Phase1Error): instance.capture_timer_rollback([unsafe])
            self.assertFalse((instance.receipt_dir / ".timer-rollback").exists())

    def test_commit_timer_clears_rollback_authority_before_cleanup(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            rollback = Path(tmp) / "rollback"; rollback.mkdir(); instance = object.__new__(orchestrator.Orchestrator); instance.timer_rollback = {"directory": str(rollback)}
            with patch.object(orchestrator.shutil, "rmtree", side_effect=OSError("injected cleanup failure")):
                result = instance.commit_timer()
            self.assertIsNone(instance.timer_rollback)
            self.assertEqual(result["status"], "committed-with-cleanup-warning")

    def test_remote_finalize_crosses_commit_point_before_receipt_write(self) -> None:
        instance = object.__new__(orchestrator.Orchestrator)
        instance.converge_tx = "phase1-test"; instance.converge_changed = True; instance.timer_rollback = {"directory": "/tmp/not-used"}; instance.result = {}
        instance.remote_call = lambda *args, **kwargs: {"transaction_id": "phase1-test", "status": "finalized"}
        calls: list[str] = []
        def commit_timer(): calls.append("commit"); instance.timer_rollback = None; return {"status": "committed", "rollback_material_removed": True}
        instance.commit_timer = commit_timer; instance.save_json = lambda *args, **kwargs: (_ for _ in ()).throw(OSError("receipt full")); instance.add_step = lambda *args, **kwargs: None
        with self.assertRaises(OSError): instance.finalize_remote()
        self.assertEqual(calls, ["commit"])
        self.assertIsNone(instance.converge_tx)
        self.assertFalse(instance.converge_changed)

    def test_qualification_reports_all_blocked_dependencies_before_stopping(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            instance = object.__new__(orchestrator.Orchestrator)
            instance.receipt_dir = Path(tmp); instance.qualification = {"schema_version": 2, "status": "running", "started_utc": orchestrator.utc_now(), "checks": []}
            instance.args = SimpleNamespace(server="debian3", mode="apply")
            instance.ensure_age = lambda: (_ for _ in ()).throw(orchestrator.Phase1Error("age unavailable"))
            instance.timer_preflight = lambda: {"status": "qualified"}
            instance.inspect_server = lambda: (_ for _ in ()).throw(orchestrator.Phase1Error("server unavailable"))
            with self.assertRaisesRegex(orchestrator.Phase1Error, "qualification did not pass"):
                instance.run_qualification()
            matrix = json.loads((Path(tmp) / "qualification.json").read_text())
            names = {item["name"]: item["status"] for item in matrix["checks"]}
            self.assertEqual(names["age-identity-and-tools"], "failed")
            self.assertEqual(names["server-inspect-readiness-nodes"], "failed")
            self.assertEqual(names["control-plane-plan"], "blocked")
            self.assertEqual(names["prechange-backup-encrypt-verify-materialize"], "blocked")
            self.assertFalse(matrix["cluster_configuration_mutated"])

    def test_successful_qualification_verifies_prechange_backup_before_return(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            calls: list[str] = []
            instance = object.__new__(orchestrator.Orchestrator)
            instance.receipt_dir = Path(tmp); instance.qualification = {"schema_version": 2, "status": "running", "started_utc": orchestrator.utc_now(), "checks": []}
            instance.args = SimpleNamespace(server="debian3", mode="apply")
            instance.ensure_age = lambda: calls.append("age") or "age1fixture"
            instance.timer_preflight = lambda: calls.append("timer") or {"status": "qualified"}
            inspect = {"cacerts_sha256": "c" * 64}
            instance.inspect_server = lambda: calls.append("inspect") or inspect
            instance.plan_control_plane = lambda value: calls.append("plan") or ("https://100.0.0.3:6443", ["100.0.0.3"], {"changed": True})
            instance.kubectl_ready = lambda: calls.append("kubectl") or {"status": "ready"}
            instance.probe_reachable_nodes = lambda endpoint, ca: calls.append("peers") or [{"status": "reachable"}]
            instance.backup_cycle = lambda recipient, purpose: calls.append(f"backup:{purpose}") or {"verified": {"status": "verified"}}
            result = instance.run_qualification()
            self.assertEqual(calls, ["age", "timer", "inspect", "plan", "kubectl", "peers", "backup:prechange"])
            self.assertEqual(result["prechange"]["verified"]["status"], "verified")
            self.assertEqual(instance.qualification["status"], "passed")

    def test_apply_cannot_converge_until_qualification_returns_success(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            calls: list[str] = []
            instance = object.__new__(orchestrator.Orchestrator)
            instance.receipt_dir = Path(tmp)
            instance.args = SimpleNamespace(mode="apply", server="debian3", transaction=None, recovery_action=None)
            instance.result = {"steps": []}; instance.qualification = {"schema_version": 2, "status": "running", "checks": []}; instance.log = lambda message: None
            instance.install_remote_helper = lambda: calls.append("helper")
            instance.run_qualification = lambda: calls.append("qualify") or {"endpoint": "https://100.0.0.3:6443", "sans": ["100.0.0.3"], "recipient": "age1", "prechange": {"verified": True}}
            instance.converge = lambda endpoint, sans: calls.append("converge") or {}
            instance.converge_changed = True; instance.converge_tx = "tx"
            post = {"service": {"ready": {"ready": True}}, "node_probe": {"ok": True, "ready_count": 1}, "cacerts_sha256": "c" * 64}
            instance.remote_call = lambda *args, **kwargs: calls.append("post-inspect") or post
            instance.save_json = lambda *args, **kwargs: None; instance.kubectl_ready = lambda: calls.append("kubectl") or {}; instance.probe_reachable_nodes = lambda *args: calls.append("peers") or []
            instance.backup_cycle = lambda recipient, purpose: calls.append(f"backup:{purpose}") or {"verified": True}
            instance.install_timer = lambda recipient: calls.append("timer") or {}; instance.finalize_remote = lambda: calls.append("finalize") or {}; instance.prune = lambda: calls.append("prune") or []
            instance.rollback_timer = lambda: None; instance.backup_tx = None; instance.rollback_remote = lambda error: None; instance.mark_failed_backup = lambda error: None
            rc = instance.run_locked()
            self.assertEqual(rc, 0)
            self.assertLess(calls.index("qualify"), calls.index("converge"))
            self.assertLess(calls.index("converge"), calls.index("backup:postchange"))
            self.assertLess(calls.index("backup:postchange"), calls.index("finalize"))

    def test_timer_rollback_failure_does_not_block_control_plane_rollback(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            calls: list[str] = []
            instance = object.__new__(orchestrator.Orchestrator)
            instance.receipt_dir = Path(tmp)
            instance.args = SimpleNamespace(mode="apply", server="debian3", transaction=None, recovery_action=None)
            instance.result = {"steps": []}
            instance.qualification = {"schema_version": 2, "status": "running", "checks": []}
            instance.log = lambda message: None
            instance.install_remote_helper = lambda: (_ for _ in ()).throw(orchestrator.Phase1Error("apply failed"))
            instance.rollback_timer = lambda: (_ for _ in ()).throw(orchestrator.Phase1Error("timer rollback failed"))
            instance.backup_tx = None
            instance.rollback_remote = lambda error: calls.append(error) or {"status": "rolled-back"}
            instance.mark_failed_backup = lambda error: None

            rc = instance.run_locked()

            self.assertEqual(rc, 2)
            self.assertEqual(calls, ["apply failed"])
            result = json.loads((Path(tmp) / "phase1-result.json").read_text())
            self.assertEqual(result["rollback"]["status"], "rolled-back")
            self.assertEqual(result["timer_rollback"]["status"], "rollback-call-failed")
            self.assertFalse(result["git_publication_allowed"])

    def test_rendered_systemd_units_validate(self) -> None:
        analyzer = shutil.which("systemd-analyze")
        if analyzer is None: self.skipTest("systemd-analyze is unavailable")
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp); install = root / "install root"; config = root / "config dir/config.json"; backup = root / "backup dir"; lock = root / "state dir"
            service_text = orchestrator.SYSTEMD_SERVICE.read_text()
            replacements = {"@INSTALL_ROOT@": orchestrator.Orchestrator.systemd_escape(install), "@CONFIG_PATH@": orchestrator.Orchestrator.systemd_escape(config), "@BACKUP_DIR@": orchestrator.Orchestrator.systemd_escape(backup), "@CONFIG_DIR@": orchestrator.Orchestrator.systemd_escape(config.parent), "@LOCK_DIR@": orchestrator.Orchestrator.systemd_escape(lock)}
            for key, value in replacements.items(): service_text = service_text.replace(key, value)
            service = root / orchestrator.SYSTEMD_SERVICE.name; timer = root / orchestrator.SYSTEMD_TIMER.name; install.mkdir(parents=True); executable = install / "phase1-control-plane.sh"; executable.write_text("#!/bin/sh\nexit 0\n"); executable.chmod(0o755); config.parent.mkdir(parents=True); config.write_text("{}\n"); backup.mkdir(parents=True); lock.mkdir(parents=True); service.write_text(service_text); timer.write_text(orchestrator.SYSTEMD_TIMER.read_text())
            runtime = root / "runtime"; runtime.mkdir(mode=0o700); environment = dict(os.environ); environment["XDG_RUNTIME_DIR"] = str(runtime)
            completed = subprocess.run([analyzer, "--user", "verify", str(service), str(timer)], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, env=environment, check=False)
            self.assertEqual(completed.returncode, 0, completed.stdout + completed.stderr)


if __name__ == "__main__":
    unittest.main()
