from __future__ import annotations

import json
import time
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import Any

from openai import (
    APIConnectionError,
    APIError,
    APITimeoutError,
    AsyncOpenAI,
    AuthenticationError,
    BadRequestError,
    NotFoundError,
    RateLimitError,
)
from tenacity import (
    retry,
    retry_if_exception_type,
    stop_after_attempt,
    wait_exponential,
)

from app.core.config import settings
from app.core.logger import logger

# Query embedding is owned by the embedding-engine (Track 1). This client handles chat only.


class LLMClientError(RuntimeError):
    """Base LLM client error."""

    def __init__(self, message: str, *, status_code: int | None = None):
        super().__init__(message)
        self.status_code = status_code


class LLMAuthError(LLMClientError):
    """Authentication failure (401)."""


class LLMRateLimitError(LLMClientError):
    """Rate limit exceeded (429)."""


class LLMTimeoutError(LLMClientError):
    """Request timed out."""


class LLMInvalidModelError(LLMClientError):
    """Invalid or unknown model."""


class LLMMalformedOutputError(LLMClientError):
    """Structured output could not be parsed."""


@dataclass(frozen=True)
class ChatCompletionResult:
    text: str
    model: str
    latency_ms: int
    prompt_tokens: int | None = None
    completion_tokens: int | None = None


@dataclass(frozen=True)
class ToolCallResult:
    id: str
    name: str
    arguments: dict[str, Any]


@dataclass(frozen=True)
class ToolCompletionResult:
    text: str
    tool_calls: list[ToolCallResult]
    model: str
    finish_reason: str
    latency_ms: int
    prompt_tokens: int | None = None
    completion_tokens: int | None = None


@dataclass(frozen=True)
class ModelInfo:
    id: str
    owned_by: str | None = None
    created: int | None = None


