"""LLM provider factory functions.

Creates OpenAI SDK clients and LangChain chat models for either
Azure OpenAI or vLLM backends based on settings.llm_provider /
settings.embedding_provider.

vLLM exposes an OpenAI-compatible API, so the openai Python SDK
works directly with base_url pointed at the vLLM server.
"""

from __future__ import annotations

import logging
from functools import lru_cache
from typing import Callable

import httpx
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI

from src.core.config import settings
from src.resilience.breakers import async_breaker_call, breaker_call, llm_breaker

logger = logging.getLogger(__name__)


# ---------------------------------------------------------------------------
# Chat clients (OpenAI SDK)
# ---------------------------------------------------------------------------

def _vllm_base_url(url: str) -> str:
    """Ensure vLLM base_url ends with /v1 (OpenAI SDK expects it)."""
    return url.rstrip("/") + "/v1" if not url.rstrip("/").endswith("/v1") else url


def _fork_safe_http_client(**overrides) -> httpx.Client:
    defaults = dict(
        http2=False,
        limits=httpx.Limits(max_connections=1, max_keepalive_connections=0),
        timeout=httpx.Timeout(30.0, connect=5.0),
    )
    defaults.update(overrides)
    return httpx.Client(**defaults)


def _fork_safe_async_http_client(**overrides) -> httpx.AsyncClient:
    defaults = dict(
        http2=False,
        limits=httpx.Limits(
            max_connections=settings.llm_concurrency,
            max_keepalive_connections=settings.llm_concurrency,
        ),
        timeout=httpx.Timeout(30.0, connect=5.0),
    )
    defaults.update(overrides)
    return httpx.AsyncClient(**defaults)


def get_chat_client(*, http_client: httpx.Client | None = None) -> OpenAI | AzureOpenAI:
    """Return a sync chat client for the configured LLM provider."""
    hc = http_client or _fork_safe_http_client()
    if settings.llm_provider == "vllm":
        base_url = _vllm_base_url(settings.llm_base_url)
        logger.info("Chat client: provider=vllm model=%s", settings.llm_model_name)
        return OpenAI(
            base_url=base_url,
            api_key=settings.llm_api_key or "dummy",
            default_headers={"X-API-Key": settings.llm_api_key} if settings.llm_api_key else {},
            http_client=hc,
        )
    logger.info("Chat client: provider=azure deployment=%s", settings.azure_openai_deployment)
    return AzureOpenAI(
        api_key=settings.azure_openai_api_key,
        api_version=settings.azure_openai_api_version,
        azure_endpoint=settings.azure_openai_endpoint,
        http_client=hc,
    )


def get_async_chat_client(
    *, http_client: httpx.AsyncClient | None = None
) -> AsyncOpenAI | AsyncAzureOpenAI:
    """Return an async chat client for the configured LLM provider."""
    hc = http_client or _fork_safe_async_http_client()
    if settings.llm_provider == "vllm":
        base_url = _vllm_base_url(settings.llm_base_url)
        logger.info("Async chat client: provider=vllm model=%s", settings.llm_model_name)
        return AsyncOpenAI(
            base_url=base_url,
            api_key=settings.llm_api_key or "dummy",
            default_headers={"X-API-Key": settings.llm_api_key} if settings.llm_api_key else {},
            http_client=hc,
        )
    logger.info("Async chat client: provider=azure deployment=%s", settings.azure_openai_deployment)
    return AsyncAzureOpenAI(
        api_key=settings.azure_openai_api_key,
        api_version=settings.azure_openai_api_version,
        azure_endpoint=settings.azure_openai_endpoint,
        http_client=hc,
    )


def get_chat_model_name() -> str:
    """Return the model/deployment name for chat.completions.create(model=...)."""
    if settings.llm_provider == "vllm":
        return settings.llm_model_name
    return settings.azure_openai_deployment


# ---------------------------------------------------------------------------
# LangChain chat service
# ---------------------------------------------------------------------------

