from __future__ import annotations

from dataclasses import dataclass
from typing import Sequence

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.models.embedding_record import EmbeddingRecord
from app.qdrant.client import get_qdrant_client
from app.qdrant.payload import embedding_record_to_qdrant_payload
from app.qdrant.point_id import qdrant_point_id


class QdrantVectorStoreError(RuntimeError):
    """Raised when Qdrant vector upsert/delete fails."""


@dataclass(frozen=True)
class VectorUpsertItem:
    record: EmbeddingRecord
    vector: list[float]


@dataclass(frozen=True)
class VectorUpsertResult:
    record_id: str
    point_id: str


@dataclass(frozen=True)
class VectorSearchHit:
    """Ranked Qdrant hit with payload metadata (no embedded text)."""

    record_id: str
    score: float
    source_type: str
    chunk_type: str
    graph_node_id: str | None = None
    file_path: str | None = None
    doc_path: str | None = None
    symbol_name: str | None = None
    section_title: str | None = None
    commit_sha: str | None = None


def _validate_vector(vector: Sequence[float], record_id: str) -> list[float]:
    if len(vector) != settings.EMBEDDING_DIMENSION:
        raise QdrantVectorStoreError(
            f"Vector for {record_id!r} has dimension {len(vector)}, "
            f"expected {settings.EMBEDDING_DIMENSION}"
        )
    return list(vector)


def _build_point(item: VectorUpsertItem) -> models.PointStruct:
    point_id = qdrant_point_id(item.record.record_id)
    return models.PointStruct(
        id=point_id,
        vector=_validate_vector(item.vector, item.record.record_id),
        payload=embedding_record_to_qdrant_payload(item.record),
    )


def upsert_vectors(
    items: list[VectorUpsertItem],
    *,
    client: QdrantClient | None = None,
    collection_name: str | None = None,
    batch_size: int | None = None,
) -> list[VectorUpsertResult]:
    """
    Upsert embedding vectors into Qdrant.

    Uses deterministic point IDs from record_id so retries overwrite the same point.
    """
    if not items:
        return []

    client = client or get_qdrant_client()
    collection_name = collection_name or settings.QDRANT_COLLECTION_NAME
    batch_size = batch_size or settings.EMBEDDING_BATCH_SIZE

    results: list[VectorUpsertResult] = []
    for offset in range(0, len(items), batch_size):
        batch = items[offset : offset + batch_size]
        points = [_build_point(item) for item in batch]
        try:
            client.upsert(
                collection_name=collection_name,
                points=points,
                wait=True,
            )
        except Exception as exc:
            raise QdrantVectorStoreError(
                f"Qdrant upsert failed for batch starting at offset {offset}: {exc}"
            ) from exc

        for item, point in zip(batch, points):
            results.append(
                VectorUpsertResult(
                    record_id=item.record.record_id,
                    point_id=str(point.id),
                )
            )

    logger.info(
        "Upserted vectors to Qdrant collection={collection} count={count}",
        collection=collection_name,
        count=len(results),
    )
    return results


def delete_vectors_by_record_ids(
    record_ids: list[str],
    *,
    client: QdrantClient | None = None,
    collection_name: str | None = None,
) -> int:
    """Delete Qdrant points by embedding record_id (deterministic point IDs)."""
    if not record_ids:
        return 0

    client = client or get_qdrant_client()
    collection_name = collection_name or settings.QDRANT_COLLECTION_NAME
    point_ids = [qdrant_point_id(record_id) for record_id in record_ids]

    try:
        client.delete(
            collection_name=collection_name,
            points_selector=models.PointIdsList(points=point_ids),
            wait=True,
        )
    except Exception as exc:
        raise QdrantVectorStoreError(
            f"Qdrant delete by record_ids failed: {exc}"
        ) from exc

    logger.info(
        "Deleted vectors from Qdrant collection={collection} count={count}",
        collection=collection_name,
        count=len(point_ids),
    )
    return len(point_ids)


def delete_vectors_by_snapshot(
    repo_id: str,
    snapshot_id: str,
    *,
    client: QdrantClient | None = None,
    collection_name: str | None = None,
) -> None:
    """Delete all Qdrant points for a repo snapshot using payload filters."""
    client = client or get_qdrant_client()
    collection_name = collection_name or settings.QDRANT_COLLECTION_NAME

    snapshot_filter = models.Filter(
        must=[
            models.FieldCondition(
                key="repo_id",
                match=models.MatchValue(value=repo_id),
            ),
            models.FieldCondition(
                key="snapshot_id",
                match=models.MatchValue(value=snapshot_id),
            ),
        ]
    )

    try:
        client.delete(
            collection_name=collection_name,
            points_selector=models.FilterSelector(filter=snapshot_filter),
            wait=True,
        )
    except Exception as exc:
        raise QdrantVectorStoreError(
            f"Qdrant delete by snapshot failed repo_id={repo_id!r} snapshot_id={snapshot_id!r}: {exc}"
        ) from exc

    logger.info(
        "Deleted snapshot vectors from Qdrant collection={collection} repo_id={repo_id} snapshot_id={snapshot_id}",
        collection=collection_name,
        repo_id=repo_id,
        snapshot_id=snapshot_id,
    )


