"""Register and resolve discovered host adapter strategies."""
from __future__ import annotations

import inspect
from collections.abc import Iterable, Mapping
from pathlib import Path

from okstra_project.dirs import okstra_home

from ..domain.host import (
    HostAdapterContractError,
    HostClaim,
    HostDescriptor,
    HostNotRegistered,
    HostResolutionContext,
    HostUnavailable,
    ProviderNotRegistered,
)
from ..ports import HostAdapter
from .factory_loader import FactoryLoadError, load_relative_factory
from .host_discovery import HostManifest, discover_host_manifests
from .provider_registry import ProviderRegistry, default_provider_registry


_ADAPTER_METHODS = (
    "detect",
    "probe",
    "interaction",
    "host_model",
    "lead_session",
    "worker_dispatch",
    "usage_accounting",
)
_PORT_METHODS = {
    "interaction": ("plan",),
    "host_model": ("resolve",),
    "lead_session": ("build_start", "build_resume"),
    "worker_dispatch": ("build_plan",),
    "usage_accounting": ("collect",),
}


class _RegisteredHostAdapter:
    def __init__(self, descriptor: HostDescriptor, adapter: object) -> None:
        self.descriptor = descriptor
        self._adapter = adapter

    def detect(self, context):
        return self._adapter.detect(context)

    def probe(self, context):
        return self._adapter.probe(context)

    def interaction(self):
        return self._adapter.interaction()

    def host_model(self):
        return self._adapter.host_model()

    def lead_session(self):
        return self._adapter.lead_session()

    def worker_dispatch(self):
        return self._adapter.worker_dispatch()

    def usage_accounting(self):
        return self._adapter.usage_accounting()


class HostAdapterRegistry:
    def __init__(
        self,
        adapters: Mapping[str, HostAdapter],
        aliases: Mapping[str, HostAdapter],
    ) -> None:
        self._adapters = dict(adapters)
        self._aliases = dict(aliases)

    @classmethod
    def from_roots(
        cls,
        roots: Iterable[Path],
        provider_registry: ProviderRegistry,
    ) -> "HostAdapterRegistry":
        entries = _load_entries(roots, provider_registry)
        adapters = _index_host_ids(entries)
        aliases = _index_aliases(entries, adapters)
        return cls(adapters, aliases)

    def resolve(self, host_id_or_alias: str) -> HostAdapter:
        normalized = host_id_or_alias.strip().lower()
        adapter = self._aliases.get(normalized)
        if adapter is not None:
            return adapter
        allowed = ", ".join(self._adapters)
        raise HostNotRegistered(
            f"unknown host {host_id_or_alias!r}. Allowed values: {allowed}"
        )

    def resolve_request(
        self,
        requested_host: str,
        context: HostResolutionContext,
    ) -> HostDescriptor:
        descriptor, _, _ = self.resolve_request_details(requested_host, context)
        return descriptor

    def resolve_request_details(
        self,
        requested_host: str,
        context: HostResolutionContext,
    ) -> tuple[HostDescriptor, str, str]:
        normalized = requested_host.strip().lower()
        if normalized and normalized != "auto":
            descriptor = self.resolve(normalized).descriptor
            return descriptor, "explicit", (
                f"Explicit runtime {descriptor.id!r} requested."
            )
        env_host = context.explicit_env_host.strip()
        if env_host and env_host.lower() != "auto":
            descriptor = self.resolve(env_host).descriptor
            return descriptor, descriptor.id, (
                f"OKSTRA_RUNTIME_HOST={env_host} selected runtime {descriptor.id!r}."
            )
        adapter, claim = self._resolve_claim(context)
        return adapter.descriptor, adapter.descriptor.id, claim.reason

    def ids(self) -> tuple[str, ...]:
        return tuple(self._adapters)

    def catalog(self) -> tuple[HostDescriptor, ...]:
        return tuple(adapter.descriptor for adapter in self._adapters.values())

    def _resolve_claim(
        self, context: HostResolutionContext
    ) -> tuple[HostAdapter, HostClaim]:
        claims = tuple(
            claim
            for adapter in self._adapters.values()
            if (claim := _detect_claim(adapter, context)) is not None
        )
        if not claims:
            raise HostUnavailable("no host adapter claimed the current session")
        highest_priority = max(claim.priority for claim in claims)
        winners = tuple(
            claim for claim in claims if claim.priority == highest_priority
        )
        if len(winners) != 1:
            host_ids = ", ".join(sorted(claim.host_id for claim in winners))
            raise HostUnavailable(f"ambiguous host adapter claims: {host_ids}")
        winner = winners[0]
        return self.resolve(winner.host_id), winner


