"""Provider and model value objects shared by the registry and adapters."""
from __future__ import annotations

from collections.abc import Callable
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Literal, Mapping, Optional, Protocol

from .role import RoleCatalogError, normalize_role, ROLE_CATALOG
from .worker_exec import WorkerWriteCapability

if TYPE_CHECKING:
    from okstra_ctl.domain.worker_exec import ExecutionStrategy


_MISSING = object()


@dataclass(frozen=True)
class ServedModelAttestation:
    """Provider-observed model identity after one execution attempt."""

    observed_model: str | None
    normalized_model_ref: str | None
    level: Literal["exact", "channel", "unknown"]
    source: Literal[
        "provider-output",
        "native-api",
        "cli-handoff",
        "unavailable",
    ]

    @classmethod
    def unknown(cls) -> ServedModelAttestation:
        return cls(None, None, "unknown", "unavailable")


ServedModelNormalizer = Callable[[str | None], ServedModelAttestation]


def unavailable_served_model(_raw_model: str | None) -> ServedModelAttestation:
    return ServedModelAttestation.unknown()


@dataclass(frozen=True)
class HostModelBinding:
    """Final provider value and host-native value for one selected model."""

    runner: str
    catalog_execution_value: str
    resolved_execution_value: str
    host_model_value: str | None
    binding_fidelity: Literal["exact", "channel", "unsupported"]
    worker_write_capability: WorkerWriteCapability | None


@dataclass(frozen=True, init=False)
class ModelSpec:
    """One canonical model plus provider-owned aliases and execution value."""

    model_id: str
    display_name: str
    execution_value: str
    aliases: tuple[str, ...]
    version_kind: Literal["pinned", "channel"]
    channel_family: str | None
    supported_roles: frozenset[str] | None
    selectable: bool
    pricing: Optional[tuple[float, float, float]]

    def __init__(
        self,
        model_id: str,
        display_name: str,
        execution_value: str | tuple[float, float, float] | None | object = _MISSING,
        aliases: tuple[str, ...] = (),
        version_kind: Literal["pinned", "channel"] = "pinned",
        channel_family: str | None = None,
        supported_roles: frozenset[str] | None = None,
        selectable: bool = True,
        pricing: Optional[tuple[float, float, float]] = None,
        *,
        in_picker: bool | None = None,
    ) -> None:
        """Accept the former display/execution form while adapters migrate."""
        is_legacy = execution_value is _MISSING or not isinstance(execution_value, str)
        if is_legacy:
            legacy_pricing = execution_value if isinstance(execution_value, tuple) else None
            canonical_id = display_name
            resolved_display = model_id
            resolved_execution = display_name
            resolved_pricing = legacy_pricing
        else:
            canonical_id = model_id
            resolved_display = display_name
            resolved_execution = execution_value
            resolved_pricing = pricing
        object.__setattr__(self, "model_id", canonical_id)
        object.__setattr__(self, "display_name", resolved_display)
        object.__setattr__(self, "execution_value", resolved_execution)
        object.__setattr__(self, "aliases", aliases)
        object.__setattr__(self, "version_kind", version_kind)
        object.__setattr__(self, "channel_family", channel_family)
        object.__setattr__(self, "supported_roles", supported_roles)
        object.__setattr__(self, "selectable", selectable if in_picker is None else in_picker)
        object.__setattr__(self, "pricing", resolved_pricing)

    @property
    def display(self) -> str:
        """Legacy display-name view."""
        return self.display_name

    @property
    def execution(self) -> str:
        """Legacy execution-value view."""
        return self.execution_value

    @property
    def in_picker(self) -> bool:
        """Legacy picker view."""
        return self.selectable


@dataclass(frozen=True)
class ExecutionNormalization:
    """One precomputed CLI execution value or its discovery-time error."""

    resolved_execution_value: str | None
    error: str = ""

    def __post_init__(self) -> None:
        if bool(self.resolved_execution_value) == bool(self.error):
            raise ValueError(
                "execution normalization requires exactly one value or error"
            )


