#!/usr/bin/env python3
"""Build a deterministic, secret-free K3s node-enrollment transaction plan.

Phase 2 intentionally performs no live mutation.  It can inspect a real
post-Tailscale candidate or replay the shipped sanitized fixture, validate the
candidate against live/tracked sources of truth, and emit the exact Phase 3
transaction/rollback plan plus atomic registry previews.
"""
from __future__ import annotations

import argparse
import json
import os
import secrets
import shlex
import sys
from pathlib import Path
from typing import Any, Callable, Mapping, Sequence

K3S_DIR = Path(__file__).resolve().parent
LIB_DIR = K3S_DIR / "lib"
if str(LIB_DIR) not in sys.path:
    sys.path.insert(0, str(LIB_DIR))

from phase1_common import (  # noqa: E402
    atomic_write_json,
    build_ssh_command,
    exclusive_lock,
    load_host_registry,
    normalize_path,
    read_json,
    resolve_host_access,
    run_command,
    sha256_file,
    utc_now,
    utc_stamp,
)
from phase2_common import (  # noqa: E402
    CandidateIdentity,
    EnrollmentLedger,
    Phase2Error,
    assert_secret_free,
    build_enrollment_plan,
    build_registry_previews,
    canonical_json_bytes,
    recovery_door_contract,
    sha256_json,
    summarize_plan,
    validate_candidate_uniqueness,
    validate_host_preflight,
)

CANDIDATE_HELPER = K3S_DIR / "remote" / "phase2-candidate.py"
SERVER_HELPER = K3S_DIR / "remote" / "phase2-server.py"


class Logger:
    def __init__(self, path: Path) -> None:
        self.path = path
        path.parent.mkdir(parents=True, exist_ok=True)

    def __call__(self, message: str) -> None:
        line = f"[phase2 {utc_now()}] {message}"
        print(line, file=sys.stderr, flush=True)
        with self.path.open("a", encoding="utf-8") as handle:
            handle.write(line + "\n")


def parse_tailscale_status(document: Mapping[str, Any]) -> list[dict[str, Any]]:
    peers: list[dict[str, Any]] = []

    def add(raw: Mapping[str, Any]) -> None:
        dns = str(raw.get("DNSName") or "").rstrip(".")
        host = str(raw.get("HostName") or "")
        ips = [str(item) for item in raw.get("TailscaleIPs", []) if isinstance(item, str)]
        if not host and dns:
            host = dns.split(".")[0]
        if not host and not dns and not ips:
            return
        peers.append(
            {
                "name": host,
                "dns_name": dns,
                "ips": sorted(set(ips)),
                "online": bool(raw.get("Online", True)),
            }
        )

    self_item = document.get("Self")
    if isinstance(self_item, Mapping):
        add(self_item)
    peer_root = document.get("Peer")
    if isinstance(peer_root, Mapping):
        for value in peer_root.values():
            if isinstance(value, Mapping):
                add(value)
    elif isinstance(peer_root, list):
        for value in peer_root:
            if isinstance(value, Mapping):
                add(value)
    dedup: dict[tuple[str, str, tuple[str, ...]], dict[str, Any]] = {}
    for peer in peers:
        key = (peer["name"], peer["dns_name"], tuple(peer["ips"]))
        dedup[key] = peer
    return sorted(dedup.values(), key=lambda item: (item["name"], item["dns_name"], item["ips"]))


def parse_cluster_nodes(document: Mapping[str, Any]) -> list[dict[str, Any]]:
    items = document.get("items")
    if not isinstance(items, list):
        raise Phase2Error("kubectl node document has no items array")
    nodes: list[dict[str, Any]] = []
    for item in items:
        if not isinstance(item, Mapping):
            continue
        metadata = item.get("metadata") if isinstance(item.get("metadata"), Mapping) else {}
        status = item.get("status") if isinstance(item.get("status"), Mapping) else {}
        labels = metadata.get("labels") if isinstance(metadata.get("labels"), Mapping) else {}
        conditions = status.get("conditions") if isinstance(status.get("conditions"), list) else []
        addresses = status.get("addresses") if isinstance(status.get("addresses"), list) else []
        ready = "Unknown"
        for condition in conditions:
            if isinstance(condition, Mapping) and condition.get("type") == "Ready":
                ready = str(condition.get("status") or "Unknown")
        internal_ip = None
        for address in addresses:
            if isinstance(address, Mapping) and address.get("type") == "InternalIP":
                internal_ip = str(address.get("address") or "")
                break
        nodes.append(
            {
                "name": str(metadata.get("name") or ""),
                "hostname": str(labels.get("kubernetes.io/hostname") or metadata.get("name") or ""),
                "internal_ip": internal_ip,
                "ready": ready,
                "uid": str(metadata.get("uid") or ""),
            }
        )
    return sorted(nodes, key=lambda item: item["name"])


