from __future__ import annotations

from collections.abc import Mapping
from dataclasses import dataclass, field
from enum import StrEnum
from typing import Any

from app.llm.contracts import LLMOutputKind, LLMTask


class LLMErrorKind(StrEnum):
    TRANSIENT = "transient"
    RATE_LIMITED = "rate_limited"
    PROVIDER_UNAVAILABLE = "provider_unavailable"
    UNSUPPORTED_CAPABILITY = "unsupported_capability"
    INVALID_CONFIGURATION = "invalid_configuration"


@dataclass(slots=True)
class LLMProviderAvailability:
    configured: bool
    enabled: bool = True
    reason: str | None = None
    is_stub: bool = False


@dataclass(slots=True)
class LLMProviderCapabilities:
    text_generation: bool
    structured_generation: bool
    native_structured_output: bool = False
    max_context_tokens: int | None = None


@dataclass(slots=True)
class LLMRouteFallbackRule:
    provider: str
    model: str
    reason: str | None = None


@dataclass(slots=True)
class LLMResolvedRoute:
    task: LLMTask
    output_kind: LLMOutputKind
    provider: str
    model: str
    fallback_rule: LLMRouteFallbackRule | None = None
    metadata: dict[str, Any] = field(default_factory=dict)


@dataclass(slots=True)
class LLMTextRequest:
    task: LLMTask | None
    prompt: str
    system_prompt: str | None = None
    model: str | None = None
    contents: Any = None
    generation_config: Any = None
    metadata: Mapping[str, Any] = field(default_factory=dict)


@dataclass(slots=True)
class LLMStructuredRequest:
    task: LLMTask | None
    prompt: str
    output_model: Any
    output_kind: LLMOutputKind
    system_prompt: str | None = None
    model: str | None = None
    metadata: Mapping[str, Any] = field(default_factory=dict)


@dataclass(slots=True)
class LLMTextResult:
    content: str
    provider: str
    model: str
    metrics: dict[str, Any] = field(default_factory=dict)
    metadata: dict[str, Any] = field(default_factory=dict)
    fallback: LLMRouteFallbackRule | None = None


@dataclass(slots=True)
class LLMStructuredResult:
    content: Any
    provider: str
    model: str
    output_kind: LLMOutputKind
    metrics: dict[str, Any] = field(default_factory=dict)
    metadata: dict[str, Any] = field(default_factory=dict)
    fallback: LLMRouteFallbackRule | None = None


class LLMProviderError(RuntimeError):
    def __init__(
        self,
        message: str,
        *,
        kind: LLMErrorKind,
        provider: str,
        model: str | None = None,
        retryable: bool = False,
        details: dict[str, Any] | None = None,
    ) -> None:
        super().__init__(message)
        self.kind = kind
        self.provider = provider
        self.model = model
        self.retryable = retryable
        self.details = details or {}

    def to_error_payload(self) -> dict[str, Any]:
        payload = {
            "class": self.__class__.__name__,
            "message": str(self),
            "kind": self.kind.value,
            "provider": self.provider,
        }
        if self.model:
            payload["model"] = self.model
        if self.details:
            payload["details"] = dict(self.details)
        return payload