def default_host_registry(
    provider_registry: ProviderRegistry | None = None,
) -> HostAdapterRegistry:
    bundled_root = Path(__file__).resolve().parents[1] / "adapters" / "hosts"
    user_root = okstra_home() / "adapters" / "hosts"
    providers = provider_registry or default_provider_registry()
    return HostAdapterRegistry.from_roots((bundled_root, user_root), providers)


def _load_entries(
    roots: Iterable[Path], provider_registry: ProviderRegistry
) -> tuple[tuple[HostManifest, HostAdapter], ...]:
    entries: list[tuple[HostManifest, HostAdapter]] = []
    for manifest in discover_host_manifests(roots):
        _validate_provider(manifest, provider_registry)
        entries.append((manifest, _load_host_adapter(manifest, provider_registry)))
    return tuple(entries)


def _validate_provider(
    manifest: HostManifest, provider_registry: ProviderRegistry
) -> None:
    if not manifest.native_provider_id:
        return
    if manifest.native_provider_id not in provider_registry.providers:
        raise ProviderNotRegistered(
            f"native provider {manifest.native_provider_id!r} is not registered"
        )


def _load_host_adapter(
    manifest: HostManifest,
    provider_registry: ProviderRegistry,
) -> HostAdapter:
    try:
        factory = load_relative_factory(manifest.path, manifest.factory_ref)
        parameters = inspect.signature(factory).parameters.values()
        accepts_registry = any(
            parameter.name == "provider_registry"
            or parameter.kind is inspect.Parameter.VAR_KEYWORD
            for parameter in parameters
        )
        adapter = (
            factory(provider_registry=provider_registry)
            if accepts_registry
            else factory()
        )
    except FactoryLoadError as exc:
        raise HostAdapterContractError(str(exc)) from exc
    except Exception as exc:
        raise HostAdapterContractError(
            f"host factory failed for {manifest.host_id!r}"
        ) from exc
    descriptor = _validated_descriptor(adapter, manifest)
    return _RegisteredHostAdapter(descriptor, adapter)


def _validated_descriptor(adapter: object, manifest: HostManifest) -> HostDescriptor:
    for method_name in _ADAPTER_METHODS:
        if not callable(getattr(adapter, method_name, None)):
            raise HostAdapterContractError(
                f"host factory returned invalid adapter {manifest.host_id!r}"
            )
    _validate_ports(adapter, manifest.host_id)
    descriptor = _copy_descriptor(getattr(adapter, "descriptor", None), manifest.host_id)
    if descriptor.id != manifest.host_id:
        raise HostAdapterContractError(f"host id mismatch for {manifest.host_id!r}")
    if descriptor.native_provider_id != manifest.native_provider_id:
        raise HostAdapterContractError(
            f"native provider mismatch for {manifest.host_id!r}"
        )
    if descriptor.required_executables != manifest.required_executables:
        raise HostAdapterContractError(
            f"required executables mismatch for {manifest.host_id!r}"
        )
    if Path(descriptor.relay_contract).resolve() != manifest.relay_contract:
        raise HostAdapterContractError(
            f"relay contract mismatch for {manifest.host_id!r}"
        )
    return descriptor


def _validate_ports(adapter: object, host_id: str) -> None:
    for accessor, required_methods in _PORT_METHODS.items():
        try:
            port = getattr(adapter, accessor)()
        except Exception as exc:
            raise HostAdapterContractError(
                f"host {host_id!r} {accessor} port could not be created"
            ) from exc
        missing = tuple(
            method for method in required_methods
            if not callable(getattr(port, method, None))
        )
        if missing:
            required = ", ".join(missing)
            raise HostAdapterContractError(
                f"host {host_id!r} {accessor} port is missing {required}"
            )