class ExecutionNormalizer(Protocol):
    """Provider-owned batch snapshot boundary for ambient CLI discovery."""

    def snapshot(
        self,
        models: tuple[ModelSpec, ...],
        roles: tuple[str, ...],
    ) -> Mapping[tuple[str, str], ExecutionNormalization]: ...


def snapshot_execution_normalizations(
    models: tuple[ModelSpec, ...],
    roles: tuple[str, ...],
    normalize: Callable[[str, str], str],
) -> Mapping[tuple[str, str], ExecutionNormalization]:
    """Evaluate a pure provider normalizer for every model/role pair."""
    outcomes: dict[tuple[str, str], ExecutionNormalization] = {}
    for model in models:
        for role in roles:
            try:
                resolved = normalize(model.execution_value, role)
            except Exception as exc:  # provider errors become immutable facts
                outcomes[(model.model_id, role)] = ExecutionNormalization(
                    None,
                    str(exc),
                )
            else:
                outcomes[(model.model_id, role)] = ExecutionNormalization(resolved)
    return outcomes


@dataclass(frozen=True)
class LeadLaunchSpec:
    """How to start a provider CLI as an interactive lead session."""

    executable: str
    model_flag: str
    prompt_flag: str = ""
    start_session_id_flag: str = ""
    resume_session_id_flag: str = ""
    sandbox_waiver: tuple[str, ...] = ()
    sandbox_waiver_note: str = ""


@dataclass(frozen=True)
class ProviderSpec:
    provider: str
    display_label: str
    models: Mapping[str, ModelSpec]
    default_models: Mapping[str, str]
    wrapper: str
    supported_roles: frozenset[str]
    execution_capabilities: frozenset[str] = frozenset()
    lead_launch: Optional[LeadLaunchSpec] = None
    exec_strategy: Optional["ExecutionStrategy"] = None
    execution_normalizer: Optional[ExecutionNormalizer] = None
    served_model_normalizer: ServedModelNormalizer = unavailable_served_model
    _legacy_role_restriction: frozenset[str] = field(
        init=False,
        repr=False,
        compare=False,
    )

    def __post_init__(self) -> None:
        legacy_roles = _normalized_legacy_roles(self.supported_roles)
        if self.execution_capabilities:
            capabilities = self.execution_capabilities
            restriction = frozenset()
        else:
            capabilities = _legacy_execution_capabilities(legacy_roles)
            restriction = legacy_roles
        object.__setattr__(self, "execution_capabilities", capabilities)
        object.__setattr__(self, "_legacy_role_restriction", restriction)

    def supports_role(self, role: str) -> bool:
        try:
            canonical_role = normalize_role(role)
        except RoleCatalogError:
            return False
        if self._legacy_role_restriction:
            return canonical_role in self._legacy_role_restriction
        required = ROLE_CATALOG.required_capabilities(canonical_role)
        return required.issubset(self.execution_capabilities)


@dataclass(frozen=True)
class RunnerResolution:
    runner: str
    wrapper: str
    reason: str


@dataclass(frozen=True)
class LeadProviderAssignment:
    provider: str
    runner: str
    wrapper: str
    reason: str


@dataclass
class ModelAssignment:
    display: str
    execution: str
    model_ref: str = ""
    version_kind: Literal["pinned", "channel"] = "pinned"


class UnknownModelError(ValueError):
    """Raised when a provider has no model alias matching a raw value."""


class UnknownProviderError(ValueError):
    """Raised when a requested provider is absent from the registry."""


def _normalized_legacy_roles(roles: frozenset[str]) -> frozenset[str]:
    normalized: set[str] = set()
    for role in roles:
        try:
            normalized.add(normalize_role(role))
        except RoleCatalogError:
            continue
    return frozenset(normalized)


def _legacy_execution_capabilities(roles: frozenset[str]) -> frozenset[str]:
    capabilities: set[str] = set()
    for role in roles:
        if role == "leader":
            capabilities.add("lead-session")
            continue
        capabilities.add("worker-artifact-io")
        if role == "analyser":
            capabilities.add("source-readonly")
        if role == "implementer":
            capabilities.add("project-mutation")
        if role in {"report-writer", "translator"}:
            capabilities.add("extended-artifact-authoring")
    return frozenset(capabilities)
