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

import io
import json
import os
import sqlite3
import sys
import tarfile
import tempfile
import unittest
from pathlib import Path

ROOT = Path(__file__).resolve().parents[3]
sys.path.insert(0, str(ROOT / "tools/k3s/lib"))
from phase1_common import (  # noqa: E402
    Phase1Error,
    absolute_path_without_symlink_resolution,
    build_ssh_command,
    exclusive_lock,
    parse_k3s_version,
    prune_backup_pairs,
    resolve_host_access,
    safe_extract_tar,
    sha256_file,
    verify_manifest,
)


class Phase1CommonTest(unittest.TestCase):
    def test_parse_k3s_version(self) -> None:
        self.assertEqual(parse_k3s_version("k3s version v1.33.4+k3s1 (abc)"), "v1.33.4+k3s1")
        with self.assertRaises(Phase1Error):
            parse_k3s_version("not k3s")

    def test_resolve_host_and_build_strict_ssh_command(self) -> None:
        registry = {
            "hosts": [
                {
                    "name": "debian3",
                    "state": "reachable",
                    "access": {
                        "tailscale_ip": {
                            "host": "100.101.104.41",
                            "port": 2222,
                            "user": "user",
                            "identity_file": "~/.ssh/id_test",
                        },
                        "tailscale_ssh": {"host": "debian3.example.ts.net", "port": 22, "user": "user"},
                    },
                }
            ]
        }
        access = resolve_host_access(registry, "debian3")
        self.assertEqual(access.host, "100.101.104.41")
        self.assertEqual(access.magic_dns, "debian3.example.ts.net")
        command = build_ssh_command(access)
        self.assertIn("StrictHostKeyChecking=yes", command)
        self.assertNotIn("StrictHostKeyChecking=accept-new", command)
        self.assertFalse(any("HostKeyAlias=" in part for part in command))
        self.assertEqual(command[-2:], ["--", "user@100.101.104.41"])

    def test_resolve_host_rejects_ssh_argument_injection(self) -> None:
        def registry(host: str, user: str) -> dict[str, object]:
            return {
                "hosts": [
                    {
                        "name": "debian3",
                        "state": "reachable",
                        "access": {"tailscale_ip": {"host": host, "user": user}},
                    }
                ]
            }

        with self.assertRaisesRegex(Phase1Error, "invalid SSH user"):
            resolve_host_access(registry("100.101.104.41", "-oProxyCommand=evil"), "debian3")
        with self.assertRaisesRegex(Phase1Error, "invalid SSH host"):
            resolve_host_access(registry("-oProxyCommand=evil", "user"), "debian3")

    def test_safe_extract_rejects_traversal_and_links(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            traversal = root / "traversal.tar"
            with tarfile.open(traversal, "w") as archive:
                info = tarfile.TarInfo("../escape")
                payload = b"x"
                info.size = len(payload)
                archive.addfile(info, io.BytesIO(payload))
            with self.assertRaises(Phase1Error):
                safe_extract_tar(traversal, root / "out1")

            linked = root / "linked.tar"
            with tarfile.open(linked, "w") as archive:
                info = tarfile.TarInfo("link")
                info.type = tarfile.SYMTYPE
                info.linkname = "/etc/passwd"
                archive.addfile(info)
            with self.assertRaises(Phase1Error):
                safe_extract_tar(linked, root / "out2")

            duplicate = root / "duplicate.tar"
            with tarfile.open(duplicate, "w") as archive:
                for payload in [b"one", b"two"]:
                    info = tarfile.TarInfo("same")
                    info.size = len(payload)
                    archive.addfile(info, io.BytesIO(payload))
            with self.assertRaises(Phase1Error):
                safe_extract_tar(duplicate, root / "out3")

    def test_manifest_verification_requires_paired_token(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            state = root / "payload/datastore/db/state.db"
            state.parent.mkdir(parents=True)
            sqlite3.connect(state).close()
            token = root / "payload/server/token"
            token.parent.mkdir(parents=True)
            token.write_text("K10abcdef::0123456789abcdef\n")
            files = []
            for path in [state, token]:
                files.append(
                    {
                        "path": path.relative_to(root).as_posix(),
                        "size": path.stat().st_size,
                        "sha256": sha256_file(path),
                    }
                )
            manifest = {
                "schema_version": 1,
                "datastore": {"type": "sqlite"},
                "server_token_path": "payload/server/token",
                "files": files,
            }
            manifest_path = root / "manifest.json"
            manifest_path.write_text(json.dumps(manifest))
            self.assertEqual(verify_manifest(root, manifest_path)["datastore"]["type"], "sqlite")
            manifest["server_token_path"] = "payload/server/missing"
            manifest_path.write_text(json.dumps(manifest))
            with self.assertRaises(Phase1Error):
                verify_manifest(root, manifest_path)

    def test_manifest_rejects_unlisted_archive_payload(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            state = root / "payload/datastore/db/state.db"
            state.parent.mkdir(parents=True)
            sqlite3.connect(state).close()
            token = root / "payload/server/token"
            token.parent.mkdir(parents=True)
            token.write_text("K10abcdef::0123456789abcdef\n")
            extra = root / "payload/unlisted"
            extra.write_text("not in manifest")
            files = [
                {
                    "path": path.relative_to(root).as_posix(),
                    "size": path.stat().st_size,
                    "sha256": sha256_file(path),
                }
                for path in [state, token]
            ]
            manifest = {
                "schema_version": 1,
                "datastore": {"type": "sqlite"},
                "server_token_path": "payload/server/token",
                "files": files,
            }
            manifest_path = root / "manifest.json"
            manifest_path.write_text(json.dumps(manifest))
            with self.assertRaises(Phase1Error):
                verify_manifest(root, manifest_path)

    def test_manifest_rejects_mode_drift(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            state = root / "payload/datastore/db/state.db"
            state.parent.mkdir(parents=True)
            sqlite3.connect(state).close()
            token = root / "payload/server/token"
            token.parent.mkdir(parents=True)
            token.write_text("K10abcdef::0123456789abcdef\n")
            state.chmod(0o600)
            token.chmod(0o600)
            files = [
                {
                    "path": path.relative_to(root).as_posix(),
                    "size": path.stat().st_size,
                    "sha256": sha256_file(path),
                    "mode": "0600",
                }
                for path in [state, token]
            ]
            manifest = {
                "schema_version": 1,
                "datastore": {"type": "sqlite"},
                "server_token_path": "payload/server/token",
                "files": files,
            }
            manifest_path = root / "manifest.json"
            manifest_path.write_text(json.dumps(manifest))
            verify_manifest(root, manifest_path)
            token.chmod(0o644)
            with self.assertRaises(Phase1Error):
                verify_manifest(root, manifest_path)

    def test_absolute_path_helper_does_not_hide_final_symlink(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            target = root / "target"
            target.write_text("secret")
            link = root / "link"
            link.symlink_to(target)
            expanded = absolute_path_without_symlink_resolution(link)
            self.assertEqual(expanded, link)
            self.assertTrue(expanded.is_symlink())

    def test_exclusive_lock_rejects_concurrent_invocation(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            lock = Path(tmp) / "phase1.lock"
            with exclusive_lock(lock):
                with self.assertRaises(Phase1Error):
                    with exclusive_lock(lock):
                        self.fail("second lock unexpectedly succeeded")

    def test_retention_prunes_archive_and_metadata_pairs_only(self) -> None:
        with tempfile.TemporaryDirectory() as tmp:
            root = Path(tmp)
            for index in range(4):
                archive = root / f"overdeck-k3s-debian3-2026080{index}T000000Z.tar.age"
                archive.write_bytes(str(index).encode())
                archive.with_suffix("").with_suffix(".json").write_text(
                    json.dumps({"status": "verified", "encrypted_sha256": sha256_file(archive)})
                )
                os.utime(archive, (index + 1, index + 1))
            unverified = root / "overdeck-k3s-debian3-20260701T000000Z.tar.age"
            unverified.write_bytes(b"retain-for-review")
            unrelated = root / "other.tar.age"
            unrelated.write_bytes(b"keep")
            removed = prune_backup_pairs(root, prefix="overdeck-k3s-debian3", retain=2)
            self.assertEqual(len(removed), 2)
            self.assertTrue(unrelated.exists())
            self.assertTrue(unverified.exists())
            self.assertEqual(len(list(root.glob("overdeck-k3s-debian3-*.tar.age"))), 3)


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