from __future__ import annotations

from dataclasses import dataclass, replace
from typing import Any

from app.config import config
from app.core.logging import get_audit_logger
from app.llm.contracts import LLMOutputKind, LLMProviderPort, LLMTask, ModelSelectionPolicy
from app.llm.models import (
    LLMErrorKind,
    LLMProviderError,
    LLMResolvedRoute,
    LLMStructuredRequest,
    LLMStructuredResult,
    LLMTextRequest,
    LLMTextResult,
)


@dataclass(frozen=True, slots=True)
class TaskRouteProfile:
    output_kind: LLMOutputKind
    provider: str
    model: str
    metadata: dict[str, Any] | None = None

    def build_route(self, *, task: LLMTask) -> LLMResolvedRoute:
        return LLMResolvedRoute(
            task=task,
            output_kind=self.output_kind,
            provider=self.provider,
            model=self.model,
            metadata=dict(self.metadata or {}),
        )


class DefaultModelSelectionPolicy(ModelSelectionPolicy):
    _ROUTES: dict[LLMTask, TaskRouteProfile] = {
        LLMTask.CLINICAL_DOCUMENT_EXTRACT: TaskRouteProfile(
            output_kind=LLMOutputKind.STRUCTURED_OBJECT,
            provider="gemini",
            model=config.GEMINI_MODEL_EXTRACT_LOW_RISK,
            metadata={"required_capability": "structured_output", "volume": "high"},
        ),
        LLMTask.PREFACTURA_PAGE_CLASSIFICATION: TaskRouteProfile(
            output_kind=LLMOutputKind.STRUCTURED_OBJECT,
            provider="gemini",
            model=config.GEMINI_MODEL_EXTRACT_LOW_RISK,
            metadata={"required_capability": "structured_output", "domain": "prefactura"},
        ),
        LLMTask.HISTORIA_CHUNK_SUMMARY: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_EXTRACT_LOW_RISK,
            metadata={"required_capability": "text_generation", "long_context": True, "volume": "high"},
        ),
        LLMTask.HISTORIA_FINAL_SUMMARY: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "text_generation", "long_context": True},
        ),
        LLMTask.SOAT_REASONING: TaskRouteProfile(
            output_kind=LLMOutputKind.STRUCTURED_OBJECT,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "structured_output", "domain": "soat"},
        ),
        LLMTask.SOAT_CODE_GENERATION: TaskRouteProfile(
            output_kind=LLMOutputKind.STRUCTURED_COLLECTION,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "structured_output", "domain": "soat"},
        ),
        LLMTask.SOAT_CHAT: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "text_generation", "long_context": True, "domain": "soat"},
        ),
        LLMTask.SOAT_MANUAL_CODE_LOOKUP: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "text_generation", "domain": "soat"},
        ),
        LLMTask.GLOSA_NOTE_GENERATION: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "text_generation", "domain": "soat"},
        ),
        LLMTask.MEDICATION_PERTINENCE: TaskRouteProfile(
            output_kind=LLMOutputKind.STRUCTURED_COLLECTION,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "structured_output", "domain": "auditoria_medicamentos"},
        ),
        LLMTask.EPICRISIS_SUMMARY_COMPOSITION: TaskRouteProfile(
            output_kind=LLMOutputKind.STRUCTURED_OBJECT,
            provider="gemini",
            model=config.GEMINI_MODEL_REASONING,
            metadata={"required_capability": "structured_output", "domain": "epicrisis"},
        ),
        LLMTask.CIE10_RESOLUTION: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_EXTRACT_LOW_RISK,
            metadata={"required_capability": "text_generation", "volume": "high"},
        ),
        LLMTask.CUPS_RESOLUTION: TaskRouteProfile(
            output_kind=LLMOutputKind.CONTROLLED_PLAIN_TEXT,
            provider="gemini",
            model=config.GEMINI_MODEL_EXTRACT_LOW_RISK,
            metadata={"required_capability": "text_generation", "volume": "high"},
        ),
    }

    def resolve(self, task: LLMTask, metadata: dict[str, Any] | None = None) -> LLMResolvedRoute:
        if task == LLMTask.CLINICAL_DOCUMENT_EXTRACT:
            return self._build_clinical_extract_route(task=task, metadata=metadata)
        return self._ROUTES[task].build_route(task=task)

    def _build_clinical_extract_route(
        self,
        *,
        task: LLMTask,
        metadata: dict[str, Any] | None,
    ) -> LLMResolvedRoute:
        profile = self._ROUTES[task]
        route = profile.build_route(task=task)
        normalized_metadata = dict(metadata or {})
        document_type = str(normalized_metadata.get("document_type") or "").strip().lower()
        risk_level = str(normalized_metadata.get("risk_level") or "low").strip().lower() or "low"
        if document_type == "factura":
            route = replace(route, model=config.GEMINI_MODEL_EXTRACT_HIGH_RISK)
            risk_level = "high"
        elif risk_level == "high":
            route = replace(route, model=config.GEMINI_MODEL_EXTRACT_HIGH_RISK)
        route.metadata.update(normalized_metadata)
        route.metadata.setdefault("risk_level", risk_level)
        return route


