"""JSON entrypoint for host catalog, resolution, and launch plans."""
from __future__ import annotations

import argparse
import json
import os
import sys
from contextlib import contextmanager
from dataclasses import asdict
from pathlib import Path

from ..agent_invocation import (
    AgentInvocationError,
    AgentModelAssignment,
    agent_model_assignment_from_payload,
    verify_agent_invocation,
)
from ..dispatch_state import DispatchError, record_verified_agent_dispatch
from ..json_boundary import JsonBoundaryError, load_owned_object

from ..application.start_run import start_run
from ..application.resume_run import resume_run
from ..domain.host import (
    HostAdapterContractError,
    HostCapabilityMismatch,
    HostNotRegistered,
    HostResolutionContext,
    HostSessionContext,
    HostUnavailable,
    LeadStartRequest,
    LeadResumeRequest,
    ProviderNotRegistered,
    ProviderUnavailable,
)
from ..registry.host_registry import HostAdapterRegistry, default_host_registry
from ..models import lead_launch_argv, lead_launch_spec


_ENTRYPOINT_ERRORS = (
    HostAdapterContractError,
    HostCapabilityMismatch,
    HostNotRegistered,
    HostUnavailable,
    ProviderNotRegistered,
    ProviderUnavailable,
    DispatchError,
    ValueError,
)


def catalog_payload(registry: HostAdapterRegistry) -> dict[str, object]:
    return {
        "schemaVersion": 1,
        "hosts": [asdict(descriptor) for descriptor in registry.catalog()],
    }


def resolve_payload(
    registry: HostAdapterRegistry,
    requested_host: str,
    context: HostResolutionContext,
) -> dict[str, object]:
    descriptor, host, reason = registry.resolve_request_details(
        requested_host, context
    )
    return {
        "ok": True,
        "requestedRuntime": requested_host.strip().lower() or "auto",
        "resolvedRuntime": descriptor.id,
        "launchMode": descriptor.launch_mode,
        "installTargets": sorted(descriptor.install_targets),
        "host": host,
        "confidence": "high",
        "reason": reason,
        "fallbackFrom": None,
        "warnings": [],
    }


def main(
    argv: list[str] | None = None,
    *,
    registry: HostAdapterRegistry | None = None,
) -> int:
    parser = _build_parser()
    args = parser.parse_args(argv)
    try:
        active_registry = registry or default_host_registry()
        payload = _command_payload(args, active_registry)
    except _ENTRYPOINT_ERRORS as exc:
        sys.stderr.write(f"{exc}\n")
        return 2
    if getattr(args, "output_format", "json") == "null":
        _write_null_launch_plan(payload)
    else:
        public_payload = {
            key: value for key, value in payload.items()
            if not key.startswith("_")
        }
        print(json.dumps(public_payload, default=_json_default, sort_keys=True))
    return 0


def _build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(prog="okstra-hosts")
    commands = parser.add_subparsers(dest="command", required=True)
    catalog = commands.add_parser("catalog")
    catalog.add_argument("--json", action="store_true")
    resolve = commands.add_parser("resolve")
    resolve.add_argument("--requested", required=True)
    resolve.add_argument("--command", dest="runtime_command", required=True)
    resolve.add_argument("--environment-host", default="")
    resolve.add_argument(
        "--capability",
        action="append",
        choices=("claude-skill-handoff", "tmux"),
        default=[],
    )
    launch = commands.add_parser("launch-plan")
    _add_session_plan_args(launch)
    launch.add_argument("--prepared-launch-json", default="")
    resume = commands.add_parser("resume-plan")
    _add_session_plan_args(resume)
    resume.add_argument("--session-id", required=True)
    resume.add_argument("--task-type", default="")
    resume.add_argument("--request-arg", action="append", default=[])
    probe = commands.add_parser("probe")
    probe.add_argument("--host", required=True)
    probe.add_argument(
        "--entry-mode",
        required=True,
        choices=("spawn-process", "current-session"),
    )
    probe.add_argument("--project-root", required=True)
    probe.add_argument("--home-dir", required=True)
    probe.add_argument("--available-function", action="append", default=[])
    return parser


def _add_session_plan_args(parser: argparse.ArgumentParser) -> None:
    parser.add_argument("--host", required=True)
    parser.add_argument(
        "--entry-mode",
        required=True,
        choices=("spawn-process", "current-session"),
    )
    parser.add_argument(
        "--format",
        dest="output_format",
        choices=("json", "null"),
        default="json",
    )
    parser.add_argument("--project-root", default=str(Path.cwd()))


