import asyncio

import pytest
from qdrant_client.http import models

from app.models.embedding_record import EmbeddingRecord
from app.qdrant import point_id, vector_store


def _sample_record(record_id: str = "emb_chk_sym_config_go_Load") -> EmbeddingRecord:
    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 _FakeQdrantClient:
    def __init__(self):
        self.points: dict[str, dict] = {}
        self.upsert_calls: list[list[models.PointStruct]] = []
        self.deleted_point_ids: list[str] = []
        self.deleted_filters: list[models.Filter] = []

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

    def delete(self, collection_name, points_selector, wait=True):
        if isinstance(points_selector, models.PointIdsList):
            for pid in points_selector.points:
                self.deleted_point_ids.append(str(pid))
                self.points.pop(str(pid), None)
            return

        if isinstance(points_selector, models.FilterSelector):
            self.deleted_filters.append(points_selector.filter)
            to_delete = []
            must_keys = {cond.key for cond in points_selector.filter.must}
            for pid, data in self.points.items():
                payload = data["payload"]
                repo_match = any(
                    cond.key == "repo_id"
                    and cond.match.value == payload.get("repo_id")
                    for cond in points_selector.filter.must
                )
                snap_match = True
                if "snapshot_id" in must_keys:
                    snap_match = any(
                        cond.key == "snapshot_id"
                        and cond.match.value == payload.get("snapshot_id")
                        for cond in points_selector.filter.must
                    )
                if repo_match and snap_match:
                    to_delete.append(pid)
            for pid in to_delete:
                del self.points[pid]


def test_upsert_vectors_uses_deterministic_point_ids():
    fake = _FakeQdrantClient()
    record = _sample_record()
    item = vector_store.VectorUpsertItem(record=record, vector=_vector())

    results = vector_store.upsert_vectors([item], client=fake)

    assert len(results) == 1
    assert results[0].record_id == record.record_id
    expected_point_id = point_id.qdrant_point_id(record.record_id)
    assert results[0].point_id == expected_point_id
    assert expected_point_id in fake.points
    assert fake.points[expected_point_id]["payload"]["record_id"] == record.record_id


def test_upsert_vectors_is_idempotent_on_retry():
    fake = _FakeQdrantClient()
    record = _sample_record()
    item = vector_store.VectorUpsertItem(record=record, vector=_vector(fill=0.1))
    item2 = vector_store.VectorUpsertItem(record=record, vector=_vector(fill=0.9))

    vector_store.upsert_vectors([item], client=fake)
    vector_store.upsert_vectors([item2], client=fake)

    assert len(fake.points) == 1
    point_key = point_id.qdrant_point_id(record.record_id)
    assert fake.points[point_key]["vector"][0] == 0.9


def test_upsert_vectors_rejects_wrong_dimension():
    fake = _FakeQdrantClient()
    item = vector_store.VectorUpsertItem(
        record=_sample_record(),
        vector=[0.1, 0.2],
    )

    with pytest.raises(vector_store.QdrantVectorStoreError, match="dimension"):
        vector_store.upsert_vectors([item], client=fake)


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

    deleted = vector_store.delete_vectors_by_record_ids([record.record_id], client=fake)

    assert deleted == 1
    assert point_id.qdrant_point_id(record.record_id) not in fake.points
    assert fake.deleted_point_ids == [point_id.qdrant_point_id(record.record_id)]


def test_delete_vectors_by_snapshot():
    fake = _FakeQdrantClient()
    record_a = _sample_record("emb_a")
    record_b = _sample_record("emb_b")
    record_b_other = EmbeddingRecord.from_code_chunk(
        record_id="emb_other",
        repo_id="AD/other-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",
    )

    vector_store.upsert_vectors(
        [
            vector_store.VectorUpsertItem(record=record_a, vector=_vector(fill=0.1)),
            vector_store.VectorUpsertItem(record=record_b, vector=_vector(fill=0.2)),
            vector_store.VectorUpsertItem(record=record_b_other, vector=_vector(fill=0.3)),
        ],
        client=fake,
    )

    vector_store.delete_vectors_by_snapshot(
        "AD/example-repo",
        "snap_abc123",
        client=fake,
    )

    assert point_id.qdrant_point_id("emb_other") in fake.points
    assert point_id.qdrant_point_id("emb_a") not in fake.points
    assert point_id.qdrant_point_id("emb_b") not in fake.points


def test_delete_vectors_by_repo():
    fake = _FakeQdrantClient()
    record_a = _sample_record("emb_a")
    record_b_other = EmbeddingRecord.from_code_chunk(
        record_id="emb_other",
        repo_id="AD/other-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",
    )

    vector_store.upsert_vectors(
        [
            vector_store.VectorUpsertItem(record=record_a, vector=_vector(fill=0.1)),
            vector_store.VectorUpsertItem(record=record_b_other, vector=_vector(fill=0.3)),
        ],
        client=fake,
    )

    vector_store.delete_vectors_by_repo("AD/example-repo", client=fake)

    assert point_id.qdrant_point_id("emb_other") in fake.points
    assert point_id.qdrant_point_id("emb_a") not in fake.points


def test_upsert_record_vectors_updates_mongo(monkeypatch):
    from app.services import qdrant_index_service

    record = _sample_record()
    item = vector_store.VectorUpsertItem(record=record, vector=_vector())
    updated: list[dict] = []

    async def _update_embed_status(record_id, status, **kwargs):
        updated.append({"record_id": record_id, "status": status, **kwargs})
        return True

    monkeypatch.setattr(
        qdrant_index_service.embedding_record_repo,
        "update_embed_status",
        _update_embed_status,
    )
    monkeypatch.setattr(
        qdrant_index_service,
        "upsert_vectors",
        lambda items, **kwargs: [
            vector_store.VectorUpsertResult(
                record_id=items[0].record.record_id,
                point_id=point_id.qdrant_point_id(items[0].record.record_id),
            )
        ],
    )

    result = asyncio.run(qdrant_index_service.upsert_record_vectors([item]))

    assert result["upserted"] == 1
    assert result["mongo_updated"] == 1
    assert updated[0]["qdrant_point_id"] == point_id.qdrant_point_id(record.record_id)
