import asyncio

import pytest
from qdrant_client.http import models

from app.models.embedding_record import EmbeddingRecord
from app.qdrant import vector_store
from app.retrieval import vector_search
from app.retrieval.models import RetrieveFilters
from app.services import embedding_service


def _sample_record(
    record_id: str = "emb_chk_sym_config_go_Load",
    *,
    source_type: str = "code",
) -> EmbeddingRecord:
    if source_type == "docs":
        return EmbeddingRecord.from_doc_chunk(
            record_id=record_id,
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            commit_sha="a1b2c3d4e5f6",
            text="Overview section text",
            doc_path="docs/architecture.md",
            section_title="Overview",
            graph_node_id="doc_chunk_overview",
        )
    return EmbeddingRecord.from_code_chunk(
        record_id=record_id,
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        commit_sha="a1b2c3d4e5f6",
        text="func Load() {}",
        file_path="internal/config/config.go",
        symbol_name="Load",
        start_line=1,
        end_line=2,
        graph_node_id="sym_config_go_Load",
    )


def _vector(dimension: int = 1536, fill: float = 0.1) -> list[float]:
    return [fill] * dimension


class _FakeQueryResponse:
    def __init__(self, points: list[models.ScoredPoint]) -> None:
        self.points = points


class _FakeQdrantClient:
    def __init__(self):
        self.points: dict[str, dict] = {}
        self.query_points_calls: list[dict] = []

    def upsert(self, collection_name, points, wait=True):
        for point in points:
            self.points[str(point.id)] = {
                "vector": point.vector,
                "payload": point.payload,
            }

    def query_points(
        self,
        collection_name,
        query,
        query_filter=None,
        limit=10,
        score_threshold=None,
        with_payload=True,
        **kwargs,
    ):
        self.query_points_calls.append(
            {
                "collection_name": collection_name,
                "query": query,
                "query_filter": query_filter,
                "limit": limit,
                "score_threshold": score_threshold,
            }
        )
        matches: list[models.ScoredPoint] = []
        for point_id, data in self.points.items():
            payload = data["payload"]
            if not _payload_matches_filter(payload, query_filter):
                continue
            score = 0.9
            if score_threshold is not None and score < score_threshold:
                continue
            matches.append(
                models.ScoredPoint(
                    id=point_id,
                    version=0,
                    score=score,
                    payload=payload,
                )
            )
        matches.sort(key=lambda item: item.score, reverse=True)
        return _FakeQueryResponse(matches[:limit])


def _payload_matches_filter(payload: dict, query_filter: models.Filter | None) -> bool:
    if query_filter is None:
        return True
    for condition in query_filter.must:
        if not isinstance(condition, models.FieldCondition):
            continue
        key = condition.key
        if isinstance(condition.match, models.MatchValue):
            if payload.get(key) != condition.match.value:
                return False
        elif isinstance(condition.match, models.MatchAny):
            if payload.get(key) not in condition.match.any:
                return False
    return True


def test_search_vectors_filters_by_repo_snapshot_and_source_type():
    fake = _FakeQdrantClient()
    code_record = _sample_record("emb_code")
    docs_record = _sample_record("emb_docs", source_type="docs")
    other_snapshot = EmbeddingRecord.from_code_chunk(
        record_id="emb_other_snap",
        repo_id="AD/example-repo",
        snapshot_id="snap_other",
        commit_sha="a1b2c3d4e5f6",
        text="func Other() {}",
        file_path="other.go",
        symbol_name="Other",
        start_line=1,
        end_line=2,
        graph_node_id="sym_other",
    )

    for record in (code_record, docs_record, other_snapshot):
        vector_store.upsert_vectors(
            [vector_store.VectorUpsertItem(record=record, vector=_vector())],
            client=fake,
        )

    hits = vector_store.search_vectors(
        _vector(),
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        source_types=["code"],
        top_k=10,
        score_threshold=0.5,
        client=fake,
    )

    assert len(hits) == 1
    assert hits[0].record_id == "emb_code"
    assert hits[0].source_type == "code"
    assert hits[0].symbol_name == "Load"
    assert fake.query_points_calls[0]["query_filter"].must[2].match.any == ["code"]


def test_search_vectors_rejects_invalid_top_k():
    fake = _FakeQdrantClient()

    with pytest.raises(vector_store.QdrantVectorStoreError, match="top_k"):
        vector_store.search_vectors(
            _vector(),
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            top_k=0,
            client=fake,
        )


def test_search_vectors_rejects_wrong_query_dimension():
    fake = _FakeQdrantClient()

    with pytest.raises(vector_store.QdrantVectorStoreError, match="dimension"):
        vector_store.search_vectors(
            [0.1, 0.2],
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            client=fake,
        )


def test_retrieve_filters_rejects_invalid_source_types():
    from pydantic import ValidationError

    with pytest.raises(ValidationError):
        RetrieveFilters(source_types=["code", "invalid"])


def test_search_by_query_text_embeds_and_searches(monkeypatch):
    fake = _FakeQdrantClient()
    record = _sample_record()
    vector_store.upsert_vectors(
        [vector_store.VectorUpsertItem(record=record, vector=_vector())],
        client=fake,
    )

    class _FakeEmbeddingService:
        async def embed_text(self, text: str) -> list[float]:
            assert text == "where is Load defined?"
            return _vector(fill=0.2)

    embedding_service.reset_embedding_service()
    monkeypatch.setattr(
        vector_search,
        "get_embedding_service",
        lambda: _FakeEmbeddingService(),
    )
    monkeypatch.setattr(vector_store, "get_qdrant_client", lambda: fake)

    result = asyncio.run(
        vector_search.search_by_query_text(
            "where is Load defined?",
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            top_k=5,
            score_threshold=0.5,
        )
    )

    assert result.embedding_model == "text-embedding-3-small"
    assert result.embedding_dimension == 1536
    assert len(result.hits) == 1
    assert result.hits[0].record_id == record.record_id
    assert result.vector_latency_ms >= 0