def _command_payload(
    args: argparse.Namespace, registry: HostAdapterRegistry
) -> dict[str, object]:
    if args.command == "catalog":
        return catalog_payload(registry)
    if args.command == "resolve":
        context = HostResolutionContext(
            args.environment_host, frozenset(args.capability)
        )
        try:
            return resolve_payload(registry, args.requested, context)
        except HostUnavailable as exc:
            raise HostUnavailable(
                _unclaimed_host_message(registry, args.runtime_command)
            ) from exc
        except HostNotRegistered as exc:
            if args.requested.strip().lower() in {"", "auto"} and context.explicit_env_host:
                allowed = "|".join(registry.ids())
                raise HostNotRegistered(
                    f"invalid OKSTRA_RUNTIME_HOST {context.explicit_env_host!r} "
                    f"(expected: {allowed})."
                ) from exc
            raise
    if args.command == "probe":
        return _probe_payload(registry, args)
    if args.command == "resume-plan":
        return _resume_plan_payload(registry, args)
    return _launch_plan_payload(registry, args)


def _probe_payload(
    registry: HostAdapterRegistry,
    args: argparse.Namespace,
) -> dict[str, object]:
    adapter = registry.resolve(args.host)
    context = HostSessionContext(
        host_id=adapter.descriptor.id,
        entry_mode=args.entry_mode,
        available_functions=frozenset(args.available_function),
        interaction_surface="preflight",
    )
    with _probe_paths(args.project_root, args.home_dir):
        readiness = adapter.probe(context)
    statuses = {check.get("status") for check in readiness.checks}
    status = "unknown" if "unavailable" in statuses else (
        "ready" if readiness.ok else "blocked"
    )
    return {
        "runtime": adapter.descriptor.id,
        "status": status,
        "checks": list(readiness.checks),
        "relayContract": adapter.descriptor.relay_contract,
    }


@contextmanager
def _probe_paths(project_root: str, home_dir: str):
    keys = ("OKSTRA_PROBE_PROJECT_ROOT", "OKSTRA_PROBE_HOME_DIR")
    previous = {key: os.environ.get(key) for key in keys}
    os.environ[keys[0]] = project_root
    os.environ[keys[1]] = home_dir
    try:
        yield
    finally:
        for key, value in previous.items():
            if value is None:
                os.environ.pop(key, None)
            else:
                os.environ[key] = value


def _unclaimed_host_message(registry: HostAdapterRegistry, command: str) -> str:
    runtime_flag = "--runtime" if command in {
        "install", "ensure-installed", "doctor"
    } else "--lead-runtime"
    allowed = "|".join(registry.ids())
    return (
        f"host unknown: set OKSTRA_RUNTIME_HOST={allowed} or pass "
        f"{runtime_flag} explicitly."
    )


def _launch_plan_payload(
    registry: HostAdapterRegistry,
    args: argparse.Namespace,
) -> dict[str, object]:
    adapter = registry.resolve(args.host)
    context = HostSessionContext(
        host_id=adapter.descriptor.id,
        entry_mode=args.entry_mode,
        available_functions=frozenset(),
        interaction_surface="terminal",
    )
    project_root, request_argv, dispatch_authority = _launch_request(
        args, adapter.descriptor.id
    )
    request = LeadStartRequest(project_root, "", request_argv)
    plan = start_run(request, context, adapter)
    if dispatch_authority is not None:
        run_manifest_path, metadata_path, _assignment = dispatch_authority
        record_verified_agent_dispatch(
            project_root=project_root,
            run_manifest_path=run_manifest_path,
            metadata_path=metadata_path,
            enforcement_mode=(
                "core-pre-dispatch"
                if plan.mode == "spawn-process"
                else "host-native-spec-link-gate"
            ),
            allow_native_core=plan.mode == "spawn-process",
        )
    return {"ok": True, "_project_root": str(project_root), **asdict(plan)}


def _resume_plan_payload(
    registry: HostAdapterRegistry,
    args: argparse.Namespace,
) -> dict[str, object]:
    adapter = registry.resolve(args.host)
    context = HostSessionContext(
        host_id=adapter.descriptor.id,
        entry_mode=args.entry_mode,
        available_functions=frozenset(),
        interaction_surface="terminal",
    )
    project_root = Path(args.project_root)
    request = LeadResumeRequest(
        project_root,
        args.task_type,
        args.session_id,
        tuple(args.request_arg),
    )
    plan = resume_run(request, context, adapter)
    return {"ok": True, "_project_root": str(project_root), **asdict(plan)}


