"""Unified, validated model lookup over registered provider adapters."""
from __future__ import annotations

import re
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING

from .domain.provider import ModelSpec, UnknownModelError
from .domain.role import ROLE_DEFINITIONS, RoleCatalogError, normalize_role
from .domain.worker_exec import WorkerWriteCapability

if TYPE_CHECKING:
    from .domain.provider import ProviderSpec
    from .registry.provider_registry import ProviderRegistry


_BUNDLED_PROVIDER_ORDER = ("claude", "antigravity", "codex", "grok", "kimi")
_TOKEN = re.compile(r"[a-z0-9][a-z0-9._-]*\Z")
_VERSION_KINDS = frozenset({"pinned", "channel"})


class ModelRef(str):
    """A canonical `<provider>/<model>` reference."""

    def __new__(cls, provider_id: str, model_id: str) -> "ModelRef":
        provider = provider_id.strip().lower()
        model = model_id.strip().lower()
        if not provider or not model:
            raise ValueError("model ref requires a provider and model id")
        value = str.__new__(cls, f"{provider}/{model}")
        value.provider_id = provider
        value.model_id = model
        return value

    @classmethod
    def parse(cls, raw: str) -> "ModelRef":
        provider, separator, model = raw.strip().partition("/")
        if not separator:
            raise ValueError(f"invalid model ref: {raw!r}")
        return cls(provider, model)


@dataclass(frozen=True)
class ResolvedModel:
    model_ref: ModelRef
    provider_id: str
    display_name: str
    execution_value: str
    aliases: tuple[str, ...]
    version_kind: str
    channel_family: str | None
    selectable: bool
    pricing: tuple[float, float, float] | None
    supported_roles: frozenset[str] | None


@dataclass(frozen=True)
class ModelAvailability:
    available: bool
    reason: str


@dataclass(frozen=True)
class ProviderRuntimeFacts:
    """Immutable provider facts required by assignment and rendering."""

    provider: str
    display_label: str
    wrapper: str
    supported_roles: frozenset[str]
    cli_worker_write_capability: WorkerWriteCapability | None
    execution_normalization_required: bool

    def supports_role(self, role: str) -> bool:
        try:
            return normalize_role(role) in self.supported_roles
        except RoleCatalogError:
            return False


@dataclass(frozen=True)
class ModelPoolValidationFailure:
    code: str
    message: str


class ModelPoolValidationError(ValueError):
    """Raised with every invalid adapter-catalog contract row."""

    def __init__(self, failures: tuple[ModelPoolValidationFailure, ...]) -> None:
        self.failures = failures
        super().__init__("; ".join(f"{row.code}: {row.message}" for row in failures))


