import pytest

from app.core.config import Settings
from app.models.enums import EmbedStatus, SourceType
from app.repositories import embedding_record_repo, indexing_run_repo


class _FakeIndexingRuns:
    def __init__(self, doc: dict):
        self.doc = doc

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

    async def update_one(self, query, update):
        if (
            self.doc.get("repo_id") != query.get("repo_id")
            or self.doc.get("snapshot_id") != query.get("snapshot_id")
        ):
            return type("Result", (), {"matched_count": 0, "modified_count": 0})()

        if "$set" in update:
            for key, value in update["$set"].items():
                if key.startswith("stages."):
                    stage_name = key.split(".", 1)[1]
                    self.doc.setdefault("stages", {})[stage_name] = value
                else:
                    self.doc[key] = value
        return type("Result", (), {"matched_count": 1, "modified_count": 1})()


class _FakeIndexingRunsDB:
    def __init__(self, doc: dict):
        self.indexing_runs = _FakeIndexingRuns(doc)


@pytest.fixture
def enabled_settings(monkeypatch):
    settings = Settings(
        _env_file=None,
        INDEXING_RUNS_ENABLED=True,
        INDEXING_RUNS_DATABASE="adpilot_repo_sync",
    )
    monkeypatch.setattr(indexing_run_repo, "settings", settings)
    return settings


@pytest.mark.asyncio
async def test_mark_embedding_running(enabled_settings, monkeypatch):
    doc = {
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "stages": {"commit_intel": "completed"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    await indexing_run_repo.mark_embedding_running("AD/example-repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "running"


@pytest.mark.asyncio
async def test_mark_embedding_running_skips_when_already_completed(enabled_settings, monkeypatch):
    doc = {
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "stages": {"embedding": "completed"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    await indexing_run_repo.mark_embedding_running("AD/example-repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "completed"


@pytest.mark.asyncio
async def test_mark_embedding_failed(enabled_settings, monkeypatch):
    doc = {
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "stages": {"embedding": "running"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    await indexing_run_repo.mark_embedding_failed("AD/example-repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "failed"


@pytest.mark.asyncio
async def test_maybe_mark_embedding_completed_without_commits(enabled_settings, monkeypatch):
    doc = {
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "expected_commits": 0,
        "stages": {"embedding": "running"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    await indexing_run_repo.maybe_mark_embedding_completed("AD/example-repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "completed"


@pytest.mark.asyncio
async def test_maybe_mark_embedding_completed_when_commits_embedded(enabled_settings, monkeypatch):
    doc = {
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "expected_commits": 2,
        "stages": {"embedding": "running"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    async def _count_commits(repo_id, snapshot_id, **kwargs):
        assert kwargs["source_type"] == SourceType.COMMIT
        assert kwargs["embed_status"] == EmbedStatus.EMBEDDED
        return 2

    async def _count_completed(repo_id, snapshot_id):
        return 2

    async def _count_deltas(repo_id, snapshot_id):
        return 2

    monkeypatch.setattr(embedding_record_repo, "count_records_by_snapshot", _count_commits)
    monkeypatch.setattr(indexing_run_repo.commit_fanin_repo, "count_completed_commits", _count_completed)
    monkeypatch.setattr(indexing_run_repo.commit_fanin_repo, "count_graph_deltas", _count_deltas)

    await indexing_run_repo.maybe_mark_embedding_completed("AD/example-repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "completed"


@pytest.mark.asyncio
async def test_maybe_mark_embedding_completed_uses_delta_backed_target(enabled_settings, monkeypatch):
    doc = {
        "repo_id": "AD/example-repo",
        "snapshot_id": "snap_abc123",
        "expected_commits": 581,
        "stages": {"embedding": "running"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    async def _count_commits(repo_id, snapshot_id, **kwargs):
        return 568

    async def _count_completed(repo_id, snapshot_id):
        return 580

    async def _count_deltas(repo_id, snapshot_id):
        return 568

    monkeypatch.setattr(embedding_record_repo, "count_records_by_snapshot", _count_commits)
    monkeypatch.setattr(indexing_run_repo.commit_fanin_repo, "count_completed_commits", _count_completed)
    monkeypatch.setattr(indexing_run_repo.commit_fanin_repo, "count_graph_deltas", _count_deltas)

    await indexing_run_repo.maybe_mark_embedding_completed("AD/example-repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "completed"


@pytest.mark.asyncio
async def test_maybe_mark_embedding_completed_waits_when_commits_pending(
    enabled_settings, monkeypatch
):
    doc = {
        "repo_id": "github:org/repo",
        "snapshot_id": "snap_abc123",
        "expected_commits": 10,
        "stages": {"embedding": "running"},
    }
    fake_db = _FakeIndexingRunsDB(doc)
    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", lambda: fake_db)

    async def _count_commits(repo_id, snapshot_id, **kwargs):
        return 0

    async def _count_completed(repo_id, snapshot_id):
        return 2

    async def _count_deltas(repo_id, snapshot_id):
        return 0

    monkeypatch.setattr(embedding_record_repo, "count_records_by_snapshot", _count_commits)
    monkeypatch.setattr(indexing_run_repo.commit_fanin_repo, "count_completed_commits", _count_completed)
    monkeypatch.setattr(indexing_run_repo.commit_fanin_repo, "count_graph_deltas", _count_deltas)

    await indexing_run_repo.maybe_mark_embedding_completed("github:org/repo", "snap_abc123")

    assert doc["stages"]["embedding"] == "running"


@pytest.mark.asyncio
async def test_disabled_indexing_runs_noop(monkeypatch):
    settings = Settings(_env_file=None, INDEXING_RUNS_ENABLED=False)
    monkeypatch.setattr(indexing_run_repo, "settings", settings)

    called = False

    def _boom():
        nonlocal called
        called = True
        raise AssertionError("should not access Mongo when disabled")

    monkeypatch.setattr(indexing_run_repo, "get_indexing_runs_db", _boom)

    await indexing_run_repo.mark_embedding_running("AD/example-repo", "snap_abc123")
    assert called is False