def _launch_request(
    args: argparse.Namespace,
    canonical_host_id: str,
) -> tuple[
    Path,
    tuple[str, ...],
    tuple[Path, Path, AgentModelAssignment] | None,
]:
    if not args.prepared_launch_json:
        return Path(args.project_root), (), None
    try:
        prepared = json.loads(args.prepared_launch_json)
    except json.JSONDecodeError as exc:
        raise ValueError("invalid prepared launch JSON") from exc
    if not isinstance(prepared, dict):
        raise ValueError("prepared launch JSON must be an object")
    if prepared.get("leadRuntime") != canonical_host_id:
        raise ValueError("prepared launch runtime does not match selected host")
    project_root = prepared.get("projectRoot")
    if not isinstance(project_root, str) or not project_root:
        raise ValueError("prepared launch is missing projectRoot")
    root = Path(project_root)
    run_manifest_value = prepared.get("runManifestPath")
    metadata_value = prepared.get("leadPromptMetadataPath")
    if not isinstance(run_manifest_value, str) or not run_manifest_value:
        raise ValueError("prepared launch is missing runManifestPath")
    if not isinstance(metadata_value, str) or not metadata_value:
        raise ValueError("prepared launch is missing leadPromptMetadataPath")
    run_manifest_path = Path(run_manifest_value)
    metadata_path = Path(metadata_value)
    try:
        manifest = load_owned_object(run_manifest_path, artifact="run manifest")
        resources = manifest["resources"]
        assignment = agent_model_assignment_from_payload(
            manifest["invocationAssignments"]["lead"]
        )
        prompt_path = root / resources["leadExecutionPromptPath"]
        canonical_metadata = root / resources["leadPromptMetadataPath"]
        metadata = load_owned_object(
            canonical_metadata, artifact="lead prompt metadata"
        )
        if metadata["prompt"]["path"] != resources["leadExecutionPromptPath"]:
            raise KeyError("lead prompt resource does not match metadata")
    except (OSError, JsonBoundaryError, KeyError, TypeError, AgentInvocationError) as exc:
        raise ValueError("lead invocation authority is incomplete") from exc
    if metadata_path.resolve(strict=False) != canonical_metadata.resolve(strict=False):
        raise ValueError("prepared lead metadata does not match run manifest")
    if prepared.get("leadProvider") != assignment.provider:
        raise ValueError("prepared lead provider does not match run manifest")
    if prepared.get("leadModelExecutionValue") != assignment.model_execution_value:
        raise ValueError("prepared lead model does not match run manifest")
    errors = verify_agent_invocation(
        metadata_path,
        project_root=root,
        expected_run_manifest_path=run_manifest_path,
        expected_assignment=assignment,
        expected_worker_id="lead",
        expected_assignment_ref="lead",
        expected_audience="lead",
    )
    if errors:
        raise ValueError("lead invocation verification failed: " + "; ".join(errors))
    session_id = prepared.get("leadSessionId")
    if not isinstance(session_id, str):
        raise ValueError("prepared launch is missing leadSessionId")
    argv = lead_launch_argv(
        assignment.provider,
        model=assignment.model_execution_value,
        session_id=session_id,
        prompt=prompt_path.read_text(encoding="utf-8"),
    )
    launch = lead_launch_spec(assignment.provider)
    prefix_size = 1 + len(launch.sandbox_waiver)
    return (
        root,
        tuple(argv[prefix_size:]),
        (run_manifest_path, canonical_metadata, assignment),
    )


def _write_null_launch_plan(payload: dict[str, object]) -> None:
    fields = [
        payload["_project_root"],
        payload["sandbox_note"],
        str(len(payload["env"])),
    ]
    fields.extend(
        f"{key}={value}" for key, value in payload["env"].items()
    )
    fields.extend(payload["argv"])
    for field in fields:
        value = str(field)
        if "\0" in value:
            raise ValueError("launch plan field contains a null byte")
        sys.stdout.write(value + "\0")


def _json_default(value: object) -> object:
    if isinstance(value, (frozenset, set)):
        return sorted(value)
    if isinstance(value, Path):
        return str(value)
    raise TypeError(f"not JSON serializable: {type(value).__name__}")


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