def _copy_descriptor(raw: object, host_id: str) -> HostDescriptor:
    try:
        descriptor = HostDescriptor(
            id=raw.id,
            aliases=raw.aliases,
            native_provider_id=raw.native_provider_id,
            required_executables=raw.required_executables,
            launch_mode=raw.launch_mode,
            install_targets=raw.install_targets,
            agent_id=raw.agent_id,
            agent_label=raw.agent_label,
            role=raw.role,
            dispatch_mode=raw.dispatch_mode,
            session_accounting=raw.session_accounting,
            has_claude_session=raw.has_claude_session,
            relay_contract=raw.relay_contract,
        )
    except (AttributeError, TypeError) as exc:
        raise HostAdapterContractError(
            f"host factory returned invalid descriptor {host_id!r}"
        ) from exc
    if not _descriptor_fields_are_valid(descriptor):
        raise HostAdapterContractError(
            f"host factory returned invalid descriptor {host_id!r}"
        )
    return descriptor


def _descriptor_fields_are_valid(descriptor: HostDescriptor) -> bool:
    normalized_ids = (descriptor.id, descriptor.native_provider_id)
    if not all(isinstance(value, str) for value in normalized_ids):
        return False
    if not descriptor.id or descriptor.id != descriptor.id.strip().lower():
        return False
    if descriptor.native_provider_id != descriptor.native_provider_id.strip().lower():
        return False
    if not isinstance(descriptor.aliases, tuple) or not all(
        isinstance(alias, str) and alias == alias.strip().lower() and alias
        for alias in descriptor.aliases
    ):
        return False
    if not isinstance(descriptor.required_executables, tuple) or not all(
        isinstance(executable, str) and executable
        for executable in descriptor.required_executables
    ):
        return False
    if descriptor.launch_mode not in {"lead", "team"}:
        return False
    if not isinstance(descriptor.install_targets, frozenset) or not all(
        isinstance(target, str) and target for target in descriptor.install_targets
    ):
        return False
    text_fields = (
        descriptor.agent_id,
        descriptor.agent_label,
        descriptor.role,
        descriptor.dispatch_mode,
        descriptor.session_accounting,
        descriptor.relay_contract,
    )
    return all(isinstance(value, str) and value for value in text_fields) and isinstance(
        descriptor.has_claude_session, bool
    )


def _index_host_ids(
    entries: tuple[tuple[HostManifest, HostAdapter], ...]
) -> dict[str, HostAdapter]:
    adapters: dict[str, HostAdapter] = {}
    for manifest, adapter in entries:
        if manifest.host_id in adapters:
            raise HostAdapterContractError(
                f"duplicate host id {manifest.host_id!r}"
            )
        adapters[manifest.host_id] = adapter
    return adapters


def _index_aliases(
    entries: tuple[tuple[HostManifest, HostAdapter], ...],
    adapters: Mapping[str, HostAdapter],
) -> dict[str, HostAdapter]:
    aliases = dict(adapters)
    for _, adapter in entries:
        for raw_alias in adapter.descriptor.aliases:
            alias = raw_alias.strip().lower()
            existing = aliases.get(alias)
            if not alias or (existing is not None and existing is not adapter):
                raise HostAdapterContractError(f"duplicate host key {alias!r}")
            aliases[alias] = adapter
    return aliases


def _detect_claim(
    adapter: HostAdapter, context: HostResolutionContext
) -> HostClaim | None:
    try:
        raw_claim = adapter.detect(context)
    except Exception as exc:
        raise HostAdapterContractError(
            f"host detector failed for {adapter.descriptor.id!r}"
        ) from exc
    if raw_claim is None:
        return None
    try:
        claim = HostClaim(raw_claim.host_id, raw_claim.priority, raw_claim.reason)
    except (AttributeError, TypeError) as exc:
        raise HostAdapterContractError(
            f"host detector returned invalid claim for {adapter.descriptor.id!r}"
        ) from exc
    if claim.host_id != adapter.descriptor.id or not isinstance(claim.priority, int):
        raise HostAdapterContractError(
            f"host detector returned invalid claim for {adapter.descriptor.id!r}"
        )
    return claim