class ModelPool:
    """One model-query surface composed from provider-owned catalog entries."""

    def __init__(self, registry: "ProviderRegistry") -> None:
        self._registry = registry
        self._providers = self._ordered_providers()
        self._provider_defaults = self._snapshot_provider_defaults()
        self._provider_facts = self._snapshot_provider_facts()
        self._models, self._aliases = self._index_models()

    @classmethod
    def from_registry(cls, registry: "ProviderRegistry") -> "ModelPool":
        pool = cls(registry)
        pool.validate()
        return pool

    def list(self, *, role: str | None = None) -> tuple[ResolvedModel, ...]:
        canonical_role = self._normalize_role(role) if role else None
        return tuple(
            model
            for model in self._models.values()
            if model.selectable and self._supports_role(model, canonical_role)
        )

    def list_catalog(self, *, role: str | None = None) -> tuple[ResolvedModel, ...]:
        canonical_role = self._normalize_role(role) if role else None
        return tuple(
            model
            for model in self._models.values()
            if self._supports_role(model, canonical_role)
        )

    def resolve(self, model_ref: str) -> ResolvedModel:
        try:
            return self._models[ModelRef.parse(model_ref)]
        except (KeyError, ValueError) as exc:
            raise UnknownModelError(f"unknown model ref {model_ref!r}") from exc

    def resolve_alias(self, provider: str, alias: str) -> ResolvedModel:
        provider_id = provider.strip().lower()
        alias_id = alias.strip().lower()
        try:
            model_ref = self._aliases[(provider_id, alias_id)]
        except KeyError as exc:
            allowed = ", ".join(self.aliases_for(provider_id))
            raise UnknownModelError(
                f"{provider} model {alias!r} is not a supported alias. "
                f"Allowed values: {allowed}"
            ) from exc
        return self.resolve(model_ref)

    def resolve_provider(self, provider_id: str) -> ProviderRuntimeFacts:
        """Return frozen assignment facts without leaking registry objects."""
        normalized = provider_id.strip().lower()
        try:
            return self._provider_facts[normalized]
        except KeyError:
            self._registry.resolve(provider_id)
            raise AssertionError("unreachable")

    def default_candidates(self, role: str) -> tuple[ResolvedModel, ...]:
        canonical_role = self._normalize_role(role)
        candidates: list[ResolvedModel] = []
        for provider_id in self._providers:
            default = self._default_for_role(
                self._provider_defaults[provider_id],
                canonical_role,
            )
            if default is not None:
                candidates.append(self.resolve_alias(provider_id, default))
        return tuple(candidates)

    def default_candidate(
        self,
        provider_id: str,
        role: str,
    ) -> ResolvedModel | None:
        """Return one provider's role default without scanning other providers."""
        canonical_role = self._normalize_role(role)
        provider = self.resolve_provider(provider_id)
        default = self._default_for_role(
            self._provider_defaults[provider.provider],
            canonical_role,
        )
        if default is None:
            return None
        return self.resolve_alias(provider.provider, default)

    def availability(
        self,
        model_ref: str,
        role: str,
        host_runtime: str | None = None,
        entry_mode: str | None = None,
    ) -> ModelAvailability:
        del host_runtime, entry_mode
        try:
            model = self.resolve(model_ref)
        except UnknownModelError:
            return ModelAvailability(False, "unknown-model")
        if not model.selectable:
            return ModelAvailability(False, "not-selectable")
        try:
            self._normalize_role(role)
        except RoleCatalogError:
            return ModelAvailability(False, "unknown-role")
        return ModelAvailability(True, "available")

    def validate(self) -> None:
        failures = self._validation_failures()
        if failures:
            raise ModelPoolValidationError(tuple(failures))

    def aliases_for(self, provider: str) -> tuple[str, ...]:
        provider_id = provider.strip().lower()
        return tuple(
            alias
            for indexed_provider, alias in self._aliases
            if indexed_provider == provider_id
        )

    def _ordered_providers(self) -> tuple[str, ...]:
        providers = self._registry.providers
        bundled = tuple(provider for provider in _BUNDLED_PROVIDER_ORDER if provider in providers)
        users = tuple(sorted(provider for provider in providers if provider not in bundled))
        return (*bundled, *users)

    def _snapshot_provider_facts(self) -> dict[str, ProviderRuntimeFacts]:
        facts: dict[str, ProviderRuntimeFacts] = {}
        canonical_roles = tuple(row.id for row in ROLE_DEFINITIONS)
        for provider_id in self._providers:
            spec = self._registry.resolve(provider_id)
            capability = (
                spec.exec_strategy.policy_support().worker_write_capability()
                if spec.exec_strategy is not None
                else None
            )
            facts[provider_id] = ProviderRuntimeFacts(
                provider=spec.provider,
                display_label=spec.display_label,
                wrapper=spec.wrapper,
                supported_roles=frozenset(
                    role for role in canonical_roles if spec.supports_role(role)
                ),
                cli_worker_write_capability=capability,
                execution_normalization_required=(
                    spec.execution_normalizer is not None
                ),
            )
        return facts

    def _snapshot_provider_defaults(self):
        return {
            provider_id: MappingProxyType(dict(
                self._registry.resolve(provider_id).default_models
            ))
            for provider_id in self._providers
        }

    def _index_models(
        self,
    ) -> tuple[dict[ModelRef, ResolvedModel], dict[tuple[str, str], ModelRef]]:
        models: dict[ModelRef, ResolvedModel] = {}
        aliases: dict[tuple[str, str], ModelRef] = {}
        for provider_id in self._providers:
            for mapping_key, spec in self._registry.resolve(provider_id).models.items():
                model_ref = self._model_ref(provider_id, spec.model_id)
                if model_ref is None:
                    continue
                models.setdefault(model_ref, self._resolved_model(model_ref, spec))
                for alias in (spec.model_id, mapping_key, *spec.aliases):
                    aliases.setdefault((provider_id, alias.strip().lower()), model_ref)
        return models, aliases

    @staticmethod
    def _resolved_model(model_ref: ModelRef, spec: ModelSpec) -> ResolvedModel:
        return ResolvedModel(
            model_ref=model_ref,
            provider_id=model_ref.provider_id,
            display_name=spec.display_name,
            execution_value=spec.execution_value,
            aliases=spec.aliases,
            version_kind=spec.version_kind,
            channel_family=spec.channel_family,
            selectable=spec.selectable,
            pricing=spec.pricing,
            supported_roles=spec.supported_roles,
        )

    def _validation_failures(self) -> list[ModelPoolValidationFailure]:
        failures: list[ModelPoolValidationFailure] = []
        seen_refs: set[ModelRef] = set()
        aliases: dict[tuple[str, str], ModelRef] = {}
        for provider_id in self._providers:
            spec = self._registry.resolve(provider_id)
            self._validate_provider_token(provider_id, failures)
            self._validate_provider_roles(spec.supported_roles, provider_id, failures)
            for role in spec.default_models:
                self._validate_role(role, provider_id, failures)
            channel_families: set[str] = set()
            pinned_families: list[str] = []
            for mapping_key, model in spec.models.items():
                model_ref = self._model_ref(provider_id, model.model_id)
                raw_ref = self._raw_model_ref(provider_id, model.model_id)
                if model_ref is not None and model_ref in seen_refs:
                    failures.append(ModelPoolValidationFailure(
                        "duplicate-model-ref", str(model_ref)
                    ))
                if model_ref is not None:
                    seen_refs.add(model_ref)
                self._validate_model_structure(
                    provider_id, model, raw_ref, channel_families, pinned_families,
                    failures,
                )
                self._validate_model_roles(model, provider_id, raw_ref, failures)
                if model_ref is None:
                    continue
                for alias in (model.model_id, mapping_key, *model.aliases):
                    key = (provider_id, alias.strip().lower())
                    previous = aliases.get(key)
                    if previous is not None and previous != model_ref:
                        failures.append(ModelPoolValidationFailure(
                            "alias-collision", f"{provider_id}/{alias}"
                        ))
                    aliases[key] = model_ref
            self._validate_pinned_families(
                provider_id, pinned_families, channel_families, failures
            )
            self._validate_defaults(provider_id, failures)
        return failures

    @staticmethod
    def _validate_provider_token(
        provider_id: str,
        failures: list[ModelPoolValidationFailure],
    ) -> None:
        if not _TOKEN.fullmatch(provider_id):
            failures.append(ModelPoolValidationFailure(
                "invalid-model-ref", f"invalid provider token: {provider_id!r}"
            ))

    @staticmethod
    def _validate_model_structure(
        provider_id: str,
        model: ModelSpec,
        model_ref: str,
        channel_families: set[str],
        pinned_families: list[str],
        failures: list[ModelPoolValidationFailure],
    ) -> None:
        if not _TOKEN.fullmatch(provider_id) or not _TOKEN.fullmatch(model.model_id):
            failures.append(ModelPoolValidationFailure(
                "invalid-model-ref", f"invalid model token: {model_ref!r}"
            ))
        if not model.execution_value.strip():
            failures.append(ModelPoolValidationFailure(
                "missing-execution-value", str(model_ref)
            ))
        ModelPool._validate_model_family(
            provider_id, model, model_ref, channel_families, pinned_families, failures
        )

    @staticmethod
    def _validate_model_family(
        provider_id: str,
        model: ModelSpec,
        model_ref: str,
        channel_families: set[str],
        pinned_families: list[str],
        failures: list[ModelPoolValidationFailure],
    ) -> None:
        if model.version_kind not in _VERSION_KINDS:
            failures.append(ModelPoolValidationFailure(
                "invalid-version-kind", f"{model_ref}: {model.version_kind}"
            ))
            return
        family = model.channel_family
        if model.version_kind == "channel" and not family:
            failures.append(ModelPoolValidationFailure(
                "missing-channel-family", str(model_ref)
            ))
            return
        if family and not _TOKEN.fullmatch(family):
            failures.append(ModelPoolValidationFailure(
                "invalid-channel-family", f"{model_ref}: {family!r}"
            ))
            return
        if not family:
            return
        if model.version_kind == "channel":
            channel_families.add(family)
        else:
            pinned_families.append(family)

    @staticmethod
    def _validate_pinned_families(
        provider_id: str,
        pinned_families: list[str],
        channel_families: set[str],
        failures: list[ModelPoolValidationFailure],
    ) -> None:
        for family in pinned_families:
            if family not in channel_families:
                failures.append(ModelPoolValidationFailure(
                    "orphan-channel-family", f"{provider_id}: {family}"
                ))

    @staticmethod
    def _validate_provider_roles(
        roles: frozenset[str], provider_id: str, failures: list[ModelPoolValidationFailure]
    ) -> None:
        for role in roles:
            ModelPool._validate_role(role, provider_id, failures)

    def _validate_model_roles(
        self,
        model: ModelSpec,
        provider_id: str,
        model_ref: str,
        failures: list[ModelPoolValidationFailure],
    ) -> None:
        provider = self._registry.resolve(provider_id)
        for role in model.supported_roles or ():
            ModelPool._validate_role(role, str(model_ref), failures)
            try:
                canonical_role = normalize_role(role)
            except RoleCatalogError:
                continue
            if not provider.supports_role(canonical_role):
                failures.append(ModelPoolValidationFailure(
                    "model-role-exceeds-provider", f"{model_ref}: {canonical_role}"
                ))

    @staticmethod
    def _validate_role(
        role: str, owner: str, failures: list[ModelPoolValidationFailure]
    ) -> None:
        try:
            normalize_role(role)
        except RoleCatalogError:
            failures.append(ModelPoolValidationFailure("unknown-role", f"{owner}: {role}"))

    def _validate_defaults(
        self,
        provider_id: str,
        failures: list[ModelPoolValidationFailure],
    ) -> None:
        provider = self._registry.resolve(provider_id)
        for role in self._supported_roles(provider):
            default = self._default_for_role(provider.default_models, role)
            if default is None:
                failures.append(ModelPoolValidationFailure(
                    "missing-default", f"{provider_id}: {role}"
                ))
                continue
            try:
                model = self.resolve_alias(provider_id, default)
            except UnknownModelError:
                failures.append(ModelPoolValidationFailure(
                    "missing-default", f"{provider_id}: {role} -> {default}"
                ))
                continue
            if not model.selectable or not self._supports_role(model, role):
                failures.append(ModelPoolValidationFailure(
                    "missing-default", f"{provider_id}: {role} -> {default}"
                ))

    @staticmethod
    def _supported_roles(spec: object) -> tuple[str, ...]:
        from .domain.role import ROLE_DEFINITIONS

        return tuple(
            definition.id
            for definition in ROLE_DEFINITIONS
            if spec.supports_role(definition.id)
        )

    @staticmethod
    def _default_for_role(defaults: object, canonical_role: str) -> str | None:
        for role, model_id in defaults.items():
            try:
                if normalize_role(role) == canonical_role:
                    return model_id
            except RoleCatalogError:
                continue
        return None

    def _supports_role(self, model: ResolvedModel, role: str | None) -> bool:
        if role is None:
            return True
        return self.availability(model.model_ref, role).available

    @staticmethod
    def _normalize_role(role: str) -> str:
        return normalize_role(role)

    @staticmethod
    def _model_ref(provider_id: str, model_id: str) -> ModelRef | None:
        try:
            return ModelRef(provider_id, model_id)
        except ValueError:
            return None

    @staticmethod
    def _raw_model_ref(provider_id: str, model_id: str) -> str:
        return f"{provider_id}/{model_id}"