class LLMClient:
    """OpenAI chat completion wrapper for RAG and Tier-2 intent."""

    RETRYABLE = (RateLimitError, APIError, APITimeoutError, APIConnectionError)

    def __init__(
        self,
        *,
        api_key: str | None = None,
        default_model: str | None = None,
        timeout_sec: int | None = None,
        max_retries: int | None = None,
        backoff_factor: float | None = None,
        client: AsyncOpenAI | None = None,
    ):
        key = api_key if api_key is not None else settings.OPENAI_API_KEY
        if not str(key).strip():
            if settings.is_local():
                key = "local-dev"
            else:
                raise LLMClientError("OPENAI_API_KEY is required outside local/dev")

        self._default_model = default_model or settings.LLM_MODEL
        self._timeout_sec = timeout_sec or settings.LLM_REQUEST_TIMEOUT_SEC
        self._max_retries = max_retries or settings.LLM_MAX_RETRIES
        self._backoff_factor = backoff_factor or settings.LLM_RETRY_BACKOFF_FACTOR
        self._client = client or AsyncOpenAI(api_key=key, timeout=self._timeout_sec)

    def _map_error(self, exc: Exception) -> LLMClientError:
        if isinstance(exc, AuthenticationError):
            return LLMAuthError(str(exc), status_code=401)
        if isinstance(exc, RateLimitError):
            return LLMRateLimitError(str(exc), status_code=429)
        if isinstance(exc, APITimeoutError):
            return LLMTimeoutError(str(exc), status_code=504)
        if isinstance(exc, NotFoundError):
            return LLMInvalidModelError(str(exc), status_code=404)
        if isinstance(exc, BadRequestError):
            message = str(exc).lower()
            if "model" in message:
                return LLMInvalidModelError(str(exc), status_code=400)
            return LLMClientError(str(exc), status_code=400)
        if isinstance(exc, LLMClientError):
            return exc
        return LLMClientError(str(exc))

    @retry(
        stop=stop_after_attempt(settings.LLM_MAX_RETRIES),
        wait=wait_exponential(
            multiplier=settings.LLM_RETRY_BACKOFF_FACTOR,
            min=2,
            max=60,
        ),
        retry=retry_if_exception_type(RETRYABLE),
        reraise=True,
    )
    async def _create_completion(
        self,
        *,
        messages: list[dict[str, Any]],
        model: str,
        response_format: dict[str, Any] | None = None,
        stream: bool = False,
        tools: list[dict[str, Any]] | None = None,
        tool_choice: str | dict[str, Any] | None = None,
    ):
        kwargs: dict[str, Any] = {
            "model": model,
            "messages": messages,
            "stream": stream,
        }
        if response_format is not None:
            kwargs["response_format"] = response_format
        if tools is not None:
            kwargs["tools"] = tools
        if tool_choice is not None:
            kwargs["tool_choice"] = tool_choice
        return await self._client.chat.completions.create(**kwargs)

    async def complete(
        self,
        messages: list[dict[str, str]],
        *,
        model: str | None = None,
        response_format: dict[str, Any] | None = None,
    ) -> ChatCompletionResult:
        resolved_model = model or self._default_model
        started = time.perf_counter()
        try:
            response = await self._create_completion(
                messages=messages,
                model=resolved_model,
                response_format=response_format,
                stream=False,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        latency_ms = int((time.perf_counter() - started) * 1000)
        choice = response.choices[0] if response.choices else None
        text = (choice.message.content or "") if choice and choice.message else ""

        usage = response.usage
        prompt_tokens = usage.prompt_tokens if usage else None
        completion_tokens = usage.completion_tokens if usage else None

        logger.info(
            "LLM complete model={model} latency_ms={latency_ms} prompt_tokens={prompt_tokens} completion_tokens={completion_tokens}",
            model=resolved_model,
            latency_ms=latency_ms,
            prompt_tokens=prompt_tokens,
            completion_tokens=completion_tokens,
        )

        return ChatCompletionResult(
            text=text,
            model=resolved_model,
            latency_ms=latency_ms,
            prompt_tokens=prompt_tokens,
            completion_tokens=completion_tokens,
        )

    async def complete_with_tools(
        self,
        messages: list[dict[str, Any]],
        tools: list[dict[str, Any]],
        *,
        model: str | None = None,
        tool_choice: str | dict[str, Any] | None = "auto",
    ) -> ToolCompletionResult:
        resolved_model = model or self._default_model
        started = time.perf_counter()
        try:
            response = await self._create_completion(
                messages=messages,
                model=resolved_model,
                tools=tools,
                tool_choice=tool_choice,
                stream=False,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        latency_ms = int((time.perf_counter() - started) * 1000)
        choice = response.choices[0] if response.choices else None
        message = choice.message if choice else None
        text = (message.content or "") if message else ""
        finish_reason = (choice.finish_reason or "") if choice else ""

        tool_calls: list[ToolCallResult] = []
        for tool_call in (message.tool_calls or []) if message else []:
            raw_args = tool_call.function.arguments if tool_call.function else ""
            try:
                arguments = json.loads(raw_args or "{}")
            except json.JSONDecodeError:
                arguments = {}
            if not isinstance(arguments, dict):
                arguments = {}
            tool_calls.append(
                ToolCallResult(
                    id=tool_call.id,
                    name=tool_call.function.name if tool_call.function else "",
                    arguments=arguments,
                )
            )

        usage = response.usage
        prompt_tokens = usage.prompt_tokens if usage else None
        completion_tokens = usage.completion_tokens if usage else None

        logger.info(
            "LLM complete_with_tools model={model} latency_ms={latency_ms} "
            "tool_calls={tool_calls} tools={tools} finish_reason={finish_reason} text_len={text_len}",
            model=resolved_model,
            latency_ms=latency_ms,
            tool_calls=len(tool_calls),
            tools=[tc.name for tc in tool_calls],
            finish_reason=finish_reason,
            text_len=len(text),
        )

        return ToolCompletionResult(
            text=text,
            tool_calls=tool_calls,
            model=resolved_model,
            finish_reason=finish_reason,
            latency_ms=latency_ms,
            prompt_tokens=prompt_tokens,
            completion_tokens=completion_tokens,
        )

    async def complete_json(
        self,
        messages: list[dict[str, str]],
        *,
        model: str | None = None,
    ) -> dict[str, Any]:
        result = await self.complete(
            messages,
            model=model,
            response_format={"type": "json_object"},
        )
        try:
            return json.loads(result.text)
        except json.JSONDecodeError as exc:
            raise LLMMalformedOutputError(
                f"LLM returned invalid JSON: {result.text[:200]}"
            ) from exc

    async def stream(
        self,
        messages: list[dict[str, str]],
        *,
        model: str | None = None,
    ) -> AsyncIterator[str]:
        resolved_model = model or self._default_model
        started = time.perf_counter()
        try:
            response = await self._create_completion(
                messages=messages,
                model=resolved_model,
                stream=True,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        async for chunk in response:
            if not chunk.choices:
                continue
            delta = chunk.choices[0].delta
            if delta and delta.content:
                yield delta.content

        latency_ms = int((time.perf_counter() - started) * 1000)
        logger.info(
            "LLM stream complete model={model} latency_ms={latency_ms}",
            model=resolved_model,
            latency_ms=latency_ms,
        )

    async def list_models(self) -> list[ModelInfo]:
        try:
            page = await self._client.models.list()
        except Exception as exc:
            raise self._map_error(exc) from exc

        return [
            ModelInfo(id=model.id, owned_by=model.owned_by, created=model.created)
            for model in page.data
        ]
