"""Capability-based strategy shared by bundled host adapters."""
from __future__ import annotations

import json
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from pathlib import Path

from okstra_ctl.adapters.dispatch import default_worker_dispatch_port
from okstra_ctl.domain.host import (
    HostClaim,
    HostDescriptor,
    HostReadiness,
    HostResolutionContext,
    HostSessionContext,
    LeadResumeRequest,
    LeadSessionPlan,
    LeadStartRequest,
    ProviderUnavailable,
)
from okstra_ctl.domain.provider import LeadLaunchSpec
from okstra_ctl.domain.wizard.interaction import (
    AnswerProtocol,
    InteractionPlan,
    WizardPrompt,
)
from okstra_ctl.ports import (
    HostModelBindingPort,
    InteractionPort,
    LeadSessionPort,
    UsageAccountingPort,
    WorkerDispatchPort,
)
from okstra_ctl.ports.host_model import FailClosedHostModelBindingPort
from okstra_ctl.registry.provider_registry import (
    ProviderRegistry,
    default_provider_registry,
)


@dataclass(frozen=True)
class HostPorts:
    interaction: InteractionPort
    host_model: HostModelBindingPort
    lead_session: LeadSessionPort
    worker_dispatch: WorkerDispatchPort
    usage_accounting: UsageAccountingPort


class PendingHostPort:
    """Placeholder replaced as each concrete port migrates behind the adapter."""


PENDING_HOST_PORT = PendingHostPort()
PLAIN_TEXT_FUNCTIONS = frozenset({"plain_text_input"})
INTERACTION_FUNCTIONS = PLAIN_TEXT_FUNCTIONS | frozenset({
    "plain_text_input",
    "native_single_select",
    "native_multi_select",
    "native_question_group",
})
_RELAY_HEADING = "## Wizard interaction relay"


@dataclass(frozen=True)
class NativePickerLimits:
    """How many questions and options a host's native picker can show.

    Counts that miss the range keep every option and fall back to numbered
    text. The defaults are Claude Code's AskUserQuestion shape; a host that
    can show more must declare that in its relay `nativeLimits`.
    """

    min_options: int = 2
    max_options: int = 4
    max_questions: int = 4

    def option_count_fits(self, count: int) -> bool:
        return self.min_options <= count <= self.max_options


def wizard_relay_contract(path: Path) -> dict[str, object]:
    """The JSON object under `## Wizard interaction relay` in a host relay."""
    body = path.read_text(encoding="utf-8")
    section = body.split(_RELAY_HEADING, 1)[1]
    encoded = section.split("```json\n", 1)[1].split("\n```", 1)[0]
    payload = json.loads(encoded)
    if not isinstance(payload, dict):
        raise ValueError(f"wizard relay in {path} is not a JSON object")
    return payload


def native_limits_from_relay(contract: Mapping[str, object]) -> NativePickerLimits:
    raw = contract.get("nativeLimits")
    if not isinstance(raw, Mapping):
        return NativePickerLimits()
    return NativePickerLimits(
        min_options=int(raw.get("minOptions", 2)),
        max_options=int(raw.get("maxOptions", 4)),
        max_questions=int(raw.get("maxQuestions", 4)),
    )


class CapabilityInteractionPort:
    def __init__(self, *, limits: NativePickerLimits | None = None) -> None:
        self._limits = limits or NativePickerLimits()

    def plan(
        self,
        prompt: WizardPrompt,
        context: HostSessionContext,
    ) -> InteractionPlan:
        functions = context.available_functions & INTERACTION_FUNCTIONS
        if prompt.kind == "pick_group":
            if (
                "native_question_group" in functions
                and self._native_group_fits(prompt, functions)
            ):
                return InteractionPlan("native-group", AnswerProtocol("group-json"))
            return InteractionPlan(
                "sequential-group",
                AnswerProtocol("numbered-group-json"),
            )
        if prompt.kind != "pick":
            return InteractionPlan("plain-text", AnswerProtocol("exact-value"))
        native_options_fit = self._native_options_fit(prompt)
        if (
            prompt.multi
            and "native_multi_select" in functions
            and native_options_fit
        ):
            return InteractionPlan(
                "native-multi",
                AnswerProtocol("exact-value", multi=True),
            )
        if (
            not prompt.multi
            and "native_single_select" in functions
            and native_options_fit
        ):
            return InteractionPlan("native-single", AnswerProtocol("exact-value"))
        kind = "numbered-multi" if prompt.multi else "numbered-single"
        return InteractionPlan(kind, AnswerProtocol("numbered", multi=prompt.multi))

    def _native_options_fit(self, prompt: WizardPrompt) -> bool:
        labels = tuple(option.label for option in prompt.options)
        return (
            self._limits.option_count_fits(len(labels))
            and len(set(labels)) == len(labels)
            and (not prompt.multi or all(", " not in label for label in labels))
        )

    def _native_group_fits(
        self, prompt: WizardPrompt, functions: frozenset[str]
    ) -> bool:
        if any(question.multi for question in prompt.questions) and (
            "native_multi_select" not in functions
        ):
            return False
        return (
            1 <= len(prompt.questions) <= self._limits.max_questions
            and all(self._native_options_fit(question) for question in prompt.questions)
        )


def numbered_interaction_port() -> CapabilityInteractionPort:
    return CapabilityInteractionPort()


def relay_interaction_port(relay_path: str | Path) -> CapabilityInteractionPort:
    contract = wizard_relay_contract(Path(relay_path))
    return CapabilityInteractionPort(limits=native_limits_from_relay(contract))


