from __future__ import annotations

from app.models.repos import RepoEntry, RepoStages
from app.repositories.repo_repository import RepoRepository


class _FakeCollection:
    def __init__(self, docs):
        self._docs = docs

    def find(self, _query):
        class _Cursor:
            def __init__(self, docs):
                self._docs = list(docs)

            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

        return _Cursor(self._docs)

    async def find_one(self, query, sort=None):
        repo_id = query.get("repo_id")
        matches = [doc for doc in self._docs if doc.get("repo_id") == repo_id]
        if "stages.graph_finalize" in query:
            matches = [
                doc
                for doc in matches
                if (doc.get("stages") or {}).get("graph_finalize") == query["stages.graph_finalize"]
            ]
        if "$or" in query:
            or_clauses = query["$or"]
            filtered = []
            for doc in matches:
                for clause in or_clauses:
                    field, condition = next(iter(clause.items()))
                    value = doc.get(field, 0)
                    op, threshold = next(iter(condition.items()))
                    if op == "$gt" and value > threshold:
                        filtered.append(doc)
                        break
            matches = filtered
        if not matches:
            return None
        if sort:
            field, direction = sort[0]
            matches.sort(key=lambda doc: doc.get(field), reverse=direction < 0)
        return matches[0]


class _FakeDB:
    def __init__(self, repos, indexing_runs):
        self.repositories = _FakeCollection(repos)
        self.indexing_runs = _FakeCollection(indexing_runs)


def test_list_repos_prefers_graph_finalized_snapshot(monkeypatch):
    repos = [{"repo_id": "ad/example", "clone_url": "https://example/repo.git"}]
    indexing_runs = [
        {
            "repo_id": "ad/example",
            "snapshot_id": "snap_empty",
            "status": "completed",
            "created_at": 2,
            "expected_code_files": 0,
            "processed_code_files": 0,
            "stages": {"graph_finalize": "completed", "embedding": None},
        },
        {
            "repo_id": "ad/example",
            "snapshot_id": "snap_good",
            "status": "completed",
            "created_at": 1,
            "expected_code_files": 77,
            "processed_code_files": 77,
            "stages": {"graph_finalize": "completed", "embedding": "completed"},
        },
    ]

    fake_db = _FakeDB(repos, indexing_runs)

    def fake_get_repo_sync_db():
        return fake_db

    monkeypatch.setattr(
        "app.repositories.repo_repository.get_repo_sync_db",
        fake_get_repo_sync_db,
    )

    import asyncio

    entries = asyncio.run(RepoRepository().list_repos())
    assert len(entries) == 1
    assert entries[0].latest_snapshot_id == "snap_good"
    assert entries[0].stages == RepoStages(embedding="completed")
