"""Embedding service for converting records to vector representations."""

import asyncio
import time
from typing import Any

import structlog
from openai import AsyncAzureOpenAI
from openai.types import CreateEmbeddingResponse

from src.core.config import settings
from src.llm import get_async_embedding_client, get_embedding_deployment
from src.llm.factory import async_create_embedding
from src.core.exceptions import IndexingError
from src.extraction.models import ExtractedRecord

logger = structlog.get_logger()


def record_to_text(record: ExtractedRecord) -> str:
    """Convert ExtractedRecord to text representation for embedding.

    Creates a text representation that captures:
    1. Content fields: "Key1: Value1. Key2: Value2. ..."
    2. Resolved symbol meanings (if present): "Key_resolved: Meaning"
    3. Source context: "From: filename, Sheet: sheet_name, Row: row_number"

    Args:
        record: Extracted record with content, headers, and source metadata

    Returns:
        Text representation suitable for embedding

    Example:
        >>> record = ExtractedRecord(
        ...     content={"Item Code": "2024-4_0019", "Report": "not applicable"},
        ...     headers=["Item Code", "Report"],
        ...     _source={
        ...         "filename": "distribution_spec.xlsx",
        ...         "sheet": "Item Mapping",
        ...         "row": 42,
        ...     }
        ... )
        >>> text = record_to_text(record)
        >>> print(text)
        Item Code: 2024-4_0019. Report: not applicable. From: distribution_spec.xlsx, Sheet: Item Mapping, Row: 42
    """
    content_parts = []

    # Add content fields (skip None values)
    for key, value in record.content.items():
        if value is not None:
            content_parts.append(f"{key}: {value}")

    # Add resolved fields if symbols were resolved
    resolved_content = getattr(record, "resolved_content", None)
    if resolved_content:
        for key, value in resolved_content.items():
            # Only add fields that end with _resolved and have non-None values
            if "_resolved" in key and value is not None:
                content_parts.append(f"{key}: {value}")

    # Add source context for citation
    source = record._source
    source_text = (
        f"From: {source['filename']}, "
        f"Sheet: {source['sheet']}, "
        f"Row: {source['row']}"
    )

    # Combine all parts
    if content_parts:
        return f"{'. '.join(content_parts)}. {source_text}"
    else:
        # Edge case: empty content (shouldn't happen but handle gracefully)
        return source_text


class EmbeddingService:
    """Service for converting text to embeddings using Azure OpenAI.

    Uses Azure OpenAI's text-embedding-3-large model to generate 3072-dimensional
    vector embeddings for semantic search.
    """

    def __init__(self) -> None:
        """Initialize embedding service configuration."""
        self.logger = structlog.get_logger()
        self.deployment = get_embedding_deployment()
        self.dimensions = settings.embedding_dimensions

    def _create_client(self):
        """Create a fresh async embedding client.

        Creates a new client instance to avoid event loop binding issues
        when called from different async contexts (e.g., Celery workers).
        """
        return get_async_embedding_client()

    async def close(self) -> None:
        """Close the underlying HTTP client to prevent resource leaks."""
        await self.client.close()

    async def embed_batch(
        self,
        texts: list[str],
        max_retries: int = 3,
    ) -> list[list[float]]:
        """Embed batch of texts using Azure OpenAI text-embedding-3-large.

        Processes up to 100 texts per API call for efficiency. Uses exponential
        backoff retry logic for transient errors.

        Args:
            texts: List of text representations (max 100 per call)
            max_retries: Number of retries for transient errors (default: 3)

        Returns:
            List of embedding vectors, each with 3072 dimensions

        Raises:
            IndexingError: If embedding fails after all retries
            ValueError: If batch size exceeds 100 texts

        Example:
            >>> service = EmbeddingService()
            >>> texts = ["Item Code: 2024-4_0019. Report: not applicable"]
            >>> embeddings = await service.embed_batch(texts)
            >>> len(embeddings)
            1
            >>> len(embeddings[0])
            3072
        """
        if len(texts) > 100:
            raise ValueError(
                f"Batch size must be ≤ 100, got {len(texts)}. "
                "Split into smaller batches."
            )

        if not texts:
            raise ValueError("Cannot embed empty batch")

        start_time = time.perf_counter()

        for attempt in range(max_retries):
            # Create fresh client for each attempt to avoid event loop issues
            # Use async context manager to properly close HTTP connections
            client = self._create_client()
            try:
                # Call embeddings API
                # vLLM may not support the dimensions parameter
                kwargs: dict = dict(model=self.deployment, input=texts)
                if settings.embedding_provider != "vllm":
                    kwargs["dimensions"] = self.dimensions
                response: CreateEmbeddingResponse = await async_create_embedding(client, **kwargs)

                # Extract embeddings from response
                embeddings = [item.embedding for item in response.data]

                # Calculate timing metrics
                duration_ms = (time.perf_counter() - start_time) * 1000

                self.logger.info(
                    "embeddings_created",
                    batch_size=len(texts),
                    dimensions=self.dimensions,
                    duration_ms=round(duration_ms, 2),
                    texts_per_second=round(len(texts) / (duration_ms / 1000), 2),
                    attempt=attempt + 1,
                )

                return embeddings

            except Exception as e:
                error_msg = str(e)
                is_last_attempt = attempt == max_retries - 1

                if is_last_attempt:
                    self.logger.error(
                        "embedding_failed_permanently",
                        batch_size=len(texts),
                        error=error_msg,
                        attempts=max_retries,
                    )
                    raise IndexingError(
                        f"Failed to embed batch after {max_retries} attempts: {error_msg}",
                        details={
                            "batch_size": len(texts),
                            "attempts": max_retries,
                            "error": error_msg,
                        },
                    ) from e

                # Exponential backoff: 1s, 2s, 4s
                backoff_seconds = 2**attempt

                self.logger.warning(
                    "embedding_attempt_failed_retrying",
                    batch_size=len(texts),
                    error=error_msg,
                    attempt=attempt + 1,
                    max_retries=max_retries,
                    backoff_seconds=backoff_seconds,
                )

                await asyncio.sleep(backoff_seconds)
            finally:
                # Ensure HTTP client is properly closed to avoid event loop binding issues
                await client.close()

        # Should never reach here due to raise in loop, but for type safety
        raise IndexingError("Embedding failed unexpectedly")

    async def embed_records(
        self,
        records: list[ExtractedRecord],
    ) -> list[list[float]]:
        """Embed a batch of ExtractedRecords.

        Convenience method that converts records to text and embeds them.

        Args:
            records: List of ExtractedRecord objects (max 100)

        Returns:
            List of embedding vectors, one per record

        Raises:
            IndexingError: If embedding fails
            ValueError: If batch size exceeds 100

        Example:
            >>> service = EmbeddingService()
            >>> records = [...]  # List of ExtractedRecord objects
            >>> embeddings = await service.embed_records(records)
        """
        if len(records) > 100:
            raise ValueError(
                f"Batch size must be ≤ 100, got {len(records)}. "
                "Split into smaller batches."
            )

        # Convert records to text
        texts = [record_to_text(record) for record in records]

        # Embed the texts
        return await self.embed_batch(texts)
