from __future__ import annotations

from qdrant_client import QdrantClient
from qdrant_client.http import models

from app.core.config import settings
from app.core.logger import logger
from app.qdrant.client import get_qdrant_client

PAYLOAD_INDEX_FIELDS = ("repo_id", "snapshot_id", "source_type", "record_id")


class QdrantBootstrapError(RuntimeError):
    """Raised when Qdrant collection bootstrap fails."""


def _collection_exists(client: QdrantClient, name: str) -> bool:
    collections = client.get_collections().collections
    return any(item.name == name for item in collections)


def _validate_existing_collection(client: QdrantClient, name: str) -> None:
    info = client.get_collection(name)
    vectors = info.config.params.vectors
    if isinstance(vectors, dict):
        vector_params = next(iter(vectors.values()))
    else:
        vector_params = vectors

    if vector_params.size != settings.EMBEDDING_DIMENSION:
        raise QdrantBootstrapError(
            f"Qdrant collection {name!r} vector size {vector_params.size} "
            f"does not match EMBEDDING_DIMENSION={settings.EMBEDDING_DIMENSION}"
        )

    if vector_params.distance != models.Distance.COSINE:
        raise QdrantBootstrapError(
            f"Qdrant collection {name!r} distance must be Cosine, got {vector_params.distance}"
        )


def _ensure_payload_indexes(client: QdrantClient, collection_name: str) -> None:
    for field_name in PAYLOAD_INDEX_FIELDS:
        try:
            client.create_payload_index(
                collection_name=collection_name,
                field_name=field_name,
                field_schema=models.PayloadSchemaType.KEYWORD,
            )
            logger.info(
                "Created Qdrant payload index collection={collection} field={field}",
                collection=collection_name,
                field=field_name,
            )
        except Exception as exc:
            if "already exists" in str(exc).lower():
                continue
            logger.warning(
                "Qdrant payload index create skipped collection={collection} field={field} error={error}",
                collection=collection_name,
                field=field_name,
                error=str(exc),
            )


def bootstrap_qdrant(client: QdrantClient | None = None) -> str:
    """
    Ensure the embeddings collection exists with the configured vector schema.

    Returns the collection name.
    """
    client = client or get_qdrant_client()
    collection_name = settings.QDRANT_COLLECTION_NAME

    if _collection_exists(client, collection_name):
        _validate_existing_collection(client, collection_name)
        logger.info(
            "Qdrant collection already exists collection={collection}",
            collection=collection_name,
        )
    else:
        client.create_collection(
            collection_name=collection_name,
            vectors_config=models.VectorParams(
                size=settings.EMBEDDING_DIMENSION,
                distance=models.Distance.COSINE,
            ),
        )
        logger.info(
            "Created Qdrant collection collection={collection} size={size}",
            collection=collection_name,
            size=settings.EMBEDDING_DIMENSION,
        )

    _ensure_payload_indexes(client, collection_name)
    return collection_name
