import asyncio

import pytest

from app.models.embedding_run import build_run_dedupe_key, build_run_id
from app.models.enums import EmbedStatus
from app.repositories import embedding_run_repo


class _FakeEmbeddingRuns:
    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 find_one_and_update(self, query, update, upsert, return_document):
        dedupe_key = query["dedupe_key"]
        existing = self.docs.get(dedupe_key, {})
        created = existing.get("created_at") or update["$setOnInsert"]["created_at"]
        merged = {
            **update.get("$setOnInsert", {}),
            **existing,
            **update["$set"],
            "created_at": created,
        }
        self.docs[dedupe_key] = merged
        return merged

    async def find_one(self, query):
        if "run_id" in query:
            return next(
                (doc for doc in self.docs.values() if doc.get("run_id") == query["run_id"]),
                None,
            )
        if "dedupe_key" in query:
            return self.docs.get(query["dedupe_key"])
        return None

    async def update_one(self, query, update):
        for doc in self.docs.values():
            if doc.get("run_id") == query.get("run_id"):
                doc.update(update["$set"])
                if "$inc" in update:
                    doc["retry_count"] = doc.get("retry_count", 0) + update["$inc"]["retry_count"]
                return _FakeUpdateResult(matched_count=1)
        return _FakeUpdateResult(matched_count=0)

    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")
            ]
        )


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

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

    def limit(self, _):
        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 _FakeDB:
    def __init__(self):
        self.embedding_runs = _FakeEmbeddingRuns()


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


def test_upsert_run_dedupes_by_event(monkeypatch):
    fake_db = _FakeDB()
    monkeypatch.setattr(embedding_run_repo, "get_db", lambda: fake_db)

    first = asyncio.run(
        embedding_run_repo.upsert_run(
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            event_id="evt_1",
        )
    )
    second = asyncio.run(
        embedding_run_repo.upsert_run(
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            event_id="evt_1",
            celery_task_id="task-123",
        )
    )

    assert first.run_id == second.run_id
    assert len(fake_db.embedding_runs.docs) == 1

    fetched = asyncio.run(embedding_run_repo.get_run(first.run_id))
    assert fetched is not None
    assert fetched.celery_task_id == "task-123"


def test_update_run_state_transitions(monkeypatch):
    fake_db = _FakeDB()
    monkeypatch.setattr(embedding_run_repo, "get_db", lambda: fake_db)

    run = asyncio.run(
        embedding_run_repo.upsert_run(
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            event_id="evt_2",
        )
    )

    asyncio.run(
        embedding_run_repo.update_run_state(
            run.run_id,
            EmbedStatus.PROCESSING,
            total_records=10,
            mark_started=True,
        )
    )
    asyncio.run(
        embedding_run_repo.update_run_state(
            run.run_id,
            EmbedStatus.COMPLETED,
            embedded_records=10,
            mark_completed=True,
        )
    )

    updated = asyncio.run(embedding_run_repo.get_run(run.run_id))
    assert updated is not None
    assert updated.status == EmbedStatus.COMPLETED
    assert updated.total_records == 10
    assert updated.embedded_records == 10
    assert updated.started_at is not None
    assert updated.completed_at is not None


def test_build_run_helpers():
    dedupe = build_run_dedupe_key("AD/example-repo", "snap_abc123", "evt_1")
    run_id = build_run_id(dedupe)
    assert dedupe.endswith(":embedding")
    assert run_id.startswith("erun_")