class ProviderLeadSessionPort:
    def __init__(
        self,
        *,
        host_id: str,
        native_provider_id: str,
        provider_registry: ProviderRegistry,
    ) -> None:
        self._host_id = host_id
        self._native_provider_id = native_provider_id
        self._provider_registry = provider_registry

    def build_start(
        self,
        request: LeadStartRequest,
        context: HostSessionContext,
    ) -> LeadSessionPlan:
        if context.entry_mode == "current-session":
            return self._current_session_plan()
        launch = self._lead_launch()
        argv = (launch.executable, *launch.sandbox_waiver, *request.argv)
        return self._spawn_plan(argv, launch.sandbox_waiver_note)

    def build_resume(
        self,
        request: LeadResumeRequest,
        context: HostSessionContext,
    ) -> LeadSessionPlan:
        if context.entry_mode == "current-session":
            return self._current_session_plan()
        launch = self._lead_launch()
        session_args = (
            (launch.resume_session_id_flag, request.session_id)
            if launch.resume_session_id_flag and request.session_id
            else ()
        )
        argv = (
            launch.executable,
            *launch.sandbox_waiver,
            *session_args,
            *request.argv,
        )
        return self._spawn_plan(argv, launch.sandbox_waiver_note)

    def _lead_launch(self) -> LeadLaunchSpec:
        if not self._native_provider_id:
            raise ProviderUnavailable(
                f"host {self._host_id!r} has no native lead provider"
            )
        provider = self._provider_registry.resolve(self._native_provider_id)
        if provider.lead_launch is None:
            raise ProviderUnavailable(
                f"provider {self._native_provider_id!r} cannot lead"
            )
        return provider.lead_launch

    def _current_session_plan(self) -> LeadSessionPlan:
        return LeadSessionPlan("current-session", (), {}, "")

    def _spawn_plan(
        self, argv: tuple[str, ...], sandbox_note: str
    ) -> LeadSessionPlan:
        return LeadSessionPlan(
            "spawn-process",
            argv,
            {"OKSTRA_RUNTIME_HOST": self._host_id},
            sandbox_note,
        )


class CapabilityHostAdapter:
    def __init__(
        self,
        descriptor: HostDescriptor,
        *,
        executable_finder: Callable[[str], str | None],
        interaction_port: InteractionPort,
        lead_session_port: LeadSessionPort,
        worker_dispatch_port: WorkerDispatchPort,
        usage_accounting_port: UsageAccountingPort,
        supported_functions: frozenset[str],
        detector: Callable[[HostResolutionContext], HostClaim | None],
        provider_registry: ProviderRegistry | None = None,
        readiness_probe: Callable[
            [HostSessionContext], tuple[Mapping[str, object], ...]
        ] | None = None,
        host_model_port: HostModelBindingPort | None = None,
    ) -> None:
        self.descriptor = descriptor
        self._executable_finder = executable_finder
        resolved_provider_registry = provider_registry or default_provider_registry()
        resolved_lead_session = lead_session_port
        if lead_session_port is PENDING_HOST_PORT:
            resolved_lead_session = ProviderLeadSessionPort(
                host_id=descriptor.id,
                native_provider_id=descriptor.native_provider_id,
                provider_registry=resolved_provider_registry,
            )
        resolved_worker_dispatch = worker_dispatch_port
        if worker_dispatch_port is PENDING_HOST_PORT:
            resolved_worker_dispatch = default_worker_dispatch_port(
                descriptor,
                resolved_provider_registry,
            )
        self._ports = HostPorts(
            interaction_port,
            host_model_port or FailClosedHostModelBindingPort(descriptor.id),
            resolved_lead_session,
            resolved_worker_dispatch,
            usage_accounting_port,
        )
        self._supported_functions = supported_functions
        self._detector = detector
        self._readiness_probe = readiness_probe or (lambda context: ())

    def detect(self, context: HostResolutionContext) -> HostClaim | None:
        return self._detector(context)

    def probe(self, context: HostSessionContext) -> HostReadiness:
        effective_functions = self._effective_functions(context)
        if context.host_id != self.descriptor.id:
            reason = (
                f"host mismatch: expected {self.descriptor.id!r}, "
                f"received {context.host_id!r}"
            )
            return HostReadiness(False, effective_functions, (), (), reason)
        if context.entry_mode == "current-session":
            return self._readiness(context, effective_functions)
        missing = tuple(
            executable
            for executable in self.descriptor.required_executables
            if self._executable_finder(executable) is None
        )
        if missing:
            reason = f"missing required executables: {', '.join(missing)}"
            return HostReadiness(False, effective_functions, missing, (), reason)
        return self._readiness(context, effective_functions)

    def interaction(self) -> InteractionPort:
        return self._ports.interaction

    def host_model(self) -> HostModelBindingPort:
        return self._ports.host_model

    def lead_session(self) -> LeadSessionPort:
        return self._ports.lead_session

    def worker_dispatch(self) -> WorkerDispatchPort:
        return self._ports.worker_dispatch

    def usage_accounting(self) -> UsageAccountingPort:
        return self._ports.usage_accounting

    def _effective_functions(
        self, context: HostSessionContext
    ) -> frozenset[str]:
        return self._supported_functions & context.available_functions

    def _readiness(
        self,
        context: HostSessionContext,
        effective_functions: frozenset[str],
    ) -> HostReadiness:
        checks = self._readiness_probe(context)
        blockers = tuple(
            check for check in checks
            if check.get("status") in {"required", "unavailable"}
        )
        if blockers:
            return HostReadiness(
                False, effective_functions, (), checks, "host not ready"
            )
        return HostReadiness(True, effective_functions, (), checks, "ready")


def no_automatic_claim(context: HostResolutionContext) -> HostClaim | None:
    """Require an explicit host instead of inferring one from installed CLIs."""
    return None
