"""Google Gemini provider via the google-genai SDK."""

from __future__ import annotations

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

from google import genai
from google.genai import errors as genai_errors
from google.genai import types
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
from app.llm.providers._translate import (
    parse_gemini_tool_calls,
    to_gemini_contents,
    to_gemini_function_declarations,
)
from app.llm.types import (
    ChatCompletionResult,
    LLMAuthError,
    LLMClientError,
    LLMInvalidModelError,
    LLMMalformedOutputError,
    LLMRateLimitError,
    LLMTimeoutError,
    ModelInfo,
    ToolCompletionResult,
)


class GeminiProvider:
    """Gemini generate_content provider."""

    RETRYABLE = (genai_errors.ServerError,)

    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: genai.Client | None = None,
    ):
        key = api_key if api_key is not None else settings.GEMINI_API_KEY
        if not str(key).strip():
            if settings.is_local():
                key = "local-dev"
            else:
                raise LLMAuthError(
                    "GEMINI_API_KEY is required outside local/dev",
                    status_code=401,
                )

        self._default_model = default_model or "gemini-2.5-flash"
        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
        # google-genai HttpOptions.timeout is milliseconds.
        http_options = types.HttpOptions(timeout=self._timeout_sec * 1000)
        self._client = client or genai.Client(api_key=key, http_options=http_options)

    def _map_error(self, exc: Exception) -> LLMClientError:
        if isinstance(exc, LLMClientError):
            return exc
        if isinstance(exc, genai_errors.ClientError):
            code = int(getattr(exc, "code", 0) or 0)
            message = str(exc)
            lowered = message.lower()
            if code in {401, 403} or "api key" in lowered or "permission" in lowered:
                return LLMAuthError(message, status_code=code or 401)
            if code == 429 or "rate" in lowered or "quota" in lowered:
                return LLMRateLimitError(message, status_code=429)
            if code == 404 or "model" in lowered and "not found" in lowered:
                return LLMInvalidModelError(message, status_code=code or 404)
            if code == 400:
                if "model" in lowered:
                    return LLMInvalidModelError(message, status_code=400)
                return LLMClientError(message, status_code=400)
            return LLMClientError(message, status_code=code or None)
        if isinstance(exc, genai_errors.ServerError):
            code = int(getattr(exc, "code", 0) or 0)
            if code == 504 or "timeout" in str(exc).lower():
                return LLMTimeoutError(str(exc), status_code=504)
            return LLMClientError(str(exc), status_code=code or 500)
        if isinstance(exc, TimeoutError):
            return LLMTimeoutError(str(exc), status_code=504)
        return LLMClientError(str(exc))

    @staticmethod
    def _extract_text(response: Any) -> str:
        text = getattr(response, "text", None)
        if text:
            return str(text)
        candidates = getattr(response, "candidates", None) or []
        parts_text: list[str] = []
        for candidate in candidates:
            content = getattr(candidate, "content", None)
            parts = getattr(content, "parts", None) or []
            for part in parts:
                part_text = getattr(part, "text", None)
                if part_text:
                    parts_text.append(str(part_text))
        return "".join(parts_text)

    @staticmethod
    def _response_parts(response: Any) -> list[Any]:
        candidates = getattr(response, "candidates", None) or []
        if not candidates:
            return []
        content = getattr(candidates[0], "content", None)
        return list(getattr(content, "parts", None) or [])

    @staticmethod
    def _map_finish_reason(response: Any) -> str:
        candidates = getattr(response, "candidates", None) or []
        if not candidates:
            return ""
        reason = getattr(candidates[0], "finish_reason", None)
        if reason is None:
            return ""
        name = getattr(reason, "name", None) or str(reason)
        lowered = str(name).lower()
        if "stop" in lowered:
            return "stop"
        if "tool" in lowered or "function" in lowered:
            return "tool_calls"
        if "max" in lowered:
            return "length"
        return lowered

    @staticmethod
    def _usage_tokens(response: Any) -> tuple[int | None, int | None]:
        usage = getattr(response, "usage_metadata", None)
        if not usage:
            return None, None
        prompt = getattr(usage, "prompt_token_count", None)
        completion = getattr(usage, "candidates_token_count", None)
        return prompt, completion

    def _build_tools_config(
        self,
        tools: list[dict[str, Any]] | None,
        *,
        response_mime_type: str | None = None,
        system: str | None = None,
        tool_choice: str | dict[str, Any] | None = None,
    ) -> types.GenerateContentConfig:
        kwargs: dict[str, Any] = {
            "automatic_function_calling": types.AutomaticFunctionCallingConfig(
                disable=True
            ),
        }
        if system:
            kwargs["system_instruction"] = system
        if response_mime_type:
            kwargs["response_mime_type"] = response_mime_type
        if tools is not None:
            declarations = [
                types.FunctionDeclaration(
                    name=item.get("name"),
                    description=item.get("description"),
                    parameters_json_schema=item.get("parameters"),
                )
                for item in to_gemini_function_declarations(tools)
            ]
            kwargs["tools"] = [types.Tool(function_declarations=declarations)]
            mode = "AUTO"
            if tool_choice == "none":
                mode = "NONE"
            elif tool_choice == "required":
                mode = "ANY"
            elif isinstance(tool_choice, dict):
                mode = "ANY"
            kwargs["tool_config"] = types.ToolConfig(
                function_calling_config=types.FunctionCallingConfig(mode=mode)
            )
        return types.GenerateContentConfig(**kwargs)

    @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 _generate(
        self,
        *,
        model: str,
        contents: list[dict[str, Any]],
        config: types.GenerateContentConfig | None = None,
    ):
        return await self._client.aio.models.generate_content(
            model=model,
            contents=contents,
            config=config,
        )

    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
        system, contents = to_gemini_contents(list(messages))
        mime = None
        if response_format and response_format.get("type") == "json_object":
            mime = "application/json"
        config = self._build_tools_config(
            None,
            response_mime_type=mime,
            system=system,
        )

        started = time.perf_counter()
        try:
            response = await self._generate(
                model=resolved_model,
                contents=contents,
                config=config,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        latency_ms = int((time.perf_counter() - started) * 1000)
        text = self._extract_text(response)
        prompt_tokens, completion_tokens = self._usage_tokens(response)

        logger.info(
            "Gemini 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
        system, contents = to_gemini_contents(messages)
        config = self._build_tools_config(
            tools,
            system=system,
            tool_choice=tool_choice,
        )

        started = time.perf_counter()
        try:
            response = await self._generate(
                model=resolved_model,
                contents=contents,
                config=config,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        latency_ms = int((time.perf_counter() - started) * 1000)
        text = self._extract_text(response)
        parts = self._response_parts(response)
        tool_calls = parse_gemini_tool_calls(parts)
        # Prefer response.function_calls when the SDK exposes it.
        if not tool_calls and getattr(response, "function_calls", None):
            tool_calls = parse_gemini_tool_calls(response.function_calls)
        finish_reason = self._map_finish_reason(response)
        if tool_calls and finish_reason in {"", "stop"}:
            finish_reason = "tool_calls"

        prompt_tokens, completion_tokens = self._usage_tokens(response)

        logger.info(
            "Gemini complete_with_tools model={model} latency_ms={latency_ms} "
            "tool_calls={tool_calls} finish_reason={finish_reason} text_len={text_len}",
            model=resolved_model,
            latency_ms=latency_ms,
            tool_calls=len(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
        system, contents = to_gemini_contents(list(messages))
        config = self._build_tools_config(None, system=system)
        started = time.perf_counter()
        try:
            stream = await self._client.aio.models.generate_content_stream(
                model=resolved_model,
                contents=contents,
                config=config,
            )
            async for chunk in stream:
                text = getattr(chunk, "text", None)
                if text:
                    yield text
        except Exception as exc:
            raise self._map_error(exc) from exc

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

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

        models: list[ModelInfo] = []
        async for model in pager:
            name = getattr(model, "name", None) or getattr(model, "id", None) or ""
            model_id = str(name).removeprefix("models/")
            if not model_id:
                continue
            models.append(ModelInfo(id=model_id, owned_by="google"))
        return models
