import asyncio

import pytest

from app.repositories import indexing_run_repo
from app.repositories.indexing_run_errors import IndexingRunNotReadyError


class _FakeIndexingRuns:
    def __init__(self):
        self.doc = {
            "repo_id": "ad/repo",
            "snapshot_id": "snap_1",
            "processed_docs_files": 0,
            "expected_docs_files": 2,
            "stages": {"docs_parse": "pending"},
        }

    async def find_one(self, query, projection=None):
        if (
            self.doc["repo_id"] == query["repo_id"]
            and self.doc["snapshot_id"] == query["snapshot_id"]
        ):
            return self.doc.copy()
        return None

    async def find_one_and_update(self, query, update, return_document=None):
        if (
            self.doc["repo_id"] != query["repo_id"]
            or self.doc["snapshot_id"] != query["snapshot_id"]
        ):
            return None
        if "$inc" in update:
            self.doc["processed_docs_files"] += update["$inc"]["processed_docs_files"]
        if "$set" in update:
            for key, value in update["$set"].items():
                if key == "stages.docs_parse":
                    self.doc.setdefault("stages", {})["docs_parse"] = value
                else:
                    self.doc[key] = value
        return self.doc.copy()

    async def update_one(self, query, update):
        if "$set" in update:
            for key, value in update["$set"].items():
                if key == "stages.docs_parse":
                    self.doc.setdefault("stages", {})["docs_parse"] = value
                else:
                    self.doc[key] = value


class _FakeClient:
    def __init__(self):
        self.indexing_runs = _FakeIndexingRuns()

    def __getitem__(self, name):
        if name == "adpilot_repo_sync":
            return self
        raise KeyError(name)


def test_increment_processed_docs_updates_stage(monkeypatch):
    fake_client = _FakeClient()
    monkeypatch.setattr(indexing_run_repo, "get_client", lambda: fake_client)
    monkeypatch.setenv("INDEXING_RUNS_ENABLED", "true")

    asyncio.run(indexing_run_repo.increment_processed_docs("ad/repo", "snap_1"))

    assert fake_client.indexing_runs.doc["processed_docs_files"] == 1
    assert fake_client.indexing_runs.doc["stages"]["docs_parse"] == "running"

    asyncio.run(indexing_run_repo.increment_processed_docs("ad/repo", "snap_1"))

    assert fake_client.indexing_runs.doc["processed_docs_files"] == 2
    assert fake_client.indexing_runs.doc["stages"]["docs_parse"] == "completed"


def test_increment_processed_docs_raises_when_run_missing(monkeypatch):
    fake_client = _FakeClient()
    monkeypatch.setattr(indexing_run_repo, "get_client", lambda: fake_client)
    monkeypatch.setenv("INDEXING_RUNS_ENABLED", "true")

    with pytest.raises(IndexingRunNotReadyError):
        asyncio.run(indexing_run_repo.increment_processed_docs("ad/repo", "missing_snap"))
