"""Unit tests for the explicit sandbox creation stub surface."""

from __future__ import annotations

import ast
from pathlib import Path

STUB_PATH = Path(__file__).parent.parent / "microsandbox" / "_microsandbox.pyi"

EXPECTED_KWARGS = [
    "image",
    "memory",
    "cpus",
    "max_memory",
    "max_cpus",
    "workdir",
    "shell",
    "security",
    "hostname",
    "user",
    "entrypoint",
    "cmd",
    "init",
    "replace",
    "replace_with_timeout",
    "max_duration",
    "idle_timeout",
    "ephemeral",
    "env",
    "labels",
    "scripts",
    "pull_policy",
    "log_level",
    "registry_auth",
    "registry_insecure",
    "registry_ca_certs",
    "volumes",
    "patches",
    "ports",
    "vsock",
    "network",
    "secrets",
    "secret_violation_action",
    "detached",
]


def _sandbox_class() -> ast.ClassDef:
    tree = ast.parse(STUB_PATH.read_text())
    for node in tree.body:
        if isinstance(node, ast.ClassDef) and node.name == "Sandbox":
            return node
    raise AssertionError("Sandbox missing from stub")


def _method(name: str) -> ast.FunctionDef | ast.AsyncFunctionDef:
    for node in _sandbox_class().body:
        if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == name:
            return node
    raise AssertionError(f"Sandbox.{name} missing from stub")


def test_create_methods_have_explicit_keyword_only_contracts() -> None:
    create = _method("create")
    connect_or_create = _method("connect_or_create")
    create_with_progress = _method("create_with_progress")

    assert isinstance(create, ast.AsyncFunctionDef)
    assert isinstance(connect_or_create, ast.AsyncFunctionDef)
    assert isinstance(create_with_progress, ast.FunctionDef)
    for method in (create, connect_or_create, create_with_progress):
        assert method.args.kwarg is None
        assert [arg.arg for arg in method.args.kwonlyargs] == EXPECTED_KWARGS
        assert all(default is not None for default in method.args.kw_defaults)


def test_create_closed_values_are_precisely_typed() -> None:
    create = _method("create")
    annotations = {
        arg.arg: ast.unparse(arg.annotation)
        for arg in create.args.kwonlyargs
        if arg.annotation is not None
    }

    assert annotations["security"] == "SecurityProfile | None"
    assert annotations["init"] == "str | InitConfig | InitOptions | None"
    assert annotations["pull_policy"] == "PullPolicy | None"
    assert annotations["log_level"] == "LogLevel | None"
    assert annotations["registry_auth"] == "RegistryAuth | None"
    assert annotations["registry_insecure"] == "bool"
    assert (
        annotations["registry_ca_certs"]
        == "list[bytes | bytearray | str | os.PathLike[str]] | None"
    )
    assert annotations["volumes"] == "Mapping[str, MountConfig] | None"
    assert annotations["patches"] == "Sequence[PatchConfig] | None"
    assert annotations["network"] == "Network | None"


def test_default_workload_methods_have_explicit_keyword_only_contracts() -> None:
    exec_default = _method("exec_default")
    exec_default_stream = _method("exec_default_stream")
    attach_default = _method("attach_default")

    for method in (exec_default, exec_default_stream, attach_default):
        assert isinstance(method, ast.AsyncFunctionDef)
        assert method.args.kwarg is None
        assert method.args.args[0].arg == "self"

    assert [arg.arg for arg in exec_default.args.kwonlyargs] == [
        "cwd",
        "user",
        "env",
        "timeout",
        "stdin",
        "tty",
        "rlimits",
    ]
    assert [arg.arg for arg in exec_default_stream.args.kwonlyargs] == [
        "cwd",
        "user",
        "env",
        "timeout",
        "stdin",
        "tty",
        "rlimits",
    ]
    assert [arg.arg for arg in attach_default.args.kwonlyargs] == [
        "cwd",
        "user",
        "env",
        "detach_keys",
    ]


def test_lifecycle_convergence_methods_are_typed() -> None:
    tree = ast.parse(STUB_PATH.read_text())
    classes = {
        node.name: node
        for node in tree.body
        if isinstance(node, ast.ClassDef) and node.name in {"Sandbox", "SandboxHandle"}
    }

    for class_name in ("Sandbox", "SandboxHandle"):
        methods = {
            node.name
            for node in classes[class_name].body
            if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
        }
        assert {"id", "wait_for_status", "restart", "destroy"} <= methods

    handle_methods = {
        node.name
        for node in classes["SandboxHandle"].body
        if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
    }
    assert "connect_or_start" in handle_methods


def test_restore_has_only_destination_options() -> None:
    restore = _method("restore")
    names = {arg.arg for arg in restore.args.kwonlyargs}
    assert {"name", "cow_memory", "forked", "disk_only", "snapshot_base", "volumes",
            "ports", "vsock", "allow_missing_resources"} <= names
    assert not names & {"image", "cmd", "replace", "detached", "from_snapshot", "network"}
    assert {"cpus", "memory", "network_policy", "max_connections", "disable_network",
            "security", "max_duration", "idle_timeout"} <= names
    assert names == {arg.arg for arg in _method("restore_with_progress").args.kwonlyargs}


def test_restore_accepts_backend_neutral_snapshot_objects() -> None:
    # Object seeds preserve an explicit remote ID instead of treating it as a host path.
    for name in ("restore", "restore_with_progress"):
        snapshot = _method(name).args.args[0]
        assert snapshot.arg == "snapshot"
        assert ast.unparse(snapshot.annotation) == (
            "Snapshot | SnapshotHandle | str | os.PathLike[str]"
        )


def test_restore_controls_preserve_optional_values_and_policy_type() -> None:
    for name in ("restore", "restore_with_progress"):
        method = _method(name)
        annotations = {arg.arg: ast.unparse(arg.annotation) for arg in method.args.kwonlyargs}
        defaults = dict(zip(
            [arg.arg for arg in method.args.kwonlyargs], method.args.kw_defaults, strict=True
        ))
        assert annotations["network_policy"] == "NetworkPolicy | None"
        assert annotations["security"] == "SecurityProfile | None"
        for option in ("cpus", "memory", "network_policy", "max_connections", "security",
                       "max_duration", "idle_timeout"):
            assert ast.literal_eval(defaults[option]) is None


def test_fork_methods_retain_branch_alias_signatures() -> None:
    classes = {
        node.name: node
        for node in ast.parse(STUB_PATH.read_text()).body
        if isinstance(node, ast.ClassDef)
    }
    for name in ("Sandbox", "SandboxHandle"):
        methods = {
            node.name: node
            for node in classes[name].body
            if isinstance(node, ast.AsyncFunctionDef)
        }
        for canonical, alias in (("fork", "branch"), ("fork_many", "branch_many")):
            assert ast.dump(methods[canonical].args) == ast.dump(methods[alias].args)
            assert ast.dump(methods[canonical].returns) == ast.dump(methods[alias].returns)
