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

import json
import os
import tempfile
import unittest
from pathlib import Path
import sys

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

from phase2_common import (
    CandidateIdentity,
    EnrollmentLedger,
    Phase2Error,
    assert_secret_free,
    build_enrollment_plan,
    build_registry_previews,
    find_secret_strings,
    recovery_door_contract,
    sha256_json,
    validate_candidate_uniqueness,
    validate_host_preflight,
    write_json_pair_atomic,
)
from phase1_common import read_json


class Phase2CommonTests(unittest.TestCase):
    def identity(self, **overrides):
        value = {
            "name": "debian4",
            "dns_name": "debian4.example.ts.net",
            "tailscale_ipv4": "100.64.44.44",
            "machine_id": "4" * 32,
            "os_id": "debian",
            "os_version_id": "12",
            "architecture": "x86_64",
            "ssh_user": "user",
            "rustdesk": "28884444",
        }
        value.update(overrides)
        return CandidateIdentity.from_mapping(value)

    def plan(self):
        identity = self.identity()
        return build_enrollment_plan(
            identity,
            control_plane={"api": {"endpoint": "https://100.101.104.41:6443", "cacerts_sha256": "1" * 64}},
            version_lock={"version": "v1.36.3+k3s1", "launcher": {"sha256": "2" * 64}},
            source_digests={"fleet.json": "3" * 64, "buildbox-hosts.json": "4" * 64},
            recovery_doors=recovery_door_contract(identity),
        )

    def test_tailscale_shared_range_is_accepted(self):
        self.assertEqual(self.identity().tailscale_ipv4, "100.64.44.44")

    def test_non_tailscale_ipv4_is_rejected(self):
        with self.assertRaises(Phase2Error):
            self.identity(tailscale_ipv4="192.168.1.8")

    def test_invalid_machine_id_is_rejected(self):
        with self.assertRaises(Phase2Error):
            self.identity(machine_id="abc")

    def test_secret_scanner_allows_token_metadata(self):
        assert_secret_free({"bootstrap_token": {"value_recorded": False, "ttl_seconds": 600}})

    def test_secret_scanner_rejects_token_value(self):
        findings = find_secret_strings({"token_value": "abcdef.0123456789abcdef"})
        self.assertTrue(findings)

    def test_plan_has_stable_digest_and_seventeen_steps(self):
        first = self.plan()
        second = self.plan()
        self.assertEqual(first, second)
        self.assertEqual(len(first["steps"]), 17)
        self.assertEqual(first["plan_sha256"], sha256_json({k: v for k, v in first.items() if k != "plan_sha256"}))
        self.assertFalse(first["live_mutation_allowed"])
        self.assertFalse(first["bootstrap_token"]["value_recorded"])

    def test_preflight_rejects_active_agent(self):
        with self.assertRaises(Phase2Error):
            validate_host_preflight({
                "systemd": True, "cgroup_v2": True, "tailscale_online": True,
                "tailscale_interface": True, "sudo_noninteractive": True, "clock_synchronized": True,
                "memory_bytes": 8 * 1024**3, "disk_free_bytes": 100 * 1024**3,
                "k3s_agent_state": "active",
            })

    def test_uniqueness_requires_one_exact_peer(self):
        identity = self.identity()
        with self.assertRaises(Phase2Error):
            validate_candidate_uniqueness(identity, {"nodes": {}}, {"hosts": []}, [], [])

    def test_registry_preview_disables_dispatch(self):
        fleet = {"schema_version": 1, "nodes": {}, "fallback": {"requires_all_unavailable": []}}
        hosts = {"schema_version": 1, "hosts": [], "orders": {"build": [], "e2e": []}}
        fleet_out, hosts_out, pair = build_registry_previews(fleet, hosts, self.identity())
        self.assertEqual(fleet_out["nodes"]["debian4"]["execution"], "none")
        self.assertEqual(pair["execution"], "none")
        self.assertNotIn("debian4", hosts_out["orders"]["build"])

    def test_ledger_is_bound_and_unexecuted(self):
        plan = self.plan()
        with tempfile.TemporaryDirectory() as tmp:
            ledger = EnrollmentLedger(Path(tmp) / "ledger.json", plan, create=True)
            self.assertEqual(ledger.next_step(), "identity-discovery")
            self.assertEqual(ledger.pending_rollback(), [])
            with self.assertRaises(Phase2Error):
                ledger.record("uniqueness-gate", "started")
            ledger.record("identity-discovery", "started")
            ledger.record("identity-discovery", "completed")
            self.assertEqual(ledger.next_step(), "uniqueness-gate")

    def test_atomic_pair_rolls_back_second_rename_failure(self):
        first_value = {"old": 1}
        second_value = {"old": 2}
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            first = root / "first.json"
            second = root / "second.json"
            first.write_text(json.dumps(first_value))
            second.write_text(json.dumps(second_value))
            real_replace = os.replace
            calls = 0
            def broken(src, dst):
                nonlocal calls
                calls += 1
                if calls == 2:
                    raise OSError("injected second publication failure")
                return real_replace(src, dst)
            import phase2_common
            phase2_common.os.replace = broken
            try:
                with self.assertRaises(OSError):
                    write_json_pair_atomic(first, {"new": 1}, second, {"new": 2})
            finally:
                phase2_common.os.replace = real_replace
            self.assertEqual(json.loads(first.read_text()), first_value)
            self.assertEqual(json.loads(second.read_text()), second_value)


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