from __future__ import annotations

import asyncio
import logging
from typing import Any

from app.core.settings import Settings

from .mlx_backend import HAS_MLX, MlxBackendMixin
from .openai_backend import OpenAiBackendMixin

logger = logging.getLogger("dscons.llm.service")


class LLMClient(
    MlxBackendMixin,
    OpenAiBackendMixin,
):
    """Multi-backend LLM client supporting local MLX and OpenAI-compatible endpoints."""

    def __init__(self, settings: Settings) -> None:
        """Initialize the client with application settings."""
        self._settings = settings
        self._timeout = settings.request_timeout_seconds
        self._model_path = settings.llm_model_name
        self._default_max_tokens = settings.llm_max_tokens
        self._provider = settings.llm_provider
        self._api_base_url = settings.llm_api_base_url.rstrip("/")
        self._api_key = settings.llm_api_key
        self._model = None
        self._tokenizer = None

        if self._provider == "auto":
            self._active_provider = "mlx" if HAS_MLX else "openai_compatible"
        else:
            self._active_provider = self._provider

    async def create_chat_completion(
        self,
        messages: list[dict[str, str]],
        *,
        model: str | None = None,
        temperature: float = 0.2,
        max_tokens: int | None = None,
        extra_body: dict[str, Any] | None = None,
        agent_code: str | None = None,
    ) -> dict[str, Any]:
        """Generate a chat completion using the configured Plug and Play provider."""
        provider_info = self.get_provider_for_agent(agent_code)
        provider_str = provider_info.get("provider", self._active_provider)
        effective_model = model or provider_info.get("model", self._model_path)

        if provider_str == "mlx":
            return await self._generate_mlx(
                messages,
                temperature=temperature,
                max_tokens=max_tokens,
                agent_code=agent_code,
            )

        from app.core.ai.domain.models import AiProviderType, AiTextRequest
        from app.core.ai.infrastructure.gateway import get_ai_gateway

        pref_provider = None
        if provider_str in ("antigravity", "antigravity_sdk"):
            pref_provider = AiProviderType.ANTIGRAVITY
        elif provider_str == "openrouter":
            pref_provider = AiProviderType.OPENROUTER
        elif provider_str == "lmstudio":
            pref_provider = AiProviderType.LMSTUDIO

        gateway = get_ai_gateway()
        req = AiTextRequest(
            messages=messages,
            canonical_model_id=effective_model,
            preferred_provider=pref_provider,
            agent_code=agent_code,
            temperature=temperature
            if temperature != 0.2
            else provider_info.get("temperature", 0.2),
            max_tokens=max_tokens
            or provider_info.get("max_tokens", self._default_max_tokens),
            extra_body=extra_body,
        )

        try:
            resp = await gateway.generate_text(req)
            return {
                "choices": [{"message": {"content": resp.content}}],
                "model": resp.model_used,
                "usage": {
                    "prompt_tokens": resp.prompt_tokens,
                    "completion_tokens": resp.completion_tokens,
                },
            }
        except Exception as e:
            logger.warning(
                "[LLM_CLIENT] EnterpriseAiGateway call failed (%s). Falling back to legacy OpenAI backend...",
                e,
            )
            return await self._generate_openai_compatible(
                messages,
                model=effective_model,
                temperature=temperature,
                max_tokens=max_tokens,
                extra_body=extra_body,
                agent_code=agent_code,
            )

    async def chat(
        self,
        messages: list[dict[str, str]],
        *,
        model: str | None = None,
        temperature: float = 0.2,
        max_tokens: int | None = None,
        extra_body: dict[str, Any] | None = None,
        agent_code: str | None = None,
    ) -> str:
        """Return the assistant text content from a chat completion response."""
        data = await self.create_chat_completion(
            messages,
            model=model,
            temperature=temperature,
            max_tokens=max_tokens,
            extra_body=extra_body,
            agent_code=agent_code,
        )

        # --- COST TRACKING ---
        if not hasattr(self, "session_cost_usd"):
            self.session_cost_usd = 0.0
            self.session_tokens = {"prompt": 0, "completion": 0}

        usage = data.get("usage", {})
        pt = usage.get("prompt_tokens", 0)
        ct = usage.get("completion_tokens", 0)
        self.session_tokens["prompt"] += pt
        self.session_tokens["completion"] += ct

        effective_model = data.get("model", model or "unknown").lower()
        if "opus" in effective_model:
            cost = (pt * 15.0 + ct * 75.0) / 1000000
        elif "fable" in effective_model:
            cost = (pt * 3.0 + ct * 15.0) / 1000000
        elif "sol" in effective_model:
            cost = (pt * 2.5 + ct * 10.0) / 1000000
        elif "deepseek-v4-pro" in effective_model:
            cost = (pt * 0.14 + ct * 0.28) / 1000000
        else:
            cost = (pt * 0.5 + ct * 1.5) / 1000000

        self.session_cost_usd += cost
        logger.info(
            f"[{effective_model}] Used {pt} prompt, {ct} comp tokens. Cost: ${cost:.6f}. Session Total: ${self.session_cost_usd:.6f}"
        )
        # ---------------------

        choices = data.get("choices", [])
        if not choices:
            return ""
        message = choices[0].get("message", {})
        content = message.get("content", "")
        if not content and "reasoning_content" in message:
            content = message.get("reasoning_content", "")
        return content if isinstance(content, str) else str(content)