class DefaultModelRouter:
    def __init__(
        self,
        *,
        providers: dict[str, LLMProviderPort],
        policy: ModelSelectionPolicy,
        service: str = "llm_router",
    ) -> None:
        self._providers = dict(providers)
        self._policy = policy
        self._service = service
        self._audit_logger = get_audit_logger()

    def generate_text(self, request: LLMTextRequest) -> LLMTextResult:
        route = self._resolve_route(request.task, metadata=dict(request.metadata or {}))
        primary_request = replace(request, model=route.model)
        self._log_route("route.select", route.provider, route.model, route=route, fallback_used=False)
        try:
            provider = self._provider_or_raise(route.provider, route.model)
            return provider.generate_text(primary_request)
        except LLMProviderError as exc:
            return self._handle_text_fallback(route=route, request=primary_request, exc=exc)

    def generate_structured(self, request: LLMStructuredRequest) -> LLMStructuredResult:
        route = self._resolve_route(request.task, metadata=dict(request.metadata or {}))
        primary_request = replace(request, model=route.model)
        self._log_route("route.select", route.provider, route.model, route=route, fallback_used=False)
        try:
            provider = self._provider_or_raise(route.provider, route.model)
            return provider.generate_structured(primary_request)
        except LLMProviderError as exc:
            return self._handle_structured_fallback(route=route, request=primary_request, exc=exc)

    def _resolve_route(self, task: LLMTask | None, *, metadata: dict[str, Any] | None = None) -> LLMResolvedRoute:
        if task is None:
            raise LLMProviderError(
                "La tarea LLM es requerida para resolver la ruta",
                kind=LLMErrorKind.INVALID_CONFIGURATION,
                provider="router",
            )
        return self._policy.resolve(task, metadata=metadata)

    def _provider_or_raise(self, provider_name: str, model: str) -> LLMProviderPort:
        provider = self._providers.get(provider_name)
        if provider is None:
            raise LLMProviderError(
                f"Proveedor {provider_name} no registrado",
                kind=LLMErrorKind.INVALID_CONFIGURATION,
                provider=provider_name,
                model=model,
            )
        availability = provider.availability()
        if not availability.configured:
            raise LLMProviderError(
                availability.reason or f"Proveedor {provider_name} no configurado",
                kind=LLMErrorKind.INVALID_CONFIGURATION,
                provider=provider_name,
                model=model,
            )
        return provider

    def _handle_text_fallback(
        self,
        *,
        route: LLMResolvedRoute,
        request: LLMTextRequest,
        exc: LLMProviderError,
    ) -> LLMTextResult:
        if not self._should_fallback(exc, route):
            raise
        fallback = route.fallback_rule
        assert fallback is not None
        self._log_route(
            "route.fallback",
            fallback.provider,
            fallback.model,
            route=route,
            fallback_used=True,
            reason=exc.kind.value,
        )
        provider = self._provider_or_raise(fallback.provider, fallback.model)
        result = provider.generate_text(replace(request, model=fallback.model))
        return replace(result, fallback=fallback)

    def _handle_structured_fallback(
        self,
        *,
        route: LLMResolvedRoute,
        request: LLMStructuredRequest,
        exc: LLMProviderError,
    ) -> LLMStructuredResult:
        if not self._should_fallback(exc, route):
            raise
        fallback = route.fallback_rule
        assert fallback is not None
        self._log_route(
            "route.fallback",
            fallback.provider,
            fallback.model,
            route=route,
            fallback_used=True,
            reason=exc.kind.value,
        )
        provider = self._provider_or_raise(fallback.provider, fallback.model)
        result = provider.generate_structured(replace(request, model=fallback.model))
        return replace(result, fallback=fallback)

    def _should_fallback(self, exc: LLMProviderError, route: LLMResolvedRoute) -> bool:
        if route.fallback_rule is None:
            return False
        return exc.kind in {
            LLMErrorKind.TRANSIENT,
            LLMErrorKind.RATE_LIMITED,
            LLMErrorKind.PROVIDER_UNAVAILABLE,
            LLMErrorKind.UNSUPPORTED_CAPABILITY,
            LLMErrorKind.INVALID_CONFIGURATION,
        }

    def _log_route(
        self,
        action: str,
        provider_name: str,
        model: str,
        *,
        route: LLMResolvedRoute,
        fallback_used: bool,
        reason: str | None = None,
    ) -> None:
        resource: dict[str, Any] = {
            "task": route.task.value,
            "output_kind": route.output_kind.value,
            "fallback_used": fallback_used,
            "selected_provider": provider_name,
            "selected_model": model,
        }
        if reason:
            resource["fallback_reason"] = reason
        self._audit_logger.llm_event(
            action=action,
            outcome="success",
            service=self._service,
            provider=provider_name,
            model=model,
            resource=resource,
        )