def get_chat_service(temperature: float | None = None):
    """Return a provider-agnostic ChatService wrapping LangChain."""
    from src.llm.chat_service import ChatService

    temp = temperature if temperature is not None else settings.azure_openai_temperature
    max_tokens = settings.azure_openai_max_tokens

    if settings.llm_provider == "vllm":
        from langchain_openai import ChatOpenAI

        base_url = _vllm_base_url(settings.llm_base_url)
        logger.info("Chat service: provider=vllm model=%s temperature=%s max_tokens=%s", settings.llm_model_name, temp, max_tokens)
        llm = ChatOpenAI(
            base_url=base_url,
            api_key=settings.llm_api_key or "dummy",
            model=settings.llm_model_name,
            temperature=temp,
            max_completion_tokens=max_tokens,
            default_headers={"X-API-Key": settings.llm_api_key} if settings.llm_api_key else {},
        )
    else:
        from langchain_openai import AzureChatOpenAI

        logger.info("Chat service: provider=azure deployment=%s temperature=%s max_tokens=%s", settings.azure_openai_deployment, temp, max_tokens)
        llm = AzureChatOpenAI(
            azure_endpoint=settings.azure_openai_endpoint,
            api_key=settings.azure_openai_api_key,
            api_version=settings.azure_openai_api_version,
            azure_deployment=settings.azure_openai_deployment,
            temperature=temp,
            max_completion_tokens=max_tokens,
        )

    return ChatService(llm)


# ---------------------------------------------------------------------------
# Embedding clients
# ---------------------------------------------------------------------------

def get_embed_model():
    """Return a LlamaIndex embedding model for the configured provider."""
    if settings.embedding_provider == "vllm":
        from llama_index.embeddings.openai import OpenAIEmbedding

        base_url = _vllm_base_url(settings.embedding_base_url)
        logger.info("Embed model: provider=vllm model=%s dimensions=%s batch_size=16", settings.embedding_model_name, settings.embedding_dimensions)
        return OpenAIEmbedding(
            api_base=base_url,
            api_key=settings.embedding_api_key or "dummy",
            model_name=settings.embedding_model_name,
            embed_batch_size=16,
        )

    from llama_index.embeddings.azure_openai import AzureOpenAIEmbedding

    logger.info("Embed model: provider=azure deployment=%s dimensions=%s batch_size=16", settings.azure_openai_embedding_deployment, settings.embedding_dimensions)
    return AzureOpenAIEmbedding(
        azure_endpoint=settings.azure_openai_endpoint,
        api_key=settings.azure_openai_api_key,
        api_version=settings.azure_openai_api_version,
        azure_deployment=settings.azure_openai_embedding_deployment,
        model=settings.azure_openai_embedding_deployment,
        dimensions=settings.embedding_dimensions,
        embed_batch_size=16,
    )


def get_async_embedding_client() -> AsyncOpenAI | AsyncAzureOpenAI:
    """Return an async embedding client for the configured provider."""
    if settings.embedding_provider == "vllm":
        base_url = _vllm_base_url(settings.embedding_base_url)
        logger.info("Async embed client: provider=vllm model=%s dimensions=%s", settings.embedding_model_name, settings.embedding_dimensions)
        return AsyncOpenAI(
            base_url=base_url,
            api_key=settings.embedding_api_key or "dummy",
        )
    logger.info("Async embed client: provider=azure deployment=%s dimensions=%s", settings.azure_openai_embedding_deployment, settings.embedding_dimensions)
    return AsyncAzureOpenAI(
        api_key=settings.azure_openai_api_key,
        api_version=settings.azure_openai_api_version,
        azure_endpoint=settings.azure_openai_endpoint,
    )


def get_embedding_deployment() -> str:
    """Return the model/deployment name for embeddings.create(model=...)."""
    if settings.embedding_provider == "vllm":
        return settings.embedding_model_name
    return settings.azure_openai_embedding_deployment


# ---------------------------------------------------------------------------
# Breaker-wrapped public call helpers (TEXTIQ-394 / SYS-NFR-20)
#
# Every LLM SDK call should go through one of these helpers so that the shared
# llm_breaker is updated on every failure. This keeps the breaker policy in
# one place and means call sites never touch pybreaker directly.
# ---------------------------------------------------------------------------


def create_chat_completion(client, **kwargs):
    """Sync ``client.chat.completions.create(**kwargs)`` routed through llm_breaker."""
    return breaker_call(llm_breaker, client.chat.completions.create, **kwargs)


