import asyncio

import pytest

from datetime import datetime

from app.models.embedding_record import EmbeddingRecord
from app.models.embedding_record_upsert import EmbeddingRecordUpsert
from app.models.enums import EmbedStatus, SourceType
from app.repositories import embedding_record_repo


class _FakeBulkResult:
    upserted_count = 1
    modified_count = 0


class _FakeEmbeddingRecords:
    def __init__(self):
        self.docs: dict[str, dict] = {}
        self.indexes: list = []

    async def create_index(self, keys, **kwargs):
        self.indexes.append((keys, kwargs))

    async def bulk_write(self, operations, ordered=False):
        for op in operations:
            record_id = op._filter["record_id"]
            existing = self.docs.get(record_id, {})
            created_at = existing.get("created_at") or op._doc["$setOnInsert"]["created_at"]
            merged = {**existing, **op._doc["$set"], "created_at": created_at}
            self.docs[record_id] = merged
        return _FakeBulkResult()

    async def find_one(self, query):
        if "record_id" in query:
            return self.docs.get(query["record_id"])
        return None

    def find(self, query):
        return _FakeCursor(
            [
                doc
                for doc in self.docs.values()
                if doc.get("repo_id") == query.get("repo_id")
                and doc.get("snapshot_id") == query.get("snapshot_id")
            ]
        )

    async def count_documents(self, query):
        return len(
            [
                doc
                for doc in self.docs.values()
                if doc.get("repo_id") == query.get("repo_id")
                and doc.get("snapshot_id") == query.get("snapshot_id")
            ]
        )

    async def update_one(self, query, update):
        record_id = query["record_id"]
        if record_id not in self.docs:
            return _FakeUpdateResult(matched_count=0)
        self.docs[record_id].update(update["$set"])
        return _FakeUpdateResult(matched_count=1)

    async def delete_many(self, query):
        ids = set(query["record_id"]["$in"])
        deleted = [key for key in list(self.docs) if key in ids]
        for key in deleted:
            del self.docs[key]
        return _FakeDeleteResult(deleted_count=len(deleted))


class _FakeCursor:
    def __init__(self, docs):
        self.docs = docs

    def skip(self, _):
        return self

    def limit(self, _):
        return self

    def sort(self, *_args, **_kwargs):
        return self

    def __aiter__(self):
        self._iter = iter(self.docs)
        return self

    async def __anext__(self):
        try:
            return next(self._iter)
        except StopIteration as exc:
            raise StopAsyncIteration from exc


class _FakeUpdateResult:
    def __init__(self, matched_count):
        self.matched_count = matched_count


class _FakeDeleteResult:
    def __init__(self, deleted_count):
        self.deleted_count = deleted_count


class _FakeDB:
    def __init__(self):
        self.embedding_records = _FakeEmbeddingRecords()


def _sample_record() -> EmbeddingRecord:
    return EmbeddingRecord.from_doc_chunk(
        record_id="emb_chk_docs_1",
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        commit_sha="a1b2c3d4e5f6",
        text="Overview text",
        doc_path="docs/architecture.md",
        section_title="Overview",
        graph_node_id="doc_chunk_1",
    )


@pytest.fixture(autouse=True)
def reset_index_flags():
    embedding_record_repo._INDEXES_READY = False
    yield


def test_upsert_operation_puts_created_at_only_on_insert():
    now = datetime.utcnow()
    record = _sample_record()
    op = embedding_record_repo._upsert_item_to_operation(
        EmbeddingRecordUpsert(record=record),
        now,
    )

    assert "created_at" not in op._doc["$set"]
    assert op._doc["$setOnInsert"]["created_at"] == now


def test_upsert_records_is_idempotent(monkeypatch):
    fake_db = _FakeDB()
    monkeypatch.setattr(embedding_record_repo, "get_db", lambda: fake_db)

    record = _sample_record()
    asyncio.run(embedding_record_repo.upsert_records([record]))
    asyncio.run(embedding_record_repo.upsert_records([record]))

    assert len(fake_db.embedding_records.docs) == 1


def test_list_and_update_record_status(monkeypatch):
    fake_db = _FakeDB()
    monkeypatch.setattr(embedding_record_repo, "get_db", lambda: fake_db)

    record = _sample_record()
    asyncio.run(embedding_record_repo.upsert_records([record]))

    stored = asyncio.run(embedding_record_repo.get_record(record.record_id))
    assert stored is not None
    assert stored.embed_status == EmbedStatus.PENDING
    assert stored.source_type == SourceType.DOCS

    updated = asyncio.run(
        embedding_record_repo.update_embed_status(
            record.record_id,
            EmbedStatus.EMBEDDED,
            qdrant_point_id=record.record_id,
        )
    )
    assert updated is True

    stored = asyncio.run(embedding_record_repo.get_record(record.record_id))
    assert stored.embed_status == EmbedStatus.EMBEDDED
    assert stored.qdrant_point_id == record.record_id
    assert stored.embedded_at is not None


def test_delete_records_by_ids(monkeypatch):
    fake_db = _FakeDB()
    monkeypatch.setattr(embedding_record_repo, "get_db", lambda: fake_db)

    record = _sample_record()
    asyncio.run(embedding_record_repo.upsert_records([record]))
    deleted = asyncio.run(
        embedding_record_repo.delete_records_by_ids([record.record_id, "missing"])
    )

    assert deleted == 1
    assert asyncio.run(embedding_record_repo.get_record(record.record_id)) is None