def delete_vectors_by_repo(
    repo_id: str,
    *,
    client: QdrantClient | None = None,
    collection_name: str | None = None,
) -> None:
    """Delete all Qdrant points for a repository using payload filters."""
    client = client or get_qdrant_client()
    collection_name = collection_name or settings.QDRANT_COLLECTION_NAME

    repo_filter = models.Filter(
        must=[
            models.FieldCondition(
                key="repo_id",
                match=models.MatchValue(value=repo_id),
            ),
        ]
    )

    try:
        client.delete(
            collection_name=collection_name,
            points_selector=models.FilterSelector(filter=repo_filter),
            wait=True,
        )
    except Exception as exc:
        raise QdrantVectorStoreError(
            f"Qdrant delete by repo failed repo_id={repo_id!r}: {exc}"
        ) from exc

    logger.info(
        "Deleted repo vectors from Qdrant collection={collection} repo_id={repo_id}",
        collection=collection_name,
        repo_id=repo_id,
    )


def _build_search_filter(
    repo_id: str,
    snapshot_id: str,
    source_types: Sequence[str] | None,
) -> models.Filter:
    must: list[models.Condition] = [
        models.FieldCondition(
            key="repo_id",
            match=models.MatchValue(value=repo_id),
        ),
        models.FieldCondition(
            key="snapshot_id",
            match=models.MatchValue(value=snapshot_id),
        ),
    ]
    if source_types:
        must.append(
            models.FieldCondition(
                key="source_type",
                match=models.MatchAny(any=list(source_types)),
            )
        )
    return models.Filter(must=must)


def _payload_str(payload: dict, key: str) -> str | None:
    value = payload.get(key)
    if value is None:
        return None
    text = str(value).strip()
    return text or None


def _scored_point_to_hit(point: models.ScoredPoint) -> VectorSearchHit:
    payload = point.payload or {}
    record_id = _payload_str(payload, "record_id")
    if not record_id:
        raise QdrantVectorStoreError(
            f"Qdrant hit point_id={point.id!r} is missing payload.record_id"
        )
    source_type = _payload_str(payload, "source_type") or ""
    chunk_type = _payload_str(payload, "chunk_type") or ""
    return VectorSearchHit(
        record_id=record_id,
        score=float(point.score),
        source_type=source_type,
        chunk_type=chunk_type,
        graph_node_id=_payload_str(payload, "graph_node_id"),
        file_path=_payload_str(payload, "file_path"),
        doc_path=_payload_str(payload, "doc_path"),
        symbol_name=_payload_str(payload, "symbol_name"),
        section_title=_payload_str(payload, "section_title"),
        commit_sha=_payload_str(payload, "commit_sha"),
    )


def search_vectors(
    query_vector: Sequence[float],
    *,
    repo_id: str,
    snapshot_id: str,
    source_types: Sequence[str] | None = None,
    top_k: int | None = None,
    score_threshold: float | None = None,
    client: QdrantClient | None = None,
    collection_name: str | None = None,
) -> list[VectorSearchHit]:
    """
    Semantic search in Qdrant scoped to repo/snapshot with optional source_type filter.
    """
    client = client or get_qdrant_client()
    collection_name = collection_name or settings.QDRANT_COLLECTION_NAME
    top_k = top_k if top_k is not None else settings.RETRIEVE_DEFAULT_TOP_K
    score_threshold = (
        score_threshold
        if score_threshold is not None
        else settings.RETRIEVE_DEFAULT_SCORE_THRESHOLD
    )

    if top_k < 1 or top_k > settings.RETRIEVE_MAX_TOP_K:
        raise QdrantVectorStoreError(
            f"top_k must be between 1 and {settings.RETRIEVE_MAX_TOP_K}, got {top_k}"
        )

    validated_vector = _validate_vector(query_vector, "query")
    query_filter = _build_search_filter(repo_id, snapshot_id, source_types)

    try:
        # qdrant-client >= 1.18 removed QdrantClient.search(); use query_points.
        if hasattr(client, "query_points"):
            response = client.query_points(
                collection_name=collection_name,
                query=validated_vector,
                query_filter=query_filter,
                limit=top_k,
                score_threshold=score_threshold,
                with_payload=True,
            )
            results = list(response.points or [])
        else:
            results = client.search(
                collection_name=collection_name,
                query_vector=validated_vector,
                query_filter=query_filter,
                limit=top_k,
                score_threshold=score_threshold,
                with_payload=True,
            )
    except Exception as exc:
        raise QdrantVectorStoreError(f"Qdrant search failed: {exc}") from exc

    hits = [_scored_point_to_hit(point) for point in results]
    logger.info(
        "Qdrant search collection={collection} repo_id={repo_id} snapshot_id={snapshot_id} "
        "hits={hits} top_k={top_k} score_threshold={score_threshold}",
        collection=collection_name,
        repo_id=repo_id,
        snapshot_id=snapshot_id,
        hits=len(hits),
        top_k=top_k,
        score_threshold=score_threshold,
    )
    return hits