async def async_create_chat_completion(client, **kwargs):
    """Async ``client.chat.completions.create(**kwargs)`` routed through llm_breaker."""
    return await async_breaker_call(llm_breaker, client.chat.completions.create, **kwargs)


async def async_create_embedding(client, **kwargs):
    """Async ``client.embeddings.create(**kwargs)`` routed through llm_breaker."""
    return await async_breaker_call(llm_breaker, client.embeddings.create, **kwargs)


# ---------------------------------------------------------------------------
# Tokenizer (for chunk-size token budgeting)
# ---------------------------------------------------------------------------

@lru_cache(maxsize=1)
def _build_tokenizer() -> Callable[[str], list]:
    """Build a `text -> list-of-tokens` callable for the configured provider.

    Returned callable's output length equals the token count (compatible with
    LlamaIndex SentenceSplitter.tokenizer and LangChain length_function=len).
    """
    if settings.embedding_provider == "vllm":
        try:
            from transformers import AutoTokenizer

            tok = AutoTokenizer.from_pretrained(settings.embedding_model_name)
            logger.info("Tokenizer: provider=vllm tokenizer=%s", settings.embedding_model_name)
            return lambda t: tok.encode(t, add_special_tokens=False)
        except Exception as e:
            logger.warning(
                "HF tokenizer load failed for %r (%s); using char/4 fallback",
                settings.embedding_model_name, e,
            )
            # Empty string -> 0 tokens. Non-empty -> ceil(len/4).
            return lambda t: [0] * (-(-len(t) // 4)) if t else []

    import tiktoken

    enc = tiktoken.get_encoding("cl100k_base")
    logger.info("Tokenizer: provider=azure tokenizer=cl100k_base")
    return lambda t: enc.encode(t)


@lru_cache(maxsize=1)
def _build_token_counter() -> Callable[[str], int]:
    tok = _build_tokenizer()
    return lambda t: len(tok(t))


@lru_cache(maxsize=4)
def _build_docling_tokenizer(max_tokens: int):
    """Build a docling-core tokenizer wrapper for HybridChunker.

    Cached per `max_tokens` so repeated `SemanticChunker(max_tokens=...)`
    instantiations in the same worker reuse the downloaded HF tokenizer.
    """
    from docling_core.transforms.chunker.tokenizer.openai import OpenAITokenizer

    if settings.embedding_provider == "vllm":
        try:
            from docling_core.transforms.chunker.tokenizer.huggingface import (
                HuggingFaceTokenizer,
            )
            from transformers import AutoTokenizer

            hf_tok = AutoTokenizer.from_pretrained(settings.embedding_model_name)
            logger.info(
                "Docling tokenizer: provider=vllm tokenizer=%s max_tokens=%s",
                settings.embedding_model_name, max_tokens,
            )
            return HuggingFaceTokenizer(tokenizer=hf_tok, max_tokens=max_tokens)
        except Exception as e:
            logger.warning(
                "HF tokenizer load failed for %r (%s); falling back to tiktoken cl100k_base",
                settings.embedding_model_name, e,
            )

    import tiktoken

    logger.info("Docling tokenizer: provider=azure tokenizer=cl100k_base max_tokens=%s", max_tokens)
    return OpenAITokenizer(
        tokenizer=tiktoken.get_encoding("cl100k_base"),
        max_tokens=max_tokens,
    )


def get_tokenizer() -> Callable[[str], list]:
    """Return a provider-aware tokenizer returning a list of token IDs.

    Output length equals token count. Compatible with LlamaIndex
    SentenceSplitter.tokenizer and LangChain length_function=len.
    """
    return _build_tokenizer()


def get_token_counter() -> Callable[[str], int]:
    """Return a provider-aware `text -> token_count` callable."""
    return _build_token_counter()


def get_docling_tokenizer(max_tokens: int = 8000):
    """Return a docling-core tokenizer wrapper for HybridChunker.

    Note: docling's HuggingFaceTokenizer counts via `tokenizer.tokenize()`
    while get_tokenizer() counts via `tokenizer.encode(add_special_tokens=False)`.
    For most HF models these agree to within a few tokens; callers that need
    byte-identical cross-path counts should use only one entry point.
    """
    return _build_docling_tokenizer(max_tokens)