def select_tailscale_candidate(peers: Sequence[Mapping[str, Any]], name: str, ssh_user: str) -> CandidateIdentity:
    matches = [
        item
        for item in peers
        if item.get("name") == name or str(item.get("dns_name") or "").split(".")[0] == name
    ]
    if len(matches) != 1:
        raise Phase2Error(f"expected exactly one Tailscale peer named {name!r}; found {len(matches)}")
    peer = matches[0]
    ipv4 = [value for value in peer.get("ips", []) if isinstance(value, str) and ":" not in value]
    if len(ipv4) != 1:
        raise Phase2Error(f"candidate {name!r} must have exactly one Tailscale IPv4; found {len(ipv4)}")
    # machine-id and OS facts are filled after the remote candidate helper runs.
    return CandidateIdentity.from_mapping(
        {
            "name": name,
            "dns_name": str(peer.get("dns_name") or name),
            "tailscale_ipv4": ipv4[0],
            "machine_id": "0" * 32,
            "os_id": "debian",
            "os_version_id": "unknown",
            "architecture": "x86_64",
            "ssh_user": ssh_user,
            "rustdesk": None,
        }
    )


class Orchestrator:
    def __init__(self, args: argparse.Namespace) -> None:
        self.args = args
        self.repo_root = normalize_path(args.repo_root)
        self.fleet_path = normalize_path(args.fleet or self.repo_root / "modules/fleet/fleet.json")
        self.hosts_path = normalize_path(
            args.host_registry or self.repo_root / "modules/workstation/claude/buildbox-hosts.json"
        )
        self.kubeconfig = normalize_path(args.kubeconfig)
        self.fixture_path = normalize_path(args.fixture) if args.fixture else None
        default_receipt = (
            Path.home() / ".local/state/overdeck/k3s-enrollment" / f"phase2-{utc_stamp()}-{os.getpid()}"
        )
        self.receipt_dir = normalize_path(args.receipt_dir or default_receipt)
        if self.receipt_dir.exists() or self.receipt_dir.is_symlink():
            raise Phase2Error(f"receipt path already exists: {self.receipt_dir}")
        self.receipt_dir.mkdir(parents=True, mode=0o700)
        (self.receipt_dir / "logs").mkdir(mode=0o700)
        (self.receipt_dir / "registry-preview").mkdir(mode=0o700)
        self.log = Logger(self.receipt_dir / "logs" / "phase2.log")
        self.lock_file = normalize_path(
            args.lock_file or Path.home() / ".local/state/overdeck" / f"k3s-enroll-{args.candidate}.lock"
        )
        self.candidate_helper_remote: str | None = None
        self.server_helper_remote: str | None = None
        self.candidate_ssh: list[str] | None = None
        self.server_ssh: list[str] | None = None
        self.result: dict[str, Any] = {
            "schema_version": 1,
            "phase": 2,
            "mode": args.mode,
            "status": "running",
            "started_utc": utc_now(),
            "candidate": args.candidate,
            "repo_root": str(self.repo_root),
            "receipt_dir": str(self.receipt_dir),
            "fixture": str(self.fixture_path) if self.fixture_path else None,
            "steps": [],
            "live_mutation_performed": False,
            "candidate_mutation_performed": False,
            "cluster_mutation_performed": False,
            "registry_mutation_performed": False,
            "phase3_authorized": False,
            "git_publication_allowed": False,
        }
        atomic_write_json(self.receipt_dir / "phase2-result.json", self.result)

    def add_step(self, name: str, status: str, **detail: Any) -> None:
        item = {"name": name, "status": status, "at_utc": utc_now()}
        item.update(detail)
        assert_secret_free(item, f"step {name}")
        self.result["steps"].append(item)
        atomic_write_json(self.receipt_dir / "phase2-result.json", self.result)

    def save_json(self, relative: str, value: Any) -> None:
        assert_secret_free(value, relative)
        atomic_write_json(self.receipt_dir / relative, value, mode=0o600)

    def _source_digests(self) -> dict[str, str]:
        return {
            "fleet.json": sha256_file(self.fleet_path),
            "buildbox-hosts.json": sha256_file(self.hosts_path),
        }

    def _load_fixture(self) -> dict[str, Any]:
        if self.fixture_path is None:
            raise Phase2Error("fixture path is not configured")
        document = read_json(self.fixture_path)
        if not isinstance(document, dict) or document.get("schema_version") != 1:
            raise Phase2Error(f"unsupported Phase 2 fixture: {self.fixture_path}")
        self.add_step("fixture-load", "passed", fixture_name=document.get("fixture_name"))
        return document

    def _candidate_ssh_command(self, dns_name: str, user: str) -> list[str]:
        return [
            "ssh",
            "-p",
            "22",
            "-o",
            "BatchMode=yes",
            "-o",
            f"ConnectTimeout={self.args.ssh_timeout}",
            "-o",
            "ServerAliveInterval=15",
            "-o",
            "ServerAliveCountMax=2",
            "-o",
            "StrictHostKeyChecking=yes",
            f"{user}@{dns_name}",
        ]

    def _upload_helper(self, ssh: Sequence[str], helper: Path, label: str) -> str:
        if not helper.is_file() or helper.is_symlink():
            raise Phase2Error(f"unsafe {label} helper: {helper}")
        remote = f"/tmp/overdeck-k3s-phase2-{label}-{os.getpid()}-{secrets.token_hex(4)}.py"
        run_command(list(ssh) + ["umask 077; cat > " + shlex.quote(remote)], input_bytes=helper.read_bytes(), timeout=60)
        run_command(list(ssh) + [shlex.join(["chmod", "0700", remote])], timeout=30)
        expected = sha256_file(helper)
        completed = run_command(list(ssh) + [shlex.join(["sha256sum", remote])], timeout=30)
        actual = str(completed.stdout).split()[0]
        if actual != expected:
            raise Phase2Error(f"{label} helper SHA-256 mismatch")
        self.add_step(f"{label}-helper-upload", "passed", sha256=expected, persistence="ephemeral")
        return remote

    def _remote_json(self, ssh: Sequence[str], helper: str, *arguments: str, sudo: bool = True) -> dict[str, Any]:
        argv = (["sudo", "-n"] if sudo else []) + ["/usr/bin/python3", helper, *arguments]
        completed = run_command(list(ssh) + [shlex.join(argv)], check=False, timeout=180)
        stdout = completed.stdout if isinstance(completed.stdout, str) else completed.stdout.decode(errors="replace")
        stderr = completed.stderr if isinstance(completed.stderr, str) else completed.stderr.decode(errors="replace")
        try:
            payload = json.loads(stdout)
        except json.JSONDecodeError as exc:
            raise Phase2Error(
                f"remote helper returned invalid JSON (exit {completed.returncode}); stderr={stderr[-1200:]!r}"
            ) from exc
        assert_secret_free(payload, "remote helper response")
        if completed.returncode != 0 or payload.get("status") == "error":
            raise Phase2Error(str(payload.get("error") or f"remote helper exited {completed.returncode}"))
        return payload

    def _collect_live(self) -> dict[str, Any]:
        tailscale_status = run_command(["tailscale", "status", "--json"], timeout=30)
        try:
            tailscale_document = json.loads(str(tailscale_status.stdout))
        except json.JSONDecodeError as exc:
            raise Phase2Error("tailscale status --json returned invalid JSON") from exc
        peers = parse_tailscale_status(tailscale_document)
        provisional = select_tailscale_candidate(peers, self.args.candidate, self.args.candidate_user)
        self.candidate_ssh = self._candidate_ssh_command(provisional.dns_name, provisional.ssh_user)
        self.candidate_helper_remote = self._upload_helper(self.candidate_ssh, CANDIDATE_HELPER, "candidate")
        candidate = self._remote_json(self.candidate_ssh, self.candidate_helper_remote, "inspect")
        identity_value = dict(candidate.get("identity") or {})
        identity_value.update(
            {
                "name": provisional.name,
                "dns_name": provisional.dns_name,
                "tailscale_ipv4": provisional.tailscale_ipv4,
                "ssh_user": provisional.ssh_user,
            }
        )
        identity = CandidateIdentity.from_mapping(identity_value)

        registry = load_host_registry(self.hosts_path)
        server_access = resolve_host_access(
            registry,
            self.args.server,
            preferred_door=self.args.server_ssh_door,
            require_reachable=True,
        )
        self.server_ssh = build_ssh_command(server_access, timeout=self.args.ssh_timeout)
        self.server_helper_remote = self._upload_helper(self.server_ssh, SERVER_HELPER, "server")
        server = self._remote_json(self.server_ssh, self.server_helper_remote, "inspect")

        nodes_result = run_command(
            ["kubectl", "--kubeconfig", str(self.kubeconfig), "get", "nodes", "-o", "json"], timeout=45
        )
        try:
            nodes_document = json.loads(str(nodes_result.stdout))
        except json.JSONDecodeError as exc:
            raise Phase2Error("kubectl get nodes returned invalid JSON") from exc
        cluster_nodes = parse_cluster_nodes(nodes_document)
        return {
            "candidate": identity.as_dict(),
            "candidate_preflight": candidate.get("preflight"),
            "tailscale_peers": peers,
            "cluster_nodes": cluster_nodes,
            "control_plane": server.get("control_plane"),
            "version_lock": server.get("version_lock"),
        }

    def _cleanup_helpers(self) -> None:
        for ssh, path, label in (
            (self.candidate_ssh, self.candidate_helper_remote, "candidate"),
            (self.server_ssh, self.server_helper_remote, "server"),
        ):
            if ssh and path:
                completed = run_command(list(ssh) + [shlex.join(["rm", "-f", path])], check=False, timeout=30)
                if completed.returncode != 0:
                    self.result.setdefault("cleanup_warnings", []).append(f"{label} helper cleanup failed")
        self.candidate_helper_remote = None
        self.server_helper_remote = None

    def run_locked(self) -> int:
        before = self._source_digests()
        try:
            if self.args.mode == "apply":
                raise Phase2Error(
                    "live enrollment is hard-disabled in Phase 2; use the reviewed Phase 3 canary package"
                )
            source = self._load_fixture() if self.fixture_path else self._collect_live()
            identity = CandidateIdentity.from_mapping(source.get("candidate") or {})
            if identity.name != self.args.candidate:
                raise Phase2Error(
                    f"candidate argument {self.args.candidate!r} does not match discovered identity {identity.name!r}"
                )
            fleet = read_json(self.fleet_path)
            hosts = read_json(self.hosts_path)
            if not isinstance(fleet, dict) or not isinstance(hosts, dict):
                raise Phase2Error("source registries must be JSON objects")
            uniqueness = validate_candidate_uniqueness(
                identity,
                fleet,
                hosts,
                source.get("cluster_nodes") or [],
                source.get("tailscale_peers") or [],
            )
            self.save_json("uniqueness.json", uniqueness)
            self.add_step("candidate-uniqueness", "passed")
            preflight = validate_host_preflight(source.get("candidate_preflight") or {})
            self.save_json("candidate-preflight.json", preflight)
            self.add_step("candidate-preflight", "passed")

            source_digests = self._source_digests()
            plan = build_enrollment_plan(
                identity,
                control_plane=source.get("control_plane") or {},
                version_lock=source.get("version_lock") or {},
                source_digests=source_digests,
                recovery_doors=recovery_door_contract(identity),
            )
            self.save_json("plan.json", plan)
            (self.receipt_dir / "PLAN.txt").write_text(summarize_plan(plan), encoding="utf-8")
            os.chmod(self.receipt_dir / "PLAN.txt", 0o600)
            ledger = EnrollmentLedger(self.receipt_dir / "transaction-ledger.json", plan, create=True)
            if ledger.next_step() != "identity-discovery" or ledger.pending_rollback():
                raise Phase2Error("new enrollment ledger did not start from a clean state")
            self.add_step("deterministic-plan", "passed", plan_sha256=plan["plan_sha256"], step_count=len(plan["steps"]))

            fleet_preview, hosts_preview, pair = build_registry_previews(fleet, hosts, identity)
            self.save_json("registry-preview/fleet.json", fleet_preview)
            self.save_json("registry-preview/buildbox-hosts.json", hosts_preview)
            self.save_json("registry-preview/pair.json", pair)
            self.add_step("registry-preview", "passed", execution="none", ordered=False)

            after = self._source_digests()
            if before != after:
                raise Phase2Error("dry-run changed a tracked source registry")
            plan_roundtrip = json.loads(canonical_json_bytes(plan).decode("utf-8"))
            if sha256_json({k: v for k, v in plan_roundtrip.items() if k != "plan_sha256"}) != plan["plan_sha256"]:
                raise Phase2Error("plan digest self-check failed")
            assert_secret_free(
                {
                    "plan": plan,
                    "fleet_preview": fleet_preview,
                    "hosts_preview": hosts_preview,
                    "pair": pair,
                },
                "Phase 2 outputs",
            )
            self.result.update(
                {
                    "status": "success",
                    "finished_utc": utc_now(),
                    "candidate_identity": identity.as_dict(),
                    "plan_sha256": plan["plan_sha256"],
                    "transaction_id": plan["transaction_id"],
                    "plan_step_count": len(plan["steps"]),
                    "plan_deterministic": True,
                    "secret_scan_passed": True,
                    "source_digests_before": before,
                    "source_digests_after": after,
                    "registry_preview": pair,
                    "live_mutation_performed": False,
                    "candidate_mutation_performed": False,
                    "cluster_mutation_performed": False,
                    "registry_mutation_performed": False,
                    "phase3_authorized": False,
                    "git_publication_allowed": True,
                    "next_phase": 3,
                }
            )
            atomic_write_json(self.receipt_dir / "phase2-result.json", self.result)
            self.log("Phase 2 dry-run plan completed without live mutation")
            return 0
        except (Exception, KeyboardInterrupt) as exc:
            error = str(exc) or type(exc).__name__
            self.log(f"failure: {error}")
            after = self._source_digests()
            self.result.update(
                {
                    "status": "failed",
                    "finished_utc": utc_now(),
                    "error": error,
                    "source_digests_before": before,
                    "source_digests_after": after,
                    "live_mutation_performed": False,
                    "candidate_mutation_performed": False,
                    "cluster_mutation_performed": False,
                    "registry_mutation_performed": False,
                    "phase3_authorized": False,
                    "git_publication_allowed": False,
                }
            )
            atomic_write_json(self.receipt_dir / "phase2-result.json", self.result)
            return 2
        finally:
            self._cleanup_helpers()
            atomic_write_json(self.receipt_dir / "phase2-result.json", self.result)

    def run(self) -> int:
        try:
            with exclusive_lock(self.lock_file):
                return self.run_locked()
        except (Exception, KeyboardInterrupt) as exc:
            error = str(exc) or type(exc).__name__
            self.result.update(
                {
                    "status": "failed",
                    "finished_utc": utc_now(),
                    "error": error,
                    "git_publication_allowed": False,
                }
            )
            atomic_write_json(self.receipt_dir / "phase2-result.json", self.result)
            self.log(f"failure before locked execution: {error}")
            return 2


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("candidate", help="Tailscale machine name, for example debian4")
    parser.add_argument("--mode", choices=["plan", "dry-run", "apply"], default="dry-run")
    parser.add_argument("--repo-root", default=str(K3S_DIR.parent.parent))
    parser.add_argument("--fixture", help="Sanitized fixture JSON; skips all live SSH/Kubernetes calls")
    parser.add_argument("--fleet")
    parser.add_argument("--host-registry")
    parser.add_argument("--server", default="debian3")
    parser.add_argument("--server-ssh-door", choices=["tailscale_ip", "tailscale_ssh", "lan"], default="tailscale_ip")
    parser.add_argument("--candidate-user", default="user")
    parser.add_argument("--kubeconfig", default=str(Path.home() / ".kube/config-buildboxes"))
    parser.add_argument("--receipt-dir")
    parser.add_argument("--lock-file")
    parser.add_argument("--ssh-timeout", type=int, default=12)
    return parser


def main(argv: Sequence[str] | None = None) -> int:
    args = build_parser().parse_args(argv)
    if args.ssh_timeout < 1 or args.ssh_timeout > 120:
        print("phase2: --ssh-timeout must be between 1 and 120", file=sys.stderr)
        return 2
    try:
        orchestrator = Orchestrator(args)
    except Exception as exc:
        print(f"phase2: {exc}", file=sys.stderr)
        return 2
    return orchestrator.run()


if __name__ == "__main__":
    raise SystemExit(main())
