"""Anthropic (Claude) provider via the official AsyncAnthropic SDK."""

from __future__ import annotations

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

from anthropic import (
    APIConnectionError,
    APIError,
    APITimeoutError,
    AsyncAnthropic,
    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
from app.llm.providers._translate import (
    parse_anthropic_tool_calls,
    to_anthropic_messages,
    to_anthropic_tools,
)
from app.llm.types import (
    ChatCompletionResult,
    LLMAuthError,
    LLMClientError,
    LLMInvalidModelError,
    LLMMalformedOutputError,
    LLMRateLimitError,
    LLMTimeoutError,
    ModelInfo,
    ToolCompletionResult,
)

_DEFAULT_MAX_TOKENS = 4096


class AnthropicProvider:
    """Anthropic Messages API provider."""

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

        self._default_model = default_model or "claude-sonnet-4-20250514"
        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._max_tokens = max_tokens
        self._client = client or AsyncAnthropic(
            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))

    @staticmethod
    def _extract_text(content: Any) -> str:
        parts: list[str] = []
        for block in content or []:
            block_type = (
                block.get("type") if isinstance(block, dict) else getattr(block, "type", None)
            )
            if block_type == "text":
                text = (
                    block.get("text") if isinstance(block, dict) else getattr(block, "text", "")
                )
                if text:
                    parts.append(str(text))
        return "".join(parts)

    @staticmethod
    def _map_tool_choice(
        tool_choice: str | dict[str, Any] | None,
    ) -> dict[str, Any] | None:
        if tool_choice is None:
            return None
        if isinstance(tool_choice, str):
            if tool_choice == "auto":
                return {"type": "auto"}
            if tool_choice == "none":
                return {"type": "none"}
            if tool_choice == "required":
                return {"type": "any"}
            return {"type": "auto"}
        if isinstance(tool_choice, dict):
            # OpenAI-style {"type":"function","function":{"name":"..."}}
            if tool_choice.get("type") == "function":
                name = (tool_choice.get("function") or {}).get("name")
                if name:
                    return {"type": "tool", "name": name}
            return tool_choice
        return None

    @staticmethod
    def _map_stop_reason(stop_reason: str | None) -> str:
        if stop_reason == "tool_use":
            return "tool_calls"
        if stop_reason == "end_turn":
            return "stop"
        return stop_reason or ""

    @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_message(
        self,
        *,
        model: str,
        messages: list[dict[str, Any]],
        system: str | None = None,
        tools: list[dict[str, Any]] | None = None,
        tool_choice: dict[str, Any] | None = None,
        max_tokens: int | None = None,
    ):
        kwargs: dict[str, Any] = {
            "model": model,
            "messages": messages,
            "max_tokens": max_tokens or self._max_tokens,
        }
        if system:
            kwargs["system"] = system
        if tools is not None:
            kwargs["tools"] = tools
        if tool_choice is not None:
            kwargs["tool_choice"] = tool_choice
        return await self._client.messages.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
        system, anthropic_messages = to_anthropic_messages(list(messages))
        if response_format and response_format.get("type") == "json_object":
            hint = "Respond with a single valid JSON object and no other text."
            system = f"{system}\n\n{hint}" if system else hint

        started = time.perf_counter()
        try:
            response = await self._create_message(
                model=resolved_model,
                messages=anthropic_messages,
                system=system,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        latency_ms = int((time.perf_counter() - started) * 1000)
        text = self._extract_text(response.content)
        usage = getattr(response, "usage", None)
        prompt_tokens = getattr(usage, "input_tokens", None) if usage else None
        completion_tokens = getattr(usage, "output_tokens", None) if usage else None

        logger.info(
            "Anthropic 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, anthropic_messages = to_anthropic_messages(messages)
        anthropic_tools = to_anthropic_tools(tools)
        mapped_choice = self._map_tool_choice(tool_choice)

        started = time.perf_counter()
        try:
            response = await self._create_message(
                model=resolved_model,
                messages=anthropic_messages,
                system=system,
                tools=anthropic_tools,
                tool_choice=mapped_choice,
            )
        except Exception as exc:
            raise self._map_error(exc) from exc

        latency_ms = int((time.perf_counter() - started) * 1000)
        text = self._extract_text(response.content)
        tool_calls = parse_anthropic_tool_calls(response.content)
        finish_reason = self._map_stop_reason(getattr(response, "stop_reason", None))

        usage = getattr(response, "usage", None)
        prompt_tokens = getattr(usage, "input_tokens", None) if usage else None
        completion_tokens = getattr(usage, "output_tokens", None) if usage else None

        logger.info(
            "Anthropic 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, anthropic_messages = to_anthropic_messages(list(messages))
        started = time.perf_counter()
        try:
            async with self._client.messages.stream(
                model=resolved_model,
                messages=anthropic_messages,
                max_tokens=self._max_tokens,
                **({"system": system} if system else {}),
            ) as stream:
                async for text in stream.text_stream:
                    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(
            "Anthropic 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.models.list()
        except Exception as exc:
            raise self._map_error(exc) from exc

        models: list[ModelInfo] = []
        async for model in pager:
            model_id = getattr(model, "id", None) or str(model)
            created = getattr(model, "created_at", None)
            if hasattr(created, "timestamp"):
                created = int(created.timestamp())
            models.append(
                ModelInfo(
                    id=model_id,
                    owned_by="anthropic",
                    created=created if isinstance(created, int) else None,
                )
            )
        return models
